diff --git a/controller/relay.go b/controller/relay.go index c97ab45b4ac4..0569ccd0479d 100644 --- a/controller/relay.go +++ b/controller/relay.go @@ -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 { @@ -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", @@ -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: @@ -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 // 允许跨域 }, diff --git a/middleware/auth.go b/middleware/auth.go index 23d933fbe0c1..2d679cef26e7 100644 --- a/middleware/auth.go +++ b/middleware/auth.go @@ -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 ") { diff --git a/middleware/distributor.go b/middleware/distributor.go index d626941456c7..eac33e986461 100644 --- a/middleware/distributor.go +++ b/middleware/distributor.go @@ -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" @@ -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 } diff --git a/relay/channel/claude/adaptor.go b/relay/channel/claude/adaptor.go index 6daf5b6f245e..5e0557419576 100644 --- a/relay/channel/claude/adaptor.go +++ b/relay/channel/claude/adaptor.go @@ -1,6 +1,7 @@ package claude import ( + "encoding/json" "errors" "fmt" "io" @@ -14,6 +15,7 @@ import ( "github.com/QuantumNous/new-api/types" "github.com/gin-gonic/gin" + "github.com/samber/lo" ) type Adaptor struct { @@ -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) { diff --git a/relay/channel/gemini/adaptor.go b/relay/channel/gemini/adaptor.go index 680c4ee484ec..69e7b8c3179d 100644 --- a/relay/channel/gemini/adaptor.go +++ b/relay/channel/gemini/adaptor.go @@ -1,6 +1,7 @@ package gemini import ( + "encoding/json" "errors" "fmt" "io" @@ -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 { @@ -238,8 +244,61 @@ 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) { @@ -247,6 +306,11 @@ 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) { + // 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") { diff --git a/relay/channel/zhipu/adaptor.go b/relay/channel/zhipu/adaptor.go index 3ed4b3596112..4a058a819b8d 100644 --- a/relay/channel/zhipu/adaptor.go +++ b/relay/channel/zhipu/adaptor.go @@ -1,6 +1,7 @@ package zhipu import ( + "encoding/json" "errors" "fmt" "io" @@ -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 } @@ -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 } diff --git a/relay/channel/zhipu_4v/adaptor.go b/relay/channel/zhipu_4v/adaptor.go index 0af8a16bbeef..bb2baee42955 100644 --- a/relay/channel/zhipu_4v/adaptor.go +++ b/relay/channel/zhipu_4v/adaptor.go @@ -1,6 +1,7 @@ package zhipu_4v import ( + "encoding/json" "errors" "fmt" "io" @@ -44,6 +45,8 @@ 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 baseURL := info.ChannelBaseUrl if baseURL == "" { baseURL = channelconstant.ChannelBaseURLs[channelconstant.ChannelTypeZhipu_v4] @@ -102,8 +105,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 + 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) { diff --git a/relay/helper/valid_request.go b/relay/helper/valid_request.go index c5477ccead65..0735b8e731ad 100644 --- a/relay/helper/valid_request.go +++ b/relay/helper/valid_request.go @@ -34,7 +34,13 @@ func GetAndValidateRequest(c *gin.Context, format types.RelayFormat) (request dt case types.RelayFormatClaude: request, err = GetAndValidateClaudeRequest(c) case types.RelayFormatOpenAIResponses: - request, err = GetAndValidateResponsesRequest(c) + if strings.EqualFold(c.Request.Method, "GET") && strings.EqualFold(c.GetHeader("Upgrade"), "websocket") { + request = &dto.OpenAIResponsesRequest{ + Model: c.Query("model"), + } + } else { + request, err = GetAndValidateResponsesRequest(c) + } case types.RelayFormatOpenAIResponsesCompaction: request, err = GetAndValidateResponsesCompactionRequest(c) diff --git a/relay/wss_responses.go b/relay/wss_responses.go new file mode 100644 index 000000000000..135136e21ed3 --- /dev/null +++ b/relay/wss_responses.go @@ -0,0 +1,430 @@ +package relay + +import ( + "bytes" + "fmt" + "io" + "net/http" + "time" + + "github.com/QuantumNous/new-api/common" + "github.com/QuantumNous/new-api/dto" + relaycommon "github.com/QuantumNous/new-api/relay/common" + "github.com/QuantumNous/new-api/service" + "github.com/QuantumNous/new-api/types" + + "github.com/gin-gonic/gin" + "github.com/gorilla/websocket" +) + +type CaptureResponseWriter struct { + gin.ResponseWriter + Body *bytes.Buffer +} + +func (w *CaptureResponseWriter) Write(b []byte) (int, error) { + return w.Body.Write(b) +} + +func (w *CaptureResponseWriter) WriteString(s string) (int, error) { + return w.Body.WriteString(s) +} + +func (w *CaptureResponseWriter) WriteHeader(statusCode int) { + // Do nothing, prevent panic on hijacked connections +} + +func (w *CaptureResponseWriter) WriteHeaderNow() { + // Do nothing +} + +func (w *CaptureResponseWriter) Flush() { + // Do nothing, prevent panic on hijacked connections +} + +func WssResponsesHelper(c *gin.Context, info *relaycommon.RelayInfo) (newAPIError *types.NewAPIError) { + if info.ClientWs == nil { + return types.NewError(fmt.Errorf("websocket connection is nil"), types.ErrorCodeGetChannelFailed, types.ErrOptionWithSkipRetry()) + } + + info.InitChannelMeta(c) + + // Set a 5-minute timeout as requested by user + timeout := 5 * time.Minute + info.ClientWs.SetReadDeadline(time.Now().Add(timeout)) + info.ClientWs.SetWriteDeadline(time.Now().Add(timeout)) + + // Default Ping handler in Gorilla responds with a Pong. + info.ClientWs.SetPingHandler(func(appData string) error { + info.ClientWs.SetReadDeadline(time.Now().Add(timeout)) + info.ClientWs.SetWriteDeadline(time.Now().Add(timeout)) + return info.ClientWs.WriteMessage(websocket.PongMessage, []byte(appData)) + }) + + responseID := "resp_" + common.GetUUID() + now := common.GetTimestamp() + seqCounter := 0 + responseOpened := false + var responsesReq dto.OpenAIResponsesRequest + + defer func() { + if newAPIError != nil && info.ClientWs != nil { + if !responseOpened { + // Ensure response.created is sent before response.failed + _ = sendWsResponseEvent(info.ClientWs, seqCounter, "response.created", gin.H{ + "response": gin.H{ + "id": responseID, + "object": "response", + "created_at": now, + "status": "in_progress", + "model": responsesReq.Model, + }, + }) + seqCounter++ + responseOpened = true + } + _ = sendWsResponseEvent(info.ClientWs, seqCounter, "response.failed", gin.H{ + "response": gin.H{ + "id": responseID, + "status": "failed", + "error": newAPIError.ToOpenAIError(), + }, + }) + seqCounter++ + } + }() + + // 1. Read the first message from WebSocket + var message []byte + var err error + for { + _, message, err = info.ClientWs.ReadMessage() + if err != nil { + return types.NewError(err, types.ErrorCodeReadRequestBodyFailed) + } + if len(message) > 0 { + break + } + } + + if err := common.Unmarshal(message, &responsesReq); err != nil { + return types.NewError(err, types.ErrorCodeInvalidRequest) + } + + // Model consistency check + if responsesReq.Model == "" { + responsesReq.Model = info.OriginModelName + } + if responsesReq.Model != info.OriginModelName { + return types.NewError(fmt.Errorf("model mismatch: expected %s, got %s", info.OriginModelName, responsesReq.Model), types.ErrorCodeInvalidRequest, types.ErrOptionWithSkipRetry()) + } + + // Required field validation + if len(responsesReq.Input) == 0 || string(responsesReq.Input) == "null" { + return types.NewError(fmt.Errorf("input is required"), types.ErrorCodeInvalidRequest) + } + + info.Request = &responsesReq + + // 2. Setup adaptor + adaptor := GetAdaptor(info.ApiType) + if adaptor == nil { + return types.NewError(fmt.Errorf("invalid api type: %d", info.ApiType), types.ErrorCodeInvalidApiType, types.ErrOptionWithSkipRetry()) + } + adaptor.Init(info) + + // Bridging + request, err := common.DeepCopy(&responsesReq) + if err != nil { + return types.NewError(err, types.ErrorCodeInvalidRequest) + } + + // Force non-streaming for the interceptor to work correctly with synchronous capture + falseVal := false + request.Stream = &falseVal + + convertedRequest, err := adaptor.ConvertOpenAIResponsesRequest(c, info, *request) + if err != nil { + return types.NewError(err, types.ErrorCodeConvertRequestFailed) + } + + jsonData, err := common.Marshal(convertedRequest) + if err != nil { + return types.NewError(err, types.ErrorCodeConvertRequestFailed) + } + + // 3. Send initial events + // event 0: response.created + respCreatedData := gin.H{ + "id": responseID, + "object": "response", + "created_at": now, + "status": "in_progress", + "background": false, + "model": responsesReq.Model, + } + if responsesReq.Temperature != nil { + respCreatedData["temperature"] = *responsesReq.Temperature + } + if responsesReq.TopP != nil { + respCreatedData["top_p"] = *responsesReq.TopP + } + + if err := sendWsResponseEvent(info.ClientWs, seqCounter, "response.created", gin.H{ + "response": respCreatedData, + }); err != nil { + return types.NewError(err, types.ErrorCodeWssWriteFailed) + } + seqCounter++ + responseOpened = true + + // event 1: response.in_progress + if err := sendWsResponseEvent(info.ClientWs, seqCounter, "response.in_progress", gin.H{ + "response": gin.H{ + "id": responseID, + "object": "response", + "created_at": now, + "status": "in_progress", + "model": responsesReq.Model, + }, + }); err != nil { + return types.NewError(err, types.ErrorCodeWssWriteFailed) + } + seqCounter++ + + // 4. Perform the actual HTTP request + resp, err := adaptor.DoRequest(c, info, io.NopCloser(bytes.NewReader(jsonData))) + if err != nil { + return types.NewOpenAIError(err, types.ErrorCodeDoRequestFailed, http.StatusInternalServerError) + } + httpResp, ok := resp.(*http.Response) + if !ok { + return types.NewError(fmt.Errorf("invalid response type from adaptor"), types.ErrorCodeDoRequestFailed) + } + defer service.CloseResponseBodyGracefully(httpResp) + + if httpResp.StatusCode != http.StatusOK { + return service.RelayErrorHandler(c.Request.Context(), httpResp, false) + } + + // Capture output + capture := &CaptureResponseWriter{ + ResponseWriter: c.Writer, + Body: bytes.NewBuffer(nil), + } + + // Ensure the adaptor knows this is NOT a streaming response for the capture to work + info.IsStream = false + + // Backup and restore context fields + originalWriter := c.Writer + originalMethod := c.Request.Method + defer func() { + c.Writer = originalWriter + c.Request.Method = originalMethod + }() + + c.Writer = capture + c.Request.Method = "POST" + + usage, newAPIError := adaptor.DoResponse(c, httpResp, info) + + // Usage handling & mandatory quota settlement + u, _ := usage.(*dto.Usage) + // Force quota consumption (pass u directly, even if nil, to allow service layer fallback) + service.PostTextConsumeQuota(c, info, u, nil) + + // Fallback for frontend display only + if u == nil { + u = &dto.Usage{ + PromptTokens: info.GetEstimatePromptTokens(), + CompletionTokens: 0, + TotalTokens: info.GetEstimatePromptTokens(), + } + } + + if newAPIError != nil { + return newAPIError + } + + usageData := gin.H{ + "total_tokens": u.TotalTokens, + "input_tokens": u.PromptTokens, + "prompt_tokens": u.PromptTokens, + "output_tokens": u.CompletionTokens, + "completion_tokens": u.CompletionTokens, + "input_tokens_details": map[string]interface{}{ + "cached_tokens": 0, + "text_tokens": u.PromptTokens, + "audio_tokens": 0, + "image_tokens": 0, + }, + "output_tokens_details": map[string]interface{}{ + "reasoning_tokens": 0, + "text_tokens": u.CompletionTokens, + "audio_tokens": 0, + }, + } + + // 5. Build and send full sequence of events to ensure client compatibility + if capture.Body.Len() > 0 { + var chatResp dto.OpenAITextResponse + if err := common.Unmarshal(capture.Body.Bytes(), &chatResp); err == nil && len(chatResp.Choices) > 0 { + content := chatResp.Choices[0].Message.StringContent() + itemID := "item_" + common.GetUUID() + + // response.output_item.added (seq 2) + if err := sendWsResponseEvent(info.ClientWs, seqCounter, "response.output_item.added", gin.H{ + "output_index": 0, + "item": gin.H{ + "id": itemID, + "object": "realtime.item", + "type": "message", + "status": "in_progress", + "role": "assistant", + "content": []any{}, + }, + }); err != nil { + return types.NewError(err, types.ErrorCodeWssWriteFailed) + } + seqCounter++ + + // response.content_part.added (seq 3) + if err := sendWsResponseEvent(info.ClientWs, seqCounter, "response.content_part.added", gin.H{ + "item_id": itemID, + "output_index": 0, + "part": gin.H{ + "type": "output_text", + "text": "", + }, + }); err != nil { + return types.NewError(err, types.ErrorCodeWssWriteFailed) + } + seqCounter++ + + // response.output_text.delta (seq 4) + if err := sendWsResponseEvent(info.ClientWs, seqCounter, "response.output_text.delta", gin.H{ + "content_index": 0, + "item_id": itemID, + "output_index": 0, + "delta": content, + }); err != nil { + return types.NewError(err, types.ErrorCodeWssWriteFailed) + } + seqCounter++ + + // response.output_text.done (seq 5) + if err := sendWsResponseEvent(info.ClientWs, seqCounter, "response.output_text.done", gin.H{ + "content_index": 0, + "item_id": itemID, + "output_index": 0, + "text": content, + }); err != nil { + return types.NewError(err, types.ErrorCodeWssWriteFailed) + } + seqCounter++ + + // response.content_part.done (seq 6) + if err := sendWsResponseEvent(info.ClientWs, seqCounter, "response.content_part.done", gin.H{ + "item_id": itemID, + "output_index": 0, + "part": gin.H{ + "type": "output_text", + "text": content, + }, + }); err != nil { + return types.NewError(err, types.ErrorCodeWssWriteFailed) + } + seqCounter++ + + // response.output_item.done (seq 7) + if err := sendWsResponseEvent(info.ClientWs, seqCounter, "response.output_item.done", gin.H{ + "output_index": 0, + "item": gin.H{ + "id": itemID, + "object": "realtime.item", + "type": "message", + "status": "completed", + "role": "assistant", + "content": []any{ + gin.H{ + "type": "output_text", + "text": content, + }, + }, + }, + }); err != nil { + return types.NewError(err, types.ErrorCodeWssWriteFailed) + } + seqCounter++ + + // response.completed (seq 8) + if err := sendWsResponseEvent(info.ClientWs, seqCounter, "response.completed", gin.H{ + "response": gin.H{ + "id": responseID, + "object": "response", + "created_at": now, + "status": "completed", + "completed_at": now, + "model": responsesReq.Model, + "output": []any{ + gin.H{ + "id": itemID, + "status": "completed", + "usage": usageData, + }, + }, + "usage": usageData, + }, + }); err != nil { + return types.NewError(err, types.ErrorCodeWssWriteFailed) + } + seqCounter++ + } else { + // Fallback for empty or unmarshalable content + _ = sendWsResponseEvent(info.ClientWs, seqCounter, "response.completed", gin.H{ + "response": gin.H{ + "id": responseID, + "object": "response", + "created_at": now, + "status": "completed", + "completed_at": now, + "model": responsesReq.Model, + "output": []any{}, + "usage": usageData, + }, + }) + seqCounter++ + } + } else { + // Terminal event for empty response + _ = sendWsResponseEvent(info.ClientWs, seqCounter, "response.completed", gin.H{ + "response": gin.H{ + "id": responseID, + "object": "response", + "created_at": now, + "status": "completed", + "completed_at": now, + "model": responsesReq.Model, + "output": []any{}, + "usage": usageData, + }, + }) + seqCounter++ + } + + // service.PostTextConsumeQuota moved up to enforce settlement + return nil +} + +func sendWsResponseEvent(ws *websocket.Conn, seq int, eventType string, data gin.H) error { + msg := gin.H{ + "type": eventType, + "sequence_number": seq, + } + for k, v := range data { + msg[k] = v + } + _ = ws.SetWriteDeadline(time.Now().Add(10 * time.Second)) + return ws.WriteJSON(msg) +} diff --git a/router/relay-router.go b/router/relay-router.go index 17a13cad7fd6..dc17bf5a1546 100644 --- a/router/relay-router.go +++ b/router/relay-router.go @@ -78,6 +78,9 @@ func SetRelayRouter(router *gin.Engine) { wsRouter.GET("/realtime", func(c *gin.Context) { controller.Relay(c, types.RelayFormatOpenAIRealtime) }) + wsRouter.GET("/responses", func(c *gin.Context) { + controller.Relay(c, types.RelayFormatOpenAIResponses) + }) } { //http router diff --git a/types/error.go b/types/error.go index 9717401ae7b2..311e5ca77a47 100644 --- a/types/error.go +++ b/types/error.go @@ -50,6 +50,7 @@ const ( ErrorCodeDoRequestFailed ErrorCode = "do_request_failed" ErrorCodeGetChannelFailed ErrorCode = "get_channel_failed" ErrorCodeGenRelayInfoFailed ErrorCode = "gen_relay_info_failed" + ErrorCodeWssWriteFailed ErrorCode = "wss_write_failed" // channel error ErrorCodeChannelNoAvailableKey ErrorCode = "channel:no_available_key"