fix(gateway): fail over on non-JSON 2xx responses

This commit is contained in:
wucm667
2026-06-15 11:04:24 +08:00
parent e34ad2b194
commit ab9987b2e2
2 changed files with 232 additions and 0 deletions
@@ -0,0 +1,177 @@
package service
import (
"bytes"
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/config"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
)
type nonJSONTempUnschedAccountRepo struct {
AccountRepository
tempUnschedCalls int
tempReason string
}
func (r *nonJSONTempUnschedAccountRepo) SetTempUnschedulable(_ context.Context, _ int64, _ time.Time, reason string) error {
r.tempUnschedCalls++
r.tempReason = reason
return nil
}
func TestHandleNonStreamingResponse_NonJSON2xxTriggersFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
body := []byte("(upstream request failed)")
resp := &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{
"Content-Type": []string{"text/plain"},
"X-Request-Id": []string{"rid-invalid-json"},
},
Body: io.NopCloser(bytes.NewReader(body)),
}
svc := &GatewayService{
cfg: &config.Config{},
rateLimitService: &RateLimitService{},
}
usage, err := svc.handleNonStreamingResponse(context.Background(), resp, c, &Account{ID: 1}, "claude-sonnet-4-6", "claude-sonnet-4-6")
require.Nil(t, usage)
var failoverErr *UpstreamFailoverError
require.True(t, errors.As(err, &failoverErr))
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
require.Equal(t, body, failoverErr.ResponseBody)
require.Equal(t, "rid-invalid-json", failoverErr.ResponseHeaders.Get("x-request-id"))
require.False(t, c.Writer.Written(), "invalid upstream response must not be committed before failover")
}
func TestHandleNonStreamingResponse_ValidJSONUnchanged(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
body := []byte(`{"id":"msg_1","type":"message","usage":{"input_tokens":12,"output_tokens":7}}`)
resp := &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(bytes.NewReader(body)),
}
svc := &GatewayService{
cfg: &config.Config{},
rateLimitService: &RateLimitService{},
}
usage, err := svc.handleNonStreamingResponse(context.Background(), resp, c, &Account{ID: 1}, "claude-sonnet-4-6", "claude-sonnet-4-6")
require.NoError(t, err)
require.NotNil(t, usage)
require.Equal(t, 12, usage.InputTokens)
require.Equal(t, 7, usage.OutputTokens)
require.JSONEq(t, string(body), rec.Body.String())
}
func TestHandleNonStreamingResponseAnthropicAPIKeyPassthrough_NonJSON2xxTriggersFailover(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
body := []byte("(upstream request failed)")
resp := &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"text/plain"}},
Body: io.NopCloser(bytes.NewReader(body)),
}
svc := &GatewayService{cfg: &config.Config{}}
usage, err := svc.handleNonStreamingResponseAnthropicAPIKeyPassthrough(context.Background(), resp, c, &Account{ID: 2})
require.Nil(t, usage)
var failoverErr *UpstreamFailoverError
require.True(t, errors.As(err, &failoverErr))
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
require.Equal(t, body, failoverErr.ResponseBody)
require.False(t, c.Writer.Written(), "invalid passthrough response must not be committed before failover")
}
func TestHandleNonStreamingResponseAnthropicAPIKeyPassthrough_ValidJSONUnchanged(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
body := []byte(`{"id":"msg_1","type":"message","usage":{"input_tokens":5,"output_tokens":3}}`)
resp := &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": []string{"application/json"}},
Body: io.NopCloser(bytes.NewReader(body)),
}
svc := &GatewayService{cfg: &config.Config{}}
usage, err := svc.handleNonStreamingResponseAnthropicAPIKeyPassthrough(context.Background(), resp, c, &Account{ID: 2})
require.NoError(t, err)
require.NotNil(t, usage)
require.Equal(t, 5, usage.InputTokens)
require.Equal(t, 3, usage.OutputTokens)
require.JSONEq(t, string(body), rec.Body.String())
}
func TestHandleNonStreamingResponse_NonJSON2xxMatchesTempUnschedulableRule(t *testing.T) {
gin.SetMode(gin.TestMode)
rec := httptest.NewRecorder()
c, _ := gin.CreateTestContext(rec)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/messages", nil)
repo := &nonJSONTempUnschedAccountRepo{}
rateLimitService := NewRateLimitService(repo, nil, &config.Config{}, nil, nil)
svc := &GatewayService{
cfg: &config.Config{},
rateLimitService: rateLimitService,
}
account := &Account{
ID: 3,
Platform: PlatformAnthropic,
Type: AccountTypeAPIKey,
Credentials: map[string]any{
"temp_unschedulable_enabled": true,
"temp_unschedulable_rules": []any{
map[string]any{
"error_code": float64(http.StatusBadGateway),
"keywords": []any{"upstream request failed"},
"duration_minutes": float64(10),
},
},
},
}
body := []byte("(upstream request failed)")
resp := &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{},
Body: io.NopCloser(bytes.NewReader(body)),
}
_, err := svc.handleNonStreamingResponse(context.Background(), resp, c, account, "claude-sonnet-4-6", "claude-sonnet-4-6")
var failoverErr *UpstreamFailoverError
require.True(t, errors.As(err, &failoverErr))
require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode)
require.Equal(t, body, failoverErr.ResponseBody)
require.Equal(t, 1, repo.tempUnschedCalls)
require.Contains(t, repo.tempReason, `"status_code":502`)
require.Contains(t, repo.tempReason, `"matched_keyword":"upstream request failed"`)
}
@@ -5822,6 +5822,51 @@ func parseClaudeUsageFromResponseBody(body []byte) *ClaudeUsage {
return usage
}
func (s *GatewayService) invalidNonStreamingJSONFailoverError(
ctx context.Context,
resp *http.Response,
account *Account,
body []byte,
parseErr error,
requestedModel ...string,
) error {
const statusCode = http.StatusBadGateway
accountID := int64(0)
accountName := ""
retryableOnSameAccount := false
if account != nil {
accountID = account.ID
accountName = account.Name
retryableOnSameAccount = account.IsPoolMode() && account.IsPoolModeRetryableStatus(statusCode)
}
logger.LegacyPrintf(
"service.gateway",
"Account %d(%s): upstream returned non-JSON 2xx response, attempting failover: status=%d request_id=%s error=%v",
accountID,
accountName,
resp.StatusCode,
resp.Header.Get("x-request-id"),
parseErr,
)
if s.rateLimitService != nil && account != nil {
if len(requestedModel) > 0 {
s.rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body, requestedModel[0])
} else {
s.rateLimitService.HandleUpstreamError(ctx, account, statusCode, resp.Header, body)
}
}
return &UpstreamFailoverError{
StatusCode: statusCode,
ResponseBody: body,
ResponseHeaders: resp.Header,
RetryableOnSameAccount: retryableOnSameAccount,
}
}
func (s *GatewayService) handleNonStreamingResponseAnthropicAPIKeyPassthrough(
ctx context.Context,
resp *http.Response,
@@ -5837,6 +5882,13 @@ func (s *GatewayService) handleNonStreamingResponseAnthropicAPIKeyPassthrough(
return nil, err
}
if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices {
var raw json.RawMessage
if err := json.Unmarshal(body, &raw); err != nil {
return nil, s.invalidNonStreamingJSONFailoverError(ctx, resp, account, body, err)
}
}
usage := parseClaudeUsageFromResponseBody(body)
writeAnthropicPassthroughResponseHeaders(c.Writer.Header(), resp.Header, s.responseHeaderFilter)
@@ -8231,6 +8283,9 @@ func (s *GatewayService) handleNonStreamingResponse(ctx context.Context, resp *h
Usage ClaudeUsage `json:"usage"`
}
if err := json.Unmarshal(body, &response); err != nil {
if resp.StatusCode >= http.StatusOK && resp.StatusCode < http.StatusMultipleChoices {
return nil, s.invalidNonStreamingJSONFailoverError(ctx, resp, account, body, err, mappedModel)
}
return nil, fmt.Errorf("parse response: %w", err)
}