mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-11 09:46: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")
|
"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
|
||||||
|
|||||||
@@ -98,8 +98,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) {
|
||||||
@@ -656,7 +655,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)
|
||||||
|
|||||||
@@ -27,8 +27,7 @@ import ("bufio"
|
|||||||
const expectedDefaultConfigsCount = 30
|
const expectedDefaultConfigsCount = 30
|
||||||
|
|
||||||
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
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -44,8 +44,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
|
||||||
@@ -93,18 +92,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)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -117,7 +116,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))
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -20,8 +20,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
|
||||||
|
|||||||
@@ -23,8 +23,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")
|
||||||
|
|||||||
@@ -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) {
|
func seedTestAuthSource(t *testing.T, dbConn *gorm.DB) {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
if err := dbConn.Create(&model.AuthSource{
|
if err := dbConn.Create(&model.AuthSource{
|
||||||
@@ -193,6 +197,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,
|
||||||
@@ -201,7 +206,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)),
|
||||||
@@ -254,7 +259,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) {
|
||||||
@@ -288,7 +293,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,
|
||||||
@@ -303,6 +308,8 @@ func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, user
|
|||||||
}
|
}
|
||||||
|
|
||||||
func setupTestDB(t *testing.T) *gorm.DB {
|
func setupTestDB(t *testing.T) *gorm.DB {
|
||||||
|
model.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)
|
||||||
@@ -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 {
|
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
|
||||||
@@ -454,6 +468,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")
|
||||||
|
model.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)
|
||||||
@@ -1094,6 +1109,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
Key: model.ConfigKeyOIDCLoginEnabled,
|
Key: model.ConfigKeyOIDCLoginEnabled,
|
||||||
Value: "false",
|
Value: "false",
|
||||||
})
|
})
|
||||||
|
model.ResetSystemConfigRAMCacheForTest()
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(model.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 {
|
||||||
@@ -1102,6 +1118,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")
|
||||||
|
model.ResetSystemConfigRAMCacheForTest()
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(model.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)
|
||||||
|
|
||||||
@@ -1113,6 +1130,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")
|
||||||
|
model.ResetSystemConfigRAMCacheForTest()
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(model.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)
|
||||||
|
|
||||||
@@ -1124,6 +1142,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")
|
||||||
|
model.ResetSystemConfigRAMCacheForTest()
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(model.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)
|
||||||
|
|
||||||
@@ -1149,6 +1168,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")
|
||||||
|
model.ResetSystemConfigRAMCacheForTest()
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(model.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{
|
||||||
@@ -1160,6 +1180,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")
|
||||||
|
model.ResetSystemConfigRAMCacheForTest()
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(model.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(model.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)
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -45,11 +45,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)
|
||||||
@@ -295,8 +293,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
|
||||||
@@ -401,8 +399,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
|
||||||
@@ -492,10 +490,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")
|
||||||
@@ -530,8 +526,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
|
||||||
@@ -594,10 +589,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)
|
||||||
|
|||||||
@@ -86,3 +86,8 @@ func Init(ctx context.Context, opts Options) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 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) {
|
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 {
|
||||||
|
|||||||
@@ -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