mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-30 06:16:37 +08:00
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:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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{}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
Reference in New Issue
Block a user