From 2c415638fd65c28d59f59dc1c7b5d623cc7331fe Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 29 Aug 2026 08:44:37 +0800 Subject: [PATCH] autoresearch iter 15: stop CORS from querying the database on every request MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit isOriginAllowed read server_address from w_system_configs for every request carrying an Origin header — one uncached primary-DB round-trip plus a split and trim loop per browser request, while sibling config reads in the storage driver are already TTL cached. Read it through the shared CacheService with the same 5s window, falling back to the database when no cache is bound. driver_http now binds CacheService in Apply the way it already binds DBService. --- .../drivers/driver_http/cache_helper.go | 36 ++++++++ .../drivers/driver_http/middlewares.go | 43 +++++++++- .../drivers/driver_http/middlewares_test.go | 85 +++++++++++++++++++ backend/plugins/drivers/driver_http/plugin.go | 13 +++ 4 files changed, 173 insertions(+), 4 deletions(-) create mode 100644 backend/plugins/drivers/driver_http/cache_helper.go 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()