From e0f230952086fbd5d35042188a783d84085decf3 Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 29 Aug 2026 11:39:13 +0800 Subject: [PATCH] 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 --- .agents/skills/new-api/SKILL.md | 20 ++++ AGENTS.md | 6 +- Makefile | 4 +- backend/cmd/app.go | 2 +- backend/core/extpoints/extpoints_test.go | 34 +++++++ backend/core/extpoints/router.go | 94 +++++++++++++++++++ backend/core/scoped_extpoints.go | 12 +++ backend/plugins/domain/auth/middleware.go | 31 ++++++ backend/plugins/domain/auth/plugin.go | 18 ++++ backend/plugins/domain/auth/service_test.go | 34 +++++++ backend/plugins/domain/user/handlers.go | 19 +++- backend/plugins/domain/user/plugin_test.go | 47 ++++++++++ backend/plugins/drivers/driver_http/config.go | 1 + backend/plugins/drivers/driver_http/engine.go | 56 ++++++----- .../drivers/driver_http/middlewares.go | 24 +++++ .../drivers/driver_http/middlewares_test.go | 38 ++++++++ backend/plugins/drivers/driver_http/plugin.go | 3 +- 17 files changed, 411 insertions(+), 32 deletions(-) diff --git a/.agents/skills/new-api/SKILL.md b/.agents/skills/new-api/SKILL.md index f56ecfb1..5093307b 100644 --- a/.agents/skills/new-api/SKILL.md +++ b/.agents/skills/new-api/SKILL.md @@ -120,6 +120,26 @@ func (p *Plugin) registerRoutes(ctx *core.Context) { } ``` +### 步骤 3:公开接口与白名单注册 (`RegisterWhitelist`) + +如果插件包含**无需登录**的公开端点(如登录、注册、人机校验、Webhooks、公开状态查询),必须在 `Apply` 中主动注册到白名单: + +```go +func (p *Plugin) Apply(ctx *core.Context) error { + // 注册公开接口白名单(支持精确路径与通配符如 /api/v1/oauth/*) + ctx.Router().RegisterWhitelist( + "/api/v1/public/ping", + "/api/v1/public/webhook/*", + ) + + // 或在子路由组中相对注册: + publicGroup := ctx.Router().Group("/api/v1/public") + publicGroup.RegisterWhitelist("/status", "/docs/*") + ... +} +``` +> 💡 **防线机制**:注册到白名单的路由在经过 `auth.RequireAuthMiddleware()` 时将自动放行,彻底消除全局/组级鉴权中间件引起的 401 Unauthorized 误拦截。 + --- ## 3. Handler 与 Service 职责划分 diff --git a/AGENTS.md b/AGENTS.md index 09cdf6e2..8997eea6 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -109,7 +109,11 @@ Strong success criteria let you loop independently. Weak criteria ("make it work - **单向服务契约调用**:调用方仅面向 `backend/core/contracts` 编程,在 `Apply` 中通过 `core.Provide[contracts.XxxService](ctx, svc)` 注册服务,通过 `core.Inject[contracts.XxxService](ctx)` 或 `ctx.Using(func(svc contracts.XxxService) { ... })` 声明式解析。 - **事件总线广播**:状态联动与解耦通信统一通过强类型事件 `ctx.Events().Emit()` 广播,由感兴趣的插件通过 `ctx.Events().On()` 订阅,消除双向依赖与循环引用。 - **扩展点自包含注册**: - - **HTTP 路由**:插件自包含在 `Apply` 中通过 `ctx.Router().Group(...)` 挂载路由与中间件,禁止跨插件散落注册。 + - **HTTP 路由与白名单机制**: + - 插件自包含在 `Apply` 中通过 `ctx.Router().Group(...)` 挂载路由与中间件,禁止跨插件散落注册。 + - **白名单机制**:`driver_http` 与微内核扩展点提供路由白名单支持(`ctx.Router().RegisterWhitelist(patterns...)`),支持精确路径与通配符(如 `/api/v1/oauth/*`)。 + - **所有权主动声明**:认证域(`auth` 插件)与各业务插件必须在 `Apply` 中主动注册其公开/免鉴权接口(如 `/api/v1/user/login`、`/api/v1/oauth/callback`、`/api/v1/cap/*` 等)。 + - **鉴权中间件放行防线**:`auth` 提供的登录鉴权中间件(`LoginRequired`)必须先执行白名单匹配并自动放行,彻底杜绝免鉴权接口被全局或组级鉴权中间件误拦截(返回 401 Unauthorized)。 - **异步与定时任务**:插件自包含在 `Apply` 中通过 `ctx.Task().Register(...)` 与 `ctx.Schedule().RegisterCron(...)` 声明。 - **静态启动配置**:插件自包含在 `Apply` 中通过 `ctx.Config().Bind("", &cfg)` 读取**自己声明**的配置,字段以 tag 表达来源:`config`(yaml 路径)、`env`(覆盖变量名)、`default`、`autoEnable`(该变量存在即置真)、`secret`(导出脱敏)。需要在 `Apply` 之前被门禁求值的键,必须在 `DeclareConfig()` 中提前声明并实现 `core.ConfigGatedPlugin`。新增基础设施 key 保持顶层命名(`redis.*`),插件私有配置归 `plugins..*`。**严禁**再造全局配置单例或在 `backend/pkg/` 读取配置。 - **动态设置**:插件自包含在 `Apply` 中通过 `ctx.Settings().Register(core.SettingSchema{...})` 声明可热更新的管理台设置模式(与上面的静态启动配置分属两层)。 diff --git a/Makefile b/Makefile index 46e21363..a6c323d7 100644 --- a/Makefile +++ b/Makefile @@ -93,14 +93,14 @@ dev-f: dev-b: @echo "==> Starting backend development server..." - go run main.go all + cd backend && go run main.go all dev: @echo "==> Starting frontend and backend development servers in parallel..." @PIDS=""; \ STATUS=0; \ ( cd frontend && pnpm dev 2>&1 | sed 's/^/[frontend] /' ) & PIDS="$$PIDS $$!"; \ - ( go run main.go all 2>&1 | sed 's/^/[backend] /' ) & PIDS="$$PIDS $$!"; \ + ( cd backend && go run main.go all 2>&1 | sed 's/^/[backend] /' ) & PIDS="$$PIDS $$!"; \ for PID in $$PIDS; do \ wait $$PID || STATUS=1; \ done; \ diff --git a/backend/cmd/app.go b/backend/cmd/app.go index 5d5321d6..179744ac 100644 --- a/backend/cmd/app.go +++ b/backend/cmd/app.go @@ -40,7 +40,7 @@ import ( const ( defaultShutdownTimeout = 15 * time.Second - defaultHTTPAddr = "127.0.0.1:3000" + defaultHTTPAddr = "127.0.0.1:8000" ) // runProfileApp prepares and runs the application for a given profile. diff --git a/backend/core/extpoints/extpoints_test.go b/backend/core/extpoints/extpoints_test.go index 1788b91a..adfc3e85 100644 --- a/backend/core/extpoints/extpoints_test.go +++ b/backend/core/extpoints/extpoints_test.go @@ -93,6 +93,40 @@ func TestRouterExtension(t *testing.T) { assert.True(t, foundUserPut) } +func TestRouterWhitelist(t *testing.T) { + r := extpoints.NewRouterRegistry() + require.NotNil(t, r) + + r.RegisterWhitelist( + "/healthz", + "/api/v1/user/login", + "/api/v1/oauth/*", + ) + + api := r.Group("/api/v1") + api.RegisterWhitelist("/cap/challenge", "/cap/redeem") + + whitelist := r.Whitelist() + assert.Contains(t, whitelist, "/healthz") + assert.Contains(t, whitelist, "/api/v1/user/login") + assert.Contains(t, whitelist, "/api/v1/oauth/*") + assert.Contains(t, whitelist, "/api/v1/cap/challenge") + assert.Contains(t, whitelist, "/api/v1/cap/redeem") + + // Exact match + assert.True(t, r.IsWhitelisted("/healthz")) + assert.True(t, r.IsWhitelisted("/api/v1/user/login")) + assert.True(t, api.IsWhitelisted("/api/v1/cap/challenge")) + + // Wildcard match + assert.True(t, r.IsWhitelisted("/api/v1/oauth/sources")) + assert.True(t, r.IsWhitelisted("/api/v1/oauth/github/authorize")) + + // Non-whitelisted + assert.False(t, r.IsWhitelisted("/api/v1/orders")) + assert.False(t, r.IsWhitelisted("/api/v1/user/profile")) +} + func TestMigrationExtension(t *testing.T) { m := extpoints.NewMigrationRegistry() require.NotNil(t, m) diff --git a/backend/core/extpoints/router.go b/backend/core/extpoints/router.go index e73db66f..2debe0bd 100644 --- a/backend/core/extpoints/router.go +++ b/backend/core/extpoints/router.go @@ -34,6 +34,9 @@ type RouterExtension interface { Middlewares() []any Unregister(method, path string) bool UnregisterByID(id uint64) bool + RegisterWhitelist(patterns ...string) + Whitelist() []string + IsWhitelisted(path string) bool } // RouterRegistry implements RouterExtension as the root route and middleware collector. @@ -42,6 +45,7 @@ type RouterRegistry struct { nextID uint64 routes []RouteDefinition middlewares []any + whitelist []string } // NewRouterRegistry creates a new root router collector. @@ -176,6 +180,40 @@ func (r *RouterRegistry) Routes() []RouteDefinition { return res } +// RegisterWhitelist adds path patterns to the whitelist. +func (r *RouterRegistry) RegisterWhitelist(patterns ...string) { + r.mu.Lock() + defer r.mu.Unlock() + for _, p := range patterns { + clean := cleanPath(p) + if clean != "" { + r.whitelist = append(r.whitelist, clean) + } + } +} + +// Whitelist returns a copy of all registered whitelist path patterns. +func (r *RouterRegistry) Whitelist() []string { + r.mu.RLock() + defer r.mu.RUnlock() + res := make([]string, len(r.whitelist)) + copy(res, r.whitelist) + return res +} + +// IsWhitelisted checks if the given path matches any registered whitelist pattern. +func (r *RouterRegistry) IsWhitelisted(path string) bool { + r.mu.RLock() + defer r.mu.RUnlock() + clean := cleanPath(path) + for _, pattern := range r.whitelist { + if MatchPathPattern(pattern, clean) { + return true + } + } + return false +} + // RouterGroup represents a scoped route group with a path prefix and group-level middlewares. type RouterGroup struct { registry *RouterRegistry @@ -293,6 +331,23 @@ func (g *RouterGroup) Middlewares() []any { return res } +// RegisterWhitelist adds path patterns under this group prefix to the whitelist. +func (g *RouterGroup) RegisterWhitelist(patterns ...string) { + for _, p := range patterns { + g.registry.RegisterWhitelist(joinPaths(g.prefix, p)) + } +} + +// Whitelist returns a copy of all registered whitelist path patterns. +func (g *RouterGroup) Whitelist() []string { + return g.registry.Whitelist() +} + +// IsWhitelisted checks if the given path matches any registered whitelist pattern. +func (g *RouterGroup) IsWhitelisted(path string) bool { + return g.registry.IsWhitelisted(path) +} + func cleanPath(p string) string { if p == "" { return "/" @@ -317,3 +372,42 @@ func joinPaths(base, relative string) string { relative = strings.TrimPrefix(relative, "/") return cleanPath(base + "/" + relative) } + +// MatchPathPattern checks if a URL path matches a pattern (supports exact match and wildcards). +func MatchPathPattern(pattern, path string) bool { + pattern = cleanPath(pattern) + path = cleanPath(path) + + if pattern == path { + return true + } + + // Suffix wildcard: /api/v1/oauth/* matches /api/v1/oauth and /api/v1/oauth/... + if strings.HasSuffix(pattern, "/*") { + prefix := strings.TrimSuffix(pattern, "/*") + if path == prefix || strings.HasPrefix(path, prefix+"/") { + return true + } + } + + // Parameter wildcard: /api/v1/oauth/*/authorize or /api/v1/oauth/:source/authorize + patternParts := strings.Split(pattern, "/") + pathParts := strings.Split(path, "/") + if len(patternParts) == len(pathParts) { + matched := true + for i, part := range patternParts { + if part == "*" || strings.HasPrefix(part, ":") { + continue + } + if part != pathParts[i] { + matched = false + break + } + } + if matched { + return true + } + } + + return false +} diff --git a/backend/core/scoped_extpoints.go b/backend/core/scoped_extpoints.go index 5db27401..9bd4d923 100644 --- a/backend/core/scoped_extpoints.go +++ b/backend/core/scoped_extpoints.go @@ -93,6 +93,18 @@ func (s *scopedRouterExtension) UnregisterByID(id uint64) bool { return s.underlying.UnregisterByID(id) } +func (s *scopedRouterExtension) RegisterWhitelist(patterns ...string) { + s.underlying.RegisterWhitelist(patterns...) +} + +func (s *scopedRouterExtension) Whitelist() []string { + return s.underlying.Whitelist() +} + +func (s *scopedRouterExtension) IsWhitelisted(path string) bool { + return s.underlying.IsWhitelisted(path) +} + // scopedTaskExtension wraps a TaskExtension to automatically register // teardown disposers on the associated Context when task handlers are declared. type scopedTaskExtension struct { diff --git a/backend/plugins/domain/auth/middleware.go b/backend/plugins/domain/auth/middleware.go index db551ef9..49ed23fa 100644 --- a/backend/plugins/domain/auth/middleware.go +++ b/backend/plugins/domain/auth/middleware.go @@ -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() diff --git a/backend/plugins/domain/auth/plugin.go b/backend/plugins/domain/auth/plugin.go index fdc35ab2..2e98a15e 100644 --- a/backend/plugins/domain/auth/plugin.go +++ b/backend/plugins/domain/auth/plugin.go @@ -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") { diff --git a/backend/plugins/domain/auth/service_test.go b/backend/plugins/domain/auth/service_test.go index c000b58f..2431ec67 100644 --- a/backend/plugins/domain/auth/service_test.go +++ b/backend/plugins/domain/auth/service_test.go @@ -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) +} diff --git a/backend/plugins/domain/user/handlers.go b/backend/plugins/domain/user/handlers.go index 3646820e..e306cb53 100644 --- a/backend/plugins/domain/user/handlers.go +++ b/backend/plugins/domain/user/handlers.go @@ -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()) } diff --git a/backend/plugins/domain/user/plugin_test.go b/backend/plugins/domain/user/plugin_test.go index 387d2757..1d471c13 100644 --- a/backend/plugins/domain/user/plugin_test.go +++ b/backend/plugins/domain/user/plugin_test.go @@ -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=") +} diff --git a/backend/plugins/drivers/driver_http/config.go b/backend/plugins/drivers/driver_http/config.go index dea5144f..df2046c1 100644 --- a/backend/plugins/drivers/driver_http/config.go +++ b/backend/plugins/drivers/driver_http/config.go @@ -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"` diff --git a/backend/plugins/drivers/driver_http/engine.go b/backend/plugins/drivers/driver_http/engine.go index cc80a8c5..2ca969a1 100644 --- a/backend/plugins/drivers/driver_http/engine.go +++ b/backend/plugins/drivers/driver_http/engine.go @@ -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 +} diff --git a/backend/plugins/drivers/driver_http/middlewares.go b/backend/plugins/drivers/driver_http/middlewares.go index f323ce60..cb15f036 100644 --- a/backend/plugins/drivers/driver_http/middlewares.go +++ b/backend/plugins/drivers/driver_http/middlewares.go @@ -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 diff --git a/backend/plugins/drivers/driver_http/middlewares_test.go b/backend/plugins/drivers/driver_http/middlewares_test.go index 8f6b9149..03b29da7 100644 --- a/backend/plugins/drivers/driver_http/middlewares_test.go +++ b/backend/plugins/drivers/driver_http/middlewares_test.go @@ -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") + } +} diff --git a/backend/plugins/drivers/driver_http/plugin.go b/backend/plugins/drivers/driver_http/plugin.go index 4fcaf4ac..83d933bc 100644 --- a/backend/plugins/drivers/driver_http/plugin.go +++ b/backend/plugins/drivers/driver_http/plugin.go @@ -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))