feat: support Responses to Chat (#5787)

* fix(openai): harden Chat-to-Responses compatibility

Add a shared Responses-to-Chat stream state machine and use it from the OpenAI relay path. Preserve assistant text alongside tool calls, bind tool argument deltas by output_index, map incomplete finish reasons, support reasoning/custom tool events, and buffer upstream SSE for non-stream Chat clients.

Add deterministic service tests and relay SSE tests for the conversion path.

Related to #5745.

* refactor: rename openaicompat to relayconvert for improved clarity

* feat(gemini): support responses request conversion

* feat: add responses to chat conversion support

* fix: harden responses chat conversion edge cases
This commit is contained in:
Calcium-Ion
2026-06-28 14:25:47 +08:00
committed by GitHub
parent 3a506f50f0
commit 2d5a041639
35 changed files with 2731 additions and 40 deletions
+15 -2
View File
@@ -104,10 +104,18 @@ func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommo
if err != nil {
return nil, err
}
if converter != dto.AdvancedCustomConverterNone {
switch converter {
case dto.AdvancedCustomConverterNone:
return a.convertOpenAICompatibleResponsesRequest(c, info, request)
case dto.AdvancedCustomConverterOpenAIResponsesToOpenAIChatCompletions:
chatReq, err := service.ResponsesRequestToChatCompletionsRequest(&request)
if err != nil {
return nil, err
}
return a.convertOpenAICompatibleRequest(c, info, chatReq)
default:
return nil, fmt.Errorf("converter %q does not support OpenAI Responses requests", converter)
}
return a.convertOpenAICompatibleResponsesRequest(c, info, request)
}
func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.EmbeddingRequest) (any, error) {
@@ -221,6 +229,11 @@ func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycom
return openai.OaiResponsesToChatStreamHandler(c, info, resp)
}
return openai.OaiResponsesToChatHandler(c, info, resp)
case dto.AdvancedCustomConverterOpenAIResponsesToOpenAIChatCompletions:
if info.IsStream {
return openai.OaiChatToResponsesStreamHandler(c, info, resp)
}
return openai.OaiChatToResponsesHandler(c, info, resp)
default:
return nil, types.NewOpenAIError(fmt.Errorf("unsupported advanced custom converter: %s", a.converter), types.ErrorCodeInvalidRequest, http.StatusBadRequest, types.ErrOptionWithSkipRetry())
}
@@ -6,6 +6,7 @@ import (
"net/url"
"testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
relaycommon "github.com/QuantumNous/new-api/relay/common"
@@ -279,6 +280,44 @@ func TestAdaptorMatchesGeminiIncomingPathTemplate(t *testing.T) {
}
}
func TestAdaptorConvertsResponsesRequestToOpenAIChatUpstream(t *testing.T) {
adaptor := &Adaptor{}
info := advancedCustomRelayInfo(&dto.AdvancedCustomConfig{
Routes: []dto.AdvancedCustomRoute{
{
IncomingPath: "/v1/responses",
UpstreamPath: "/v1/chat/completions",
Converter: dto.AdvancedCustomConverterOpenAIResponsesToOpenAIChatCompletions,
},
},
})
info.RelayMode = relayconstant.RelayModeResponses
info.RequestURLPath = "/v1/responses"
c := advancedCustomGinContext("/v1/responses")
converted, err := adaptor.ConvertOpenAIResponsesRequest(c, info, dto.OpenAIResponsesRequest{
Model: "gpt-test",
Instructions: mustAdvancedCustomRawMessage(t, "system rules"),
Input: mustAdvancedCustomRawMessage(t, "hello"),
})
require.NoError(t, err)
chatReq, ok := converted.(*dto.GeneralOpenAIRequest)
require.True(t, ok)
assert.Equal(t, "gpt-test", chatReq.Model)
require.Len(t, chatReq.Messages, 2)
assert.Equal(t, "system", chatReq.Messages[0].Role)
assert.Equal(t, "system rules", chatReq.Messages[0].StringContent())
assert.Equal(t, "user", chatReq.Messages[1].Role)
assert.Equal(t, "hello", chatReq.Messages[1].StringContent())
requestURL, err := adaptor.GetRequestURL(info)
require.NoError(t, err)
parsedURL, err := url.Parse(requestURL)
require.NoError(t, err)
assert.Equal(t, "/v1/chat/completions", parsedURL.Path)
}
func advancedCustomRelayInfo(config *dto.AdvancedCustomConfig) *relaycommon.RelayInfo {
return &relaycommon.RelayInfo{
RelayFormat: types.RelayFormatOpenAI,
@@ -302,3 +341,10 @@ func advancedCustomGinContext(path string) *gin.Context {
c.Request.Header.Set("Content-Type", "application/json")
return c
}
func mustAdvancedCustomRawMessage(t *testing.T, value any) []byte {
t.Helper()
raw, err := common.Marshal(value)
require.NoError(t, err)
return raw
}
+44
View File
@@ -5,6 +5,8 @@ import (
"testing"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/service"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
@@ -12,6 +14,48 @@ func commonPointer[T any](value T) *T {
return &value
}
func TestResponseOpenAI2ClaudeToolUseInputIsObject(t *testing.T) {
tests := []struct {
name string
args string
want map[string]interface{}
}{
{name: "object", args: `{"q":"x"}`, want: map[string]interface{}{"q": "x"}},
{name: "empty", args: "", want: map[string]interface{}{}},
{name: "invalid", args: "{", want: map[string]interface{}{}},
{name: "null", args: "null", want: map[string]interface{}{}},
{name: "array", args: `["x"]`, want: map[string]interface{}{}},
{name: "string", args: `"x"`, want: map[string]interface{}{}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
msg := dto.Message{Role: "assistant"}
msg.SetToolCalls([]dto.ToolCallRequest{
{
ID: "call_1",
Type: "function",
Function: dto.FunctionRequest{
Name: "lookup",
Arguments: tt.args,
},
},
})
resp := service.ResponseOpenAI2Claude(&dto.OpenAITextResponse{
Id: "chatcmpl_1",
Model: "gpt-test",
Choices: []dto.OpenAITextResponseChoice{
{Message: msg, FinishReason: "tool_calls"},
},
}, nil)
require.Len(t, resp.Content, 1)
assert.Equal(t, "tool_use", resp.Content[0].Type)
assert.Equal(t, tt.want, resp.Content[0].Input)
})
}
}
func TestFormatClaudeResponseInfo_MessageStart(t *testing.T) {
claudeInfo := &ClaudeResponseInfo{
Usage: &dto.Usage{},
+19 -2
View File
@@ -12,6 +12,7 @@ import (
"github.com/QuantumNous/new-api/relay/channel/openai"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/constant"
"github.com/QuantumNous/new-api/service/relayconvert"
"github.com/QuantumNous/new-api/setting/model_setting"
"github.com/QuantumNous/new-api/setting/reasoning"
"github.com/QuantumNous/new-api/types"
@@ -238,8 +239,17 @@ func (a *Adaptor) ConvertEmbeddingRequest(c *gin.Context, info *relaycommon.Rela
}
func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error) {
// TODO implement me
return nil, errors.New("not implemented")
request, err := preprocessGeminiOpenAIResponsesRequest(request)
if err != nil {
return nil, err
}
chatRequest, err := relayconvert.ResponsesRequestToChatCompletionsRequest(&request)
if err != nil {
return nil, err
}
return a.ConvertOpenAIRequest(c, info, chatRequest)
}
func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (any, error) {
@@ -247,6 +257,13 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
}
func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
if info.RelayMode == constant.RelayModeResponses {
if info.IsStream {
return GeminiResponsesStreamHandler(c, info, resp)
}
return GeminiResponsesHandler(c, info, resp)
}
if info.RelayMode == constant.RelayModeGemini {
if strings.Contains(info.RequestURLPath, ":embedContent") ||
strings.Contains(info.RequestURLPath, ":batchEmbedContents") {
+99
View File
@@ -0,0 +1,99 @@
package gemini
import (
"strings"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/dto"
)
const (
geminiResponsesInputTypeCustomToolCall = "custom_tool_call"
geminiResponsesInputTypeCustomToolCallOutput = "custom_tool_call_output"
geminiResponsesInputTypeFunctionCallOutput = "function_call_output"
)
func preprocessGeminiOpenAIResponsesRequest(request dto.OpenAIResponsesRequest) (dto.OpenAIResponsesRequest, error) {
tools, err := filterGeminiResponsesTools(request.Tools)
if err != nil {
return request, err
}
request.Tools = tools
input, err := filterGeminiResponsesInput(request.Input)
if err != nil {
return request, err
}
request.Input = input
return request, nil
}
func filterGeminiResponsesTools(raw []byte) ([]byte, error) {
if !geminiRawJSONPresent(raw) || common.GetJsonType(raw) != "array" {
return raw, nil
}
var tools []map[string]any
if err := common.Unmarshal(raw, &tools); err != nil {
return nil, err
}
filtered := make([]map[string]any, 0, len(tools))
for _, tool := range tools {
if strings.TrimSpace(common.Interface2String(tool["type"])) != "function" {
// TODO: Support Responses custom/freeform tools when Gemini has a safe equivalent representation.
continue
}
filtered = append(filtered, tool)
}
if len(filtered) == 0 {
return nil, nil
}
return common.Marshal(filtered)
}
func filterGeminiResponsesInput(raw []byte) ([]byte, error) {
if !geminiRawJSONPresent(raw) || common.GetJsonType(raw) != "array" {
return raw, nil
}
var items []map[string]any
if err := common.Unmarshal(raw, &items); err != nil {
return nil, err
}
skippedCustomCallIDs := make(map[string]struct{})
for _, item := range items {
if strings.TrimSpace(common.Interface2String(item["type"])) != geminiResponsesInputTypeCustomToolCall {
continue
}
if callID := strings.TrimSpace(common.Interface2String(item["call_id"])); callID != "" {
skippedCustomCallIDs[callID] = struct{}{}
}
}
filtered := make([]map[string]any, 0, len(items))
for _, item := range items {
itemType := strings.TrimSpace(common.Interface2String(item["type"]))
switch itemType {
case geminiResponsesInputTypeCustomToolCall, geminiResponsesInputTypeCustomToolCallOutput:
// TODO: Support Responses custom/freeform tool calls once Gemini can preserve their semantics.
continue
case geminiResponsesInputTypeFunctionCallOutput:
if _, ok := skippedCustomCallIDs[strings.TrimSpace(common.Interface2String(item["call_id"]))]; ok {
continue
}
}
filtered = append(filtered, item)
}
return common.Marshal(filtered)
}
func geminiRawJSONPresent(raw []byte) bool {
if len(raw) == 0 {
return false
}
return common.GetJsonType(raw) != "null"
}
@@ -0,0 +1,176 @@
package gemini
import (
"testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/dto"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/tidwall/gjson"
)
func TestConvertOpenAIResponsesRequestToGeminiInstructionsAndInput(t *testing.T) {
got := mustConvertResponsesToGemini(t, dto.OpenAIResponsesRequest{
Model: "gemini-test",
Instructions: mustGeminiRawMessage(t, "system rules"),
Input: mustGeminiRawMessage(t, "hello"),
})
require.NotNil(t, got.SystemInstructions)
require.Len(t, got.SystemInstructions.Parts, 1)
assert.Equal(t, "system rules", got.SystemInstructions.Parts[0].Text)
require.Len(t, got.Contents, 1)
assert.Equal(t, "user", got.Contents[0].Role)
require.Len(t, got.Contents[0].Parts, 1)
assert.Equal(t, "hello", got.Contents[0].Parts[0].Text)
}
func TestConvertOpenAIResponsesRequestToGeminiFunctionToolAndChoice(t *testing.T) {
got := mustConvertResponsesToGemini(t, dto.OpenAIResponsesRequest{
Model: "gemini-test",
Input: mustGeminiRawMessage(t, "lookup weather"),
Tools: mustGeminiRawMessage(t, []map[string]any{
{
"type": "function",
"name": "lookup",
"description": "Lookup data",
"parameters": map[string]any{
"type": "object",
"properties": map[string]any{
"q": map[string]any{"type": "string"},
},
},
},
{"type": "custom", "name": "freeform"},
}),
ToolChoice: mustGeminiRawMessage(t, map[string]any{
"type": "function",
"name": "lookup",
}),
})
tools := got.GetTools()
require.Len(t, tools, 1)
assert.Equal(t, "lookup", gjson.GetBytes(got.Tools, "0.functionDeclarations.0.name").String())
assert.Equal(t, "Lookup data", gjson.GetBytes(got.Tools, "0.functionDeclarations.0.description").String())
require.NotNil(t, got.ToolConfig)
require.NotNil(t, got.ToolConfig.FunctionCallingConfig)
assert.Equal(t, dto.FunctionCallingConfigMode("ANY"), got.ToolConfig.FunctionCallingConfig.Mode)
assert.Equal(t, []string{"lookup"}, got.ToolConfig.FunctionCallingConfig.AllowedFunctionNames)
}
func TestConvertOpenAIResponsesRequestToGeminiFunctionCallConversation(t *testing.T) {
got := mustConvertResponsesToGemini(t, dto.OpenAIResponsesRequest{
Model: "gemini-test",
Input: mustGeminiRawMessage(t, []map[string]any{
{
"role": "assistant",
"content": []map[string]any{
{"type": "output_text", "text": "I will call."},
},
},
{
"type": "function_call",
"call_id": "call_1",
"name": "lookup",
"arguments": map[string]any{"q": "x"},
},
{
"type": "function_call_output",
"call_id": "call_1",
"output": map[string]any{"ok": true},
},
}),
Tools: mustGeminiRawMessage(t, []map[string]any{
{"type": "function", "name": "lookup", "parameters": map[string]any{"type": "object"}},
}),
})
require.Len(t, got.Contents, 2)
assert.Equal(t, "model", got.Contents[0].Role)
require.Len(t, got.Contents[0].Parts, 2)
require.NotNil(t, got.Contents[0].Parts[0].FunctionCall)
assert.Equal(t, "lookup", got.Contents[0].Parts[0].FunctionCall.FunctionName)
assert.Equal(t, map[string]interface{}{"q": "x"}, got.Contents[0].Parts[0].FunctionCall.Arguments)
assert.Equal(t, "I will call.", got.Contents[0].Parts[1].Text)
assert.Equal(t, "user", got.Contents[1].Role)
require.Len(t, got.Contents[1].Parts, 1)
require.NotNil(t, got.Contents[1].Parts[0].FunctionResponse)
assert.Equal(t, "lookup", got.Contents[1].Parts[0].FunctionResponse.Name)
assert.Equal(t, map[string]interface{}{"ok": true}, got.Contents[1].Parts[0].FunctionResponse.Response)
}
func TestConvertOpenAIResponsesRequestToGeminiSkipsCustomToolCalls(t *testing.T) {
got := mustConvertResponsesToGemini(t, dto.OpenAIResponsesRequest{
Model: "gemini-test",
Input: mustGeminiRawMessage(t, []map[string]any{
{
"role": "assistant",
"content": []map[string]any{
{"type": "output_text", "text": "before custom"},
},
},
{
"type": "custom_tool_call",
"call_id": "call_custom",
"name": "apply_patch",
"input": "patch body",
},
{
"type": "custom_tool_call_output",
"call_id": "call_custom",
"output": "ok",
},
{
"type": "function_call_output",
"call_id": "call_custom",
"output": "legacy custom output",
},
{
"role": "user",
"content": "next turn",
},
}),
Tools: mustGeminiRawMessage(t, []map[string]any{
{"type": "custom", "name": "apply_patch"},
{"type": "unknown", "name": "unknown"},
}),
})
assert.Empty(t, got.GetTools())
require.Len(t, got.Contents, 2)
assert.Equal(t, "model", got.Contents[0].Role)
require.Len(t, got.Contents[0].Parts, 1)
assert.Equal(t, "before custom", got.Contents[0].Parts[0].Text)
assert.Nil(t, got.Contents[0].Parts[0].FunctionCall)
assert.Equal(t, "user", got.Contents[1].Role)
require.Len(t, got.Contents[1].Parts, 1)
assert.Equal(t, "next turn", got.Contents[1].Parts[0].Text)
assert.Nil(t, got.Contents[1].Parts[0].FunctionResponse)
}
func mustConvertResponsesToGemini(t *testing.T, req dto.OpenAIResponsesRequest) *dto.GeminiChatRequest {
t.Helper()
info := &relaycommon.RelayInfo{
OriginModelName: req.Model,
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: req.Model,
},
}
got, err := (&Adaptor{}).ConvertOpenAIResponsesRequest(nil, info, req)
require.NoError(t, err)
geminiReq, ok := got.(*dto.GeminiChatRequest)
require.True(t, ok)
return geminiReq
}
func mustGeminiRawMessage(t *testing.T, value any) []byte {
t.Helper()
raw, err := common.Marshal(value)
require.NoError(t, err)
return raw
}
+162
View File
@@ -0,0 +1,162 @@
package gemini
import (
"errors"
"fmt"
"io"
"net/http"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/logger"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/service/relayconvert"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
)
func GeminiResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
defer service.CloseResponseBodyGracefully(resp)
responseBody, err := io.ReadAll(resp.Body)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
logger.LogDebug(c, "Gemini responses response body: %s", responseBody)
var geminiResponse dto.GeminiChatResponse
if err := common.Unmarshal(responseBody, &geminiResponse); err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
if len(geminiResponse.Candidates) == 0 {
usage := buildUsageFromGeminiMetadata(geminiResponse.UsageMetadata, info.GetEstimatePromptTokens())
if geminiResponse.PromptFeedback != nil && geminiResponse.PromptFeedback.BlockReason != nil {
common.SetContextKey(c, constant.ContextKeyAdminRejectReason, fmt.Sprintf("gemini_block_reason=%s", *geminiResponse.PromptFeedback.BlockReason))
return &usage, types.NewOpenAIError(
errors.New("request blocked by Gemini API: "+*geminiResponse.PromptFeedback.BlockReason),
types.ErrorCodePromptBlocked,
http.StatusBadRequest,
)
}
common.SetContextKey(c, constant.ContextKeyAdminRejectReason, "gemini_empty_candidates")
return &usage, types.NewOpenAIError(
errors.New("empty response from Gemini API"),
types.ErrorCodeEmptyResponse,
http.StatusInternalServerError,
)
}
chatResp := responseGeminiChat2OpenAI(c, &geminiResponse)
chatResp.Model = info.UpstreamModelName
usage := buildUsageFromGeminiMetadata(geminiResponse.UsageMetadata, info.GetEstimatePromptTokens())
chatResp.Usage = usage
responsesResp, responsesUsage, err := service.ChatCompletionsResponseToResponsesResponse(chatResp, helper.GetResponseID(c))
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
if responsesUsage == nil || responsesUsage.TotalTokens == 0 {
responsesResp.Usage = relayconvert.UsageFromChatUsage(&usage)
}
responseBody, err = common.Marshal(responsesResp)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError)
}
service.IOCopyBytesGracefully(c, resp, responseBody)
return &usage, nil
}
func GeminiResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
responseID := helper.GetResponseID(c)
created := common.GetTimestamp()
state := relayconvert.NewChatToResponsesStreamState(responseID, info.UpstreamModelName)
state.Created = created
finishReason := constant.FinishReasonStop
toolCallIndexByChoice := make(map[int]map[string]int)
nextToolCallIndexByChoice := make(map[int]int)
var streamErr *types.NewAPIError
sendEvent := func(event relayconvert.ChatToResponsesStreamEvent) bool {
data, err := common.Marshal(event.Payload)
if err != nil {
streamErr = types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError)
return false
}
helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data))
return true
}
sendChunk := func(chunk *dto.ChatCompletionsStreamResponse) bool {
events, err := relayconvert.ChatCompletionsStreamChunkToResponsesEvents(chunk, state)
if err != nil {
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
return false
}
for _, event := range events {
if !sendEvent(event) {
return false
}
}
return true
}
usage, err := geminiStreamHandler(c, info, resp, func(data string, geminiResponse *dto.GeminiChatResponse) bool {
response, isStop := streamResponseGeminiChat2OpenAI(geminiResponse)
response.Id = responseID
response.Created = created
response.Model = info.UpstreamModelName
if response.IsToolCall() {
finishReason = constant.FinishReasonToolCalls
}
for choiceIdx := range response.Choices {
choiceKey := response.Choices[choiceIdx].Index
for toolIdx := range response.Choices[choiceIdx].Delta.ToolCalls {
tool := &response.Choices[choiceIdx].Delta.ToolCalls[toolIdx]
if tool.ID == "" {
continue
}
indexByID := toolCallIndexByChoice[choiceKey]
if indexByID == nil {
indexByID = make(map[string]int)
toolCallIndexByChoice[choiceKey] = indexByID
}
if idx, ok := indexByID[tool.ID]; ok {
tool.SetIndex(idx)
continue
}
idx := nextToolCallIndexByChoice[choiceKey]
nextToolCallIndexByChoice[choiceKey] = idx + 1
indexByID[tool.ID] = idx
tool.SetIndex(idx)
}
}
if !sendChunk(response) {
return false
}
if isStop {
return sendChunk(helper.GenerateStopResponse(responseID, created, info.UpstreamModelName, finishReason))
}
return true
})
if err != nil {
return usage, err
}
if streamErr != nil {
return nil, streamErr
}
if usage != nil {
state.Usage = relayconvert.UsageFromChatUsage(usage)
}
for _, event := range relayconvert.FinalizeChatCompletionsStreamToResponses(state) {
if !sendEvent(event) {
return nil, streamErr
}
}
return usage, nil
}
@@ -0,0 +1,205 @@
package gemini
import (
"bytes"
"errors"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
relaycommon "github.com/QuantumNous/new-api/relay/common"
relayconstant "github.com/QuantumNous/new-api/relay/constant"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestGeminiResponsesHandlerReturnsOpenAIResponsesJSON(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
c.Set(common.RequestIdKey, "gemini-responses-test")
info := newGeminiResponsesRelayInfo(false)
payload := dto.GeminiChatResponse{
Candidates: []dto.GeminiChatCandidate{
{
Content: dto.GeminiChatContent{
Role: "model",
Parts: []dto.GeminiPart{
{Text: "hello"},
},
},
},
},
UsageMetadata: dto.GeminiUsageMetadata{
PromptTokenCount: 2,
CandidatesTokenCount: 3,
TotalTokenCount: 5,
},
}
body, err := common.Marshal(payload)
require.NoError(t, err)
usage, newAPIError := GeminiResponsesHandler(c, info, &http.Response{
Body: io.NopCloser(bytes.NewReader(body)),
})
require.Nil(t, newAPIError)
require.NotNil(t, usage)
assert.Equal(t, 2, usage.PromptTokens)
assert.Equal(t, 3, usage.CompletionTokens)
got := recorder.Body.String()
assert.Contains(t, got, `"object":"response"`)
assert.Contains(t, got, `"status":"completed"`)
assert.Contains(t, got, `"type":"output_text"`)
assert.Contains(t, got, `"text":"hello"`)
assert.Contains(t, got, `"input_tokens":2`)
assert.Contains(t, got, `"output_tokens":3`)
assert.NotContains(t, got, `"choices"`)
assert.NotContains(t, got, `"candidates"`)
}
func TestGeminiResponsesHandlerClosesBodyOnReadError(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
c.Set(common.RequestIdKey, "gemini-responses-read-error-test")
body := &failingReadCloser{}
usage, newAPIError := GeminiResponsesHandler(c, newGeminiResponsesRelayInfo(false), &http.Response{Body: body})
require.Nil(t, usage)
require.NotNil(t, newAPIError)
assert.True(t, body.closed)
}
func TestGeminiResponsesStreamHandlerReturnsOpenAIResponsesSSE(t *testing.T) {
gin.SetMode(gin.TestMode)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
c.Set(common.RequestIdKey, "gemini-responses-stream-test")
oldStreamingTimeout := constant.StreamingTimeout
constant.StreamingTimeout = 300
t.Cleanup(func() { constant.StreamingTimeout = oldStreamingTimeout })
info := newGeminiResponsesRelayInfo(true)
first := dto.GeminiChatResponse{
Candidates: []dto.GeminiChatCandidate{
{
Content: dto.GeminiChatContent{
Role: "model",
Parts: []dto.GeminiPart{
{Text: "hello"},
},
},
},
},
UsageMetadata: dto.GeminiUsageMetadata{
PromptTokenCount: 2,
CandidatesTokenCount: 3,
TotalTokenCount: 5,
},
}
stop := "STOP"
final := dto.GeminiChatResponse{
Candidates: []dto.GeminiChatCandidate{
{
FinishReason: &stop,
Content: dto.GeminiChatContent{
Role: "model",
Parts: []dto.GeminiPart{{Text: ""}},
},
},
},
UsageMetadata: dto.GeminiUsageMetadata{
PromptTokenCount: 2,
CandidatesTokenCount: 3,
TotalTokenCount: 5,
},
}
firstData, err := common.Marshal(first)
require.NoError(t, err)
finalData, err := common.Marshal(final)
require.NoError(t, err)
streamBody := strings.Join([]string{
"data: " + string(firstData),
"",
"data: " + string(finalData),
"",
"data: [DONE]",
"",
}, "\n")
usage, newAPIError := GeminiResponsesStreamHandler(c, info, &http.Response{
Body: io.NopCloser(strings.NewReader(streamBody)),
})
require.Nil(t, newAPIError)
require.NotNil(t, usage)
assert.Equal(t, 5, usage.TotalTokens)
got := recorder.Body.String()
assert.Equal(t, "text/event-stream", recorder.Header().Get("Content-Type"))
assert.Contains(t, got, `event: response.created`)
assert.Contains(t, got, `event: response.output_text.delta`)
assert.Contains(t, got, `"delta":"hello"`)
assert.Contains(t, got, `event: response.completed`)
assert.Contains(t, got, `"input_tokens":2`)
assert.Contains(t, got, `"output_tokens":3`)
assert.NotContains(t, got, `"choices"`)
assert.NotContains(t, got, `"candidates"`)
requireOrderedGeminiResponsesSubstrings(t, got,
`event: response.created`,
`event: response.output_item.added`,
`event: response.output_text.delta`,
`event: response.output_text.done`,
`event: response.completed`,
)
}
func newGeminiResponsesRelayInfo(isStream bool) *relaycommon.RelayInfo {
return &relaycommon.RelayInfo{
IsStream: isStream,
RelayMode: relayconstant.RelayModeResponses,
RelayFormat: types.RelayFormatOpenAIResponses,
RequestURLPath: "/v1/responses",
DisablePing: true,
OriginModelName: "gemini-test",
ChannelMeta: &relaycommon.ChannelMeta{
UpstreamModelName: "gemini-test",
},
}
}
type failingReadCloser struct {
closed bool
}
func (r *failingReadCloser) Read([]byte) (int, error) {
return 0, errors.New("read failed")
}
func (r *failingReadCloser) Close() error {
r.closed = true
return nil
}
func requireOrderedGeminiResponsesSubstrings(t *testing.T, s string, parts ...string) {
t.Helper()
offset := 0
for _, part := range parts {
idx := strings.Index(s[offset:], part)
require.NotEqualf(t, -1, idx, "missing %q after byte offset %d", part, offset)
offset += idx + len(part)
}
}
+5 -5
View File
@@ -14,7 +14,7 @@ import (
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/service/openaicompat"
"github.com/QuantumNous/new-api/service/relayconvert"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
@@ -78,7 +78,7 @@ func OaiResponsesToChatBufferedStreamHandler(c *gin.Context, info *relaycommon.R
}
defer service.CloseResponseBodyGracefully(resp)
accumulator := openaicompat.NewResponsesBufferedAccumulator()
accumulator := relayconvert.NewResponsesBufferedAccumulator()
var finalResponse *dto.OpenAIResponsesResponse
var streamErr *types.NewAPIError
@@ -184,7 +184,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
responseId := helper.GetResponseID(c)
createAt := time.Now().Unix()
state := openaicompat.NewResponsesToChatStreamState(info.UpstreamModelName, false)
state := relayconvert.NewResponsesToChatStreamState(info.UpstreamModelName, false)
state.ID = responseId
state.Created = createAt
streamErr := (*types.NewAPIError)(nil)
@@ -243,7 +243,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
return
}
chunks, err := openaicompat.ResponsesStreamEventToChatChunks(&streamResp, state)
chunks, err := relayconvert.ResponsesStreamEventToChatChunks(&streamResp, state)
if err != nil {
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
sr.Stop(streamErr)
@@ -270,7 +270,7 @@ func OaiResponsesToChatStreamHandler(c *gin.Context, info *relaycommon.RelayInfo
if info.RelayFormat == types.RelayFormatClaude && info.ClaudeConvertInfo != nil {
info.ClaudeConvertInfo.Usage = usage
}
for _, chunk := range openaicompat.FinalizeResponsesToChatStream(state) {
for _, chunk := range relayconvert.FinalizeResponsesToChatStream(state) {
if !sendChatChunk(chunk) {
return nil, streamErr
}
@@ -116,6 +116,58 @@ func TestOaiResponsesToChatBufferedStreamHandlerReturnsJSONFromSSE(t *testing.T)
require.Contains(t, got, `"finish_reason":"tool_calls"`)
}
func TestOaiChatToResponsesStreamHandlerConvertsSSEOrderAndUsage(t *testing.T) {
oldMode := gin.Mode()
gin.SetMode(gin.TestMode)
t.Cleanup(func() { gin.SetMode(oldMode) })
oldTimeout := constant.StreamingTimeout
constant.StreamingTimeout = 30
t.Cleanup(func() { constant.StreamingTimeout = oldTimeout })
body := strings.Join([]string{
`data: {"id":"chatcmpl_1","object":"chat.completion.chunk","created":1710000000,"model":"gpt-test","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}`,
`data: {"id":"chatcmpl_1","object":"chat.completion.chunk","created":1710000000,"model":"gpt-test","choices":[{"index":0,"delta":{"content":"hello"},"finish_reason":null}]}`,
`data: {"id":"chatcmpl_1","object":"chat.completion.chunk","created":1710000000,"model":"gpt-test","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"id":"call_1","type":"function","function":{"name":"lookup"}}]},"finish_reason":null}]}`,
`data: {"id":"chatcmpl_1","object":"chat.completion.chunk","created":1710000000,"model":"gpt-test","choices":[{"index":0,"delta":{"tool_calls":[{"index":0,"function":{"arguments":"{\"q\":\"x\"}"}}]},"finish_reason":null}]}`,
`data: {"id":"chatcmpl_1","object":"chat.completion.chunk","created":1710000000,"model":"gpt-test","choices":[{"index":0,"delta":{},"finish_reason":"tool_calls"}]}`,
`data: {"id":"chatcmpl_1","object":"chat.completion.chunk","created":1710000000,"model":"gpt-test","choices":[],"usage":{"prompt_tokens":2,"completion_tokens":3,"total_tokens":5}}`,
`data: [DONE]`,
``,
}, "\n")
c, recorder, resp, info := newResponsesChatTestContext(t, body, true)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
usage, err := OaiChatToResponsesStreamHandler(c, info, resp)
require.Nil(t, err)
require.NotNil(t, usage)
require.Equal(t, 2, usage.PromptTokens)
require.Equal(t, 3, usage.CompletionTokens)
require.Equal(t, 5, usage.TotalTokens)
got := recorder.Body.String()
require.Equal(t, "text/event-stream", recorder.Header().Get("Content-Type"))
require.Contains(t, got, `event: response.created`)
require.Contains(t, got, `event: response.output_text.delta`)
require.Contains(t, got, `"delta":"hello"`)
require.Contains(t, got, `event: response.function_call_arguments.delta`)
require.Contains(t, got, `"delta":"{\"q\":\"x\"}"`)
require.Contains(t, got, `event: response.completed`)
require.Contains(t, got, `"input_tokens":2`)
require.Contains(t, got, `"output_tokens":3`)
requireOrderedSubstrings(t, got,
`event: response.created`,
`event: response.output_item.added`,
`event: response.output_text.delta`,
`event: response.output_item.added`,
`event: response.function_call_arguments.delta`,
`event: response.output_text.done`,
`event: response.function_call_arguments.done`,
`event: response.completed`,
)
}
func requireOrderedSubstrings(t *testing.T, s string, parts ...string) {
t.Helper()
+131
View File
@@ -0,0 +1,131 @@
package openai
import (
"fmt"
"io"
"net/http"
"github.com/QuantumNous/new-api/common"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/logger"
relaycommon "github.com/QuantumNous/new-api/relay/common"
"github.com/QuantumNous/new-api/relay/helper"
"github.com/QuantumNous/new-api/service"
"github.com/QuantumNous/new-api/service/relayconvert"
"github.com/QuantumNous/new-api/types"
"github.com/gin-gonic/gin"
)
func OaiChatToResponsesHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
if resp == nil || resp.Body == nil {
return nil, types.NewOpenAIError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse, http.StatusInternalServerError)
}
defer service.CloseResponseBodyGracefully(resp)
body, err := io.ReadAll(resp.Body)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeReadResponseBodyFailed, http.StatusInternalServerError)
}
var chatResp dto.OpenAITextResponse
if err := common.Unmarshal(body, &chatResp); err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
if oaiError := chatResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
return nil, types.WithOpenAIError(*oaiError, resp.StatusCode)
}
responseID := helper.GetResponseID(c)
responsesResp, usage, err := service.ChatCompletionsResponseToResponsesResponse(&chatResp, responseID)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeBadResponseBody, http.StatusInternalServerError)
}
if usage == nil || usage.TotalTokens == 0 {
text := service.ExtractOutputTextFromResponses(responsesResp)
usage = service.ResponseText2Usage(c, text, info.UpstreamModelName, info.GetEstimatePromptTokens())
responsesResp.Usage = relayconvert.UsageFromChatUsage(usage)
}
responseBody, err := common.Marshal(responsesResp)
if err != nil {
return nil, types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError)
}
service.IOCopyBytesGracefully(c, resp, responseBody)
return usage, nil
}
func OaiChatToResponsesStreamHandler(c *gin.Context, info *relaycommon.RelayInfo, resp *http.Response) (*dto.Usage, *types.NewAPIError) {
if resp == nil || resp.Body == nil {
return nil, types.NewOpenAIError(fmt.Errorf("invalid response"), types.ErrorCodeBadResponse, http.StatusInternalServerError)
}
defer service.CloseResponseBodyGracefully(resp)
responseID := helper.GetResponseID(c)
state := relayconvert.NewChatToResponsesStreamState(responseID, info.UpstreamModelName)
streamErr := (*types.NewAPIError)(nil)
sendEvent := func(event relayconvert.ChatToResponsesStreamEvent) bool {
data, err := common.Marshal(event.Payload)
if err != nil {
streamErr = types.NewOpenAIError(err, types.ErrorCodeJsonMarshalFailed, http.StatusInternalServerError)
return false
}
helper.ResponseChunkData(c, dto.ResponsesStreamResponse{Type: event.Type}, string(data))
return true
}
helper.StreamScannerHandler(c, resp, info, func(data string, sr *helper.StreamResult) {
if streamErr != nil {
sr.Stop(streamErr)
return
}
var errorResp dto.OpenAITextResponse
if err := common.UnmarshalJsonStr(data, &errorResp); err == nil {
if oaiError := errorResp.GetOpenAIError(); oaiError != nil && oaiError.Type != "" {
streamErr = types.WithOpenAIError(*oaiError, resp.StatusCode)
sr.Stop(streamErr)
return
}
}
var chunk dto.ChatCompletionsStreamResponse
if err := common.UnmarshalJsonStr(data, &chunk); err != nil {
logger.LogError(c, "failed to unmarshal chat stream response: "+err.Error())
sr.Error(err)
return
}
events, err := relayconvert.ChatCompletionsStreamChunkToResponsesEvents(&chunk, state)
if err != nil {
streamErr = types.NewOpenAIError(err, types.ErrorCodeBadResponse, http.StatusInternalServerError)
sr.Stop(streamErr)
return
}
for _, event := range events {
if !sendEvent(event) {
sr.Stop(streamErr)
return
}
}
})
if streamErr != nil {
return nil, streamErr
}
usage := state.Usage
if usage == nil || usage.TotalTokens == 0 {
usage = service.ResponseText2Usage(c, state.UsageText(), info.UpstreamModelName, info.GetEstimatePromptTokens())
state.Usage = relayconvert.UsageFromChatUsage(usage)
}
for _, event := range relayconvert.FinalizeChatCompletionsStreamToResponses(state) {
if !sendEvent(event) {
return nil, streamErr
}
}
return usage, nil
}
+5 -1
View File
@@ -146,7 +146,7 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad
httpResp = resp.(*http.Response)
clientStream := info.IsStream
upstreamStream := strings.HasPrefix(httpResp.Header.Get("Content-Type"), "text/event-stream")
upstreamStream := isResponsesEventStreamContentType(httpResp.Header.Get("Content-Type"))
info.IsStream = clientStream || upstreamStream
if httpResp.StatusCode != http.StatusOK {
newApiErr := service.RelayErrorHandler(c.Request.Context(), httpResp, false)
@@ -179,3 +179,7 @@ func chatCompletionsViaResponses(c *gin.Context, info *relaycommon.RelayInfo, ad
}
return usage, nil
}
func isResponsesEventStreamContentType(contentType string) bool {
return strings.Contains(strings.ToLower(contentType), "text/event-stream")
}
@@ -0,0 +1,26 @@
package relay
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestIsResponsesEventStreamContentType(t *testing.T) {
tests := []struct {
name string
contentType string
want bool
}{
{name: "plain", contentType: "text/event-stream", want: true},
{name: "mixed case with charset", contentType: "Text/Event-Stream; charset=utf-8", want: true},
{name: "json", contentType: "application/json", want: false},
{name: "empty", contentType: "", want: false},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
assert.Equal(t, tt.want, isResponsesEventStreamContentType(tt.contentType))
})
}
}