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
This commit is contained in:
ryan
2026-06-18 12:05:56 +08:00
parent e5b3a60f73
commit dd991909af
13 changed files with 240 additions and 48 deletions
@@ -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
+2 -3
View File
@@ -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)
@@ -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
+7 -8
View File
@@ -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))
+1 -2
View File
@@ -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
+1 -2
View File
@@ -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
+1 -2
View File
@@ -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")
+24 -3
View File
@@ -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)
+16 -7
View File
@@ -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
}
+8 -15
View File
@@ -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)
+5
View File
@@ -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{}
}
+5 -2
View File
@@ -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 {
+168
View File
@@ -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())
}