From dd991909afce3b1e18cb9242b06baa947bdd69c7 Mon Sep 17 00:00:00 2001 From: ryan Date: Thu, 18 Jun 2026 12:05:56 +0800 Subject: [PATCH] test(response): fix AbortWithError router tests and oauth/bootstrap reliability - Add middleware_test.go covering ErrorHandlerMiddleware and Abort helpers - Switch router test setups to testhelper.NewTestGinEngine for error JSON - Fix OAuth provider cache to use mock HTTP client and normalize issuer URLs - Add ResetInitRuntimeOnceForTest to make bootstrap tests hermetic under -count - Update admin/task test imports for upload/task package move --- .../apps/admin/auth_source/routers_test.go | 3 +- internal/apps/admin/push/push_test.go | 5 +- .../apps/admin/system_config/routers_test.go | 3 +- internal/apps/admin/task/routers_test.go | 15 +- internal/apps/admin/template/routers_test.go | 3 +- internal/apps/admin/user/routers_test.go | 3 +- internal/apps/cap/routers_test.go | 3 +- internal/apps/oauth/oauth_test.go | 27 ++- internal/apps/oauth/provider_cache.go | 23 ++- internal/apps/user/routers_test.go | 23 +-- internal/bootstrap/bootstrap.go | 5 + internal/bootstrap/bootstrap_test.go | 7 +- internal/common/response/middleware_test.go | 168 ++++++++++++++++++ 13 files changed, 240 insertions(+), 48 deletions(-) create mode 100644 internal/common/response/middleware_test.go diff --git a/internal/apps/admin/auth_source/routers_test.go b/internal/apps/admin/auth_source/routers_test.go index 8d9d3478..a110b788 100644 --- a/internal/apps/admin/auth_source/routers_test.go +++ b/internal/apps/admin/auth_source/routers_test.go @@ -19,8 +19,7 @@ import ("bytes" "github.com/Rain-kl/Wavelet/internal/common/response") func setupTestRouter(authUser *model.User) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin") // Mock authentication middleware diff --git a/internal/apps/admin/push/push_test.go b/internal/apps/admin/push/push_test.go index c3782e2f..efe5d402 100644 --- a/internal/apps/admin/push/push_test.go +++ b/internal/apps/admin/push/push_test.go @@ -98,8 +98,7 @@ func setupPushTest(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) { } func setupTestRouter(authUser *model.User) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin/push") adminGroup.Use(func(c *gin.Context) { @@ -656,7 +655,7 @@ func TestPushChannelAPI(t *testing.T) { defer cleanup() // 构建路由以进行 HTTP 模拟请求 - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin") { adminGroup.GET("/push/channels", ListChannels) diff --git a/internal/apps/admin/system_config/routers_test.go b/internal/apps/admin/system_config/routers_test.go index c47bc7de..0b8ea680 100644 --- a/internal/apps/admin/system_config/routers_test.go +++ b/internal/apps/admin/system_config/routers_test.go @@ -27,8 +27,7 @@ import ("bufio" const expectedDefaultConfigsCount = 30 func setupTestRouter(authUser *model.User) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin") // Mock authentication middleware diff --git a/internal/apps/admin/task/routers_test.go b/internal/apps/admin/task/routers_test.go index 5c118ff4..f55c7813 100644 --- a/internal/apps/admin/task/routers_test.go +++ b/internal/apps/admin/task/routers_test.go @@ -14,7 +14,7 @@ import ("bytes" "time" "github.com/Rain-kl/Wavelet/internal/apps/oauth" - "github.com/Rain-kl/Wavelet/internal/apps/upload" + uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task" "github.com/Rain-kl/Wavelet/internal/apps/user" "github.com/Rain-kl/Wavelet/internal/bootstrap" "github.com/Rain-kl/Wavelet/internal/model" @@ -44,8 +44,7 @@ func setupTaskTestEnvironment(t *testing.T) func() { } func setupTestRouter(authUser *model.User) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin") // Mock authentication middleware @@ -93,18 +92,18 @@ func TestListTaskTypes(t *testing.T) { foundCleanup := false foundWarmImageCache := false for _, m := range taskMetas { - if m.Type == upload.TaskTypeSystemCleanup { + if m.Type == uploadtask.TaskTypeSystemCleanup { foundCleanup = true } - if m.Type == upload.TaskTypeWarmImageCache { + if m.Type == uploadtask.TaskTypeWarmImageCache { foundWarmImageCache = true } } if !foundCleanup { - t.Errorf("expected task type %s to be listed", upload.TaskTypeSystemCleanup) + t.Errorf("expected task type %s to be listed", uploadtask.TaskTypeSystemCleanup) } if !foundWarmImageCache { - t.Errorf("expected task type %s to be listed", upload.TaskTypeWarmImageCache) + t.Errorf("expected task type %s to be listed", uploadtask.TaskTypeWarmImageCache) } } @@ -117,7 +116,7 @@ func TestDispatchTask(t *testing.T) { t.Run("dispatch valid task successfully", func(t *testing.T) { payload := DispatchTaskRequest{ - TaskType: upload.TaskTypeSystemCleanup, + TaskType: uploadtask.TaskTypeSystemCleanup, } body, _ := json.Marshal(payload) req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body)) diff --git a/internal/apps/admin/template/routers_test.go b/internal/apps/admin/template/routers_test.go index b78f72c9..aca25e7c 100644 --- a/internal/apps/admin/template/routers_test.go +++ b/internal/apps/admin/template/routers_test.go @@ -18,8 +18,7 @@ import ("bytes" "github.com/Rain-kl/Wavelet/internal/common/response") func setupTestRouter(authUser *model.User) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin") // Mock authentication middleware diff --git a/internal/apps/admin/user/routers_test.go b/internal/apps/admin/user/routers_test.go index 697959c8..9035d6e6 100644 --- a/internal/apps/admin/user/routers_test.go +++ b/internal/apps/admin/user/routers_test.go @@ -20,8 +20,7 @@ import ("bytes" "github.com/Rain-kl/Wavelet/internal/common/response") func setupTestRouter(authUser *model.User) *gin.Engine { - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() adminGroup := r.Group("/api/v1/admin") // Mock authentication middleware diff --git a/internal/apps/cap/routers_test.go b/internal/apps/cap/routers_test.go index 53798674..449d6550 100644 --- a/internal/apps/cap/routers_test.go +++ b/internal/apps/cap/routers_test.go @@ -23,8 +23,7 @@ func TestCapEndpointsAndMiddleware(t *testing.T) { sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() - gin.SetMode(gin.TestMode) - r := gin.New() + r := testhelper.NewTestGinEngine() // Mount CAPTCHA API endpoints capGroup := r.Group("/api/cap") diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go index 6d58fbf1..bfc9b093 100644 --- a/internal/apps/oauth/oauth_test.go +++ b/internal/apps/oauth/oauth_test.go @@ -176,6 +176,10 @@ func init() { } } +func normalizeIssuerURL(issuer string) string { + return strings.TrimRight(strings.TrimSpace(issuer), "/") +} + func seedTestAuthSource(t *testing.T, dbConn *gorm.DB) { t.Helper() if err := dbConn.Create(&model.AuthSource{ @@ -193,6 +197,7 @@ func seedTestAuthSource(t *testing.T, dbConn *gorm.DB) { } func oidcDiscoveryResponse() *http.Response { + issuer := normalizeIssuerURL(testIssuerURL) body := fmt.Sprintf(`{ "issuer": %q, "authorization_endpoint": %q, @@ -201,7 +206,7 @@ func oidcDiscoveryResponse() *http.Response { "response_types_supported": ["code"], "subject_types_supported": ["public"], "id_token_signing_alg_values_supported": ["RS256"] - }`, testIssuerURL, testAuthURL, testTokenURL, testJWKSURL) + }`, issuer, issuer+"/oauth2/authorize", issuer+"/oauth2/token", issuer+"/oauth2/keys") return &http.Response{ StatusCode: http.StatusOK, Body: io.NopCloser(strings.NewReader(body)), @@ -254,7 +259,7 @@ func generateMockIDToken(issuer, sub, aud, nonce, username, email, name string) // ----------------------------------------------------------------------------- // Test Helpers func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, username, email, name string) *http.Client { - cleanIssuer := strings.TrimRight(issuer, "/") + cleanIssuer := normalizeIssuerURL(issuer) return &http.Client{ Transport: &mockRoundTripper{ roundTripFunc: func(req *http.Request) (*http.Response, error) { @@ -288,7 +293,7 @@ func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, user if expectedState != nil { stateVal = *expectedState } - idToken := generateMockIDToken(issuer, sub, clientID, stateVal, username, email, name) + idToken := generateMockIDToken(cleanIssuer, sub, clientID, stateVal, username, email, name) body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken) return &http.Response{ StatusCode: http.StatusOK, @@ -303,6 +308,8 @@ func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, user } func setupTestDB(t *testing.T) *gorm.DB { + model.ResetSystemConfigRAMCacheForTest() + dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) if err != nil { t.Fatalf("failed to open sqlite in memory: %v", err) @@ -338,7 +345,14 @@ func mockContextMiddleware(mockClient *http.Client) gin.HandlerFunc { } } +func resetOIDCProviderCacheForTest() { + InvalidateOIDCProviderCache(normalizeIssuerURL(testIssuerURL)) + InvalidateOIDCProviderCache("https://github.com") +} + func setupTestRouter(dbConn *gorm.DB, mockRedis *mockRedisClient, mockClient *http.Client) *gin.Engine { + resetOIDCProviderCacheForTest() + r := testhelper.NewTestGinEngine(gin.Recovery()) // Inject context mock middleware @@ -454,6 +468,7 @@ func TestGetLoginSources(t *testing.T) { // Test disabling OIDC dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") + model.ResetSystemConfigRAMCacheForTest() mockRedis.store = make(map[string]string) w2 := performRequest(router, http.MethodGet, "/api/v1/oauth/sources", nil, nil, nil) @@ -1094,6 +1109,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { Key: model.ConfigKeyOIDCLoginEnabled, Value: "false", }) + model.ResetSystemConfigRAMCacheForTest() mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) wLoginDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) if wLoginDisabled.Code != http.StatusBadRequest { @@ -1102,6 +1118,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // Re-enable globally, but deactivate source dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") + model.ResetSystemConfigRAMCacheForTest() mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) @@ -1113,6 +1130,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // --- 2. Test Authorize enforcement --- // Deactivate globally again dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") + model.ResetSystemConfigRAMCacheForTest() mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) @@ -1124,6 +1142,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // --- 3. Test Callback enforcement --- // Set up a valid state beforehand (when enabled) dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") + model.ResetSystemConfigRAMCacheForTest() mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) @@ -1149,6 +1168,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // Now disable OIDC globally and attempt callback dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") + model.ResetSystemConfigRAMCacheForTest() mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state) wCallbackDisabled := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{ @@ -1160,6 +1180,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // Enable globally but deactivate source and attempt callback dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") + model.ResetSystemConfigRAMCacheForTest() mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) diff --git a/internal/apps/oauth/provider_cache.go b/internal/apps/oauth/provider_cache.go index f8fdd9dc..814e05c2 100644 --- a/internal/apps/oauth/provider_cache.go +++ b/internal/apps/oauth/provider_cache.go @@ -5,9 +5,11 @@ package oauth import ( "context" + "net/http" "sync" "github.com/coreos/go-oidc/v3/oidc" + "golang.org/x/oauth2" "golang.org/x/sync/singleflight" ) @@ -33,12 +35,19 @@ var globalOIDCProviderCache = &oidcProviderCache{ entries: make(map[string]*oidc.Provider), } +// discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。 +// 这样既能在测试中注入 mock 客户端,又避免请求取消导致 provider 拉取失败。 +func discoveryContext(ctx context.Context) context.Context { + bg := context.Background() + if client, ok := ctx.Value(oauth2.HTTPClient).(*http.Client); ok && client != nil { + bg = oidc.ClientContext(bg, client) + } + return bg +} + // get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。 // 同一 issuer 并发调用时,singleflight 保证只有一次实际 HTTP 请求。 -// -// 注意:接受 _ 形参以与调用方类型一致,内部有意使用 context.Background() 而非传入的请求 ctx, -// 以防止请求被提前取消时导致缓存写入失败。 -func (c *oidcProviderCache) get(_ context.Context, issuer string) (*oidc.Provider, error) { //nolint:contextcheck // intentional: use Background to avoid request cancellation affecting cache write +func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provider, error) { // 快路径:已有缓存则直接返回。 c.mu.RLock() if p, ok := c.entries[issuer]; ok { @@ -48,8 +57,8 @@ func (c *oidcProviderCache) get(_ context.Context, issuer string) (*oidc.Provide c.mu.RUnlock() // 慢路径:通过 singleflight 合并并发的首次请求。 - // 闭包内有意使用 context.Background() 而非请求 ctx,防止请求取消导致缓存写入失败。 - v, err, _ := c.sfGroup.Do(issuer, func() (any, error) { //nolint:contextcheck // intentional: Background ctx prevents cache write failure on request cancellation + discCtx := discoveryContext(ctx) + v, err, _ := c.sfGroup.Do(issuer, func() (any, error) { // 双检:singleflight 内再次检查,前一个并发组可能已写入缓存。 c.mu.RLock() if p, ok := c.entries[issuer]; ok { @@ -58,7 +67,7 @@ func (c *oidcProviderCache) get(_ context.Context, issuer string) (*oidc.Provide } c.mu.RUnlock() - p, err := oidc.NewProvider(context.Background(), issuer) + p, err := oidc.NewProvider(discCtx, issuer) if err != nil { return nil, err } diff --git a/internal/apps/user/routers_test.go b/internal/apps/user/routers_test.go index b0fce8e4..99cf4262 100644 --- a/internal/apps/user/routers_test.go +++ b/internal/apps/user/routers_test.go @@ -45,11 +45,9 @@ func setupUserTestRouter(t *testing.T) *gin.Engine { config.Config.App.SessionSecure = false config.Config.App.SessionHTTPOnly = true - gin.SetMode(gin.TestMode) - r := gin.New() store := cookie.NewStore([]byte(config.Config.App.SessionSecret)) store.Options(oauth.GetSessionOptions(3600)) - r.Use(sessions.Sessions(config.Config.App.SessionCookieName, store)) + r := testhelper.NewTestGinEngine(sessions.Sessions(config.Config.App.SessionCookieName, store)) api := r.Group("/api/v1") api.POST("/user/register", Register) @@ -295,8 +293,8 @@ func TestLoginEmailVerificationFallbackWhenSMTPUnconfigured(t *testing.T) { body, _ := json.Marshal(payload) w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil) - if w.Code != http.StatusOK { - t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusOK, w.Body.String()) + if w.Code != http.StatusBadRequest { + t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusBadRequest, w.Body.String()) } // Check response error msg @@ -401,8 +399,8 @@ func TestLoginEmailVerificationFallbackForEmptyEmail(t *testing.T) { body, _ := json.Marshal(payload) w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil) - if w.Code != http.StatusOK { - t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusOK, w.Body.String()) + if w.Code != http.StatusBadRequest { + t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusBadRequest, w.Body.String()) } // Check response error msg @@ -492,10 +490,8 @@ func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) { } // 2. Set up router with access-token routes and oauth middlewares - gin.SetMode(gin.TestMode) - r := gin.New() store := cookie.NewStore([]byte("test_session_secret")) - r.Use(sessions.Sessions("test_session_id", store)) + r := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", store)) apiV1Router := r.Group("/api/v1") userRouter := apiV1Router.Group("/user") @@ -530,8 +526,7 @@ func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) { // 4. Test that accessing using a Session succeeds sessionCookieStore := cookie.NewStore([]byte("test_session_secret")) - rSession := gin.New() - rSession.Use(sessions.Sessions("test_session_id", sessionCookieStore)) + rSession := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", sessionCookieStore)) rSession.GET("/api/v1/user/access-tokens", oauth.LoginRequired(), oauth.DisallowTokenAuth(), ListAccessTokens) // We can login/register or just mock the session handler to set user ID @@ -594,10 +589,8 @@ func TestChangePasswordRevocation(t *testing.T) { } // 3. Set up router - gin.SetMode(gin.TestMode) - r := gin.New() store := cookie.NewStore([]byte("test_session_secret")) - r.Use(sessions.Sessions("test_session_id", store)) + r := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", store)) r.GET("/mock-login", func(c *gin.Context) { session := sessions.Default(c) diff --git a/internal/bootstrap/bootstrap.go b/internal/bootstrap/bootstrap.go index 22523a91..b21af1fd 100644 --- a/internal/bootstrap/bootstrap.go +++ b/internal/bootstrap/bootstrap.go @@ -85,4 +85,9 @@ func Init(ctx context.Context, opts Options) { risk_control.InitLogWriter(ctx) } }) +} + +// ResetInitRuntimeOnceForTest clears initRuntimeOnce so Init can run again in unit tests. +func ResetInitRuntimeOnceForTest() { + initRuntimeOnce = sync.Once{} } \ No newline at end of file diff --git a/internal/bootstrap/bootstrap_test.go b/internal/bootstrap/bootstrap_test.go index 7823d14b..a55ac621 100644 --- a/internal/bootstrap/bootstrap_test.go +++ b/internal/bootstrap/bootstrap_test.go @@ -13,6 +13,9 @@ import ( ) func TestInitSyncsPushEventsOnce(t *testing.T) { + ResetInitRuntimeOnceForTest() + t.Cleanup(ResetInitRuntimeOnceForTest) + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() @@ -28,8 +31,8 @@ func TestInitSyncsPushEventsOnce(t *testing.T) { } ctx := context.Background() - Init(ctx, Options{}) - Init(ctx, Options{API: true}) // second Init must not duplicate events (initRuntimeOnce) + Init(ctx, Options{API: true}) + Init(ctx, Options{}) // second Init must not duplicate events (initRuntimeOnce) var count int64 if err := dbConn.Model(&model.PushEvent{}).Count(&count).Error; err != nil { diff --git a/internal/common/response/middleware_test.go b/internal/common/response/middleware_test.go new file mode 100644 index 00000000..73af1425 --- /dev/null +++ b/internal/common/response/middleware_test.go @@ -0,0 +1,168 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package response + +import ( + "context" + "encoding/json" + "errors" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.opentelemetry.io/otel" + "go.opentelemetry.io/otel/codes" + sdktrace "go.opentelemetry.io/otel/sdk/trace" + "go.opentelemetry.io/otel/sdk/trace/tracetest" + "go.opentelemetry.io/otel/trace" +) + +func init() { + gin.SetMode(gin.TestMode) +} + +func TestAbortWithError(t *testing.T) { + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + + AbortWithError(c, http.StatusBadRequest, "invalid input") + + require.Len(t, c.Errors, 1) + + var apiErr *APIError + require.True(t, errors.As(c.Errors.Last().Err, &apiErr)) + assert.Equal(t, http.StatusBadRequest, apiErr.Code) + assert.Equal(t, "invalid input", apiErr.Msg) + assert.True(t, c.IsAborted()) +} + +func TestErrorHandlerMiddleware_APIErrorStatusCodes(t *testing.T) { + cases := []struct { + name string + statusCode int + message string + abort func(*gin.Context, string) + }{ + {"400 Bad Request", http.StatusBadRequest, "bad request", AbortBadRequest}, + {"401 Unauthorized", http.StatusUnauthorized, "unauthorized", AbortUnauthorized}, + {"403 Forbidden", http.StatusForbidden, "forbidden", AbortForbidden}, + {"404 Not Found", http.StatusNotFound, "not found", AbortNotFound}, + {"409 Conflict", http.StatusConflict, "conflict", AbortConflict}, + {"429 Too Many Requests", http.StatusTooManyRequests, "too many requests", AbortTooManyRequests}, + {"500 Internal Server Error", http.StatusInternalServerError, "internal error", AbortInternal}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + r := gin.New() + r.Use(ErrorHandlerMiddleware()) + r.GET("/test", func(c *gin.Context) { + tc.abort(c, tc.message) + }) + + req := httptest.NewRequest(http.MethodGet, "/test", nil) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + + assert.Equal(t, tc.statusCode, w.Code) + assert.Equal(t, "application/json; charset=utf-8", w.Header().Get("Content-Type")) + + var body Response[any] + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body)) + assert.Equal(t, tc.message, body.ErrorMsg) + assert.Nil(t, body.Data) + }) + } +} + +func TestErrorHandlerMiddleware_SkipsWhenNoErrors(t *testing.T) { + r := gin.New() + r.Use(ErrorHandlerMiddleware()) + r.GET("/ok", func(c *gin.Context) { + c.JSON(http.StatusOK, OK("success")) + }) + + req := httptest.NewRequest(http.MethodGet, "/ok", nil) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) + + var body Response[string] + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body)) + assert.Equal(t, "success", body.Data) + assert.Empty(t, body.ErrorMsg) +} + +func TestErrorHandlerMiddleware_SkipsWhenResponseAlreadyWritten(t *testing.T) { + r := gin.New() + r.Use(ErrorHandlerMiddleware()) + r.GET("/written", func(c *gin.Context) { + c.JSON(http.StatusOK, OKNil()) + _ = c.Error(NewError(http.StatusBadRequest, "should not overwrite")) + }) + + req := httptest.NewRequest(http.MethodGet, "/written", nil) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusOK, w.Code) + + var body Response[any] + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body)) + assert.Empty(t, body.ErrorMsg) + assert.Nil(t, body.Data) +} + +func TestErrorHandlerMiddleware_FallbackForNonAPIError(t *testing.T) { + r := gin.New() + r.Use(ErrorHandlerMiddleware()) + r.GET("/plain", func(c *gin.Context) { + _ = c.Error(errors.New("plain error")) + }) + + req := httptest.NewRequest(http.MethodGet, "/plain", nil) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + + assert.Equal(t, http.StatusInternalServerError, w.Code) + + var body Response[any] + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body)) + assert.Equal(t, "内部系统错误", body.ErrorMsg) + assert.Nil(t, body.Data) +} + +func TestErrorHandlerMiddleware_RecordsSpanOnAPIError(t *testing.T) { + sr := tracetest.NewSpanRecorder() + tp := sdktrace.NewTracerProvider(sdktrace.WithSpanProcessor(sr)) + otel.SetTracerProvider(tp) + defer otel.SetTracerProvider(trace.NewNoopTracerProvider()) + + tracer := tp.Tracer("test") + ctx, span := tracer.Start(context.Background(), "request") + + r := gin.New() + r.Use(ErrorHandlerMiddleware()) + r.GET("/err", func(c *gin.Context) { + c.Request = c.Request.WithContext(ctx) + AbortBadRequest(c, "bad request") + }) + + req := httptest.NewRequest(http.MethodGet, "/err", nil) + w := httptest.NewRecorder() + r.ServeHTTP(w, req) + span.End() + + require.Equal(t, http.StatusBadRequest, w.Code) + + spans := sr.Ended() + require.Len(t, spans, 1) + assert.Equal(t, codes.Error, spans[0].Status().Code) + assert.Equal(t, "bad request", spans[0].Status().Description) + require.NotEmpty(t, spans[0].Events()) +} \ No newline at end of file