From cdac1f8a451927ab367816ffb11912c2fb6af2cd Mon Sep 17 00:00:00 2001 From: ryan Date: Sat, 20 Jun 2026 10:17:50 +0800 Subject: [PATCH] =?UTF-8?q?perf(cache):=20=E4=B8=89=E5=B1=82=E7=BC=93?= =?UTF-8?q?=E5=AD=98=E6=A1=86=E6=9E=B6=E8=A1=A5=E5=BC=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 新增 cache-framework skill,规范 RAM→Redis→DB 读路径、失效与 pub/sub - 上传元数据 Otter+Redis 缓存与多节点失效;Auth Source 缓存与 pub/sub - ListSystemConfigsByKeys 补 Redis 层;上传统计单事务;登录/Token 缓存预热 - cleanup 任务补 upload meta 失效钩子 --- .agent/skills/cache-framework/SKILL.md | 217 +++++++++++++ AGENTS.md | 1 + internal/apps/admin/auth_source/routers.go | 5 + internal/apps/oauth/auth_source_resolver.go | 8 +- internal/apps/oauth/handler_callback.go | 2 + internal/apps/oauth/oauth_test.go | 35 +++ internal/apps/upload/cache/meta_cache.go | 151 +++++++++ internal/apps/upload/cache/meta_cache_test.go | 287 ++++++++++++++++++ internal/apps/upload/filesrv/file_server.go | 7 +- .../apps/upload/filesrv/file_server_test.go | 4 + internal/apps/upload/ingest/helpers.go | 15 +- internal/apps/upload/ingest/ingest_test.go | 46 +++ internal/apps/upload/ingest/remove.go | 23 +- internal/apps/upload/stats/stats_counter.go | 7 +- .../apps/upload/stats/stats_counter_test.go | 31 ++ internal/apps/upload/task/cleanup.go | 2 + internal/apps/user/logics.go | 3 + internal/apps/user/routers.go | 2 + internal/repository/auth_source_cache.go | 269 ++++++++++++++++ internal/repository/auth_source_cache_test.go | 240 +++++++++++++++ internal/repository/system_config.go | 16 + internal/repository/system_config_test.go | 146 +++++++++ internal/repository/upload.go | 14 +- 23 files changed, 1511 insertions(+), 20 deletions(-) create mode 100644 .agent/skills/cache-framework/SKILL.md create mode 100644 internal/apps/upload/cache/meta_cache.go create mode 100644 internal/apps/upload/cache/meta_cache_test.go create mode 100644 internal/repository/auth_source_cache.go create mode 100644 internal/repository/auth_source_cache_test.go create mode 100644 internal/repository/system_config_test.go diff --git a/.agent/skills/cache-framework/SKILL.md b/.agent/skills/cache-framework/SKILL.md new file mode 100644 index 00000000..becf7240 --- /dev/null +++ b/.agent/skills/cache-framework/SKILL.md @@ -0,0 +1,217 @@ +--- +name: "cache-framework" +description: "Wavelet 项目专用:当新增或修改业务缓存(RAM/Redis/DB 三层读路径)、缓存失效、多节点 pub/sub 同步、或评估高频读是否应接入缓存时必须使用。本技能说明系统标准缓存框架、参考实现、禁止写法与分布式一致性要求。" +--- + +# 系统三层缓存框架 + +开始前阅读根目录 `AGENTS.md`(含 **Skill 关联索引**)。Wavelet 标准读路径为 **本地 RAM → Redis → PostgreSQL**(由快到慢),不是 DB 优先。 + +详细性能背景见 `docs/PERFORMANCE.md`。 + +## 关联 Skill + +| 关联 | 何时一并阅读 | +| :--- | :--- | +| [database-migration](../database-migration/SKILL.md) | 缓存对象对应新表/列/索引,或 seed 变更 | +| [new-setting](../new-setting/SKILL.md) | 系统配置类缓存(`GetSystemConfigByKey`、`ListSystemConfigsByKeys`) | +| [file-upload](../file-upload/SKILL.md) | 上传元数据 `upload:meta:{id}`、ingest/remove/cleanup 失效钩子 | +| [clickhouse-batchwriter](../clickhouse-batchwriter/SKILL.md) | 分析写入走 batchwriter,**不要**用本技能模式缓存 CH flush 队列 | +| [new-api](../new-api/SKILL.md) | 在 Handler 层接入 `GetXxxCached` 或评估高频读 | +| [new-async-task](../new-async-task/SKILL.md) | Worker/定时任务变更数据后必须 `Invalidate*`(如 `system:cleanup`) | + +## 标准模式(金标准) + +参考:`internal/repository/system_config_cache.go` + `GetSystemConfigByKey` / `ListSystemConfigsByKeys`。 + +| 层级 | 技术 | 职责 | +| :--- | :--- | :--- | +| L1 本地 | `pkg/cache/ram`(Otter v2) | 进程内热数据,最低延迟 | +| L2 共享 | Redis `db.GetJSON` / `SetJSON` / `HSetJSON` + `db.PrefixedKey` | 跨节点共享,带 TTL 或写穿 | +| L3 权威 | PostgreSQL via `db.DB(ctx)` | 唯一数据源 | + +### 读路径模板 + +```go +func GetThingCached(ctx context.Context, key string) (Thing, error) { + ensureThingCacheListener() // 订阅 pub/sub,仅 sync.Once + + if v, ok := thingRAM.GetIfPresent(key); ok { + return cloneThing(v), nil + } + if db.Redis != nil { + var v Thing + if err := db.GetJSON(ctx, redisKey(key), &v); err == nil { + thingRAM.Set(key, cloneThing(v)) + return v, nil + } + } + v, err := loadThingFromDB(ctx, key) + if err != nil { + return Thing{}, err + } + populateThingCache(ctx, v) // 回写 RAM + Redis + return v, nil +} +``` + +### 写穿(populate) + +DB miss 或业务创建成功后,**必须**回写上层: + +```go +func populateThingCache(ctx context.Context, v Thing) { + thingRAM.Set(v.Key, cloneThing(v)) + if db.Redis != nil { + _ = db.SetJSON(ctx, redisKey(v.Key), v, cacheTTL) + } +} +``` + +### 失效(Invalidate)— 分布式必做三步 + +数据变更(Admin 更新、软删除、状态迁移)时: + +1. **本机 RAM** — `thingRAM.Invalidate(key)` 或 `InvalidateAll()` +2. **Redis** — `Del` / `HDel` 对应 key +3. **pub/sub 广播** — 通知**其他节点**清除 RAM(Redis 已由写节点清掉) + +```go +func InvalidateThingCache(ctx context.Context, key string) error { + ensureThingCacheListener() + thingRAM.Invalidate(key) + if db.Redis != nil { + if err := db.Redis.Del(ctx, db.PrefixedKey(redisKey(key))).Err(); err != nil { + return err + } + publishThingRAMInvalidation(ctx, key) // 只广播 RAM 失效 + } + return nil +} +``` + +### pub/sub 监听模板 + +```go +const thingInvalidationChannel = "domain:thing_invalidation" + +func startThingCacheInvalidationListener() { + if db.Redis == nil { + return + } + go func() { + pubsub := db.Redis.Subscribe(context.Background(), thingInvalidationChannel) + defer func() { _ = pubsub.Close() }() + for msg := range pubsub.Channel() { + // 解析 payload,Invalidate RAM;勿重复 Del Redis + thingRAM.Invalidate(parsedKey) + } + }() +} +``` + +- 使用 `sync.Once` 启动监听;**`ensureListener` 必须在 `db.Redis == nil` 时直接 return,不可消费 Once**(否则测试或 Redis 晚初始化时监听器永不启动)。 +- 测试可提供 `StopThingCacheListener` + 重置 `Once`(参考 `StopUploadMetaCacheListener`、`StopAuthSourceCacheListener`)。 +- 其他节点收到消息后**只清 RAM**,不再删 Redis。 + +## 现有实现速查 + +| 域 | 文件 | L1 | L2 | pub/sub | +| :--- | :--- | :--- | :--- | :--- | +| 系统配置 | `repository/system_config_cache.go` | Otter | Redis Hash | `system:config_invalidation` ✅ | +| CAPTCHA 运行时 | `apps/cap/runtime_settings.go` | atomic.Pointer | (借配置 Redis) | 订阅 `system:config_invalidation` ✅ | +| 上传元数据 | `apps/upload/cache/meta_cache.go` | Otter | Redis JSON | `upload:meta_invalidation` ✅ | +| 上传访问白名单 | `apps/upload/cache/access_cache.go` | 进程内 TTL | (借配置读路径) | `upload:file_access_invalidation` ✅ | +| Auth Source | `repository/auth_source_cache.go` | Otter | Redis JSON | `oauth:auth_source_invalidation` ✅ | +| OAuth 用户/Token | `apps/oauth/cache.go` | 自研 map | Redis JSON | ❌ 无 pub/sub(历史债) | +| 推送渠道 | `repository/push_channel.go` | 无 | Redis JSON | ❌ 仅 Redis Del | +| Storage 驱动 | `internal/storage/storage.go` | RWMutex 快照 | — | `storage:config_invalidation` ✅ | + +## 新增缓存工作流 + +1. **判定是否需要缓存**:高频读、低变更、可容忍短暂 TTL;写路径必须能统一失效。 +2. **选型 L1**:优先 `pkg/cache/ram.MustNew`;**禁止**自研 `map+mutex+TTL`,除非有充分理由并文档说明。 +3. **选型 L2**:小对象 `SetJSON`;配置类多条目用 Redis Hash(`HSetJSON`)。 +4. **定义 Redis key**:小写蛇形,带业务前缀(`upload:meta:{id}`);统一 `db.PrefixedKey`。 +5. **实现 Invalidate + pub/sub**:凡多实例部署可读的 RAM 缓存**必须**有失效广播。 +6. **挂载变更钩子**:在所有 DB 变更入口调用 Invalidate(含 Worker/定时任务,不只 HTTP Handler)。 +7. **测试**: + - RAM hit / Redis hit / DB fallback + - Invalidate 清 L1+L2 + - pub/sub 触发他机 RAM 失效(可用 miniredis Publish 模拟) + - `Reset*RAMCacheForTest` 仅清本机 RAM +8. 运行 `go test` 相关包 + `make code-check`。 + +## 变更钩子清单(上传元数据示例) + +| 入口 | 动作 | +| :--- | :--- | +| `ingest.persistUploadRecord` 创建成功 | `SetUploadMetaCache` | +| `ingest.Remove` / `RemoveOwned` | `InvalidateUploadMetaCache` | +| `task/cleanup.go` 软删除 pending 文件 | `InvalidateUploadMetaCache` | +| 直接 `repository.SoftDeleteUpload` | **禁止** — 必须走 `upload.Remove` | + +## 禁止写法 + +```go +// ❌ 自研 L1,与 pkg/cache/ram 重复 +var mu sync.RWMutex +var items = map[uint64]entry{} + +// ❌ 只清本机 RAM + Redis,无 pub/sub(多节点 RAM 脏读) +func Invalidate(ctx context.Context, id uint64) { + localDelete(id) + redis.Del(...) +} + +// ❌ DB 变更后忘记 Worker 路径 +// cleanup 任务删了 upload 行,但未 InvalidateUploadMetaCache + +// ❌ 在 Handler 里直接查 DB,绕过已有 GetXxxCached + +// ❌ Redis key 不用 PrefixedKey(多环境共 Redis 时冲突) + +// ❌ 在 init() 里启动 pub/sub 监听 — 与 bootstrap 规范冲突;用 sync.Once 懒启动 +``` + +## 特殊场景 + +### 敏感字段(ClientSecret) + +模型 `json:"-"` 时,Redis DTO 用独立 `*RedisRecord` struct 显式序列化字段(见 `auth_source_cache.go`)。 + +### 批量读配置 + +批量接口必须与单 key 一致走 Redis(`ListSystemConfigsByKeys` 在 RAM miss 后逐 key `HGetJSON`,再 DB `IN`)。 + +### 仅进程内、短 TTL、配置衍生 + +可用进程内快照 + 订阅上游 pub/sub(`access_cache.go`、`cap/runtime_settings.go`),不必强行 Redis L2。 + +### OAuth 用户/Token + +沿用 `oauth/cache.go`;新增逻辑调用 `SetCachedUser` / `SetCachedToken` 预热,变更调用 `InvalidateCachedUser` / `InvalidateCachedToken`。 + +## 验证清单 + +```bash +go test ./internal/repository/... ./internal/apps/upload/cache/... +make code-check +``` + +- [ ] L1 使用 `pkg/cache/ram`(或已文档化的例外) +- [ ] 读路径:RAM → Redis → DB +- [ ] 写穿 populate 在 DB load / 创建成功后 +- [ ] Invalidate:RAM + Redis + Publish +- [ ] `ensureListener` + pub/sub 清他机 RAM +- [ ] 所有变更入口(含 Worker)已挂钩 +- [ ] 测试含 Invalidate 与 pub/sub + +## 相关文件 + +- L1 引擎:`pkg/cache/ram/cache.go` +- DB/Redis 助手:`internal/db/redis.go`(`GetJSON`, `SetJSON`, `HGetJSON`, `PrefixedKey`) +- 金标准:`internal/repository/system_config_cache.go` +- 上传元数据:`internal/apps/upload/cache/meta_cache.go` +- Auth Source:`internal/repository/auth_source_cache.go` +- 性能文档:`docs/PERFORMANCE.md` \ No newline at end of file diff --git a/AGENTS.md b/AGENTS.md index 2194dea3..df44c4dd 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -42,6 +42,7 @@ | `database-migration` | 数据库表结构变更、goose SQL 迁移(PG/SQLite/ClickHouse)、seed 数据 | | `clickhouse-batchwriter` | ClickHouse 批量写入、`internal/db/batchwriter` 接入、分析表异步 flush、背压与写入路径改造 | | `file-upload` | 业务上传文件、Worker 程序化摄取、`upload.Ingest` 策略选型、文件访问与 `w_uploads` / 统计排查 | +| `cache-framework` | 新增或修改业务缓存(RAM/Redis/DB 三层读路径)、缓存失效、多节点 pub/sub 同步、评估高频读是否应接入缓存 | | `push-notification` | 系统通知推送事件、统一触发器投递、带消息推送的业务功能 | | `release-guide` | 根据自上一正式版本 Tag 以来的提交整理 Version Bump 提交信息以触发双语 Release | | `shadcn` | 添加、修改或组合 shadcn/ui 组件 | diff --git a/internal/apps/admin/auth_source/routers.go b/internal/apps/admin/auth_source/routers.go index d18e2c8e..fc0339ba 100644 --- a/internal/apps/admin/auth_source/routers.go +++ b/internal/apps/admin/auth_source/routers.go @@ -13,6 +13,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/admin" "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/gin-gonic/gin" "github.com/Rain-kl/Wavelet/internal/common/response" @@ -91,6 +92,7 @@ func CreateAuthSource(c *gin.Context) { response.AbortBadRequest(c, err.Error()) return } + _ = repository.InvalidateAuthSourceCache(c.Request.Context()) source.Sanitize() c.JSON(http.StatusOK, response.OK(source)) } @@ -150,6 +152,7 @@ func UpdateAuthSource(c *gin.Context) { oauth.InvalidateOIDCProviderCache(normalizeIssuer(existing.OpenIDDiscoveryURL)) } oauth.InvalidateOIDCProviderCache(normalizeIssuer(req.OpenIDDiscoveryURL)) + _ = repository.InvalidateAuthSourceCache(c.Request.Context()) updated, err := model.GetAuthSourceByID(c.Request.Context(), id) if err != nil { @@ -191,6 +194,7 @@ func ToggleAuthSource(c *gin.Context) { response.AbortBadRequest(c, err.Error()) return } + _ = repository.InvalidateAuthSourceCache(c.Request.Context()) c.JSON(http.StatusOK, response.OKNil()) } @@ -216,6 +220,7 @@ func DeleteAuthSource(c *gin.Context) { response.AbortBadRequest(c, err.Error()) return } + _ = repository.InvalidateAuthSourceCache(c.Request.Context()) c.JSON(http.StatusOK, response.OKNil()) } diff --git a/internal/apps/oauth/auth_source_resolver.go b/internal/apps/oauth/auth_source_resolver.go index 2f157c60..e996c4be 100644 --- a/internal/apps/oauth/auth_source_resolver.go +++ b/internal/apps/oauth/auth_source_resolver.go @@ -25,16 +25,16 @@ func isOIDCLoginEnabled(ctx context.Context) bool { func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) { name := strings.TrimSpace(strings.ToLower(sourceName)) if name == "" { - sources, err := model.GetActiveAuthSources(ctx) + sources, err := repository.GetActiveAuthSourcesCached(ctx) if err != nil { return nil, err } if len(sources) == 0 { return nil, errors.New(errNoActiveAuthSource) } - return &sources[0], nil + return repository.GetAuthSourceByNameCached(ctx, sources[0].Name) } - return model.GetAuthSourceByName(ctx, name) + return repository.GetAuthSourceByNameCached(ctx, name) } func activeLoginSources(ctx context.Context) []AuthSourceView { @@ -43,7 +43,7 @@ func activeLoginSources(ctx context.Context) []AuthSourceView { return nil } - dbSources, err := model.GetActiveAuthSources(ctx) + dbSources, err := repository.GetActiveAuthSourcesCached(ctx) if err != nil { return nil } diff --git a/internal/apps/oauth/handler_callback.go b/internal/apps/oauth/handler_callback.go index d80ff250..9e1a46eb 100644 --- a/internal/apps/oauth/handler_callback.go +++ b/internal/apps/oauth/handler_callback.go @@ -179,6 +179,8 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth return } + SetCachedUser(ctx, user.ID, &user) + logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP()) listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP()) diff --git a/internal/apps/oauth/oauth_test.go b/internal/apps/oauth/oauth_test.go index e97e296d..9ff0badd 100644 --- a/internal/apps/oauth/oauth_test.go +++ b/internal/apps/oauth/oauth_test.go @@ -95,6 +95,28 @@ func (m *mockRedisClient) Del(ctx context.Context, keys ...string) *redis.IntCmd return cmd } +func (m *mockRedisClient) Scan(ctx context.Context, cursor uint64, match string, count int64) *redis.ScanCmd { + cmd := redis.NewScanCmd(ctx, nil, cursor, match, count) + var keys []string + for key := range m.store { + if redisMatchPattern(key, match) { + keys = append(keys, key) + } + } + cmd.SetVal(keys, 0) + return cmd +} + +func redisMatchPattern(key, pattern string) bool { + if pattern == "" || pattern == "*" { + return true + } + if strings.HasSuffix(pattern, "*") { + return strings.HasPrefix(key, strings.TrimSuffix(pattern, "*")) + } + return key == pattern +} + func (m *mockRedisClient) HSet(ctx context.Context, key string, values ...interface{}) *redis.IntCmd { cmd := redis.NewIntCmd(ctx) if len(values) >= 2 { @@ -129,6 +151,12 @@ func (m *mockRedisClient) HGet(ctx context.Context, key string, field string) *r return cmd } +func (m *mockRedisClient) Publish(ctx context.Context, channel string, message interface{}) *redis.IntCmd { + cmd := redis.NewIntCmd(ctx) + cmd.SetVal(1) + return cmd +} + func (m *mockRedisClient) Subscribe(ctx context.Context, channels ...string) *redis.PubSub { return redis.NewClient(&redis.Options{ Addr: "127.0.0.1:0", @@ -310,6 +338,7 @@ func newMockOIDCClient(issuer, clientID string, expectedState *string, sub, user func setupTestDB(t *testing.T) *gorm.DB { repository.ResetSystemConfigRAMCacheForTest() + repository.ResetAuthSourceRAMCacheForTest() dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) if err != nil { @@ -635,6 +664,7 @@ func TestAuthorize(t *testing.T) { // Case 2: Inactive Source Authorize dbConn.Model(&model.AuthSource{}).Where("id = ?", 101).Update("is_active", false) + _ = repository.InvalidateAuthSourceCache(context.Background()) w2 := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize", nil, nil, nil) if w2.Code != http.StatusBadRequest { t.Errorf("expected 400 for inactive source, got %d", w2.Code) @@ -1122,6 +1152,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { repository.ResetSystemConfigRAMCacheForTest() mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) + _ = repository.InvalidateAuthSourceCache(context.Background()) wSourceInactive := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) if wSourceInactive.Code != http.StatusBadRequest { @@ -1134,6 +1165,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { repository.ResetSystemConfigRAMCacheForTest() mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) + _ = repository.InvalidateAuthSourceCache(context.Background()) wAuthDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/"+testSourceName+"/authorize", nil, nil, nil) if wAuthDisabled.Code != http.StatusBadRequest { @@ -1146,6 +1178,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { repository.ResetSystemConfigRAMCacheForTest() mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled) dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) + _ = repository.InvalidateAuthSourceCache(context.Background()) wLogin := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) @@ -1187,6 +1220,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // Since callback deletes state, we need to generate state again dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true) + _ = repository.InvalidateAuthSourceCache(context.Background()) wLogin2 := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil) _ = json.Unmarshal(wLogin2.Body.Bytes(), &loginUrlResp) parsedURL, _ = url.Parse(loginUrlResp.Data.AuthorizeURL) @@ -1201,6 +1235,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) { // Deactivate source dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false) + _ = repository.InvalidateAuthSourceCache(context.Background()) reqBody2 := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state) wCallbackSourceInactive := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody2), map[string]string{ "Content-Type": "application/json", diff --git a/internal/apps/upload/cache/meta_cache.go b/internal/apps/upload/cache/meta_cache.go new file mode 100644 index 00000000..b5b9d606 --- /dev/null +++ b/internal/apps/upload/cache/meta_cache.go @@ -0,0 +1,151 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package cache + +import ( + "context" + "encoding/json" + "fmt" + "sync" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/pkg/cache/ram" +) + +const ( + uploadMetaRedisCacheTTL = 30 * 60 // seconds + uploadMetaRAMMaximumSize = 4096 + uploadMetaInvalidationChan = "upload:meta_invalidation" +) + +type uploadMetaInvalidationMessage struct { + ID uint64 `json:"id"` +} + +var ( + uploadMetaRAM = ram.MustNew[uint64, model.Upload](ram.Options{MaximumSize: uploadMetaRAMMaximumSize}) + uploadMetaListenerOnce sync.Once + uploadMetaListenerCtx context.Context + uploadMetaListenerCancel context.CancelFunc +) + +func uploadMetaRedisKey(id uint64) string { + return fmt.Sprintf("upload:meta:%d", id) +} + +func cloneUpload(upload model.Upload) model.Upload { + return upload +} + +func ensureUploadMetaCacheListener() { + if db.Redis == nil { + return + } + uploadMetaListenerOnce.Do(startUploadMetaCacheInvalidationListener) +} + +func startUploadMetaCacheInvalidationListener() { + uploadMetaListenerCtx, uploadMetaListenerCancel = context.WithCancel(context.Background()) + + go func() { + pubsub := db.Redis.Subscribe(uploadMetaListenerCtx, uploadMetaInvalidationChan) + defer func() { + _ = pubsub.Close() + }() + + go func() { + <-uploadMetaListenerCtx.Done() + _ = pubsub.Close() + }() + + for msg := range pubsub.Channel() { + var payload uploadMetaInvalidationMessage + if err := json.Unmarshal([]byte(msg.Payload), &payload); err != nil || payload.ID == 0 { + uploadMetaRAM.InvalidateAll() + continue + } + uploadMetaRAM.Invalidate(payload.ID) + } + }() +} + +func publishUploadMetaRAMInvalidation(ctx context.Context, id uint64) { + if db.Redis == nil { + return + } + payload, err := json.Marshal(uploadMetaInvalidationMessage{ID: id}) + if err != nil { + return + } + _ = db.Redis.Publish(ctx, uploadMetaInvalidationChan, payload).Err() +} + +// GetUploadByID loads upload metadata from RAM, Redis, or the database. +func GetUploadByID(ctx context.Context, id uint64) (model.Upload, error) { + ensureUploadMetaCacheListener() + + if upload, ok := uploadMetaRAM.GetIfPresent(id); ok { + return cloneUpload(upload), nil + } + + key := uploadMetaRedisKey(id) + if db.Redis != nil { + var upload model.Upload + if err := db.GetJSON(ctx, key, &upload); err == nil { + uploadMetaRAM.Set(id, cloneUpload(upload)) + return upload, nil + } + } + + var upload model.Upload + if err := db.DB(ctx). + Where("id = ? AND status IN (?, ?)", id, model.UploadStatusPending, model.UploadStatusUsed). + First(&upload).Error; err != nil { + return model.Upload{}, err + } + + SetUploadMetaCache(ctx, &upload) + return upload, nil +} + +// SetUploadMetaCache populates RAM and Redis upload metadata caches. +func SetUploadMetaCache(ctx context.Context, upload *model.Upload) { + ensureUploadMetaCacheListener() + + if upload == nil { + return + } + + cloned := cloneUpload(*upload) + uploadMetaRAM.Set(upload.ID, cloned) + if db.Redis != nil { + _ = db.SetJSON(ctx, uploadMetaRedisKey(upload.ID), cloned, uploadMetaRedisCacheTTL) + } +} + +// InvalidateUploadMetaCache clears RAM and Redis upload metadata caches and notifies peer nodes. +func InvalidateUploadMetaCache(ctx context.Context, id uint64) { + ensureUploadMetaCacheListener() + + uploadMetaRAM.Invalidate(id) + if db.Redis != nil { + _ = db.Redis.Del(ctx, db.PrefixedKey(uploadMetaRedisKey(id))).Err() + publishUploadMetaRAMInvalidation(ctx, id) + } +} + +// ResetUploadMetaCacheForTest clears the in-process upload metadata RAM cache. +func ResetUploadMetaCacheForTest() { + uploadMetaRAM.InvalidateAll() +} + +// StopUploadMetaCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard. +func StopUploadMetaCacheListener() { + if uploadMetaListenerCancel != nil { + uploadMetaListenerCancel() + uploadMetaListenerCancel = nil + } + uploadMetaListenerOnce = sync.Once{} +} diff --git a/internal/apps/upload/cache/meta_cache_test.go b/internal/apps/upload/cache/meta_cache_test.go new file mode 100644 index 00000000..56408d66 --- /dev/null +++ b/internal/apps/upload/cache/meta_cache_test.go @@ -0,0 +1,287 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package cache + +import ( + "context" + "encoding/json" + "testing" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/testhelper" + "gorm.io/gorm" +) + +func init() { + testhelper.RegisterCleanup(func() { + StopUploadMetaCacheListener() + ResetUploadMetaCacheForTest() + }) +} + +func seedUpload(t *testing.T, dbConn *gorm.DB, upload model.Upload) { + t.Helper() + if err := dbConn.Create(&upload).Error; err != nil { + t.Fatalf("create upload: %v", err) + } +} + +func TestGetUploadByIDLoadsFromDBAndPopulatesCache(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + ResetUploadMetaCacheForTest() + + ctx := context.Background() + upload := model.Upload{ + ID: 91001, + UserID: 1, + FileName: "cached.png", + FilePath: "cached.png", + FileSize: 12, + MimeType: "image/png", + Extension: "png", + Type: "avatar", + Status: model.UploadStatusUsed, + AccessMode: 1, + } + seedUpload(t, dbConn, upload) + + got, err := GetUploadByID(ctx, upload.ID) + if err != nil { + t.Fatalf("GetUploadByID: %v", err) + } + if got.ID != upload.ID || got.FileName != upload.FileName { + t.Fatalf("unexpected upload: %+v", got) + } + + var redisUpload model.Upload + if err := db.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err != nil { + t.Fatalf("redis cache miss after DB load: %v", err) + } + if redisUpload.ID != upload.ID { + t.Fatalf("redis upload id mismatch: got=%d want=%d", redisUpload.ID, upload.ID) + } + + if err := dbConn.Delete(&model.Upload{}, upload.ID).Error; err != nil { + t.Fatalf("delete upload from db: %v", err) + } + + gotCached, err := GetUploadByID(ctx, upload.ID) + if err != nil { + t.Fatalf("GetUploadByID from RAM cache: %v", err) + } + if gotCached.ID != upload.ID { + t.Fatalf("expected RAM cache hit for upload %d", upload.ID) + } +} + +func TestGetUploadByIDReadsFromRedisWhenRAMEmpty(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + ResetUploadMetaCacheForTest() + + ctx := context.Background() + upload := model.Upload{ + ID: 91002, + UserID: 1, + FileName: "redis.png", + FilePath: "redis.png", + FileSize: 8, + MimeType: "image/png", + Extension: "png", + Type: "avatar", + Status: model.UploadStatusPending, + AccessMode: 0, + } + seedUpload(t, dbConn, upload) + SetUploadMetaCache(ctx, &upload) + ResetUploadMetaCacheForTest() + + if err := dbConn.Delete(&model.Upload{}, upload.ID).Error; err != nil { + t.Fatalf("delete upload from db: %v", err) + } + + got, err := GetUploadByID(ctx, upload.ID) + if err != nil { + t.Fatalf("GetUploadByID from redis: %v", err) + } + if got.ID != upload.ID || got.FileName != upload.FileName { + t.Fatalf("unexpected upload from redis: %+v", got) + } +} + +func TestInvalidateUploadMetaCacheClearsRAMAndRedis(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + ResetUploadMetaCacheForTest() + + ctx := context.Background() + upload := model.Upload{ + ID: 91003, + UserID: 1, + FileName: "invalidate.png", + FilePath: "invalidate.png", + FileSize: 4, + MimeType: "image/png", + Extension: "png", + Type: "avatar", + Status: model.UploadStatusUsed, + AccessMode: 1, + } + seedUpload(t, dbConn, upload) + SetUploadMetaCache(ctx, &upload) + + InvalidateUploadMetaCache(ctx, upload.ID) + + var redisUpload model.Upload + if err := db.GetJSON(ctx, uploadMetaRedisKey(upload.ID), &redisUpload); err == nil { + t.Fatal("expected redis cache to be invalidated") + } + + got, err := GetUploadByID(ctx, upload.ID) + if err != nil { + t.Fatalf("GetUploadByID after invalidate should reload from DB: %v", err) + } + if got.ID != upload.ID { + t.Fatalf("unexpected upload reloaded from DB: %+v", got) + } +} + +func TestUploadMetaInvalidationPubSubClearsPeerRAM(t *testing.T) { + StopUploadMetaCacheListener() + defer StopUploadMetaCacheListener() + + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + ResetUploadMetaCacheForTest() + + ctx := context.Background() + upload := model.Upload{ + ID: 91006, + UserID: 1, + FileName: "pubsub.png", + FilePath: "pubsub.png", + FileSize: 4, + MimeType: "image/png", + Extension: "png", + Type: "avatar", + Status: model.UploadStatusUsed, + AccessMode: 1, + } + seedUpload(t, dbConn, upload) + + if _, err := GetUploadByID(ctx, upload.ID); err != nil { + t.Fatalf("GetUploadByID: %v", err) + } + time.Sleep(50 * time.Millisecond) // allow pub/sub listener to subscribe + if err := dbConn.Delete(&model.Upload{}, upload.ID).Error; err != nil { + t.Fatalf("delete upload from db: %v", err) + } + if _, err := GetUploadByID(ctx, upload.ID); err != nil { + t.Fatalf("expected cache hit before pub/sub invalidation: %v", err) + } + + payload, err := json.Marshal(uploadMetaInvalidationMessage{ID: upload.ID}) + if err != nil { + t.Fatalf("marshal invalidation payload: %v", err) + } + if err := db.Redis.Publish(ctx, uploadMetaInvalidationChan, string(payload)).Err(); err != nil { + t.Fatalf("publish invalidation: %v", err) + } + + deadline := time.Now().Add(2 * time.Second) + ramCleared := false + for time.Now().Before(deadline) { + if _, ok := uploadMetaRAM.GetIfPresent(upload.ID); !ok { + ramCleared = true + break + } + time.Sleep(20 * time.Millisecond) + } + if !ramCleared { + t.Fatal("expected peer RAM cache to be cleared by pub/sub") + } + + if err := db.Redis.Del(ctx, db.PrefixedKey(uploadMetaRedisKey(upload.ID))).Err(); err != nil { + t.Fatalf("delete redis cache: %v", err) + } + if _, err := GetUploadByID(ctx, upload.ID); err == nil { + t.Fatal("expected cache miss after pub/sub RAM eviction and redis delete") + } +} + +func TestGetUploadByIDSkipsDeletedUploads(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + ResetUploadMetaCacheForTest() + + ctx := context.Background() + upload := model.Upload{ + ID: 91004, + UserID: 1, + FileName: "deleted.png", + FilePath: "deleted.png", + FileSize: 4, + MimeType: "image/png", + Extension: "png", + Type: "avatar", + Status: model.UploadStatusDeleted, + AccessMode: 1, + } + seedUpload(t, dbConn, upload) + + if _, err := GetUploadByID(ctx, upload.ID); err == nil { + t.Fatal("expected error for deleted upload") + } +} + +func TestGetUploadByIDWorksWithRedisDisabled(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + ResetUploadMetaCacheForTest() + + redisClient := db.Redis + db.Redis = nil + t.Cleanup(func() { + db.Redis = redisClient + StopUploadMetaCacheListener() + }) + + ctx := context.Background() + upload := model.Upload{ + ID: 91005, + UserID: 1, + FileName: "ram-only.png", + FilePath: "ram-only.png", + FileSize: 6, + MimeType: "image/png", + Extension: "png", + Type: "avatar", + Status: model.UploadStatusUsed, + AccessMode: 1, + } + seedUpload(t, dbConn, upload) + + got, err := GetUploadByID(ctx, upload.ID) + if err != nil { + t.Fatalf("GetUploadByID without redis: %v", err) + } + if got.ID != upload.ID { + t.Fatalf("unexpected upload: %+v", got) + } + + if err := dbConn.Delete(&model.Upload{}, upload.ID).Error; err != nil { + t.Fatalf("delete upload from db: %v", err) + } + + gotCached, err := GetUploadByID(ctx, upload.ID) + if err != nil { + t.Fatalf("GetUploadByID from RAM without redis: %v", err) + } + if gotCached.ID != upload.ID { + t.Fatal("expected RAM cache hit when redis is disabled") + } +} \ No newline at end of file diff --git a/internal/apps/upload/filesrv/file_server.go b/internal/apps/upload/filesrv/file_server.go index 0599318e..573b48cd 100644 --- a/internal/apps/upload/filesrv/file_server.go +++ b/internal/apps/upload/filesrv/file_server.go @@ -22,7 +22,6 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/upload/util" "github.com/Rain-kl/Wavelet/internal/common" "github.com/Rain-kl/Wavelet/internal/common/response" - "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/diskcache" "github.com/Rain-kl/Wavelet/internal/model" @@ -96,10 +95,8 @@ func GetUploadRecordByID(c *gin.Context) (*model.Upload, error) { return nil, err } - var upload model.Upload - if err := db.DB(c.Request.Context()). - Where("id = ? AND status IN (?, ?)", uploadID, model.UploadStatusPending, model.UploadStatusUsed). - First(&upload).Error; err != nil { + upload, err := cache.GetUploadByID(c.Request.Context(), uploadID) + if err != nil { return nil, err } diff --git a/internal/apps/upload/filesrv/file_server_test.go b/internal/apps/upload/filesrv/file_server_test.go index 8c33561d..b56cca3c 100644 --- a/internal/apps/upload/filesrv/file_server_test.go +++ b/internal/apps/upload/filesrv/file_server_test.go @@ -34,6 +34,10 @@ import ( "gorm.io/gorm" ) +func init() { + testhelper.RegisterCleanup(cache.ResetUploadMetaCacheForTest) +} + func TestServeFileByIDAccessControl(t *testing.T) { dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() diff --git a/internal/apps/upload/ingest/helpers.go b/internal/apps/upload/ingest/helpers.go index 80216670..00580a90 100644 --- a/internal/apps/upload/ingest/helpers.go +++ b/internal/apps/upload/ingest/helpers.go @@ -11,9 +11,11 @@ import ( "strings" "time" + uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" + "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/db/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" @@ -99,7 +101,7 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i } func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey string) error { - if err := repository.CreateUpload(ctx, upload); err != nil { + if err := createUploadWithStats(ctx, upload); err != nil { _, backend, backendErr := storage.Active(ctx) if backendErr == nil { if deleteErr := backend.Delete(ctx, objectKey); deleteErr != nil { @@ -108,10 +110,19 @@ func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey st } return err } - uploadstats.RecordUploadStatsAdd(ctx, upload) + uploadcache.SetUploadMetaCache(ctx, upload) return nil } +func createUploadWithStats(ctx context.Context, upload *model.Upload) error { + return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := repository.CreateUploadTx(tx, upload); err != nil { + return err + } + return uploadstats.ApplyUploadStatsDeltaTx(tx, upload, 1) + }) +} + func createDedupRecord(ctx context.Context, existing model.Upload, req Request) (Result, error) { accessMode := resolveAccessMode(req.Type, req.AccessMode) newUpload := model.Upload{ diff --git a/internal/apps/upload/ingest/ingest_test.go b/internal/apps/upload/ingest/ingest_test.go index 5663bd92..d70bedc7 100644 --- a/internal/apps/upload/ingest/ingest_test.go +++ b/internal/apps/upload/ingest/ingest_test.go @@ -183,6 +183,52 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) { } } +func TestCreateUploadWithStatsRollsBackOnCreateFailure(t *testing.T) { + dbConn, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + ctx := context.Background() + + existing := model.Upload{ + ID: 99001, + UserID: 1001, + FileName: "existing.png", + FilePath: "uploads/existing.png", + FileSize: 64, + MimeType: "image/png", + Extension: "png", + Type: "generic", + Status: model.UploadStatusUsed, + CreatedAt: time.Now(), + } + if err := dbConn.Create(&existing).Error; err != nil { + t.Fatalf("seed upload failed: %v", err) + } + + duplicate := &model.Upload{ + ID: existing.ID, + UserID: 1002, + FileName: "duplicate.png", + FilePath: "uploads/duplicate.png", + FileSize: 128, + MimeType: "image/png", + Extension: "png", + Type: "generic", + Status: model.UploadStatusUsed, + CreatedAt: time.Now(), + } + if err := createUploadWithStats(ctx, duplicate); err == nil { + t.Fatal("createUploadWithStats with duplicate ID expected error") + } + + stats, err := loadTotalStats(ctx) + if err != nil { + t.Fatalf("loadTotalStats returned error: %v", err) + } + if stats.TotalCount != 0 || stats.TotalSize != 0 { + t.Fatalf("loadTotalStats() = count %d size %d, want zero after rolled-back stats", stats.TotalCount, stats.TotalSize) + } +} + func TestRemoveDecrementsStats(t *testing.T) { _, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() diff --git a/internal/apps/upload/ingest/remove.go b/internal/apps/upload/ingest/remove.go index 80b0a0fb..99071d68 100644 --- a/internal/apps/upload/ingest/remove.go +++ b/internal/apps/upload/ingest/remove.go @@ -6,9 +6,12 @@ package ingest import ( "context" + uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" + "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" + "gorm.io/gorm" ) // Remove soft-deletes an upload and decrements incremental stats. @@ -17,8 +20,7 @@ func Remove(ctx context.Context, uploadID uint64) (model.Upload, error) { if err != nil { return model.Upload{}, err } - uploadstats.RecordUploadStatsRemove(ctx, &upload) - if err := repository.SoftDeleteUpload(ctx, &upload); err != nil { + if err := softDeleteUploadWithStats(ctx, &upload); err != nil { return model.Upload{}, err } upload.Status = model.UploadStatusDeleted @@ -34,10 +36,23 @@ func RemoveOwned(ctx context.Context, userID, uploadID uint64) (model.Upload, er if upload.UserID != userID { return model.Upload{}, ErrForbidden } - uploadstats.RecordUploadStatsRemove(ctx, &upload) - if err := repository.SoftDeleteUpload(ctx, &upload); err != nil { + if err := softDeleteUploadWithStats(ctx, &upload); err != nil { return model.Upload{}, err } upload.Status = model.UploadStatusDeleted return upload, nil } + +func softDeleteUploadWithStats(ctx context.Context, upload *model.Upload) error { + statsSnapshot := *upload + if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := repository.SoftDeleteUploadTx(tx, upload); err != nil { + return err + } + return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1) + }); err != nil { + return err + } + uploadcache.InvalidateUploadMetaCache(ctx, upload.ID) + return nil +} \ No newline at end of file diff --git a/internal/apps/upload/stats/stats_counter.go b/internal/apps/upload/stats/stats_counter.go index 3f2ead66..810dc875 100644 --- a/internal/apps/upload/stats/stats_counter.go +++ b/internal/apps/upload/stats/stats_counter.go @@ -37,7 +37,7 @@ func RebuildUploadStats(ctx context.Context) error { } for i := range uploads { - if err := applyUploadStatsDeltaTx(tx, &uploads[i], 1); err != nil { + if err := ApplyUploadStatsDeltaTx(tx, &uploads[i], 1); err != nil { return err } } @@ -50,11 +50,12 @@ func applyUploadStatsDelta(ctx context.Context, upload *model.Upload, sign int64 return nil } return db.DB(ctx).Transaction(func(tx *gorm.DB) error { - return applyUploadStatsDeltaTx(tx, upload, sign) + return ApplyUploadStatsDeltaTx(tx, upload, sign) }) } -func applyUploadStatsDeltaTx(tx *gorm.DB, upload *model.Upload, sign int64) error { +// ApplyUploadStatsDeltaTx applies incremental upload stats within an existing transaction. +func ApplyUploadStatsDeltaTx(tx *gorm.DB, upload *model.Upload, sign int64) error { if upload == nil || !isActiveUploadStatus(upload.Status) || sign == 0 { return nil } diff --git a/internal/apps/upload/stats/stats_counter_test.go b/internal/apps/upload/stats/stats_counter_test.go index 5e3760bb..02d6c7b4 100644 --- a/internal/apps/upload/stats/stats_counter_test.go +++ b/internal/apps/upload/stats/stats_counter_test.go @@ -11,8 +11,39 @@ import ( "github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/testhelper" + "gorm.io/gorm" ) +func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) { + _, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + ctx := context.Background() + + upload := &model.Upload{ + ID: 42002, + FileSize: 256, + MimeType: "image/jpeg", + Extension: "jpg", + Type: "avatar", + Status: model.UploadStatusUsed, + CreatedAt: time.Now(), + } + + if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + return ApplyUploadStatsDeltaTx(tx, upload, 1) + }); err != nil { + t.Fatalf("ApplyUploadStatsDeltaTx returned error: %v", err) + } + + stats, err := loadUploadStats(ctx) + if err != nil { + t.Fatalf("loadUploadStats returned error: %v", err) + } + if stats.TotalCount != 1 || stats.TotalSize != 256 { + t.Fatalf("unexpected total stats: count=%d size=%d", stats.TotalCount, stats.TotalSize) + } +} + func TestApplyUploadStatsAddAndRemove(t *testing.T) { _, _, cleanup := testhelper.SetupTestEnvironment(t) defer cleanup() diff --git a/internal/apps/upload/task/cleanup.go b/internal/apps/upload/task/cleanup.go index beabc622..0d70e35d 100644 --- a/internal/apps/upload/task/cleanup.go +++ b/internal/apps/upload/task/cleanup.go @@ -10,6 +10,7 @@ import ( "fmt" "time" + uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache" "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" @@ -100,6 +101,7 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas } uploadstats.RecordUploadStatsRemove(ctx, &u) + uploadcache.InvalidateUploadMetaCache(ctx, u.ID) totalDeleted++ lastID = u.ID } diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index 4d6f0ce7..c99bb0c1 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -362,6 +362,7 @@ func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) erro if err := db.DB(ctx).Create(record).Error; err != nil { return errors.New("创建令牌失败,请稍后再试") } + oauth.SetCachedToken(ctx, record.TokenHash, record) return nil } @@ -405,6 +406,8 @@ func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *mo return "", nil, errors.New("轮换令牌失败,请稍后再试") } + oauth.SetCachedToken(ctx, tokenRecord.TokenHash, &tokenRecord) + return newTokenStr, &tokenRecord, nil } diff --git a/internal/apps/user/routers.go b/internal/apps/user/routers.go index b3903e79..d2aa9a42 100644 --- a/internal/apps/user/routers.go +++ b/internal/apps/user/routers.go @@ -168,6 +168,8 @@ func Login(c *gin.Context) { return } + oauth.SetCachedUser(ctx, user.ID, user) + logger.InfoF(ctx, "[LoginAudit] successful login for user: %s, ID: %d, IP: %s", user.Username, user.ID, c.ClientIP()) listener.EmitAdminLoggedIn(ctx, user, c.ClientIP()) diff --git a/internal/repository/auth_source_cache.go b/internal/repository/auth_source_cache.go new file mode 100644 index 00000000..6650b127 --- /dev/null +++ b/internal/repository/auth_source_cache.go @@ -0,0 +1,269 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "fmt" + "strings" + "sync" + "time" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/pkg/cache/ram" +) + +const ( + authSourceActiveRedisKey = "oauth:auth_sources:active" + authSourceByNameRedisKeyFmt = "oauth:auth_sources:by_name:%s" + authSourceByNameRedisPattern = "oauth:auth_sources:by_name:*" + authSourceActiveRAMKey = "active" + authSourceCacheTTL = time.Hour + authSourceRAMMaximumSize = 64 + authSourceInvalidationChannel = "oauth:auth_source_invalidation" +) + +// authSourceRedisRecord persists full auth source credentials in Redis. +type authSourceRedisRecord struct { + ID uint64 `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + DisplayName string `json:"display_name"` + IsActive bool `json:"is_active"` + ClientID string `json:"client_id"` + ClientSecret string `json:"client_secret"` + OpenIDDiscoveryURL string `json:"openid_discovery_url"` + Scopes string `json:"scopes"` + IconURL string `json:"icon_url"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + ClientSecretConfigured bool `json:"client_secret_configured"` +} + +var ( + authSourceActiveRAM = ram.MustNew[string, []model.AuthSource](ram.Options{MaximumSize: authSourceRAMMaximumSize}) + authSourceByNameRAM = ram.MustNew[string, model.AuthSource](ram.Options{MaximumSize: authSourceRAMMaximumSize}) + authSourceListenerOnce sync.Once + authSourceListenerCtx context.Context + authSourceListenerCancel context.CancelFunc +) + +func cloneAuthSources(sources []model.AuthSource) []model.AuthSource { + if len(sources) == 0 { + return nil + } + cloned := make([]model.AuthSource, len(sources)) + copy(cloned, sources) + return cloned +} + +func cloneAuthSource(source model.AuthSource) model.AuthSource { + return source +} + +func normalizeAuthSourceName(name string) string { + return strings.TrimSpace(strings.ToLower(name)) +} + +func authSourceByNameRedisKey(name string) string { + return fmt.Sprintf(authSourceByNameRedisKeyFmt, normalizeAuthSourceName(name)) +} + +func authSourceToRedisRecord(source model.AuthSource) authSourceRedisRecord { + return authSourceRedisRecord{ + ID: source.ID, + Name: source.Name, + Type: source.Type, + DisplayName: source.DisplayName, + IsActive: source.IsActive, + ClientID: source.ClientID, + ClientSecret: source.ClientSecret, + OpenIDDiscoveryURL: source.OpenIDDiscoveryURL, + Scopes: source.Scopes, + IconURL: source.IconURL, + CreatedAt: source.CreatedAt, + UpdatedAt: source.UpdatedAt, + ClientSecretConfigured: source.ClientSecretConfigured, + } +} + +func redisRecordToAuthSource(record authSourceRedisRecord) model.AuthSource { + return model.AuthSource{ + ID: record.ID, + Name: record.Name, + Type: record.Type, + DisplayName: record.DisplayName, + IsActive: record.IsActive, + ClientID: record.ClientID, + ClientSecret: record.ClientSecret, + OpenIDDiscoveryURL: record.OpenIDDiscoveryURL, + Scopes: record.Scopes, + IconURL: record.IconURL, + CreatedAt: record.CreatedAt, + UpdatedAt: record.UpdatedAt, + ClientSecretConfigured: record.ClientSecretConfigured, + } +} + +func ensureAuthSourceCacheListener() { + if db.Redis == nil { + return + } + authSourceListenerOnce.Do(startAuthSourceCacheInvalidationListener) +} + +func startAuthSourceCacheInvalidationListener() { + authSourceListenerCtx, authSourceListenerCancel = context.WithCancel(context.Background()) + + go func() { + pubsub := db.Redis.Subscribe(authSourceListenerCtx, authSourceInvalidationChannel) + defer func() { + _ = pubsub.Close() + }() + + go func() { + <-authSourceListenerCtx.Done() + _ = pubsub.Close() + }() + + for range pubsub.Channel() { + authSourceActiveRAM.InvalidateAll() + authSourceByNameRAM.InvalidateAll() + } + }() +} + +func publishAuthSourceRAMInvalidation(ctx context.Context) { + if db.Redis == nil { + return + } + _ = db.Redis.Publish(ctx, authSourceInvalidationChannel, "reset").Err() +} + +func populateActiveAuthSourceCache(ctx context.Context, sources []model.AuthSource) { + cloned := cloneAuthSources(sources) + authSourceActiveRAM.Set(authSourceActiveRAMKey, cloned) + if db.Redis != nil { + _ = db.SetJSON(ctx, authSourceActiveRedisKey, cloned, authSourceCacheTTL) + } +} + +func populateAuthSourceByNameCache(ctx context.Context, name string, source *model.AuthSource) { + if source == nil { + return + } + cloned := cloneAuthSource(*source) + authSourceByNameRAM.Set(normalizeAuthSourceName(name), cloned) + if db.Redis != nil { + record := authSourceToRedisRecord(cloned) + _ = db.SetJSON(ctx, authSourceByNameRedisKey(name), record, authSourceCacheTTL) + } +} + +// GetActiveAuthSourcesCached returns active auth sources from RAM, Redis, or the database. +func GetActiveAuthSourcesCached(ctx context.Context) ([]model.AuthSource, error) { + ensureAuthSourceCacheListener() + + if sources, ok := authSourceActiveRAM.GetIfPresent(authSourceActiveRAMKey); ok { + return cloneAuthSources(sources), nil + } + + if db.Redis != nil { + var sources []model.AuthSource + if err := db.GetJSON(ctx, authSourceActiveRedisKey, &sources); err == nil { + populateActiveAuthSourceCache(ctx, sources) + return cloneAuthSources(sources), nil + } + } + + sources, err := model.GetActiveAuthSources(ctx) + if err != nil { + return nil, err + } + populateActiveAuthSourceCache(ctx, sources) + return cloneAuthSources(sources), nil +} + +// GetAuthSourceByNameCached returns an auth source by name from RAM, Redis, or the database. +func GetAuthSourceByNameCached(ctx context.Context, name string) (*model.AuthSource, error) { + ensureAuthSourceCacheListener() + + normalized := normalizeAuthSourceName(name) + if normalized == "" { + return model.GetAuthSourceByName(ctx, name) + } + + if source, ok := authSourceByNameRAM.GetIfPresent(normalized); ok { + cloned := cloneAuthSource(source) + return &cloned, nil + } + + if db.Redis != nil { + var record authSourceRedisRecord + if err := db.GetJSON(ctx, authSourceByNameRedisKey(name), &record); err == nil { + source := redisRecordToAuthSource(record) + populateAuthSourceByNameCache(ctx, name, &source) + cloned := cloneAuthSource(source) + return &cloned, nil + } + } + + source, err := model.GetAuthSourceByName(ctx, name) + if err != nil { + return nil, err + } + populateAuthSourceByNameCache(ctx, name, source) + cloned := cloneAuthSource(*source) + return &cloned, nil +} + +// InvalidateAuthSourceCache clears active and per-name auth source caches from RAM and Redis. +func InvalidateAuthSourceCache(ctx context.Context) error { + ensureAuthSourceCacheListener() + + authSourceActiveRAM.InvalidateAll() + authSourceByNameRAM.InvalidateAll() + + if db.Redis == nil { + return nil + } + + if err := db.Redis.Del(ctx, db.PrefixedKey(authSourceActiveRedisKey)).Err(); err != nil { + return err + } + + pattern := db.PrefixedKey(authSourceByNameRedisPattern) + iter := db.Redis.Scan(ctx, 0, pattern, 0).Iterator() + var keys []string + for iter.Next(ctx) { + keys = append(keys, iter.Val()) + } + if err := iter.Err(); err != nil { + return err + } + if len(keys) > 0 { + if err := db.Redis.Del(ctx, keys...).Err(); err != nil { + return err + } + } + + publishAuthSourceRAMInvalidation(ctx) + return nil +} + +// StopAuthSourceCacheListener stops the Redis Pub/Sub subscription listener and resets the sync.Once guard. +func StopAuthSourceCacheListener() { + if authSourceListenerCancel != nil { + authSourceListenerCancel() + authSourceListenerCancel = nil + } + authSourceListenerOnce = sync.Once{} +} + +// ResetAuthSourceRAMCacheForTest clears only the process-local RAM cache. +func ResetAuthSourceRAMCacheForTest() { + authSourceActiveRAM.InvalidateAll() + authSourceByNameRAM.InvalidateAll() +} diff --git a/internal/repository/auth_source_cache_test.go b/internal/repository/auth_source_cache_test.go new file mode 100644 index 00000000..2f69c359 --- /dev/null +++ b/internal/repository/auth_source_cache_test.go @@ -0,0 +1,240 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/glebarez/sqlite" + "github.com/redis/go-redis/v9" + "github.com/redis/go-redis/v9/maintnotifications" + "gorm.io/gorm" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" +) + +func setupAuthSourceCacheTest(t *testing.T) (*gorm.DB, *miniredis.Miniredis, func()) { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + if err != nil { + t.Fatalf("failed to open in-memory SQLite db: %v", err) + } + if err := sqliteDB.AutoMigrate(&model.AuthSource{}); err != nil { + t.Fatalf("failed to migrate auth sources: %v", err) + } + + miniRedis, err := miniredis.Run() + if err != nil { + t.Fatalf("failed to start miniredis: %v", err) + } + + db.SetDB(sqliteDB) + db.Redis = redis.NewClient(&redis.Options{ + Addr: miniRedis.Addr(), + MaintNotificationsConfig: &maintnotifications.Config{ + Mode: maintnotifications.ModeDisabled, + }, + }) + + ResetAuthSourceRAMCacheForTest() + + cleanup := func() { + StopAuthSourceCacheListener() + ResetAuthSourceRAMCacheForTest() + db.Redis.Close() + miniRedis.Close() + db.Redis = nil + } + return sqliteDB, miniRedis, cleanup +} + +func TestGetActiveAuthSourcesCached_LoadsFromRedisBeforeDB(t *testing.T) { + dbConn, _, cleanup := setupAuthSourceCacheTest(t) + defer cleanup() + ctx := context.Background() + + if err := InvalidateAuthSourceCache(ctx); err != nil { + t.Fatalf("InvalidateAuthSourceCache() error = %v", err) + } + + source := model.AuthSource{ + Name: "cached-source", + Type: model.AuthSourceTypeOIDC, + DisplayName: "Cached Source", + IsActive: true, + ClientID: "client-id", + ClientSecret: "client-secret", + OpenIDDiscoveryURL: "https://issuer.example.com", + } + if err := model.CreateAuthSource(ctx, &source); err != nil { + t.Fatalf("CreateAuthSource() error = %v", err) + } + + warmed, err := GetActiveAuthSourcesCached(ctx) + if err != nil { + t.Fatalf("GetActiveAuthSourcesCached() warm error = %v", err) + } + if len(warmed) == 0 || warmed[0].Name != source.Name { + t.Fatalf("GetActiveAuthSourcesCached() warm = %#v, want source %q", warmed, source.Name) + } + + if err := dbConn.Delete(&model.AuthSource{}, "id = ?", source.ID).Error; err != nil { + t.Fatalf("Delete(auth source) error = %v", err) + } + + ResetAuthSourceRAMCacheForTest() + + cached, err := GetActiveAuthSourcesCached(ctx) + if err != nil { + t.Fatalf("GetActiveAuthSourcesCached() cached error = %v", err) + } + if len(cached) == 0 || cached[0].Name != source.Name { + t.Fatalf("GetActiveAuthSourcesCached() = %#v, want redis-backed source %q", cached, source.Name) + } +} + +func TestGetAuthSourceByNameCached_LoadsFromRedisBeforeDB(t *testing.T) { + dbConn, _, cleanup := setupAuthSourceCacheTest(t) + defer cleanup() + ctx := context.Background() + + if err := InvalidateAuthSourceCache(ctx); err != nil { + t.Fatalf("InvalidateAuthSourceCache() error = %v", err) + } + + source := model.AuthSource{ + Name: "by-name-source", + Type: model.AuthSourceTypeOIDC, + DisplayName: "By Name Source", + IsActive: true, + ClientID: "client-id", + ClientSecret: "client-secret", + OpenIDDiscoveryURL: "https://issuer.example.com", + } + if err := model.CreateAuthSource(ctx, &source); err != nil { + t.Fatalf("CreateAuthSource() error = %v", err) + } + + warmed, err := GetAuthSourceByNameCached(ctx, source.Name) + if err != nil { + t.Fatalf("GetAuthSourceByNameCached() warm error = %v", err) + } + if warmed.Name != source.Name || warmed.ClientSecret != source.ClientSecret { + t.Fatalf("GetAuthSourceByNameCached() warm = %#v, want %#v", warmed, source) + } + + if err := dbConn.Delete(&model.AuthSource{}, "id = ?", source.ID).Error; err != nil { + t.Fatalf("Delete(auth source) error = %v", err) + } + + ResetAuthSourceRAMCacheForTest() + + cached, err := GetAuthSourceByNameCached(ctx, source.Name) + if err != nil { + t.Fatalf("GetAuthSourceByNameCached() cached error = %v", err) + } + if cached.Name != source.Name || cached.ClientSecret != source.ClientSecret { + t.Fatalf("GetAuthSourceByNameCached() = %#v, want redis-backed source %#v", cached, source) + } +} + +func TestInvalidateAuthSourceCache_ClearsRedisKeys(t *testing.T) { + _, _, cleanup := setupAuthSourceCacheTest(t) + defer cleanup() + ctx := context.Background() + + if err := InvalidateAuthSourceCache(ctx); err != nil { + t.Fatalf("InvalidateAuthSourceCache() initial error = %v", err) + } + + source := model.AuthSource{ + Name: "invalidate-source", + Type: model.AuthSourceTypeOIDC, + DisplayName: "Invalidate Source", + IsActive: true, + ClientID: "client-id", + ClientSecret: "client-secret", + OpenIDDiscoveryURL: "https://issuer.example.com", + } + if err := model.CreateAuthSource(ctx, &source); err != nil { + t.Fatalf("CreateAuthSource() error = %v", err) + } + if _, err := GetActiveAuthSourcesCached(ctx); err != nil { + t.Fatalf("GetActiveAuthSourcesCached() warm error = %v", err) + } + if _, err := GetAuthSourceByNameCached(ctx, source.Name); err != nil { + t.Fatalf("GetAuthSourceByNameCached() warm error = %v", err) + } + + if err := InvalidateAuthSourceCache(ctx); err != nil { + t.Fatalf("InvalidateAuthSourceCache() error = %v", err) + } + + activeExists, err := db.Redis.Exists(ctx, db.PrefixedKey(authSourceActiveRedisKey)).Result() + if err != nil { + t.Fatalf("Exists(active key) error = %v", err) + } + if activeExists != 0 { + t.Fatalf("active redis key still exists after invalidation") + } + + byNameExists, err := db.Redis.Exists(ctx, db.PrefixedKey(authSourceByNameRedisKey(source.Name))).Result() + if err != nil { + t.Fatalf("Exists(by-name key) error = %v", err) + } + if byNameExists != 0 { + t.Fatalf("by-name redis key still exists after invalidation") + } +} + +func TestAuthSourceInvalidationPubSubClearsPeerRAM(t *testing.T) { + dbConn, _, cleanup := setupAuthSourceCacheTest(t) + defer cleanup() + ctx := context.Background() + + source := model.AuthSource{ + Name: "pubsub-source", + Type: model.AuthSourceTypeOIDC, + DisplayName: "PubSub Source", + IsActive: true, + ClientID: "client-id", + ClientSecret: "client-secret", + OpenIDDiscoveryURL: "https://issuer.example.com", + } + if err := model.CreateAuthSource(ctx, &source); err != nil { + t.Fatalf("CreateAuthSource() error = %v", err) + } + + if _, err := GetActiveAuthSourcesCached(ctx); err != nil { + t.Fatalf("GetActiveAuthSourcesCached() error = %v", err) + } + if err := dbConn.Delete(&model.AuthSource{}, "id = ?", source.ID).Error; err != nil { + t.Fatalf("Delete(auth source) error = %v", err) + } + if _, err := GetActiveAuthSourcesCached(ctx); err != nil { + t.Fatalf("expected RAM cache hit before pub/sub invalidation: %v", err) + } + + if err := db.Redis.Publish(ctx, authSourceInvalidationChannel, "reset").Err(); err != nil { + t.Fatalf("publish invalidation: %v", err) + } + + deadline := time.Now().Add(500 * time.Millisecond) + for time.Now().Before(deadline) { + if _, ok := authSourceActiveRAM.GetIfPresent(authSourceActiveRAMKey); !ok { + break + } + time.Sleep(10 * time.Millisecond) + } + if _, ok := authSourceActiveRAM.GetIfPresent(authSourceActiveRAMKey); ok { + t.Fatal("expected peer RAM cache to be cleared by pub/sub") + } +} \ No newline at end of file diff --git a/internal/repository/system_config.go b/internal/repository/system_config.go index 9d20d708..ac8641ca 100644 --- a/internal/repository/system_config.go +++ b/internal/repository/system_config.go @@ -75,6 +75,22 @@ func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]mod missing = append(missing, key) } + if len(missing) > 0 && db.Redis != nil { + stillMissing := make([]string, 0, len(missing)) + for _, key := range missing { + var sc model.SystemConfig + if err := db.HGetJSON(ctx, SystemConfigRedisHashKey, key, &sc); err == nil { + systemConfigRAMCache.Set(key, cloneSystemConfig(sc)) + result[key] = sc + continue + } else if !errors.Is(err, redis.Nil) { + return nil, err + } + stillMissing = append(stillMissing, key) + } + missing = stillMissing + } + if len(missing) == 0 { return result, nil } diff --git a/internal/repository/system_config_test.go b/internal/repository/system_config_test.go new file mode 100644 index 00000000..8dc601a3 --- /dev/null +++ b/internal/repository/system_config_test.go @@ -0,0 +1,146 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "testing" + + "github.com/Rain-kl/Wavelet/internal/db" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/alicebob/miniredis/v2" + "github.com/glebarez/sqlite" + "github.com/redis/go-redis/v9" + "github.com/redis/go-redis/v9/maintnotifications" + "gorm.io/gorm" +) + +func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) { + t.Helper() + + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ + DisableForeignKeyConstraintWhenMigrating: true, + }) + if err != nil { + t.Fatalf("gorm.Open(sqlite) error = %v", err) + } + if err := sqliteDB.AutoMigrate(&model.SystemConfig{}); err != nil { + t.Fatalf("AutoMigrate(SystemConfig) error = %v", err) + } + + siteConfig := model.SystemConfig{ + Key: model.ConfigKeySiteName, + Value: "Wavelet", + Type: "system", + Description: "系统平台的展示名称", + } + if err := sqliteDB.Create(&siteConfig).Error; err != nil { + t.Fatalf("Create(site_name) error = %v", err) + } + + mr, err := miniredis.Run() + if err != nil { + t.Fatalf("miniredis.Run() error = %v", err) + } + redisClient := redis.NewClient(&redis.Options{ + Addr: mr.Addr(), + MaintNotificationsConfig: &maintnotifications.Config{ + Mode: maintnotifications.ModeDisabled, + }, + }) + + previousRedis := db.Redis + db.SetDB(sqliteDB) + db.Redis = redisClient + + cleanup := func() { + StopSystemConfigCacheListener() + ResetSystemConfigRAMCacheForTest() + db.SetDB(nil) + db.Redis = previousRedis + _ = redisClient.Close() + mr.Close() + } + + return sqliteDB, cleanup +} + +func TestListSystemConfigsByKeys_EmptyKeys(t *testing.T) { + result, err := ListSystemConfigsByKeys(context.Background(), nil) + if err != nil { + t.Fatalf("ListSystemConfigsByKeys(nil) error = %v", err) + } + if len(result) != 0 { + t.Fatalf("ListSystemConfigsByKeys(nil) = %#v, want empty map", result) + } +} + +func TestListSystemConfigsByKeys_LoadsFromRedisBeforeDB(t *testing.T) { + dbConn, cleanup := setupSystemConfigTest(t) + defer cleanup() + ctx := context.Background() + + ResetSystemConfigRAMCacheForTest() + if err := InvalidateAllSystemConfigCaches(ctx); err != nil { + t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err) + } + + warm, err := GetSystemConfigByKey(ctx, model.ConfigKeySiteName) + if err != nil { + t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err) + } + if warm.Value != "Wavelet" { + t.Fatalf("GetSystemConfigByKey(site_name).Value = %q, want %q", warm.Value, "Wavelet") + } + + if err := dbConn.Model(&model.SystemConfig{}). + Where("key = ?", model.ConfigKeySiteName). + Update("value", "db_only_value").Error; err != nil { + t.Fatalf("Update(site_name) error = %v", err) + } + + ResetSystemConfigRAMCacheForTest() + + configs, err := ListSystemConfigsByKeys(ctx, []string{model.ConfigKeySiteName}) + if err != nil { + t.Fatalf("ListSystemConfigsByKeys(site_name) error = %v", err) + } + + sc, ok := configs[model.ConfigKeySiteName] + if !ok { + t.Fatal("ListSystemConfigsByKeys(site_name) missing site_name entry") + } + if sc.Value != "Wavelet" { + t.Fatalf("ListSystemConfigsByKeys(site_name).Value = %q, want redis value %q", sc.Value, "Wavelet") + } +} + +func TestListSystemConfigsByKeys_PopulatesRAMFromRedis(t *testing.T) { + _, cleanup := setupSystemConfigTest(t) + defer cleanup() + ctx := context.Background() + + ResetSystemConfigRAMCacheForTest() + if err := InvalidateAllSystemConfigCaches(ctx); err != nil { + t.Fatalf("InvalidateAllSystemConfigCaches() error = %v", err) + } + + if _, err := GetSystemConfigByKey(ctx, model.ConfigKeySiteName); err != nil { + t.Fatalf("GetSystemConfigByKey(site_name) warm error = %v", err) + } + + ResetSystemConfigRAMCacheForTest() + + if _, err := ListSystemConfigsByKeys(ctx, []string{model.ConfigKeySiteName}); err != nil { + t.Fatalf("ListSystemConfigsByKeys(site_name) error = %v", err) + } + + cached, ok := systemConfigRAMCache.GetIfPresent(model.ConfigKeySiteName) + if !ok { + t.Fatal("expected RAM cache to be populated after redis hit") + } + if cached.Value != "Wavelet" { + t.Fatalf("RAM cache value = %q, want %q", cached.Value, "Wavelet") + } +} \ No newline at end of file diff --git a/internal/repository/upload.go b/internal/repository/upload.go index 355758e7..34c8a96f 100644 --- a/internal/repository/upload.go +++ b/internal/repository/upload.go @@ -65,7 +65,12 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (model.Upload, error) { // SoftDeleteUpload marks an upload as deleted. // External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this. func SoftDeleteUpload(ctx context.Context, upload *model.Upload) error { - return db.DB(ctx).Model(upload).Update("status", model.UploadStatusDeleted).Error + return SoftDeleteUploadTx(db.DB(ctx), upload) +} + +// SoftDeleteUploadTx marks an upload as deleted within an existing transaction. +func SoftDeleteUploadTx(tx *gorm.DB, upload *model.Upload) error { + return tx.Model(upload).Update("status", model.UploadStatusDeleted).Error } // UpdateUpload applies partial field updates to an upload record. @@ -100,7 +105,12 @@ func FindReusableUploadByHash(ctx context.Context, hash string, size int64) (mod // CreateUpload persists a new upload record. // External modules must use upload.Ingest; only internal/apps/upload may call this. func CreateUpload(ctx context.Context, upload *model.Upload) error { - return db.DB(ctx).Create(upload).Error + return CreateUploadTx(db.DB(ctx), upload) +} + +// CreateUploadTx persists a new upload record within an existing transaction. +func CreateUploadTx(tx *gorm.DB, upload *model.Upload) error { + return tx.Create(upload).Error } // ListUploadsByIDs returns active uploads matching the given IDs.