mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-10-01 06:36:38 +08:00
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:
@@ -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 职责划分
|
||||
|
||||
@@ -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("<prefix>", &cfg)` 读取**自己声明**的配置,字段以 tag 表达来源:`config`(yaml 路径)、`env`(覆盖变量名)、`default`、`autoEnable`(该变量存在即置真)、`secret`(导出脱敏)。需要在 `Apply` 之前被门禁求值的键,必须在 `DeclareConfig()` 中提前声明并实现 `core.ConfigGatedPlugin`。新增基础设施 key 保持顶层命名(`redis.*`),插件私有配置归 `plugins.<name>.*`。**严禁**再造全局配置单例或在 `backend/pkg/` 读取配置。
|
||||
- **动态设置**:插件自包含在 `Apply` 中通过 `ctx.Settings().Register(core.SettingSchema{...})` 声明可热更新的管理台设置模式(与上面的静态启动配置分属两层)。
|
||||
|
||||
@@ -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; \
|
||||
|
||||
+1
-1
@@ -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.
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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"`
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
Reference in New Issue
Block a user