mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-09-21 14:19:18 +08:00
fix(gateway): fail over on non-JSON 2xx responses
This commit is contained in:
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user