mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-28 05:46:36 +08:00
perf(cache): 三层缓存框架补强
- 新增 cache-framework skill,规范 RAM→Redis→DB 读路径、失效与 pub/sub - 上传元数据 Otter+Redis 缓存与多节点失效;Auth Source 缓存与 pub/sub - ListSystemConfigsByKeys 补 Redis 层;上传统计单事务;登录/Token 缓存预热 - cleanup 任务补 upload meta 失效钩子
This commit is contained in:
@@ -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`
|
||||||
@@ -42,6 +42,7 @@
|
|||||||
| `database-migration` | 数据库表结构变更、goose SQL 迁移(PG/SQLite/ClickHouse)、seed 数据 |
|
| `database-migration` | 数据库表结构变更、goose SQL 迁移(PG/SQLite/ClickHouse)、seed 数据 |
|
||||||
| `clickhouse-batchwriter` | ClickHouse 批量写入、`internal/db/batchwriter` 接入、分析表异步 flush、背压与写入路径改造 |
|
| `clickhouse-batchwriter` | ClickHouse 批量写入、`internal/db/batchwriter` 接入、分析表异步 flush、背压与写入路径改造 |
|
||||||
| `file-upload` | 业务上传文件、Worker 程序化摄取、`upload.Ingest` 策略选型、文件访问与 `w_uploads` / 统计排查 |
|
| `file-upload` | 业务上传文件、Worker 程序化摄取、`upload.Ingest` 策略选型、文件访问与 `w_uploads` / 统计排查 |
|
||||||
|
| `cache-framework` | 新增或修改业务缓存(RAM/Redis/DB 三层读路径)、缓存失效、多节点 pub/sub 同步、评估高频读是否应接入缓存 |
|
||||||
| `push-notification` | 系统通知推送事件、统一触发器投递、带消息推送的业务功能 |
|
| `push-notification` | 系统通知推送事件、统一触发器投递、带消息推送的业务功能 |
|
||||||
| `release-guide` | 根据自上一正式版本 Tag 以来的提交整理 Version Bump 提交信息以触发双语 Release |
|
| `release-guide` | 根据自上一正式版本 Tag 以来的提交整理 Version Bump 提交信息以触发双语 Release |
|
||||||
| `shadcn` | 添加、修改或组合 shadcn/ui 组件 |
|
| `shadcn` | 添加、修改或组合 shadcn/ui 组件 |
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/apps/admin"
|
"github.com/Rain-kl/Wavelet/internal/apps/admin"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
"github.com/gin-gonic/gin"
|
"github.com/gin-gonic/gin"
|
||||||
|
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"github.com/Rain-kl/Wavelet/internal/common/response"
|
||||||
@@ -91,6 +92,7 @@ func CreateAuthSource(c *gin.Context) {
|
|||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
_ = repository.InvalidateAuthSourceCache(c.Request.Context())
|
||||||
source.Sanitize()
|
source.Sanitize()
|
||||||
c.JSON(http.StatusOK, response.OK(source))
|
c.JSON(http.StatusOK, response.OK(source))
|
||||||
}
|
}
|
||||||
@@ -150,6 +152,7 @@ func UpdateAuthSource(c *gin.Context) {
|
|||||||
oauth.InvalidateOIDCProviderCache(normalizeIssuer(existing.OpenIDDiscoveryURL))
|
oauth.InvalidateOIDCProviderCache(normalizeIssuer(existing.OpenIDDiscoveryURL))
|
||||||
}
|
}
|
||||||
oauth.InvalidateOIDCProviderCache(normalizeIssuer(req.OpenIDDiscoveryURL))
|
oauth.InvalidateOIDCProviderCache(normalizeIssuer(req.OpenIDDiscoveryURL))
|
||||||
|
_ = repository.InvalidateAuthSourceCache(c.Request.Context())
|
||||||
|
|
||||||
updated, err := model.GetAuthSourceByID(c.Request.Context(), id)
|
updated, err := model.GetAuthSourceByID(c.Request.Context(), id)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -191,6 +194,7 @@ func ToggleAuthSource(c *gin.Context) {
|
|||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
_ = repository.InvalidateAuthSourceCache(c.Request.Context())
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -216,6 +220,7 @@ func DeleteAuthSource(c *gin.Context) {
|
|||||||
response.AbortBadRequest(c, err.Error())
|
response.AbortBadRequest(c, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
_ = repository.InvalidateAuthSourceCache(c.Request.Context())
|
||||||
c.JSON(http.StatusOK, response.OKNil())
|
c.JSON(http.StatusOK, response.OKNil())
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -25,16 +25,16 @@ func isOIDCLoginEnabled(ctx context.Context) bool {
|
|||||||
func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
|
func resolveAuthSource(ctx context.Context, sourceName string) (*model.AuthSource, error) {
|
||||||
name := strings.TrimSpace(strings.ToLower(sourceName))
|
name := strings.TrimSpace(strings.ToLower(sourceName))
|
||||||
if name == "" {
|
if name == "" {
|
||||||
sources, err := model.GetActiveAuthSources(ctx)
|
sources, err := repository.GetActiveAuthSourcesCached(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
if len(sources) == 0 {
|
if len(sources) == 0 {
|
||||||
return nil, errors.New(errNoActiveAuthSource)
|
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 {
|
func activeLoginSources(ctx context.Context) []AuthSourceView {
|
||||||
@@ -43,7 +43,7 @@ func activeLoginSources(ctx context.Context) []AuthSourceView {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
dbSources, err := model.GetActiveAuthSources(ctx)
|
dbSources, err := repository.GetActiveAuthSourcesCached(ctx)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -179,6 +179,8 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
|
|||||||
return
|
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())
|
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())
|
listener.EmitAdminLoggedIn(ctx, &user, c.ClientIP())
|
||||||
|
|||||||
@@ -95,6 +95,28 @@ func (m *mockRedisClient) Del(ctx context.Context, keys ...string) *redis.IntCmd
|
|||||||
return cmd
|
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 {
|
func (m *mockRedisClient) HSet(ctx context.Context, key string, values ...interface{}) *redis.IntCmd {
|
||||||
cmd := redis.NewIntCmd(ctx)
|
cmd := redis.NewIntCmd(ctx)
|
||||||
if len(values) >= 2 {
|
if len(values) >= 2 {
|
||||||
@@ -129,6 +151,12 @@ func (m *mockRedisClient) HGet(ctx context.Context, key string, field string) *r
|
|||||||
return cmd
|
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 {
|
func (m *mockRedisClient) Subscribe(ctx context.Context, channels ...string) *redis.PubSub {
|
||||||
return redis.NewClient(&redis.Options{
|
return redis.NewClient(&redis.Options{
|
||||||
Addr: "127.0.0.1:0",
|
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 {
|
func setupTestDB(t *testing.T) *gorm.DB {
|
||||||
repository.ResetSystemConfigRAMCacheForTest()
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
|
repository.ResetAuthSourceRAMCacheForTest()
|
||||||
|
|
||||||
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
dbConn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -635,6 +664,7 @@ func TestAuthorize(t *testing.T) {
|
|||||||
|
|
||||||
// Case 2: Inactive Source Authorize
|
// Case 2: Inactive Source Authorize
|
||||||
dbConn.Model(&model.AuthSource{}).Where("id = ?", 101).Update("is_active", false)
|
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)
|
w2 := performRequest(router, http.MethodGet, "/api/v1/oauth/github/authorize", nil, nil, nil)
|
||||||
if w2.Code != http.StatusBadRequest {
|
if w2.Code != http.StatusBadRequest {
|
||||||
t.Errorf("expected 400 for inactive source, got %d", w2.Code)
|
t.Errorf("expected 400 for inactive source, got %d", w2.Code)
|
||||||
@@ -1122,6 +1152,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
repository.ResetSystemConfigRAMCacheForTest()
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
||||||
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false)
|
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)
|
wSourceInactive := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
||||||
if wSourceInactive.Code != http.StatusBadRequest {
|
if wSourceInactive.Code != http.StatusBadRequest {
|
||||||
@@ -1134,6 +1165,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
repository.ResetSystemConfigRAMCacheForTest()
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
||||||
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true)
|
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)
|
wAuthDisabled := performRequest(router, http.MethodGet, "/api/v1/oauth/"+testSourceName+"/authorize", nil, nil, nil)
|
||||||
if wAuthDisabled.Code != http.StatusBadRequest {
|
if wAuthDisabled.Code != http.StatusBadRequest {
|
||||||
@@ -1146,6 +1178,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
repository.ResetSystemConfigRAMCacheForTest()
|
repository.ResetSystemConfigRAMCacheForTest()
|
||||||
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
mockRedis.Del(context.Background(), db.PrefixedKey(repository.SystemConfigRedisHashKey)+":"+model.ConfigKeyOIDCLoginEnabled)
|
||||||
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true)
|
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)
|
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
|
// Since callback deletes state, we need to generate state again
|
||||||
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", true)
|
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)
|
wLogin2 := performRequest(router, http.MethodGet, "/api/v1/oauth/login?source="+testSourceName, nil, nil, nil)
|
||||||
_ = json.Unmarshal(wLogin2.Body.Bytes(), &loginUrlResp)
|
_ = json.Unmarshal(wLogin2.Body.Bytes(), &loginUrlResp)
|
||||||
parsedURL, _ = url.Parse(loginUrlResp.Data.AuthorizeURL)
|
parsedURL, _ = url.Parse(loginUrlResp.Data.AuthorizeURL)
|
||||||
@@ -1201,6 +1235,7 @@ func TestOIDCPolicyEnforcement(t *testing.T) {
|
|||||||
|
|
||||||
// Deactivate source
|
// Deactivate source
|
||||||
dbConn.Model(&model.AuthSource{}).Where("name = ?", testSourceName).Update("is_active", false)
|
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)
|
reqBody2 := fmt.Sprintf(`{"state":"%s","code":"test_auth_code"}`, state)
|
||||||
wCallbackSourceInactive := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody2), map[string]string{
|
wCallbackSourceInactive := performRequest(router, http.MethodPost, "/api/v1/oauth/callback", []byte(reqBody2), map[string]string{
|
||||||
"Content-Type": "application/json",
|
"Content-Type": "application/json",
|
||||||
|
|||||||
+151
@@ -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{}
|
||||||
|
}
|
||||||
+287
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -22,7 +22,6 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common"
|
"github.com/Rain-kl/Wavelet/internal/common"
|
||||||
"github.com/Rain-kl/Wavelet/internal/common/response"
|
"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/diskcache"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
|
|
||||||
@@ -96,10 +95,8 @@ func GetUploadRecordByID(c *gin.Context) (*model.Upload, error) {
|
|||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
var upload model.Upload
|
upload, err := cache.GetUploadByID(c.Request.Context(), uploadID)
|
||||||
if err := db.DB(c.Request.Context()).
|
if err != nil {
|
||||||
Where("id = ? AND status IN (?, ?)", uploadID, model.UploadStatusPending, model.UploadStatusUsed).
|
|
||||||
First(&upload).Error; err != nil {
|
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -34,6 +34,10 @@ import (
|
|||||||
"gorm.io/gorm"
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func init() {
|
||||||
|
testhelper.RegisterCleanup(cache.ResetUploadMetaCacheForTest)
|
||||||
|
}
|
||||||
|
|
||||||
func TestServeFileByIDAccessControl(t *testing.T) {
|
func TestServeFileByIDAccessControl(t *testing.T) {
|
||||||
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|||||||
@@ -11,9 +11,11 @@ import (
|
|||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
||||||
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
||||||
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
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/db/idgen"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"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 {
|
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)
|
_, backend, backendErr := storage.Active(ctx)
|
||||||
if backendErr == nil {
|
if backendErr == nil {
|
||||||
if deleteErr := backend.Delete(ctx, objectKey); deleteErr != 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
|
return err
|
||||||
}
|
}
|
||||||
uploadstats.RecordUploadStatsAdd(ctx, upload)
|
uploadcache.SetUploadMetaCache(ctx, upload)
|
||||||
return nil
|
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) {
|
func createDedupRecord(ctx context.Context, existing model.Upload, req Request) (Result, error) {
|
||||||
accessMode := resolveAccessMode(req.Type, req.AccessMode)
|
accessMode := resolveAccessMode(req.Type, req.AccessMode)
|
||||||
newUpload := model.Upload{
|
newUpload := model.Upload{
|
||||||
|
|||||||
@@ -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) {
|
func TestRemoveDecrementsStats(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|||||||
@@ -6,9 +6,12 @@ package ingest
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
|
||||||
|
uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
|
||||||
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
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/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/repository"
|
"github.com/Rain-kl/Wavelet/internal/repository"
|
||||||
|
"gorm.io/gorm"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Remove soft-deletes an upload and decrements incremental stats.
|
// 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 {
|
if err != nil {
|
||||||
return model.Upload{}, err
|
return model.Upload{}, err
|
||||||
}
|
}
|
||||||
uploadstats.RecordUploadStatsRemove(ctx, &upload)
|
if err := softDeleteUploadWithStats(ctx, &upload); err != nil {
|
||||||
if err := repository.SoftDeleteUpload(ctx, &upload); err != nil {
|
|
||||||
return model.Upload{}, err
|
return model.Upload{}, err
|
||||||
}
|
}
|
||||||
upload.Status = model.UploadStatusDeleted
|
upload.Status = model.UploadStatusDeleted
|
||||||
@@ -34,10 +36,23 @@ func RemoveOwned(ctx context.Context, userID, uploadID uint64) (model.Upload, er
|
|||||||
if upload.UserID != userID {
|
if upload.UserID != userID {
|
||||||
return model.Upload{}, ErrForbidden
|
return model.Upload{}, ErrForbidden
|
||||||
}
|
}
|
||||||
uploadstats.RecordUploadStatsRemove(ctx, &upload)
|
if err := softDeleteUploadWithStats(ctx, &upload); err != nil {
|
||||||
if err := repository.SoftDeleteUpload(ctx, &upload); err != nil {
|
|
||||||
return model.Upload{}, err
|
return model.Upload{}, err
|
||||||
}
|
}
|
||||||
upload.Status = model.UploadStatusDeleted
|
upload.Status = model.UploadStatusDeleted
|
||||||
return upload, nil
|
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
|
||||||
|
}
|
||||||
@@ -37,7 +37,7 @@ func RebuildUploadStats(ctx context.Context) error {
|
|||||||
}
|
}
|
||||||
|
|
||||||
for i := range uploads {
|
for i := range uploads {
|
||||||
if err := applyUploadStatsDeltaTx(tx, &uploads[i], 1); err != nil {
|
if err := ApplyUploadStatsDeltaTx(tx, &uploads[i], 1); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -50,11 +50,12 @@ func applyUploadStatsDelta(ctx context.Context, upload *model.Upload, sign int64
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
|
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 {
|
if upload == nil || !isActiveUploadStatus(upload.Status) || sign == 0 {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -11,8 +11,39 @@ import (
|
|||||||
"github.com/Rain-kl/Wavelet/internal/db"
|
"github.com/Rain-kl/Wavelet/internal/db"
|
||||||
"github.com/Rain-kl/Wavelet/internal/model"
|
"github.com/Rain-kl/Wavelet/internal/model"
|
||||||
"github.com/Rain-kl/Wavelet/internal/testhelper"
|
"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) {
|
func TestApplyUploadStatsAddAndRemove(t *testing.T) {
|
||||||
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
_, _, cleanup := testhelper.SetupTestEnvironment(t)
|
||||||
defer cleanup()
|
defer cleanup()
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
|
uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
|
||||||
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
|
||||||
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
|
||||||
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
|
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)
|
uploadstats.RecordUploadStatsRemove(ctx, &u)
|
||||||
|
uploadcache.InvalidateUploadMetaCache(ctx, u.ID)
|
||||||
totalDeleted++
|
totalDeleted++
|
||||||
lastID = u.ID
|
lastID = u.ID
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -362,6 +362,7 @@ func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) erro
|
|||||||
if err := db.DB(ctx).Create(record).Error; err != nil {
|
if err := db.DB(ctx).Create(record).Error; err != nil {
|
||||||
return errors.New("创建令牌失败,请稍后再试")
|
return errors.New("创建令牌失败,请稍后再试")
|
||||||
}
|
}
|
||||||
|
oauth.SetCachedToken(ctx, record.TokenHash, record)
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -405,6 +406,8 @@ func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *mo
|
|||||||
return "", nil, errors.New("轮换令牌失败,请稍后再试")
|
return "", nil, errors.New("轮换令牌失败,请稍后再试")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
oauth.SetCachedToken(ctx, tokenRecord.TokenHash, &tokenRecord)
|
||||||
|
|
||||||
return newTokenStr, &tokenRecord, nil
|
return newTokenStr, &tokenRecord, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -168,6 +168,8 @@ func Login(c *gin.Context) {
|
|||||||
return
|
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())
|
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())
|
listener.EmitAdminLoggedIn(ctx, user, c.ClientIP())
|
||||||
|
|||||||
@@ -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()
|
||||||
|
}
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -75,6 +75,22 @@ func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]mod
|
|||||||
missing = append(missing, key)
|
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 {
|
if len(missing) == 0 {
|
||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -65,7 +65,12 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (model.Upload, error) {
|
|||||||
// SoftDeleteUpload marks an upload as deleted.
|
// SoftDeleteUpload marks an upload as deleted.
|
||||||
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
|
// 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 {
|
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.
|
// 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.
|
// CreateUpload persists a new upload record.
|
||||||
// External modules must use upload.Ingest; only internal/apps/upload may call this.
|
// External modules must use upload.Ingest; only internal/apps/upload may call this.
|
||||||
func CreateUpload(ctx context.Context, upload *model.Upload) error {
|
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.
|
// ListUploadsByIDs returns active uploads matching the given IDs.
|
||||||
|
|||||||
Reference in New Issue
Block a user