feat(router): add whitelist mechanism for http driver and auth plugin

- implement route whitelist registration and wildcard matching in RouterExtension
- add cookie store session fallback when Redis is disabled in driver_http
- actively register public auth endpoints to whitelist in auth plugin
- update user handlers to persist session and clear cookie on logout
- document router whitelist mechanism in AGENTS.md and new-api skill
This commit is contained in:
ryan
2026-08-29 11:39:13 +08:00
parent 53ae3007d0
commit e0f2309520
17 changed files with 411 additions and 32 deletions
+31
View File
@@ -5,6 +5,7 @@ package auth
import (
"Wavelet/core/contracts"
"Wavelet/core/extpoints"
"Wavelet/pkg/ginutil"
"Wavelet/pkg/response"
"Wavelet/pkg/trace"
@@ -12,10 +13,35 @@ import (
"crypto/sha256"
"encoding/hex"
"errors"
"sync"
"github.com/gin-gonic/gin"
)
var (
whitelistMu sync.RWMutex
whitelist []string
)
// RegisterWhitelist registers route patterns that bypass mandatory authentication.
func RegisterWhitelist(patterns ...string) {
whitelistMu.Lock()
defer whitelistMu.Unlock()
whitelist = append(whitelist, patterns...)
}
// IsWhitelisted checks if the specified path matches the auth whitelist.
func IsWhitelisted(path string) bool {
whitelistMu.RLock()
defer whitelistMu.RUnlock()
for _, pattern := range whitelist {
if extpoints.MatchPathPattern(pattern, path) {
return true
}
}
return false
}
func hashToken(token string) string {
h := sha256.New()
h.Write([]byte(token))
@@ -112,6 +138,11 @@ func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) {
// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session
func LoginRequired() gin.HandlerFunc {
return func(c *gin.Context) {
if IsWhitelisted(c.Request.URL.Path) {
c.Next()
return
}
_, span := trace.Start(c.Request.Context(), "LoginRequired")
defer span.End()
+18
View File
@@ -122,6 +122,24 @@ func (p *Plugin) Apply(ctx *core.Context) error {
core.Provide[contracts.AuthService](ctx, p.authSvc)
core.Provide[contracts.AuthRegistry](ctx, p.authRegistry)
// 2.1 Register Public / Auth Whitelist Endpoints
publicEndpoints := []string{
"/api/v1/oauth/sources",
"/api/v1/oauth/login",
"/api/v1/oauth/*/authorize",
"/api/v1/oauth/:source/authorize",
"/api/v1/oauth/callback",
"/api/v1/user/login",
"/api/v1/user/register",
"/api/v1/user/send-email-code",
"/api/v1/cap/challenge",
"/api/v1/cap/redeem",
"/healthz",
"/metrics",
}
RegisterWhitelist(publicEndpoints...)
ctx.Router().RegisterWhitelist(publicEndpoints...)
// 3. Register HTTP Routes
oauthGroup := ctx.Router().Group("/api/v1/oauth")
{
@@ -365,3 +365,37 @@ func TestLoginStateContextParityWithLegacyImplementation(t *testing.T) {
}
})
}
func TestAuthWhitelistMiddleware(t *testing.T) {
gin.SetMode(gin.TestMode)
ctx := core.NewContext(context.Background())
p := auth.New()
require.NoError(t, p.Apply(ctx))
svc, err := core.Inject[contracts.AuthService](ctx)
require.NoError(t, err)
mw, ok := svc.RequireAuthMiddleware().(gin.HandlerFunc)
require.True(t, ok)
engine := newSessionEngine()
engine.Use(mw)
engine.POST("/api/v1/user/login", func(c *gin.Context) {
c.JSON(http.StatusOK, response.OK("login-ok"))
})
engine.GET("/api/v1/secret-profile", func(c *gin.Context) {
c.JSON(http.StatusOK, response.OK("profile-ok"))
})
// 1. Whitelisted route /api/v1/user/login passes through without auth
w1 := httptest.NewRecorder()
req1, _ := http.NewRequest(http.MethodPost, "/api/v1/user/login", nil)
engine.ServeHTTP(w1, req1)
assert.Equal(t, http.StatusOK, w1.Code)
// 2. Non-whitelisted route /api/v1/secret-profile gets 401 Unauthorized
w2 := httptest.NewRecorder()
req2, _ := http.NewRequest(http.MethodGet, "/api/v1/secret-profile", nil)
engine.ServeHTTP(w2, req2)
assert.Equal(t, http.StatusUnauthorized, w2.Code)
}
+17 -2
View File
@@ -79,7 +79,9 @@ func Login(c *gin.Context) {
sess := sessions.Default(c)
sess.Set(contracts.AuthUserIDKey, user.ID)
sess.Set(contracts.AuthUserNameKey, user.Username)
_ = sess.Save()
if err := sess.Save(); err != nil {
logger.ErrorF(c.Request.Context(), "save session failed on login: %v", err)
}
c.JSON(http.StatusOK, response.OK(user))
}
@@ -107,14 +109,27 @@ func Register(c *gin.Context) {
return
}
sess := sessions.Default(c)
sess.Set(contracts.AuthUserIDKey, newUser.ID)
sess.Set(contracts.AuthUserNameKey, newUser.Username)
if err := sess.Save(); err != nil {
logger.ErrorF(c.Request.Context(), "save session failed on register: %v", err)
}
c.JSON(http.StatusOK, response.OK(newUser))
}
// Logout logs out the current session.
func Logout(c *gin.Context) {
sess := sessions.Default(c)
sess.Options(sessions.Options{
Path: "/",
MaxAge: -1,
})
sess.Clear()
_ = sess.Save()
if err := sess.Save(); err != nil {
logger.ErrorF(c.Request.Context(), "clear session failed on logout: %v", err)
}
c.JSON(http.StatusOK, response.OKNil())
}
@@ -8,10 +8,16 @@ import (
"Wavelet/core/contracts"
"Wavelet/pkg/idgen"
"Wavelet/plugins/domain/user"
"bytes"
"context"
"net/http"
"net/http/httptest"
"path/filepath"
"testing"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -111,3 +117,44 @@ func TestUserPluginUnit(t *testing.T) {
assert.GreaterOrEqual(t, total, int64(1))
assert.NotEmpty(t, list)
}
func TestUserLoginHTTPHandler(t *testing.T) {
ctx := core.NewContext(context.Background())
ctx.Config().SetSource(core.NewMapSource(nil))
require.NoError(t, ctx.Config().Resolve())
testDB := setupTestDB(t)
dbPlugin := database.New(database.WithDB(testDB))
require.NoError(t, dbPlugin.Apply(ctx))
p := user.New()
require.NoError(t, p.Apply(ctx))
userSvc, err := core.Inject[contracts.UserService](ctx)
require.NoError(t, err)
_, err = userSvc.CreateUser(context.Background(), contracts.CreateUserRequest{
Username: "admin",
Password: "Password123!",
Email: "admin@example.com",
})
require.NoError(t, err)
r := gin.New()
cookieStore := cookie.NewStore([]byte("test-session-secret"))
r.Use(sessions.Sessions("wavelet_session", cookieStore))
r.POST("/api/v1/user/login", user.Login)
reqBody := `{"username":"admin","password":"Password123!"}`
req, _ := http.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewBufferString(reqBody))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
r.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
assert.Contains(t, w.Body.String(), `"username":"admin"`)
setCookie := w.Header().Get("Set-Cookie")
assert.NotEmpty(t, setCookie)
assert.Contains(t, setCookie, "wavelet_session=")
}
@@ -18,6 +18,7 @@ type httpAppConfig struct {
}
type httpRedisConfig struct {
Enabled bool `config:"enabled" env:"REDIS_ENABLED" default:"false"`
Addrs []string `config:"addrs" env:"REDIS_ADDR"`
Username string `config:"username" env:"REDIS_USERNAME"`
Password string `config:"password" env:"REDIS_PASSWORD" secret:"true"`
+31 -25
View File
@@ -9,6 +9,7 @@ import (
"strconv"
"github.com/gin-contrib/sessions"
"github.com/gin-contrib/sessions/cookie"
"github.com/gin-contrib/sessions/redis"
"github.com/gin-gonic/gin"
"go.opentelemetry.io/contrib/instrumentation/github.com/gin-gonic/gin/otelgin"
@@ -33,36 +34,12 @@ func BuildEngineWithConfig(appCfg httpAppConfig, redisCfg httpRedisConfig) (*gin
r.Use(gin.Recovery())
r.Use(corsMiddleware())
addrs := redisCfg.Addrs
sessionAddr := "localhost:6379"
if len(addrs) > 0 {
sessionAddr = addrs[0]
}
sessionSecret := appCfg.SessionSecret
if sessionSecret == "" {
sessionSecret = "wavelet-default-session-secret"
}
sessionStore, err := redis.NewStoreWithDB(
redisCfg.MinIdleConn,
"tcp",
sessionAddr,
redisCfg.Username,
redisCfg.Password,
strconv.Itoa(redisCfg.DB),
[]byte(sessionSecret),
)
if err != nil {
return nil, err
}
// 设置 Session Redis Key 前缀
if redisCfg.KeyPrefix != "" {
if err := redis.SetKeyPrefix(sessionStore, redisCfg.KeyPrefix+"session:"); err != nil {
log.Printf("[API] set session key prefix failed: %v\n", err)
}
}
sessionStore := initSessionStore(sessionSecret, redisCfg)
sessionCookieName := appCfg.SessionCookieName
if sessionCookieName == "" {
@@ -95,3 +72,32 @@ func BuildEngineWithConfig(appCfg httpAppConfig, redisCfg httpRedisConfig) (*gin
return r, nil
}
func initSessionStore(sessionSecret string, redisCfg httpRedisConfig) sessions.Store {
if !redisCfg.Enabled || len(redisCfg.Addrs) == 0 {
return cookie.NewStore([]byte(sessionSecret))
}
sessionAddr := redisCfg.Addrs[0]
store, err := redis.NewStoreWithDB(
redisCfg.MinIdleConn,
"tcp",
sessionAddr,
redisCfg.Username,
redisCfg.Password,
strconv.Itoa(redisCfg.DB),
[]byte(sessionSecret),
)
if err != nil {
log.Printf("[driver_http] init redis session store failed, fallback to cookie store: %v\n", err)
return cookie.NewStore([]byte(sessionSecret))
}
if redisCfg.KeyPrefix != "" {
if err := redis.SetKeyPrefix(store, redisCfg.KeyPrefix+"session:"); err != nil {
log.Printf("[API] set session key prefix failed: %v\n", err)
}
}
return store
}
@@ -5,6 +5,7 @@
package driver_http
import (
"Wavelet/core/extpoints"
"Wavelet/pkg/logger"
"Wavelet/pkg/response"
"context"
@@ -24,8 +25,31 @@ import (
var (
apiPrefixMu sync.RWMutex
apiPrefix = "/api/v1"
whitelistMu sync.RWMutex
whitelistPatterns []string
)
// SetWhitelist configures global whitelist patterns for HTTP routes.
func SetWhitelist(patterns []string) {
whitelistMu.Lock()
defer whitelistMu.Unlock()
whitelistPatterns = make([]string, len(patterns))
copy(whitelistPatterns, patterns)
}
// IsPathWhitelisted checks if the given path matches any registered whitelist pattern.
func IsPathWhitelisted(path string) bool {
whitelistMu.RLock()
defer whitelistMu.RUnlock()
for _, pattern := range whitelistPatterns {
if extpoints.MatchPathPattern(pattern, path) {
return true
}
}
return false
}
func setAPIPrefix(prefix string) {
if prefix == "" {
return
@@ -12,6 +12,7 @@ import (
"testing"
"time"
"github.com/gin-contrib/sessions"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
@@ -205,3 +206,40 @@ func TestCORSAllowedOriginReadsConfigOncePerCacheWindow(t *testing.T) {
t.Errorf("expected 1 config load across 3 requests, got %d", cache.loads)
}
}
func TestBuildEngineWithCookieFallback(t *testing.T) {
gin.SetMode(gin.TestMode)
appCfg := httpAppConfig{
SessionSecret: "test-secret",
SessionCookieName: "wavelet_session",
SessionAge: 3600,
}
redisCfg := httpRedisConfig{
Enabled: false,
}
engine, err := BuildEngineWithConfig(appCfg, redisCfg)
if err != nil {
t.Fatalf("expected nil error with cookie fallback, got: %v", err)
}
engine.GET("/set-session", func(c *gin.Context) {
sess := sessions.Default(c)
sess.Set("test_user_id", uint64(12345))
_ = sess.Save()
c.String(http.StatusOK, "ok")
})
req, _ := http.NewRequest(http.MethodGet, "/set-session", nil)
w := httptest.NewRecorder()
engine.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("expected 200 OK, got %d", w.Code)
}
setCookie := w.Header().Get("Set-Cookie")
if setCookie == "" {
t.Fatal("expected Set-Cookie header in response, got none")
}
}
@@ -177,12 +177,13 @@ func (p *Plugin) Start(ctx context.Context) error {
var err error
p.engine, err = BuildEngineWithConfig(appCfg, redisCfg)
if err != nil {
p.engine = gin.New()
p.engine, _ = BuildEngine()
}
}
// Mount routes collected in Context RouterExtension
if p.coreCtx != nil && p.coreCtx.Router() != nil {
SetWhitelist(p.coreCtx.Router().Whitelist())
for _, rd := range p.coreCtx.Router().Routes() {
allHandlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers))