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
20 changes: 16 additions & 4 deletions controller/relay.go
Original file line number Diff line number Diff line change
Expand Up @@ -75,7 +75,7 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
ws *websocket.Conn
)

if relayFormat == types.RelayFormatOpenAIRealtime {
if relayFormat == types.RelayFormatOpenAIRealtime || (relayFormat == types.RelayFormatOpenAIResponses && c.Request.Method == http.MethodGet) {
var err error
ws, err = upgrader.Upgrade(c.Writer, c.Request, nil)
if err != nil {
Expand All @@ -90,8 +90,14 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
logger.LogError(c, fmt.Sprintf("relay error: %s", newAPIError.Error()))
newAPIError.SetMessage(common.MessageWithRequestId(newAPIError.Error(), requestId))
switch relayFormat {
case types.RelayFormatOpenAIRealtime:
helper.WssError(c, ws, newAPIError.ToOpenAIError())
case types.RelayFormatOpenAIRealtime, types.RelayFormatOpenAIResponses:
if ws != nil {
helper.WssError(c, ws, newAPIError.ToOpenAIError())
} else {
c.JSON(newAPIError.StatusCode, gin.H{
"error": newAPIError.ToOpenAIError(),
})
}
case types.RelayFormatClaude:
c.JSON(newAPIError.StatusCode, gin.H{
"type": "error",
Expand Down Expand Up @@ -211,6 +217,12 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
switch relayFormat {
case types.RelayFormatOpenAIRealtime:
newAPIError = relay.WssHelper(c, relayInfo)
case types.RelayFormatOpenAIResponses:
if relayInfo.ClientWs != nil {
newAPIError = relay.WssResponsesHelper(c, relayInfo)
} else {
newAPIError = relayHandler(c, relayInfo)
}
case types.RelayFormatClaude:
newAPIError = relay.ClaudeHelper(c, relayInfo)
case types.RelayFormatGemini:
Expand Down Expand Up @@ -242,7 +254,7 @@ func Relay(c *gin.Context, relayFormat types.RelayFormat) {
}

var upgrader = websocket.Upgrader{
Subprotocols: []string{"realtime"}, // WS 握手支持的协议,如果有使用 Sec-WebSocket-Protocol,则必须在此声明对应的 Protocol TODO add other protocol
Subprotocols: []string{"realtime", "openai-beta.responses-v1"}, // WS 握手支持的协议
CheckOrigin: func(r *http.Request) bool {
return true // 允许跨域
},
Expand Down
6 changes: 6 additions & 0 deletions middleware/auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -311,6 +311,12 @@ func TokenAuth() func(c *gin.Context) {
c.Request.Header.Set("Authorization", "Bearer "+xGoogKey)
}
}
if strings.HasPrefix(c.Request.URL.Path, "/v1/responses") {
apiKey := c.Query("api_key")
if apiKey != "" {
c.Request.Header.Set("Authorization", "Bearer "+apiKey)
}
}
key := c.Request.Header.Get("Authorization")
parts := make([]string, 0)
if strings.HasPrefix(key, "Bearer ") || strings.HasPrefix(key, "bearer ") {
Expand Down
27 changes: 27 additions & 0 deletions middleware/distributor.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ import (
"github.com/QuantumNous/new-api/constant"
"github.com/QuantumNous/new-api/dto"
"github.com/QuantumNous/new-api/i18n"
"github.com/QuantumNous/new-api/logger"
"github.com/QuantumNous/new-api/model"
relayconstant "github.com/QuantumNous/new-api/relay/constant"
"github.com/QuantumNous/new-api/service"
Expand Down Expand Up @@ -339,6 +340,32 @@ func getModelRequest(c *gin.Context) (*ModelRequest, bool, error) {
if strings.HasPrefix(c.Request.URL.Path, "/v1/responses/compact") && modelRequest.Model != "" {
modelRequest.Model = ratio_setting.WithCompactModelSuffix(modelRequest.Model)
}
if strings.HasPrefix(c.Request.URL.Path, "/v1/responses") && modelRequest.Model == "" {
// logger.LogInfo(c, "DEBUG: headers: " + fmt.Sprintf("%v", c.Request.Header))
modelRequest.Model = c.Query("model")
if modelRequest.Model == "" {
protocol := c.GetHeader("Sec-WebSocket-Protocol")
if protocol != "" {
parts := strings.Split(protocol, ",")
for _, part := range parts {
part = strings.TrimSpace(part)
if strings.HasPrefix(part, "openai-model.") {
modelRequest.Model = strings.TrimPrefix(part, "openai-model.")
break
}
// Fallback: If it's a common model name pattern but not prefixed
if !strings.HasPrefix(part, "openai-") && !strings.Contains(part, "realtime") && strings.Contains(part, "-") {
modelRequest.Model = part
}
}
}
}
// 終極保底:如果真的什麼都找不到,強制使用這個預設值以避免 400 錯誤
if modelRequest.Model == "" {
modelRequest.Model = "gpt-5.3-codex"
logger.LogInfo(c, "WebSocket responses: model completely missing, fallback to gpt-5.3-codex")
}
}
return &modelRequest, shouldSelectChannel, nil
}

Expand Down
53 changes: 51 additions & 2 deletions relay/channel/claude/adaptor.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package claude

import (
"encoding/json"
"errors"
"fmt"
"io"
Expand All @@ -14,6 +15,7 @@ import (
"github.com/QuantumNous/new-api/types"

"github.com/gin-gonic/gin"
"github.com/samber/lo"
)

type Adaptor struct {
Expand Down Expand Up @@ -108,8 +110,55 @@ 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")
// Bridge Responses API to standard OpenAI Chat format, then pass through Claude converter.
oaiReq := &dto.GeneralOpenAIRequest{
Model: request.Model,
Stream: lo.ToPtr(false),
}

if request.MaxOutputTokens != nil {
oaiReq.MaxTokens = request.MaxOutputTokens
}
if request.Temperature != nil {
oaiReq.Temperature = request.Temperature
}
if request.TopP != nil {
oaiReq.TopP = request.TopP
}

// Instructions -> System Message
if len(request.Instructions) > 0 {
var instrStr string
if err := json.Unmarshal(request.Instructions, &instrStr); err == nil && instrStr != "" {
oaiReq.Messages = append(oaiReq.Messages, dto.Message{
Role: "system",
Content: instrStr,
})
}
}

// Input -> User Messages
if len(request.Input) > 0 {
inputs := request.ParseInput()
var contentParts []dto.MediaContent
for _, inp := range inputs {
if inp.Type == "input_text" {
contentParts = append(contentParts, dto.MediaContent{Type: "text", Text: inp.Text})
}
}
if len(contentParts) == 1 {
oaiReq.Messages = append(oaiReq.Messages, dto.Message{
Role: "user",
Content: contentParts[0].Text,
})
} else if len(contentParts) > 1 {
msg := dto.Message{Role: "user"}
msg.SetMediaContent(contentParts)
oaiReq.Messages = append(oaiReq.Messages, msg)
}
}

return a.ConvertOpenAIRequest(c, info, oaiReq)
}

func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (any, error) {
Expand Down
68 changes: 66 additions & 2 deletions relay/channel/gemini/adaptor.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package gemini

import (
"encoding/json"
"errors"
"fmt"
"io"
Expand Down Expand Up @@ -161,6 +162,11 @@ func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) {
}

action := "generateContent"
if info.RelayMode == constant.RelayModeResponses {
// Force non-streaming for Gemini via New API for stability with Responses API.
info.IsStream = false
}

if info.IsStream {
action = "streamGenerateContent?alt=sse"
if info.RelayMode == constant.RelayModeGemini {
Expand Down Expand Up @@ -238,15 +244,73 @@ 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")
// Bridge Responses API to standard OpenAI Chat format
oaiReq := &dto.GeneralOpenAIRequest{
Model: request.Model,
Stream: lo.ToPtr(false),
}

if request.MaxOutputTokens != nil {
oaiReq.MaxTokens = request.MaxOutputTokens
}
if request.Temperature != nil {
oaiReq.Temperature = request.Temperature
}
if request.TopP != nil {
oaiReq.TopP = request.TopP
}

// Convert instructions into a system message
if len(request.Instructions) > 0 {
var instrStr string
if err := json.Unmarshal(request.Instructions, &instrStr); err == nil && instrStr != "" {
oaiReq.Messages = append(oaiReq.Messages, dto.Message{
Role: "system",
Content: instrStr,
})
}
}

// Convert input into user messages
if len(request.Input) > 0 {
inputs := request.ParseInput()
var contentParts []dto.MediaContent
for _, inp := range inputs {
switch inp.Type {
case "input_text":
contentParts = append(contentParts, dto.MediaContent{Type: "text", Text: inp.Text})
case "input_image":
contentParts = append(contentParts, dto.MediaContent{
Type: "image_url",
ImageUrl: &dto.MessageImageUrl{Url: inp.ImageUrl},
})
}
}
if len(contentParts) == 1 && contentParts[0].Type == "text" {
oaiReq.Messages = append(oaiReq.Messages, dto.Message{
Role: "user",
Content: contentParts[0].Text,
})
} else if len(contentParts) > 0 {
msg := dto.Message{Role: "user"}
msg.SetMediaContent(contentParts)
oaiReq.Messages = append(oaiReq.Messages, msg)
}
}

return a.ConvertOpenAIRequest(c, info, oaiReq)
}

func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, requestBody io.Reader) (any, error) {
return channel.DoApiRequest(a, c, info, requestBody)
}

func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
// Force non-streaming for Responses API stability
if info.RelayMode == constant.RelayModeResponses {
info.IsStream = false
}

if info.RelayMode == constant.RelayModeGemini {
if strings.Contains(info.RequestURLPath, ":embedContent") ||
strings.Contains(info.RequestURLPath, ":batchEmbedContents") {
Expand Down
65 changes: 55 additions & 10 deletions relay/channel/zhipu/adaptor.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package zhipu

import (
"encoding/json"
"errors"
"fmt"
"io"
Expand Down Expand Up @@ -43,10 +44,9 @@ func (a *Adaptor) Init(info *relaycommon.RelayInfo) {
}

func (a *Adaptor) GetRequestURL(info *relaycommon.RelayInfo) (string, error) {
// Force non-streaming for GLM via New API for stability with Responses API.
info.IsStream = false
method := "invoke"
if info.IsStream {
method = "sse-invoke"
}
return fmt.Sprintf("%s/api/paas/v3/model-api/%s/%s", info.ChannelBaseUrl, info.UpstreamModelName, method), nil
}

Expand Down Expand Up @@ -81,16 +81,61 @@ func (a *Adaptor) DoRequest(c *gin.Context, info *relaycommon.RelayInfo, request
}

func (a *Adaptor) ConvertOpenAIResponsesRequest(c *gin.Context, info *relaycommon.RelayInfo, request dto.OpenAIResponsesRequest) (any, error) {
// TODO implement me
return nil, errors.New("not implemented")
// Bridge Responses API to standard OpenAI Chat format
oaiReq := &dto.GeneralOpenAIRequest{
Model: request.Model,
Stream: lo.ToPtr(false),
}

if request.MaxOutputTokens != nil {
oaiReq.MaxTokens = request.MaxOutputTokens
}
if request.Temperature != nil {
oaiReq.Temperature = request.Temperature
}
if request.TopP != nil {
oaiReq.TopP = request.TopP
}

// Instructions -> System Message
if len(request.Instructions) > 0 {
var instrStr string
if err := json.Unmarshal(request.Instructions, &instrStr); err == nil && instrStr != "" {
oaiReq.Messages = append(oaiReq.Messages, dto.Message{
Role: "system",
Content: instrStr,
})
}
}

// Input -> User Messages
if len(request.Input) > 0 {
inputs := request.ParseInput()
var contentParts []dto.MediaContent
for _, inp := range inputs {
if inp.Type == "input_text" {
contentParts = append(contentParts, dto.MediaContent{Type: "text", Text: inp.Text})
}
}
if len(contentParts) == 1 {
oaiReq.Messages = append(oaiReq.Messages, dto.Message{
Role: "user",
Content: contentParts[0].Text,
})
} else if len(contentParts) > 1 {
msg := dto.Message{Role: "user"}
msg.SetMediaContent(contentParts)
oaiReq.Messages = append(oaiReq.Messages, msg)
}
}

return a.ConvertOpenAIRequest(c, info, oaiReq)
}

func (a *Adaptor) DoResponse(c *gin.Context, resp *http.Response, info *relaycommon.RelayInfo) (usage any, err *types.NewAPIError) {
if info.IsStream {
usage, err = zhipuStreamHandler(c, info, resp)
} else {
usage, err = zhipuHandler(c, info, resp)
}
// Force non-streaming handler
info.IsStream = false
usage, err = zhipuHandler(c, info, resp)
return
}

Expand Down
Loading