diff --git a/backend/plugins/drivers/driver_http/cache_helper.go b/backend/plugins/drivers/driver_http/cache_helper.go new file mode 100644 index 00000000..a7da8372 --- /dev/null +++ b/backend/plugins/drivers/driver_http/cache_helper.go @@ -0,0 +1,36 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package driver_http + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "context" + "sync" +) + +var ( + cacheMu sync.RWMutex + cacheSvc contracts.CacheService +) + +func setCacheService(s contracts.CacheService) { + cacheMu.Lock() + defer cacheMu.Unlock() + cacheSvc = s +} + +// getCache resolves the cache service from a micro-kernel context, falling back +// to the instance bound during plugin registration. +func getCache(ctx context.Context) contracts.CacheService { + if c, ok := ctx.(*core.Context); ok && c != nil { + if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { + return s + } + } + cacheMu.RLock() + s := cacheSvc + cacheMu.RUnlock() + return s +} diff --git a/backend/plugins/drivers/driver_http/middlewares.go b/backend/plugins/drivers/driver_http/middlewares.go index b3de2ded..e07d995a 100644 --- a/backend/plugins/drivers/driver_http/middlewares.go +++ b/backend/plugins/drivers/driver_http/middlewares.go @@ -70,13 +70,48 @@ func loggerMiddleware() gin.HandlerFunc { } } -func isOriginAllowed(ctx context.Context, origin string) bool { - var val string +const ( + // serverAddressConfigKey 是 admin 域声明的系统配置键(跨插件不能直接引用其常量)。 + serverAddressConfigKey = "server_address" + // serverAddressCacheKey / serverAddressCacheTTL 把 CORS 允许来源的读取从每请求 + // 一次主库查询降为每 TTL 一次,TTL 与存储驱动的系统配置缓存保持一致。 + serverAddressCacheKey = "driver_http:cors:server_address" + serverAddressCacheTTL = 5 * time.Second +) + +// loadServerAddress reads the configured server address straight from system configs. +func loadServerAddress(ctx context.Context) (string, error) { db := getDB(ctx) if db == nil { - return false + return "", nil } - if err := db.Table("w_system_configs").Where("key = ?", "server_address").Pluck("value", &val).Error; err != nil || val == "" { + var val string + if err := db.Table("w_system_configs").Where("key = ?", serverAddressConfigKey).Pluck("value", &val).Error; err != nil { + return "", err + } + return val, nil +} + +// serverAddress returns the configured server address, served from the shared +// cache so CORS does not reach the database on every request. +func serverAddress(ctx context.Context) (string, error) { + cacheSvc := getCache(ctx) + if cacheSvc == nil { + return loadServerAddress(ctx) + } + + var val string + if err := cacheSvc.GetOrSet(ctx, serverAddressCacheKey, &val, serverAddressCacheTTL, func() (any, error) { + return loadServerAddress(ctx) + }); err != nil { + return "", err + } + return val, nil +} + +func isOriginAllowed(ctx context.Context, origin string) bool { + val, err := serverAddress(ctx) + if err != nil || val == "" { return false } allowedOrigins := strings.Split(val, ",") diff --git a/backend/plugins/drivers/driver_http/middlewares_test.go b/backend/plugins/drivers/driver_http/middlewares_test.go index 739ec029..8f6b9149 100644 --- a/backend/plugins/drivers/driver_http/middlewares_test.go +++ b/backend/plugins/drivers/driver_http/middlewares_test.go @@ -4,11 +4,13 @@ package driver_http import ( + "Wavelet/core/contracts" "Wavelet/pkg/testhelper" "context" "net/http" "net/http/httptest" "testing" + "time" "github.com/gin-gonic/gin" "gorm.io/gorm" @@ -120,3 +122,86 @@ func TestCORSMiddleware(t *testing.T) { } }) } + +// memoryCache 是一个可工作的内存缓存,用于统计 loader(即数据库查询)实际执行次数。 +type memoryCache struct { + values map[string]string + loads int +} + +func (c *memoryCache) Get(_ context.Context, key string, target any) error { + v, ok := c.values[key] + if !ok { + return contracts.ErrCacheMiss + } + dst, ok := target.(*string) + if !ok { + return contracts.ErrCacheMiss + } + *dst = v + return nil +} + +func (c *memoryCache) Set(_ context.Context, key string, value any, _ time.Duration) error { + v, ok := value.(string) + if !ok { + return nil + } + c.values[key] = v + return nil +} + +func (c *memoryCache) Delete(_ context.Context, key string) error { + delete(c.values, key) + return nil +} + +func (c *memoryCache) GetOrSet(_ context.Context, key string, target any, _ time.Duration, loader func() (any, error)) error { + if err := c.Get(context.Background(), key, target); err == nil { + return nil + } + raw, err := loader() + if err != nil { + return err + } + c.loads++ + v, ok := raw.(string) + if !ok { + return nil + } + c.values[key] = v + if dst, ok := target.(*string); ok { + *dst = v + } + return nil +} + +func (c *memoryCache) Invalidate(ctx context.Context, key string) error { + return c.Delete(ctx, key) +} + +// TestCORSAllowedOriginReadsConfigOncePerCacheWindow 回归:CORS 的来源校验不得在 +// 每个请求上都查询主库,缓存有效期内 loader 只应执行一次。 +func TestCORSAllowedOriginReadsConfigOncePerCacheWindow(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + setDBService(&mockDBService{db: dbConn}) + cache := &memoryCache{values: map[string]string{}} + setCacheService(cache) + defer func() { + setCacheService(nil) + setDBService(nil) + cleanup() + }() + + ctx := context.Background() + const origin = "http://localhost:8000" // testhelper 预置的 server_address + for range 3 { + if !isOriginAllowed(ctx, origin) { + t.Fatalf("origin %q should be allowed", origin) + } + } + + if cache.loads != 1 { + t.Errorf("expected 1 config load across 3 requests, got %d", cache.loads) + } +} diff --git a/backend/plugins/drivers/driver_http/plugin.go b/backend/plugins/drivers/driver_http/plugin.go index 24b18b56..fe543458 100644 --- a/backend/plugins/drivers/driver_http/plugin.go +++ b/backend/plugins/drivers/driver_http/plugin.go @@ -110,6 +110,19 @@ func (p *Plugin) Apply(ctx *core.Context) error { return nil }) + // Bind CacheService from Context + if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { + setCacheService(cache) + } else { + core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { + setCacheService(cache) + }) + } + ctx.OnDispose(func() error { + setCacheService(nil) + return nil + }) + ctx.OnDispose(func() error { shutdownCtx, cancel := context.WithTimeout(context.Background(), p.shutdownTimeout) defer cancel()