From ab9987b2e22d532abdda7091c57afe403ec299f1 Mon Sep 17 00:00:00 2001 From: wucm667 Date: Mon, 15 Jun 2026 11:04:24 +0800 Subject: [PATCH] fix(gateway): fail over on non-JSON 2xx responses --- .../gateway_non_streaming_response_test.go | 177 ++++++++++++++++++ backend/internal/service/gateway_service.go | 55 ++++++ 2 files changed, 232 insertions(+) create mode 100644 backend/internal/service/gateway_non_streaming_response_test.go diff --git a/backend/internal/service/gateway_non_streaming_response_test.go b/backend/internal/service/gateway_non_streaming_response_test.go new file mode 100644 index 0000000000..2416e3e0d0 --- /dev/null +++ b/backend/internal/service/gateway_non_streaming_response_test.go @@ -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"`) +} diff --git a/backend/internal/service/gateway_service.go b/backend/internal/service/gateway_service.go index 82b57f3041..cc0f45c1e0 100644 --- a/backend/internal/service/gateway_service.go +++ b/backend/internal/service/gateway_service.go @@ -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) }