mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
fix(openai): 修复 Chat 非流式缓冲读取错误未触发故障转移
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
@@ -422,7 +423,7 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse(
|
||||
|
||||
finalResponse, usage, acc, err := s.readOpenAICompatBufferedTerminal(resp, "openai chat_completions buffered", requestID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil, s.newOpenAICompatBufferedReadFailoverError(c, account, resp, requestID, err)
|
||||
}
|
||||
|
||||
if finalResponse == nil {
|
||||
@@ -518,6 +519,47 @@ func (s *OpenAIGatewayService) handleChatBufferedStreamingResponse(
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *OpenAIGatewayService) newOpenAICompatBufferedReadFailoverError(
|
||||
c *gin.Context,
|
||||
account *Account,
|
||||
resp *http.Response,
|
||||
requestID string,
|
||||
err error,
|
||||
) error {
|
||||
var readErr *openAICompatBufferedReadError
|
||||
if !errors.As(err, &readErr) || readErr == nil || errors.Is(readErr.cause, bufio.ErrTooLong) {
|
||||
return err
|
||||
}
|
||||
var requestContext context.Context
|
||||
if c != nil && c.Request != nil {
|
||||
requestContext = c.Request.Context()
|
||||
}
|
||||
if !shouldClassifyOpenAIUpstreamStreamReadError(readErr.cause, requestContext) {
|
||||
return err
|
||||
}
|
||||
|
||||
classifiedErr := newOpenAIUpstreamStreamReadError(readErr.cause)
|
||||
code, message, ok := OpenAIUpstreamStreamReadErrorDetails(classifiedErr)
|
||||
if !ok {
|
||||
return err
|
||||
}
|
||||
payload, _ := json.Marshal(gin.H{
|
||||
"error": gin.H{
|
||||
"type": "upstream_error",
|
||||
"code": code,
|
||||
"message": message,
|
||||
},
|
||||
})
|
||||
var responseHeaders http.Header
|
||||
if resp != nil {
|
||||
responseHeaders = resp.Header
|
||||
}
|
||||
failoverErr := s.newOpenAIStreamFailoverError(c, account, false, requestID, payload, message, responseHeaders)
|
||||
// 保留稳定错误码,确保重试耗尽后客户端和错误透传规则仍能识别传输故障。
|
||||
failoverErr.ResponseBody = payload
|
||||
return failoverErr
|
||||
}
|
||||
|
||||
// handleChatStreamingResponse reads Responses SSE events from upstream,
|
||||
// converts each to Chat Completions SSE chunks, and writes them to the client.
|
||||
func (s *OpenAIGatewayService) handleChatStreamingResponse(
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tidwall/gjson"
|
||||
)
|
||||
|
||||
type openAICompatBufferedReadErrorCloser struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (r *openAICompatBufferedReadErrorCloser) Read([]byte) (int, error) { return 0, r.err }
|
||||
func (r *openAICompatBufferedReadErrorCloser) Close() error { return nil }
|
||||
|
||||
func TestChatCompletionsBufferedResponsesReadErrorReturnsFailover(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
readErrors := []struct {
|
||||
name string
|
||||
err error
|
||||
expectedCode string
|
||||
}{
|
||||
{name: "unexpected_eof", err: io.ErrUnexpectedEOF, expectedCode: OpenAIUpstreamStreamReadErrorCode},
|
||||
{name: "http2_reset", err: errors.New("stream error: stream ID 7; INTERNAL_ERROR; received from peer"), expectedCode: OpenAIUpstreamHTTP2StreamErrorCode},
|
||||
}
|
||||
|
||||
for _, readError := range readErrors {
|
||||
t.Run(readError.name, func(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
|
||||
resp := &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}, "X-Request-Id": []string{"upstream-rid"}},
|
||||
Body: &openAICompatBufferedReadErrorCloser{err: readError.err},
|
||||
}
|
||||
account := &Account{ID: 40, Name: "openai-oauth", Platform: PlatformOpenAI}
|
||||
|
||||
result, err := (&OpenAIGatewayService{}).handleChatBufferedStreamingResponse(
|
||||
resp, c, account, "gpt-5.6-sol", "gpt-5.6-sol", "gpt-5.6-sol", time.Now(),
|
||||
)
|
||||
|
||||
require.Error(t, err)
|
||||
require.Nil(t, result)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.ErrorAs(t, err, &failoverErr)
|
||||
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
|
||||
require.Equal(t, "upstream-rid", failoverErr.ResponseHeaders.Get("x-request-id"))
|
||||
require.Equal(t, readError.expectedCode, gjson.GetBytes(failoverErr.ResponseBody, "error.code").String())
|
||||
require.Empty(t, rec.Body.String())
|
||||
require.False(t, c.Writer.Written())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestChatCompletionsBufferedResponsesReadErrorDoesNotFailoverAfterClientCancel(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
requestContext, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil).WithContext(requestContext)
|
||||
resp := &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: &openAICompatBufferedReadErrorCloser{err: io.ErrUnexpectedEOF},
|
||||
}
|
||||
|
||||
result, err := (&OpenAIGatewayService{}).handleChatBufferedStreamingResponse(
|
||||
resp,
|
||||
c,
|
||||
&Account{ID: 40, Name: "openai-oauth", Platform: PlatformOpenAI},
|
||||
"gpt-5.6-sol",
|
||||
"gpt-5.6-sol",
|
||||
"gpt-5.6-sol",
|
||||
time.Now(),
|
||||
)
|
||||
|
||||
require.Error(t, err)
|
||||
require.Nil(t, result)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.NotErrorAs(t, err, &failoverErr)
|
||||
require.Empty(t, rec.Body.String())
|
||||
}
|
||||
|
||||
func TestChatCompletionsBufferedResponsesOversizedLineDoesNotFailover(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", nil)
|
||||
resp := &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: &openAICompatBufferedReadErrorCloser{err: bufio.ErrTooLong},
|
||||
}
|
||||
|
||||
result, err := (&OpenAIGatewayService{}).handleChatBufferedStreamingResponse(
|
||||
resp,
|
||||
c,
|
||||
&Account{ID: 40, Name: "openai-oauth", Platform: PlatformOpenAI},
|
||||
"gpt-5.6-sol",
|
||||
"gpt-5.6-sol",
|
||||
"gpt-5.6-sol",
|
||||
time.Now(),
|
||||
)
|
||||
|
||||
require.ErrorIs(t, err, bufio.ErrTooLong)
|
||||
require.Nil(t, result)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.NotErrorAs(t, err, &failoverErr)
|
||||
}
|
||||
|
||||
func TestAnthropicBufferedResponsesReadErrorKeepsExistingBehavior(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
|
||||
resp := &http.Response{
|
||||
StatusCode: http.StatusOK,
|
||||
Header: http.Header{"Content-Type": []string{"text/event-stream"}},
|
||||
Body: &openAICompatBufferedReadErrorCloser{err: io.ErrUnexpectedEOF},
|
||||
}
|
||||
|
||||
result, err := (&OpenAIGatewayService{}).handleAnthropicBufferedStreamingResponse(
|
||||
resp,
|
||||
c,
|
||||
&Account{ID: 40, Name: "openai-oauth", Platform: PlatformOpenAI},
|
||||
"gpt-5.6-sol",
|
||||
"gpt-5.6-sol",
|
||||
"gpt-5.6-sol",
|
||||
time.Now(),
|
||||
)
|
||||
|
||||
require.ErrorIs(t, err, io.ErrUnexpectedEOF)
|
||||
require.Equal(t, io.ErrUnexpectedEOF, err, "Messages 路径必须保持原始读取错误,不引入 Chat failover 包装")
|
||||
require.Nil(t, result)
|
||||
var failoverErr *UpstreamFailoverError
|
||||
require.NotErrorAs(t, err, &failoverErr)
|
||||
}
|
||||
@@ -553,6 +553,10 @@ func (s *OpenAIGatewayService) handleAnthropicBufferedStreamingResponse(
|
||||
|
||||
finalResponse, usage, acc, err := s.readOpenAICompatBufferedTerminal(resp, "openai messages buffered", requestID)
|
||||
if err != nil {
|
||||
var readErr *openAICompatBufferedReadError
|
||||
if errors.As(err, &readErr) && readErr != nil {
|
||||
return nil, readErr.cause
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -675,6 +679,15 @@ func isOpenAICompatDoneSentinelLine(line string) bool {
|
||||
return ok && strings.TrimSpace(payload) == "[DONE]"
|
||||
}
|
||||
|
||||
// openAICompatBufferedReadError 只标记错误发生在上游响应体读取阶段;
|
||||
// 具体端点自行决定是否允许重放请求,避免共享读取器扩大重试范围。
|
||||
type openAICompatBufferedReadError struct {
|
||||
cause error
|
||||
}
|
||||
|
||||
func (e *openAICompatBufferedReadError) Error() string { return e.cause.Error() }
|
||||
func (e *openAICompatBufferedReadError) Unwrap() error { return e.cause }
|
||||
|
||||
func (s *OpenAIGatewayService) readOpenAICompatBufferedTerminal(
|
||||
resp *http.Response,
|
||||
logPrefix string,
|
||||
@@ -783,7 +796,7 @@ func (s *OpenAIGatewayService) readOpenAICompatBufferedTerminal(
|
||||
zap.String("request_id", requestID),
|
||||
)
|
||||
}
|
||||
return nil, usage, acc, ev.err
|
||||
return nil, usage, acc, &openAICompatBufferedReadError{cause: ev.err}
|
||||
}
|
||||
|
||||
if isOpenAICompatDoneSentinelLine(ev.line) {
|
||||
|
||||
Reference in New Issue
Block a user