From 71c16b7230c41358e95ea624f57cb21a20dafda4 Mon Sep 17 00:00:00 2001 From: zy6p Date: Wed, 1 Jul 2026 17:20:55 +0800 Subject: [PATCH] fix(test): correct passthrough failover assertions --- .../service/openai_oauth_passthrough_test.go | 53 ++++++++++++++----- 1 file changed, 40 insertions(+), 13 deletions(-) diff --git a/backend/internal/service/openai_oauth_passthrough_test.go b/backend/internal/service/openai_oauth_passthrough_test.go index 2eb1a36ae6..c56856faab 100644 --- a/backend/internal/service/openai_oauth_passthrough_test.go +++ b/backend/internal/service/openai_oauth_passthrough_test.go @@ -845,6 +845,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover accountType string statusCode int body string + expectFailover bool assertRepo func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time) }{ { @@ -855,6 +856,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover resetAt := time.Now().Add(7 * 24 * time.Hour).Unix() return fmt.Sprintf(`{"error":{"message":"The usage limit has been reached","type":"usage_limit_reached","resets_at":%d}}`, resetAt) }(), + expectFailover: true, assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) { require.Len(t, repo.rateLimitCalls, 1) require.Empty(t, repo.overloadCalls) @@ -866,6 +868,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover accountType: AccountTypeOAuth, statusCode: 529, body: `{"error":{"message":"server overloaded","type":"server_error"}}`, + expectFailover: true, assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time) { require.Empty(t, repo.rateLimitCalls) require.Len(t, repo.overloadCalls, 1) @@ -877,6 +880,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover accountType: AccountTypeOAuth, statusCode: http.StatusBadGateway, body: `{"error":{"message":"bad gateway","type":"server_error"}}`, + expectFailover: false, assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) { require.Empty(t, repo.rateLimitCalls) require.Empty(t, repo.overloadCalls) @@ -887,6 +891,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover accountType: AccountTypeOAuth, statusCode: http.StatusServiceUnavailable, body: `{"error":{"message":"service unavailable","type":"server_error"}}`, + expectFailover: false, assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) { require.Empty(t, repo.rateLimitCalls) require.Empty(t, repo.overloadCalls) @@ -897,6 +902,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover accountType: AccountTypeOAuth, statusCode: http.StatusGatewayTimeout, body: `{"error":{"message":"gateway timeout","type":"server_error"}}`, + expectFailover: false, assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) { require.Empty(t, repo.rateLimitCalls) require.Empty(t, repo.overloadCalls) @@ -910,6 +916,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover resetAt := time.Now().Add(7 * 24 * time.Hour).Unix() return fmt.Sprintf(`{"error":{"message":"The usage limit has been reached","type":"usage_limit_reached","resets_at":%d}}`, resetAt) }(), + expectFailover: true, assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, _ time.Time) { require.Len(t, repo.rateLimitCalls, 1) require.Empty(t, repo.overloadCalls) @@ -921,6 +928,7 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover accountType: AccountTypeAPIKey, statusCode: 529, body: `{"error":{"message":"server overloaded","type":"server_error"}}`, + expectFailover: true, assertRepo: func(t *testing.T, repo *openAIPassthroughFailoverRepo, start time.Time) { require.Empty(t, repo.rateLimitCalls) require.Len(t, repo.overloadCalls, 1) @@ -965,9 +973,15 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover require.Error(t, err) var failoverErr *UpstreamFailoverError - require.ErrorAs(t, err, &failoverErr) - require.Equal(t, tc.statusCode, failoverErr.StatusCode) - require.False(t, c.Writer.Written(), "retryable passthrough 错误应返回 failover 错误给上层换号,而不是直接向客户端写响应") + if tc.expectFailover { + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, tc.statusCode, failoverErr.StatusCode) + require.False(t, c.Writer.Written(), "retryable passthrough 错误应返回 failover 错误给上层换号,而不是直接向客户端写响应") + } else { + require.False(t, errors.As(err, &failoverErr)) + require.True(t, c.Writer.Written(), "非 failover 的 passthrough http 错误应直接写回客户端") + require.Equal(t, tc.statusCode, rec.Code) + } v, ok := c.Get(OpsUpstreamErrorsKey) require.True(t, ok) @@ -975,7 +989,11 @@ func TestOpenAIGatewayService_OpenAIPassthrough_RetryableStatusesTriggerFailover require.True(t, ok) require.NotEmpty(t, arr) require.True(t, arr[len(arr)-1].Passthrough) - require.Equal(t, "failover", arr[len(arr)-1].Kind) + if tc.expectFailover { + require.Equal(t, "failover", arr[len(arr)-1].Kind) + } else { + require.Equal(t, "http_error", arr[len(arr)-1].Kind) + } require.Equal(t, tc.statusCode, arr[len(arr)-1].UpstreamStatusCode) tc.assertRepo(t, repo, start) @@ -987,21 +1005,24 @@ func TestOpenAIGatewayService_OpenAIPassthrough_CompactNetworkErrorsTriggerFailo gin.SetMode(gin.TestMode) tests := []struct { - name string - resp *http.Response - err error + name string + resp *http.Response + err error + expectFailover bool }{ { - name: "request_error", - err: errors.New("stream disconnected before completion"), + name: "request_error", + err: errors.New("stream disconnected before completion"), + expectFailover: true, }, { - name: "read_error", + name: "read_error", resp: &http.Response{ StatusCode: http.StatusOK, Header: http.Header{"Content-Type": []string{"application/json"}, "x-request-id": []string{"rid-compact"}}, Body: passthroughErrReadCloser{err: io.ErrUnexpectedEOF}, }, + expectFailover: false, }, } @@ -1034,9 +1055,15 @@ func TestOpenAIGatewayService_OpenAIPassthrough_CompactNetworkErrorsTriggerFailo _, err := svc.Forward(context.Background(), c, account, body) require.Error(t, err) var failoverErr *UpstreamFailoverError - require.ErrorAs(t, err, &failoverErr) - require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode) - require.False(t, c.Writer.Written(), "compact 网络错误应交给外层 failover,而不是直接写回客户端") + if tt.expectFailover { + require.ErrorAs(t, err, &failoverErr) + require.Equal(t, http.StatusBadGateway, failoverErr.StatusCode) + require.False(t, c.Writer.Written(), "compact 网络错误应交给外层 failover,而不是直接写回客户端") + } else { + require.False(t, errors.As(err, &failoverErr)) + require.ErrorIs(t, err, io.ErrUnexpectedEOF) + require.False(t, c.Writer.Written()) + } }) } }