merge main into repository-context-enabled-background

This commit is contained in:
ryan
2026-06-18 12:13:45 +08:00
10 changed files with 237 additions and 42 deletions
@@ -18,8 +18,7 @@ import ("bytes"
"github.com/Rain-kl/Wavelet/internal/common/response") "github.com/Rain-kl/Wavelet/internal/common/response")
func setupTestRouter(authUser *model.User) *gin.Engine { func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode) r := testhelper.NewTestGinEngine()
r := gin.New()
adminGroup := r.Group("/api/v1/admin") adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware // Mock authentication middleware
+2 -3
View File
@@ -97,8 +97,7 @@ func setupPushTest(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) {
} }
func setupTestRouter(authUser *model.User) *gin.Engine { func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode) r := testhelper.NewTestGinEngine()
r := gin.New()
adminGroup := r.Group("/api/v1/admin/push") adminGroup := r.Group("/api/v1/admin/push")
adminGroup.Use(func(c *gin.Context) { adminGroup.Use(func(c *gin.Context) {
@@ -655,7 +654,7 @@ func TestPushChannelAPI(t *testing.T) {
defer cleanup() defer cleanup()
// 构建路由以进行 HTTP 模拟请求 // 构建路由以进行 HTTP 模拟请求
r := gin.New() r := testhelper.NewTestGinEngine()
adminGroup := r.Group("/api/v1/admin") adminGroup := r.Group("/api/v1/admin")
{ {
adminGroup.GET("/push/channels", ListChannels) adminGroup.GET("/push/channels", ListChannels)
+7 -8
View File
@@ -14,7 +14,7 @@ import ("bytes"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth" "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/apps/user"
"github.com/Rain-kl/Wavelet/internal/bootstrap" "github.com/Rain-kl/Wavelet/internal/bootstrap"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
@@ -43,8 +43,7 @@ func setupTaskTestEnvironment(t *testing.T) func() {
} }
func setupTestRouter(authUser *model.User) *gin.Engine { func setupTestRouter(authUser *model.User) *gin.Engine {
gin.SetMode(gin.TestMode) r := testhelper.NewTestGinEngine()
r := gin.New()
adminGroup := r.Group("/api/v1/admin") adminGroup := r.Group("/api/v1/admin")
// Mock authentication middleware // Mock authentication middleware
@@ -92,18 +91,18 @@ func TestListTaskTypes(t *testing.T) {
foundCleanup := false foundCleanup := false
foundWarmImageCache := false foundWarmImageCache := false
for _, m := range taskMetas { for _, m := range taskMetas {
if m.Type == upload.TaskTypeSystemCleanup { if m.Type == uploadtask.TaskTypeSystemCleanup {
foundCleanup = true foundCleanup = true
} }
if m.Type == upload.TaskTypeWarmImageCache { if m.Type == uploadtask.TaskTypeWarmImageCache {
foundWarmImageCache = true foundWarmImageCache = true
} }
} }
if !foundCleanup { 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 { if !foundWarmImageCache {
t.Errorf("expected task type %s to be listed", upload.TaskTypeWarmImageCache) t.Errorf("expected task type %s to be listed", uploadtask.TaskTypeWarmImageCache)
} }
} }
@@ -116,7 +115,7 @@ func TestDispatchTask(t *testing.T) {
t.Run("dispatch valid task successfully", func(t *testing.T) { t.Run("dispatch valid task successfully", func(t *testing.T) {
payload := DispatchTaskRequest{ payload := DispatchTaskRequest{
TaskType: upload.TaskTypeSystemCleanup, TaskType: uploadtask.TaskTypeSystemCleanup,
} }
body, _ := json.Marshal(payload) body, _ := json.Marshal(payload)
req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body)) req, _ := http.NewRequest("POST", "/api/v1/admin/tasks/dispatch", bytes.NewBuffer(body))
+1 -2
View File
@@ -24,8 +24,7 @@ func TestCapEndpointsAndMiddleware(t *testing.T) {
sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t) sqliteDB, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup() defer cleanup()
gin.SetMode(gin.TestMode) r := testhelper.NewTestGinEngine()
r := gin.New()
// Mount CAPTCHA API endpoints // Mount CAPTCHA API endpoints
capGroup := r.Group("/api/cap") capGroup := r.Group("/api/cap")
+24 -3
View File
@@ -177,6 +177,10 @@ func init() {
} }
} }
func normalizeIssuerURL(issuer string) string {
return strings.TrimRight(strings.TrimSpace(issuer), "/")
}
func seedTestAuthSource(t *testing.T, dbConn *gorm.DB) { func seedTestAuthSource(t *testing.T, dbConn *gorm.DB) {
t.Helper() t.Helper()
if err := dbConn.Create(&model.AuthSource{ if err := dbConn.Create(&model.AuthSource{
@@ -194,6 +198,7 @@ func seedTestAuthSource(t *testing.T, dbConn *gorm.DB) {
} }
func oidcDiscoveryResponse() *http.Response { func oidcDiscoveryResponse() *http.Response {
issuer := normalizeIssuerURL(testIssuerURL)
body := fmt.Sprintf(`{ body := fmt.Sprintf(`{
"issuer": %q, "issuer": %q,
"authorization_endpoint": %q, "authorization_endpoint": %q,
@@ -202,7 +207,7 @@ func oidcDiscoveryResponse() *http.Response {
"response_types_supported": ["code"], "response_types_supported": ["code"],
"subject_types_supported": ["public"], "subject_types_supported": ["public"],
"id_token_signing_alg_values_supported": ["RS256"] "id_token_signing_alg_values_supported": ["RS256"]
}`, testIssuerURL, testAuthURL, testTokenURL, testJWKSURL) }`, issuer, issuer+"/oauth2/authorize", issuer+"/oauth2/token", issuer+"/oauth2/keys")
return &http.Response{ return &http.Response{
StatusCode: http.StatusOK, StatusCode: http.StatusOK,
Body: io.NopCloser(strings.NewReader(body)), Body: io.NopCloser(strings.NewReader(body)),
@@ -255,7 +260,7 @@ func generateMockIDToken(issuer, sub, aud, nonce, username, email, name string)
// ----------------------------------------------------------------------------- // -----------------------------------------------------------------------------
// Test Helpers // Test Helpers
func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, username, email, name string) *http.Client { func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, username, email, name string) *http.Client {
cleanIssuer := strings.TrimRight(issuer, "/") cleanIssuer := normalizeIssuerURL(issuer)
return &http.Client{ return &http.Client{
Transport: &mockRoundTripper{ Transport: &mockRoundTripper{
roundTripFunc: func(req *http.Request) (*http.Response, error) { roundTripFunc: func(req *http.Request) (*http.Response, error) {
@@ -289,7 +294,7 @@ func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, user
if expectedState != nil { if expectedState != nil {
stateVal = *expectedState 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) body := fmt.Sprintf(`{"access_token":"mock_access_token","token_type":"Bearer","expires_in":3600,"id_token":"%s"}`, idToken)
return &http.Response{ return &http.Response{
StatusCode: http.StatusOK, StatusCode: http.StatusOK,
@@ -304,6 +309,8 @@ func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, user
} }
func setupTestDB(t *testing.T) *gorm.DB { func setupTestDB(t *testing.T) *gorm.DB {
repository.ResetSystemConfigRAMCacheForTest()
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
if err != nil { if err != nil {
t.Fatalf("failed to open sqlite in memory: %v", err) t.Fatalf("failed to open sqlite in memory: %v", err)
@@ -339,7 +346,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 { func setupTestRouter(dbConn *gorm.DB, mockRedis *mockRedisClient, mockClient *http.Client) *gin.Engine {
resetOIDCProviderCacheForTest()
r := testhelper.NewTestGinEngine(gin.Recovery()) r := testhelper.NewTestGinEngine(gin.Recovery())
// Inject context mock middleware // Inject context mock middleware
@@ -455,6 +469,7 @@ func TestGetLoginSources(t *testing.T) {
// Test disabling OIDC // Test disabling OIDC
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false")
repository.ResetSystemConfigRAMCacheForTest()
mockRedis.store = make(map[string]string) mockRedis.store = make(map[string]string)
w2 := performRequest(router, http.MethodGet, "/api/v1/oauth/sources", nil, nil, nil) w2 := performRequest(router, http.MethodGet, "/api/v1/oauth/sources", nil, nil, nil)
@@ -1095,6 +1110,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
Key: model.ConfigKeyOIDCLoginEnabled, Key: model.ConfigKeyOIDCLoginEnabled,
Value: "false", Value: "false",
}) })
repository.ResetSystemConfigRAMCacheForTest()
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
wLoginDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) wLoginDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
if wLoginDisabled.Code != http.StatusBadRequest { if wLoginDisabled.Code != http.StatusBadRequest {
@@ -1103,6 +1119,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
// Re-enable globally, but deactivate source // Re-enable globally, but deactivate source
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true")
repository.ResetSystemConfigRAMCacheForTest()
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false)
@@ -1114,6 +1131,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
// --- 2. Test Authorize enforcement --- // --- 2. Test Authorize enforcement ---
// Deactivate globally again // Deactivate globally again
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false")
repository.ResetSystemConfigRAMCacheForTest()
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true)
@@ -1125,6 +1143,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
// --- 3. Test Callback enforcement --- // --- 3. Test Callback enforcement ---
// Set up a valid state beforehand (when enabled) // Set up a valid state beforehand (when enabled)
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true")
repository.ResetSystemConfigRAMCacheForTest()
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true)
@@ -1150,6 +1169,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
// Now disable OIDC globally and attempt callback // Now disable OIDC globally and attempt callback
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false") dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "false")
repository.ResetSystemConfigRAMCacheForTest()
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state) reqBody := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state)
wCallbackDisabled := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{ wCallbackDisabled := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody), map[string]string{
@@ -1161,6 +1181,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
// Enable globally but deactivate source and attempt callback // Enable globally but deactivate source and attempt callback
dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true") dbConn.Model(&model.SystemConfig{}).Where("key = ?", model.ConfigKeyOIDCLoginEnabled).Update("value", "true")
repository.ResetSystemConfigRAMCacheForTest()
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false)
+16 -7
View File
@@ -5,9 +5,11 @@ package oauth
import ( import (
"context" "context"
"net/http"
"sync" "sync"
"github.com/coreos/go-oidc/v3/oidc" "github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2"
"golang.org/x/sync/singleflight" "golang.org/x/sync/singleflight"
) )
@@ -33,12 +35,19 @@ var globalOIDCProviderCache = &oidcProviderCache{
entries: make(map[string]*oidc.Provider), 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 获取并写入缓存。 // get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。
// 同一 issuer 并发调用时,singleflight 保证只有一次实际 HTTP 请求。 // 同一 issuer 并发调用时,singleflight 保证只有一次实际 HTTP 请求。
// func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provider, error) {
// 注意:接受 _ 形参以与调用方类型一致,内部有意使用 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
// 快路径:已有缓存则直接返回。 // 快路径:已有缓存则直接返回。
c.mu.RLock() c.mu.RLock()
if p, ok := c.entries[issuer]; ok { if p, ok := c.entries[issuer]; ok {
@@ -48,8 +57,8 @@ func (c *oidcProviderCache) get(_ context.Context, issuer string) (*oidc.Provide
c.mu.RUnlock() c.mu.RUnlock()
// 慢路径:通过 singleflight 合并并发的首次请求。 // 慢路径:通过 singleflight 合并并发的首次请求。
// 闭包内有意使用 context.Background() 而非请求 ctx,防止请求取消导致缓存写入失败。 discCtx := discoveryContext(ctx)
v, err, _ := c.sfGroup.Do(issuer, func() (any, error) { //nolint:contextcheck // intentional: Background ctx prevents cache write failure on request cancellation v, err, _ := c.sfGroup.Do(issuer, func() (any, error) {
// 双检:singleflight 内再次检查,前一个并发组可能已写入缓存。 // 双检:singleflight 内再次检查,前一个并发组可能已写入缓存。
c.mu.RLock() c.mu.RLock()
if p, ok := c.entries[issuer]; ok { if p, ok := c.entries[issuer]; ok {
@@ -58,7 +67,7 @@ func (c *oidcProviderCache) get(_ context.Context, issuer string) (*oidc.Provide
} }
c.mu.RUnlock() c.mu.RUnlock()
p, err := oidc.NewProvider(context.Background(), issuer) p, err := oidc.NewProvider(discCtx, issuer)
if err != nil { if err != nil {
return nil, err return nil, err
} }
+8 -15
View File
@@ -48,11 +48,9 @@ func setupUserTestRouter(t *testing.T) *gin.Engine {
config.Config.App.SessionSecure = false config.Config.App.SessionSecure = false
config.Config.App.SessionHTTPOnly = true config.Config.App.SessionHTTPOnly = true
gin.SetMode(gin.TestMode)
r := gin.New()
store := cookie.NewStore([]byte(config.Config.App.SessionSecret)) store := cookie.NewStore([]byte(config.Config.App.SessionSecret))
store.Options(oauth.GetSessionOptions(3600)) 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 := r.Group("/api/v1")
api.POST("/user/register", Register) api.POST("/user/register", Register)
@@ -298,8 +296,8 @@ func TestLoginEmailVerificationFallbackWhenSMTPUnconfigured(t *testing.T) {
body, _ := json.Marshal(payload) body, _ := json.Marshal(payload)
w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil) w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil)
if w.Code != http.StatusOK { if w.Code != http.StatusBadRequest {
t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusOK, w.Body.String()) t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusBadRequest, w.Body.String())
} }
// Check response error msg // Check response error msg
@@ -404,8 +402,8 @@ func TestLoginEmailVerificationFallbackForEmptyEmail(t *testing.T) {
body, _ := json.Marshal(payload) body, _ := json.Marshal(payload)
w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil) w := performUserRequest(router, http.MethodPost, "/api/v1/user/login", body, nil)
if w.Code != http.StatusOK { if w.Code != http.StatusBadRequest {
t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusOK, w.Body.String()) t.Fatalf("Login(%q) status = %d, want %d. Body: %s", username, w.Code, http.StatusBadRequest, w.Body.String())
} }
// Check response error msg // Check response error msg
@@ -495,10 +493,8 @@ func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) {
} }
// 2. Set up router with access-token routes and oauth middlewares // 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")) 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") apiV1Router := r.Group("/api/v1")
userRouter := apiV1Router.Group("/user") userRouter := apiV1Router.Group("/user")
@@ -533,8 +529,7 @@ func TestAccessTokenEndpointsDisallowTokenAuth(t *testing.T) {
// 4. Test that accessing using a Session succeeds // 4. Test that accessing using a Session succeeds
sessionCookieStore := cookie.NewStore([]byte("test_session_secret")) sessionCookieStore := cookie.NewStore([]byte("test_session_secret"))
rSession := gin.New() rSession := testhelper.NewTestGinEngine(sessions.Sessions("test_session_id", sessionCookieStore))
rSession.Use(sessions.Sessions("test_session_id", sessionCookieStore))
rSession.GET("/api/v1/user/access-tokens", oauth.LoginRequired(), oauth.DisallowTokenAuth(), ListAccessTokens) 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 // We can login/register or just mock the session handler to set user ID
@@ -597,10 +592,8 @@ func TestChangePasswordRevocation(t *testing.T) {
} }
// 3. Set up router // 3. Set up router
gin.SetMode(gin.TestMode)
r := gin.New()
store := cookie.NewStore([]byte("test_session_secret")) 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) { r.GET("/mock-login", func(c *gin.Context) {
session := sessions.Default(c) session := sessions.Default(c)
+5
View File
@@ -85,4 +85,9 @@ func Init(ctx context.Context, opts Options) {
risk_control.InitLogWriter(ctx) 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) { func TestInitSyncsPushEventsOnce(t *testing.T) {
ResetInitRuntimeOnceForTest()
t.Cleanup(ResetInitRuntimeOnceForTest)
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup() defer cleanup()
@@ -28,8 +31,8 @@ func TestInitSyncsPushEventsOnce(t *testing.T) {
} }
ctx := context.Background() ctx := context.Background()
Init(ctx, Options{}) Init(ctx, Options{API: true})
Init(ctx, Options{API: true}) // second Init must not duplicate events (initRuntimeOnce) Init(ctx, Options{}) // second Init must not duplicate events (initRuntimeOnce)
var count int64 var count int64
if err := dbConn.Model(&model.PushEvent{}).Count(&count).Error; err != nil { 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())
}