diff --git a/.agent/skills/new-api/SKILL.md b/.agent/skills/new-api/SKILL.md index 38034e18..d4d2a32b 100644 --- a/.agent/skills/new-api/SKILL.md +++ b/.agent/skills/new-api/SKILL.md @@ -115,7 +115,7 @@ internal/ ## 核心开发步骤 (Step-by-Step Flow) ### 步骤 1:数据库定义与迁移 -如果自定义功能涉及新表或字段,请参考 [database-migration](../database-migration/SKILL.md) 技能,在 `internal/infra/persistence/migrator/goose/` 目录下编写迁移文件并在 `internal/model/` 中定义 GORM 数据模型。 +如果自定义功能涉及新表或字段,请参考 [database-migration](../database-migration/SKILL.md) 技能,在 `internal/infra/persistence/migrator/goose/` 目录下编写迁移文件,在 `internal/model/` 中定义 GORM 实体(无 CRUD / 无 DB 访问),并在 `internal/repository/` 中实现数据访问(**repository 为唯一持久化入口**)。 ### 步骤 2:在模块内实现业务逻辑 (`logics.go` / `service.go`) 业务逻辑逻辑应当实现于 `internal/apps/custom/` 目录下: diff --git a/.agent/skills/new-async-task/SKILL.md b/.agent/skills/new-async-task/SKILL.md index c59e5268..d5d7463a 100644 --- a/.agent/skills/new-async-task/SKILL.md +++ b/.agent/skills/new-async-task/SKILL.md @@ -19,7 +19,8 @@ description: "Wavelet 项目专用:新增或修改 Asynq 异步任务、后台 - `internal/infra/task/worker/worker.go`:Worker 路由和队列 - `internal/infra/task/scheduler/scheduler.go`:定时调度 - `internal/apps/admin/task/routers.go`:Admin 任务 API -- `internal/model/task_execution.go`:执行记录和日志持久化 +- `internal/model/task_execution.go`:执行记录实体与 DTO +- `internal/repository/task_execution.go`:执行记录和日志持久化 需要模板时阅读 [references/CODE-EXAMPLES.md](references/CODE-EXAMPLES.md)。 @@ -42,7 +43,7 @@ description: "Wavelet 项目专用:新增或修改 Asynq 异步任务、后台 - 成功返回 `&task.TaskResult{Message: ..., Detail: ...}`。 - 失败返回 error,由任务框架处理状态和重试。 - 不要吞掉关键错误。 -- 复杂 SQL 放到 `internal/model/` 或模块内的业务服务层(如 `internal/apps//service.go` 或 `logics.go`)。 +- 持久化只通过 `internal/repository/`(唯一入口);业务编排放模块内 `logics.go` / `service.go`。`internal/model` 仅实体/DTO,禁止 CRUD 与 DB 访问。 ### 注册 diff --git a/.agent/skills/new-setting/SKILL.md b/.agent/skills/new-setting/SKILL.md index 99b9d60e..12a88416 100644 --- a/.agent/skills/new-setting/SKILL.md +++ b/.agent/skills/new-setting/SKILL.md @@ -14,7 +14,7 @@ description: "Wavelet 项目专用:当新增或修改启动时设置、数据 Wavelet 当前有两套设置入口: - 启动时设置:来自 `config.yaml` 或环境变量,适合进程启动前必须确定、通常不热更新的基础配置。 -- 系统设置:保存于数据库 `system_configs`,经 `model.SystemConfig` 和 Redis hash 缓存读取,支持运行时热更新。管理入口是 `/admin/system` 和 `/admin/settings`。 +- 系统设置:保存于数据库 `system_configs`,经 `model.SystemConfig` 实体(key 常量在 model)与 `repository` 读取层(含 Redis hash 缓存)访问,支持运行时热更新。管理入口是 `/admin/system` 和 `/admin/settings`。 系统设置分三种使用语义: @@ -30,7 +30,8 @@ Wavelet 当前有两套设置入口: 修改前快速查看这些文件,确认当前实现没有漂移: -- `internal/model/system_configs.go`: 配置 key 常量、`SystemConfig` 模型、`GetByKey`、`GetBoolByKey`、`GetIntByKey`、`GetDecimalByKey` 等读取方法。 +- `internal/model/system_configs.go`: 配置 key 常量(`ConfigKey*`)、`SystemConfig` 实体与字段语义;**不含**持久化读取 API。 +- `internal/repository/system_config.go`: 配置读取与缓存(`GetSystemConfigByKey`、`GetBoolByKey`、`GetIntByKey`、`GetDecimalByKey`、`ListVisibleSystemConfigs` 等)。 - `internal/infra/persistence/migrator/goose/postgres/*.sql` 和 `internal/infra/persistence/migrator/goose/sqlite/*.sql`: `system_configs` 表结构、初始化 seed、后续升级迁移。 - `internal/infra/persistence/migrator/migrator.go`: goose 迁移入口和 PostgreSQL/SQLite 方言选择。 - `internal/testhelper/test_helper.go`: Go 测试用默认系统配置 seed。 @@ -61,12 +62,13 @@ Wavelet 当前有两套设置入口: - 如果相关 Go 包测试依赖默认配置,同步 `internal/testhelper/test_helper.go` 的 `seedDefaultConfigs` 和公共 key 列表。 3. 读取配置。 - - 后端业务代码优先使用 `model.GetBoolByKey`、`model.GetIntByKey`、`model.GetDecimalByKey` 或 `SystemConfig.GetByKey`。 + - 后端业务代码通过 `internal/repository` 读取:`repository.GetBoolByKey`、`repository.GetIntByKey`、`repository.GetDecimalByKey` 或 `repository.GetSystemConfigByKey`;key 常量仍用 `model.ConfigKey*`。 + - 禁止新增或调用 `model.Get*ByKey` / `model.ListVisibleSystemConfigs` 等数据访问 API(model 无 CRUD)。 - 运行时可热更新的规则不要放进 `config.Config`;启动时设置才走 `internal/infra/config/model.go` 和 `config.example.yaml`。 - 不要在 handler 或业务代码里直接读 `os.Getenv()`。 4. 如果前端需要未登录或全局消费,暴露为公共可见配置。 - - 把该配置的 `visibility` 设为 `1`,`GetPublicConfig` 会通过 `model.ListVisibleSystemConfigs` 返回所有可见 key/value。 + - 把该配置的 `visibility` 设为 `1`,`GetPublicConfig` 会通过 `repository.ListVisibleSystemConfigs` 返回所有可见 key/value。 - `/api/v1/config/public` 的 `data` 是动态对象:后端返回 `map[string]string`,前端类型是 `Record`。 - 前端读取时按配置 key 访问,必要时在消费侧把字符串转换为 boolean/number/JSON。 - 检查使用方的 query key,更新后需要 invalidate `["public-config"]`。 @@ -99,9 +101,9 @@ Wavelet 当前有两套设置入口: ### 布尔公共设置 -- model key:`ConfigKeyFeatureEnabled = "feature_enabled"` +- model key:`ConfigKeyFeatureEnabled = "feature_enabled"`(定义在 `internal/model`) - goose SQL 默认值:`value='false'`,`type` 按语义选 `"system"` 或 `"business"`,`visibility=1`。 -- 后端读取:`model.GetBoolByKey(ctx, model.ConfigKeyFeatureEnabled)`。 +- 后端读取:`repository.GetBoolByKey(ctx, model.ConfigKeyFeatureEnabled)`。 - 公共响应:`/api/v1/config/public` 的 `data.feature_enabled` 为字符串 `"true"` 或 `"false"`。 - 前端图形控件:`Switch`,保存时写 `"true"` / `"false"`。 @@ -109,13 +111,13 @@ Wavelet 当前有两套设置入口: - model key:`ConfigKeyMaxSomething = "max_something"`。 - goose SQL 默认值:例如 `"5"`,`type` 通常为 `"business"`,只有前端公共消费时才设 `visibility=1`。 -- 后端读取:`model.GetIntByKey` 或 `model.GetDecimalByKey`。 +- 后端读取:`repository.GetIntByKey` 或 `repository.GetDecimalByKey`。 - 前端图形控件:`Input type="number"` 或合适的 shadcn 数值控件;保存前做最小必要校验,错误用 toast。 ### JSON 设置 - 默认值使用合法 JSON,例如 `"{}"` 或 `"[]"`。 -- 在 model 或 service 层提供解析函数,像 `GetMenuDisplayConfig` 一样把 JSON 解析错误包装成清晰错误。 +- 在 repository 或业务 logics 中提供解析函数,像 `repository.GetMenuDisplayConfig` 一样把 JSON 解析错误包装成清晰错误;不要在 model 中做 IO。 - 前端不要直接拼接 JSON 字符串;用 `JSON.stringify` 写入,用类型化对象在组件中操作。 ## 验证 @@ -125,7 +127,7 @@ Wavelet 当前有两套设置入口: - 新增或修改系统配置默认值、visibility 或公共配置读取:至少运行相关 Go 包测试,例如: ```bash -go test ./internal/model ./internal/apps/config ./internal/apps/admin/system_config +go test ./internal/repository ./internal/apps/config ./internal/apps/admin/system_config ``` - 新增 goose 迁移后,至少用当前数据库方言跑一次迁移;如果 SQL 同时改了 PostgreSQL 和 SQLite,尽量覆盖两种方言。涉及 schema/seed 的任务还应遵循 database-migration skill。 diff --git a/AGENTS.md b/AGENTS.md index 36c0d838..3cc98003 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -59,6 +59,13 @@ - 在完成代码开发后或者 git 提交前必须运行 `make format` 格式化代码。 - 需要缓存或文件管理能力时,必须复用现有平台实现,禁止在业务包中自行创建缓存目录、直接管理缓存文件或重复封装存储后端。 - 文件摄取必须通过 `upload.Ingest`(`upload.PolicyCreate` / `PolicyDedupNewRecord` / `PolicyResolveExisting`);删除必须通过 `upload.Remove` 或 `upload.RemoveOwned`。禁止业务模块直接调用 `repository.CreateUpload` / `repository.SoftDeleteUpload`,禁止 `db.Create(&model.Upload{})` 旁路写 `w_uploads`。 +- **`internal/model` 与 `internal/repository` 分层(硬规则)**: + - `internal/model/`:仅 GORM 实体、表名、配置 key、查询 DTO、无 IO 领域规则(如密码哈希校验、字段规范化)。**禁止**在 model 中调用 `db.DB` / Redis / ClickHouse,**禁止** `import internal/repository`。 + - 实体上允许仅 mutate 自身字段的 GORM hook(如 `AfterFind(*gorm.DB)`),**禁止**在 hook 内再发起 DB/缓存查询。 + - `internal/repository/`:唯一持久化入口(CRUD、事务、缓存、分析查询)。apps / logics / task 框架通过 repository 访问数据,**禁止**在 Handler 内直接写复杂 SQL。 + - apps 不得为业务 CRUD 直接调用 `db.DB`;必须走 repository(管理端 SQL 控制台、infra 内部实现等例外可保留)。 + - 依赖方向只能是 `apps → repository → model` 与 `repository → infra/persistence`;**禁止** `model → repository`。 + - 新增代码不得再增加 `model.Get/List/Create/Update/Delete*(ctx…)` 类数据访问 API;存量迁移按域收敛至 repository。 - 禁止在 `init()` 中注册跨模块集成(任务 Handler、推送内置事件、域事件监听器、任务完成钩子)。统一通过 `internal/platform/bootstrap` 在 `internal/cmd` 入口显式装配。 - `internal/router/router.go` 的 `Serve()` 仅负责 HTTP 路由与中间件,禁止在其中执行 `SyncEvents`、`InitLogWriter` 等进程级运行时初始化。 - 核心业务模块(如 `oauth`、`user`)禁止直接 `import` `internal/apps/admin/push` 或 `custom_events` 触发通知;应通过 `internal/listener` 发射域事件,由 push 模块在 bootstrap 阶段订阅。 @@ -115,7 +122,8 @@ - `internal/router/`:唯一的 HTTP 路由注册点。 - `internal/apps/`:按功能(Feature-based)组织的 HTTP Handler、中间件、内部服务与模块逻辑。移除全局 service 层,模块内部业务逻辑(如验证码业务逻辑管理器 `internal/apps/cap/manager.go`)均收敛于各自模块中;管理端模块位于 `internal/apps/admin/`。 - `internal/apps/upload/`:上传记录、文件访问控制、本地/S3 文件响应、下载及图片 WebP 压缩。业务应复用 `upload.Ingest` / `upload.Remove` 与 `GET /f/:id` 文件服务,不直接操作底层 storage 或旁路写 `w_uploads`。 -- `internal/model/`:GORM 实体和模型级业务方法。 +- `internal/model/`:GORM 实体、表映射、配置 key、查询 DTO 与无 IO 领域规则;不含数据库访问。 +- `internal/repository/`:数据访问层(平台与业务域 CRUD、缓存、ClickHouse 分析读写);唯一持久化入口。 - `internal/infra/persistence/`:PostgreSQL、Redis、ClickHouse、GORM 日志、ID 生成和 goose SQL 迁移的布线。 - `internal/infra/diskcache/`:平台级磁盘字节缓存,通过 `diskcache.GetGlobalCache()` 提供 TTL、最大空间限制、LRU 淘汰、清空、状态统计和配置热更新。写入时使用 `DefaultExpiration`(全局默认 TTL)、正数 `time.Duration`(业务 TTL)或 `NoExpiration`(无 TTL,仍受空间限制和 LRU 淘汰)。 - `internal/infra/objectstore/`:S3 兼容对象存储适配,提供对象上传、读取、删除、CDN/代理读取及远端对象本地缓存。 @@ -303,9 +311,10 @@ func doSomething(c *gin.Context) { response.AbortBadRequest(c, "...") } 数据库操作: -- 简单查询可以直接从 model 层使用 GORM。 -- 管理员代码应首选 `db.DB(ctx)` 以获得链路追踪感知的 DB 访问。 -- 不要在 Handler 中放置复杂的 SQL;将其移至 `internal/model/` 或模块内的业务服务层(如 `internal/apps//service.go` 或 `logics.go`)。 +- **持久化只通过 `internal/repository`**(或 analytics 子包)。apps / logics 不要直接 `db.DB(ctx).Where...` 拼复杂查询;简单事务编排可在 logics 中调用多个 repository 方法。 +- repository 内管理员/业务查询应使用 `db.DB(ctx)` 以获得链路追踪感知的 DB 访问。 +- 不要在 Handler 中放置 SQL;复杂查询放 `internal/repository/`,业务编排放 `internal/apps//logics.go`(或 `service.go`)。 +- `internal/model` 只定义实体与无 IO 规则,不访问数据库。 - 在 `internal/infra/persistence/migrator/goose/` 下使用 goose SQL 迁移;不要添加基于 GORM AutoMigrate 的 Schema 升级。 - 不要创建物理数据库外键。改为关系字段添加显式索引。 - 数据库默认值必须与 Go 模型零值(`nil`、`0`、`false`、`""`)匹配,以避免意外的插入。 diff --git a/Makefile b/Makefile index f611433f..ea06e7d0 100644 --- a/Makefile +++ b/Makefile @@ -38,6 +38,12 @@ build-embedded: main.go code-check: + @echo "==> Architecture guards..." + @command -v rg >/dev/null 2>&1 || { echo 'error: rg (ripgrep) is required for architecture guards' >&2; exit 1; } + @if rg -n 'db\.DB\(|db\.Redis' internal/model --glob '*.go' -g '!*_test.go' ; then \ + echo 'error: internal/model must not access db.DB or db.Redis (non-test code)' >&2; \ + exit 1; \ + fi golangci-lint run cd frontend && pnpm tsc --noEmit --jsx preserve && npx eslint . --max-warnings 0 diff --git a/docs/changelog/index.md b/docs/changelog/index.md index 21fd0f14..dcc92e5e 100644 --- a/docs/changelog/index.md +++ b/docs/changelog/index.md @@ -22,6 +22,12 @@ sidebar: false ## [unreleased] +### 改进 + +- 统一数据访问分层:业务持久化经 `internal/repository`,`internal/model` 仅保留实体与无 IO 领域规则,避免双轨 CRUD 与职责混淆。 +- 构建检查增加 `internal/model` 禁止直接访问数据库/Redis 的架构守卫,并收敛 model 与 repository 的错误文案定义边界。 + + ## [v3.4.3] - 2026-07-24 ### 新增 diff --git a/docs/design/index.md b/docs/design/index.md index ffb78642..b003caa8 100644 --- a/docs/design/index.md +++ b/docs/design/index.md @@ -94,9 +94,9 @@ OpenFlare 已收敛为**单 monorepo**(Go 模块 `github.com/Rain-kl/Wavelet` | `internal/apps/openflare/` | OpenFlare 控制面业务域(`routers.go` + `logics.go`) | | `internal/apps/{admin,user,oauth,upload,cap,...}/` | Wavelet 平台能力(用户、认证、任务、推送等) | | `internal/apps/openflare/{agent,relay,flared}/` | **Server 侧**边缘协议处理器(鉴权、心跳、WS) | -| `internal/model/` | GORM 实体(`openflare_*.go` + 平台模型) | +| `internal/model/` | GORM 实体 / DTO / 无 IO 领域规则(`openflare_*.go` + 平台模型);**不含** DB 访问 | | `internal/infra/persistence/migrator/goose/` | goose SQL 迁移(PostgreSQL / SQLite / ClickHouse) | -| `internal/repository/` | 平台域数据访问层 | +| `internal/repository/` | 数据访问层(平台 + OpenFlare 业务 CRUD、缓存、ClickHouse 分析读写);**唯一**持久化入口 | | `internal/infra/task/` | Asynq 异步任务(Worker + Scheduler) | | `internal/infra/config/` | Viper 配置加载 | | `internal/shared/` | 统一 API 响应封装(`response/`) | diff --git a/docs/plan/20260724-model-repository-layering.md b/docs/plan/20260724-model-repository-layering.md new file mode 100644 index 00000000..00c25d30 --- /dev/null +++ b/docs/plan/20260724-model-repository-layering.md @@ -0,0 +1,44 @@ +# model / repository 分层治理 + +说明:将 `internal/model` 收敛为无 IO 实体层,`internal/repository` 作为唯一持久化入口。 + +--- + +## 1. 目标与背景 (Goal & Context) + +* **需求背景**:已提交的 AGENTS 曾允许 model 直接使用 GORM,导致 OpenFlare 业务 CRUD 与平台 repository 双轨并存。工作区目标分层与代码不一致,接手成本高。 +* **开发范围 (Scope)**: + * 固化规范:`model` 仅实体 / DTO / 无 IO 规则;`repository` 唯一持久化入口。 + * 将 `internal/model` 中现有 `db.DB` / Redis / ClickHouse store 适配迁入 `internal/repository`。 + * 全量更新 call site:`model.Get/List/Create…` → `repository.…`。 + * 编译通过 + `make format` + 相关单测。 +* **Out of Scope**:不改表结构、不改 API 契约、不做业务行为变更;不强制一次重写所有测试风格。 + +## 2. 设计与决策 + +* **分层**: + * `apps → repository → model` + * `repository → infra/persistence`(及 `repository/analytics`) + * **禁止** `model → repository`、**禁止** model 内 `db.DB` / Redis / ClickHouse +* **迁移策略**:按文件拆分 package-level IO 函数至同名 `repository` 文件;类型与纯函数留在 model;store 适配整文件迁入 repository。 +* **命名**:repository 函数保持原导出名,降低 call site 改动面。 + +## 3. 具体修改文件清单 + +### 规范 + +* #### [MODIFY] `AGENTS.md` +* #### [MODIFY] `docs/design/index.md` +* #### [MODIFY] `docs/plan/index.md`(登记本计划) + +### 后端 + +* #### [MODIFY] `internal/model/*.go`(剥离 IO) +* #### [NEW/MODIFY] `internal/repository/*.go`(承接 CRUD / store) +* #### [MODIFY] `internal/apps/**`、`internal/infra/task/**` 等 call site + +## 4. 验证计划 + +* `go test ./internal/model/... ./internal/repository/...` +* 关键包 `go build ./...` +* `make format` / `make code-check`(在可接受时间内) diff --git a/docs/plan/index.md b/docs/plan/index.md index 22e4a1fb..79788151 100644 --- a/docs/plan/index.md +++ b/docs/plan/index.md @@ -24,6 +24,8 @@ ## 已完成的计划 +* [model / repository 分层治理](./20260724-model-repository-layering.md):model 无 IO;repository 唯一持久化;已完成 OpenFlare/平台 CRUD 迁入 repository。 + * [Pages 项目部署源与 GitHub Releases 自动更新 V2](./20260719-pages-source-sync-v2.md):已完成 Remote URL / GitHub Release 来源、不可变部署、自动检查更新与安全回滚,并预留独立仓库构建 Provider 边界;生产环境验收边界见计划内验证记录。 ## 使用建议 diff --git a/internal/apps/admin/auth_source/routers.go b/internal/apps/admin/auth_source/routers.go index 0366af08..b4bbf430 100644 --- a/internal/apps/admin/auth_source/routers.go +++ b/internal/apps/admin/auth_source/routers.go @@ -49,7 +49,7 @@ type ToggleAuthSourceRequest struct { // @Failure 500 {object} response.Any "内部错误" // @Router /api/v1/admin/auth-sources [get] func ListAuthSources(c *gin.Context) { - sources, err := model.GetAuthSources(c.Request.Context()) + sources, err := repository.GetAuthSources(c.Request.Context()) if err != nil { response.AbortInternal(c, err.Error()) return @@ -88,7 +88,7 @@ func CreateAuthSource(c *gin.Context) { Scopes: req.Scopes, IconURL: req.IconURL, } - if err := model.CreateAuthSource(c.Request.Context(), &source); err != nil { + if err := repository.CreateAuthSource(c.Request.Context(), &source); err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -126,7 +126,7 @@ func UpdateAuthSource(c *gin.Context) { } // 记录更新前的 Discovery URL,以便更新成功后清除旧缓存条目。 - existing, _ := model.GetAuthSourceByID(c.Request.Context(), id) + existing, _ := repository.GetAuthSourceByID(c.Request.Context(), id) source := model.AuthSource{ ID: id, @@ -141,7 +141,7 @@ func UpdateAuthSource(c *gin.Context) { IconURL: req.IconURL, } keepSecret := source.ClientSecret == "" - if err := model.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil { + if err := repository.UpdateAuthSource(c.Request.Context(), &source, keepSecret); err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -154,7 +154,7 @@ func UpdateAuthSource(c *gin.Context) { oauth.InvalidateOIDCProviderCache(normalizeIssuer(req.OpenIDDiscoveryURL)) _ = repository.InvalidateAuthSourceCache(c.Request.Context()) - updated, err := model.GetAuthSourceByID(c.Request.Context(), id) + updated, err := repository.GetAuthSourceByID(c.Request.Context(), id) if err != nil { response.AbortInternal(c, err.Error()) return @@ -190,7 +190,7 @@ func ToggleAuthSource(c *gin.Context) { return } - if err := model.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil { + if err := repository.ToggleAuthSource(c.Request.Context(), id, req.IsActive); err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -216,7 +216,7 @@ func DeleteAuthSource(c *gin.Context) { response.AbortBadRequest(c, err.Error()) return } - if err := model.DeleteAuthSource(c.Request.Context(), id); err != nil { + if err := repository.DeleteAuthSource(c.Request.Context(), id); err != nil { response.AbortBadRequest(c, err.Error()) return } @@ -229,7 +229,7 @@ func parseSourceID(c *gin.Context) (uint64, error) { if raw == "" { return 0, errors.New(admin.InvalidAuthSourceID) } - source, err := model.GetAuthSourceByName(c.Request.Context(), raw) + source, err := repository.GetAuthSourceByName(c.Request.Context(), raw) if err == nil { return source.ID, nil } diff --git a/internal/apps/admin/logs/routers.go b/internal/apps/admin/logs/routers.go index 7f77d520..c1b5454d 100644 --- a/internal/apps/admin/logs/routers.go +++ b/internal/apps/admin/logs/routers.go @@ -16,7 +16,7 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/admin" "github.com/Rain-kl/Wavelet/internal/infra/config" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/gin-gonic/gin" @@ -163,10 +163,7 @@ func buildAccessLogFilter(ctx context.Context, c *gin.Context) (analyticsrepo.Ac username := c.Query("username") if username != "" { - var userIDs []uint64 - err := db.DB(ctx).Model(&model.User{}). - Where("username LIKE ?", "%"+username+"%"). - Pluck("id", &userIDs).Error + userIDs, err := repository.ListUserIDsByUsernameContains(ctx, username) if err != nil { return filter, fmt.Errorf("查询用户信息失败: %w", err) } @@ -215,8 +212,7 @@ func enrichAccessLogsWithUsers(ctx context.Context, list []accessLogItem) { } userMap := make(map[uint64]struct{ Username, Nickname string }) - var users []model.User - if err := db.DB(ctx).Where("id IN ?", userIDs).Find(&users).Error; err == nil { + if users, err := repository.ListUsersByIDs(ctx, userIDs); err == nil { for _, u := range users { userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname} } @@ -402,8 +398,7 @@ func GetLogsAnalytics(c *gin.Context) { Username string Nickname string }) - var users []model.User - if errProfile := db.DB(ctx).Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil { + if users, errProfile := repository.ListUsersByIDs(ctx, userIDs); errProfile == nil { for _, u := range users { userProfileMap[u.ID] = struct { Username string diff --git a/internal/apps/admin/system_config/logics.go b/internal/apps/admin/system_config/logics.go index cabba62f..db1f09f3 100644 --- a/internal/apps/admin/system_config/logics.go +++ b/internal/apps/admin/system_config/logics.go @@ -10,7 +10,6 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" @@ -72,7 +71,7 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR req.Value = validatedVal } - if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := repository.RunInTransaction(ctx, func(tx *gorm.DB) error { updates := map[string]any{ "description": req.Description, } @@ -84,7 +83,7 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR updates["value"] = req.Value config.Value = req.Value } - if err := tx.Model(&config).Updates(updates).Error; err != nil { + if err := repository.UpdateSystemConfigFieldsTx(tx, &config, updates); err != nil { return err } resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value) @@ -116,13 +115,12 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate( return } - if err := tx.Model(&model.TaskExecution{}). - Where("task_type = ? AND status = ?", "storage:migrate", model.TaskExecutionStatusFailed). - Updates(map[string]any{ - "status": model.TaskExecutionStatusSucceeded, - "result": "存储配置直接更新,故障迁移任务自动标记为已解决", - "finished_at": time.Now(), - }).Error; err != nil { + if err := repository.MarkFailedTaskExecutionsSucceededTx( + tx, + "storage:migrate", + "存储配置直接更新,故障迁移任务自动标记为已解决", + time.Now(), + ); err != nil { logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err) } } diff --git a/internal/apps/admin/system_config/routers.go b/internal/apps/admin/system_config/routers.go index 4a5ca3b2..05392ce0 100644 --- a/internal/apps/admin/system_config/routers.go +++ b/internal/apps/admin/system_config/routers.go @@ -18,7 +18,6 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/cap" "github.com/Rain-kl/Wavelet/internal/apps/upload" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/shared/response" @@ -325,10 +324,8 @@ func validateAndMergeStorageConfig(ctx context.Context, value string, currentCon func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error { if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver { - var uploadCount int64 - if err := db.DB(ctx).Model(&model.Upload{}). - Where("status != ?", model.UploadStatusDeleted). - Count(&uploadCount).Error; err != nil { + uploadCount, err := repository.CountActiveUploads(ctx) + if err != nil { return fmt.Errorf("检查存量文件失败: %w", err) } if uploadCount > 0 { diff --git a/internal/apps/admin/task/routers.go b/internal/apps/admin/task/routers.go index f631c469..5459f7e1 100644 --- a/internal/apps/admin/task/routers.go +++ b/internal/apps/admin/task/routers.go @@ -11,6 +11,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/admin" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/infra/task/scheduler" @@ -119,7 +121,7 @@ func ListTaskExecutions(c *gin.Context) { } } - executions, total, err := model.ListTaskExecutions(c.Request.Context(), req) + executions, total, err := repository.ListTaskExecutions(c.Request.Context(), req) if err != nil { response.AbortInternal(c, err.Error()) return @@ -153,7 +155,7 @@ func GetTaskExecution(c *gin.Context) { return } - execution, err := model.GetTaskExecutionByID(c.Request.Context(), id) + execution, err := repository.GetTaskExecutionByID(c.Request.Context(), id) if err != nil { response.AbortNotFound(c, TaskNotFound) return @@ -211,7 +213,7 @@ func RetryTask(c *gin.Context) { // @Failure 403 {object} response.Any "无管理员权限" // @Router /api/v1/admin/tasks/schedules [get] func ListSchedules(c *gin.Context) { - schedules, err := model.ListSchedules(c.Request.Context()) + schedules, err := repository.ListSchedules(c.Request.Context()) if err != nil { response.AbortInternal(c, err.Error()) return @@ -289,7 +291,7 @@ func CreateSchedule(c *gin.Context) { IsActive: *req.IsActive, } - if err := model.CreateSchedule(c.Request.Context(), schedule); err != nil { + if err := repository.CreateSchedule(c.Request.Context(), schedule); err != nil { response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err)) return } @@ -341,7 +343,7 @@ func UpdateSchedule(c *gin.Context) { } // 检查定时任务是否存在 - schedule, err := model.GetScheduleByID(c.Request.Context(), id) + schedule, err := repository.GetScheduleByID(c.Request.Context(), id) if err != nil { response.AbortNotFound(c, ScheduleNotFound) return @@ -381,7 +383,7 @@ func UpdateSchedule(c *gin.Context) { schedule.Payload = string(validated) schedule.IsActive = *req.IsActive - if err := model.UpdateSchedule(c.Request.Context(), schedule); err != nil { + if err := repository.UpdateSchedule(c.Request.Context(), schedule); err != nil { response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err)) return } @@ -422,7 +424,7 @@ func DeleteSchedule(c *gin.Context) { response.AbortBadRequest(c, "无效的定时任务ID") return } - schedule, err := model.GetScheduleByID(c.Request.Context(), id) + schedule, err := repository.GetScheduleByID(c.Request.Context(), id) if err != nil { response.AbortNotFound(c, ScheduleNotFound) return @@ -432,7 +434,7 @@ func DeleteSchedule(c *gin.Context) { return } - if err := model.DeleteSchedule(c.Request.Context(), id); err != nil { + if err := repository.DeleteSchedule(c.Request.Context(), id); err != nil { response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err)) return } diff --git a/internal/apps/admin/task/routers_test.go b/internal/apps/admin/task/routers_test.go index ee7b7283..4eb14854 100644 --- a/internal/apps/admin/task/routers_test.go +++ b/internal/apps/admin/task/routers_test.go @@ -14,6 +14,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/oauth" uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task" "github.com/Rain-kl/Wavelet/internal/apps/user" @@ -154,8 +156,8 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) { Payload: "{}", IsActive: true, } - require.NoError(t, model.CreateSchedule(ctx, internalSchedule)) - require.NoError(t, model.CreateSchedule(ctx, publicSchedule)) + require.NoError(t, repository.CreateSchedule(ctx, internalSchedule)) + require.NoError(t, repository.CreateSchedule(ctx, publicSchedule)) req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/tasks/schedules", nil) w := httptest.NewRecorder() @@ -215,7 +217,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) { Cron: "0 * * * *", IsActive: true, } - require.NoError(t, model.CreateSchedule(ctx, schedule)) + require.NoError(t, repository.CreateSchedule(ctx, schedule)) isActive := false body, err := json.Marshal(UpdateScheduleRequest{ Name: "尝试修改内部排程", @@ -235,7 +237,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) { router.ServeHTTP(w, req) assert.Equal(t, http.StatusBadRequest, w.Code) - unchanged, err := model.GetScheduleByID(ctx, schedule.ID) + unchanged, err := repository.GetScheduleByID(ctx, schedule.ID) require.NoError(t, err) assert.Equal(t, "系统内部排程", unchanged.Name) assert.Equal(t, testInternalOnlyTaskType, unchanged.TaskType) @@ -249,7 +251,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) { Cron: "0 * * * *", IsActive: true, } - require.NoError(t, model.CreateSchedule(ctx, schedule)) + require.NoError(t, repository.CreateSchedule(ctx, schedule)) isActive := true body, err := json.Marshal(UpdateScheduleRequest{ Name: "尝试切入内部任务", @@ -269,7 +271,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) { router.ServeHTTP(w, req) assert.Equal(t, http.StatusBadRequest, w.Code) - unchanged, err := model.GetScheduleByID(ctx, schedule.ID) + unchanged, err := repository.GetScheduleByID(ctx, schedule.ID) require.NoError(t, err) assert.Equal(t, "公开排程", unchanged.Name) assert.Equal(t, uploadtask.TaskTypeSystemCleanup, unchanged.TaskType) @@ -282,7 +284,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) { Cron: "*/5 * * * *", IsActive: true, } - require.NoError(t, model.CreateSchedule(ctx, schedule)) + require.NoError(t, repository.CreateSchedule(ctx, schedule)) req := httptest.NewRequest( http.MethodDelete, fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID), @@ -296,7 +298,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) { var resp response.Any require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) assert.Equal(t, InvalidTaskType, resp.ErrorMsg) - preserved, err := model.GetScheduleByID(ctx, schedule.ID) + preserved, err := repository.GetScheduleByID(ctx, schedule.ID) require.NoError(t, err) assert.Equal(t, testInternalOnlyTaskType, preserved.TaskType) }) @@ -320,7 +322,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) { Cron: "0 * * * *", IsActive: false, } - require.NoError(t, model.CreateSchedule(ctx, schedule)) + require.NoError(t, repository.CreateSchedule(ctx, schedule)) req := httptest.NewRequest( http.MethodDelete, fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID), @@ -331,7 +333,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) { router.ServeHTTP(w, req) assert.Equal(t, http.StatusOK, w.Code) - _, err := model.GetScheduleByID(ctx, schedule.ID) + _, err := repository.GetScheduleByID(ctx, schedule.ID) assert.Error(t, err) }) } @@ -470,7 +472,7 @@ func TestListTaskExecutions(t *testing.T) { {TaskID: "exec_003", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", Retryable: true, MaxRetry: 3}, } for _, r := range records { - err := model.CreateTaskExecution(ctx, r) + err := repository.CreateTaskExecution(ctx, r) require.NoError(t, err) } @@ -581,7 +583,7 @@ func TestGetTaskExecution(t *testing.T) { MaxRetry: 3, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) t.Run("get existing execution", func(t *testing.T) { @@ -646,7 +648,7 @@ func TestRetryTask(t *testing.T) { StartedAt: &now, FinishedAt: &now, } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID) @@ -666,7 +668,7 @@ func TestRetryTask(t *testing.T) { assert.True(t, ok) assert.NotEmpty(t, newTaskID) - newExecution, err := model.GetTaskExecutionByTaskID(ctx, newTaskID) + newExecution, err := repository.GetTaskExecutionByTaskID(ctx, newTaskID) require.NoError(t, err) assert.Equal(t, 1, newExecution.RetryCount) assert.Equal(t, "retry", newExecution.TriggeredBy) @@ -682,7 +684,7 @@ func TestRetryTask(t *testing.T) { MaxRetry: 3, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID) @@ -702,7 +704,7 @@ func TestRetryTask(t *testing.T) { Retryable: false, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID) diff --git a/internal/apps/admin/user/logics.go b/internal/apps/admin/user/logics.go index fe278b1a..20bd6904 100644 --- a/internal/apps/admin/user/logics.go +++ b/internal/apps/admin/user/logics.go @@ -10,7 +10,6 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/apps/oauth" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" @@ -41,7 +40,7 @@ func updateUserStatus(ctx context.Context, id uint64, active bool) error { var tokens []model.AccessToken if !active { - _ = db.DB(ctx).Where("user_id = ?", id).Find(&tokens).Error + tokens, _ = repository.ListAccessTokensByUserID(ctx, id) } err = repository.UpdateUserActive(ctx, id, active) @@ -68,8 +67,7 @@ func deleteUser(ctx context.Context, currentUserID, targetID uint64) error { return errors.New(cannotDelete) } - var tokens []model.AccessToken - _ = db.DB(ctx).Where("user_id = ?", targetID).Find(&tokens).Error + tokens, _ := repository.ListAccessTokensByUserID(ctx, targetID) err = repository.DeleteUserWithRelations(ctx, targetID) if err == nil { @@ -181,7 +179,7 @@ func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam needRevokeTokens := (param.Password != "") || (targetUser.IsAdmin && !param.IsAdmin) var tokens []model.AccessToken if needRevokeTokens { - _ = db.DB(ctx).Where("user_id = ?", param.ID).Find(&tokens).Error + tokens, _ = repository.ListAccessTokensByUserID(ctx, param.ID) } // 更新字段 diff --git a/internal/apps/oauth/handler_callback.go b/internal/apps/oauth/handler_callback.go index f7ee72da..ed3dda69 100644 --- a/internal/apps/oauth/handler_callback.go +++ b/internal/apps/oauth/handler_callback.go @@ -130,12 +130,12 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS response.AbortUnauthorized(c, shared.UnAuthorized) return } - var user model.User - if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil { + user, err := repository.GetUserByID(ctx, userID) + if err != nil { response.AbortInternal(c, err.Error()) return } - if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ + if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{ AuthSourceID: source.ID, UserID: user.ID, ExternalID: userInfo.Sub, @@ -146,7 +146,7 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS return } user.LastLoginAt = time.Now() - _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error + _ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt) c.JSON(http.StatusOK, response.OK(buildCallbackResult(&user, "bound"))) } @@ -154,13 +154,15 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) { var user model.User - account, err := model.FindExternalAccount(ctx, source.ID, userInfo.Sub) + account, err := repository.FindExternalAccount(ctx, source.ID, userInfo.Sub) switch { case err == nil: - if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil { - response.AbortInternal(c, err.Error()) + loaded, loadErr := repository.GetUserByID(ctx, account.UserID) + if loadErr != nil { + response.AbortInternal(c, loadErr.Error()) return } + user = loaded case errors.Is(err, gorm.ErrRecordNotFound): newUser, ok := handleCallbackRegister(ctx, c, source, userInfo) if !ok { @@ -173,7 +175,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth } user.LastLoginAt = time.Now() - _ = db.DB(ctx).Model(&user).Update("last_login_at", user.LastLoginAt).Error + _ = repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt) if err := setLoginSession(ctx, c, &user); err != nil { response.AbortInternal(c, err.Error()) return @@ -209,11 +211,11 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A userInfo.Username = username var user model.User - if err := user.CreateUser(ctx, db.DB(ctx), userInfo); err != nil { + if err := repository.CreateUserFromOAuth(ctx, &user, userInfo); err != nil { response.AbortInternal(c, err.Error()) return model.User{}, false } - if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ + if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{ AuthSourceID: source.ID, UserID: user.ID, ExternalID: userInfo.Sub, diff --git a/internal/apps/oauth/handler_external_accounts.go b/internal/apps/oauth/handler_external_accounts.go index bd7f9b02..aa922c1d 100644 --- a/internal/apps/oauth/handler_external_accounts.go +++ b/internal/apps/oauth/handler_external_accounts.go @@ -8,7 +8,8 @@ import ( "strconv" "strings" - "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/shared" "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/gin-gonic/gin" @@ -26,7 +27,7 @@ import ( // @Router /api/v1/oauth/external-accounts [get] func ListExternalAccounts(c *gin.Context) { userID := GetUserIDFromContext(c) - accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID) + accounts, err := repository.ListExternalAccountsByUserID(c.Request.Context(), userID) if err != nil { response.AbortInternal(c, err.Error()) return @@ -57,7 +58,7 @@ func DeleteExternalAccount(c *gin.Context) { response.AbortBadRequest(c, errInvalidExternalAccountBindingID) return } - if err := model.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil { + if err := repository.DeleteExternalAccountForUser(c.Request.Context(), id, userID); err != nil { response.AbortBadRequest(c, err.Error()) return } diff --git a/internal/apps/oauth/middlewares.go b/internal/apps/oauth/middlewares.go index d15d5e8b..6a20fa8c 100644 --- a/internal/apps/oauth/middlewares.go +++ b/internal/apps/oauth/middlewares.go @@ -8,8 +8,8 @@ import ( "context" "errors" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/shared" "github.com/Rain-kl/Wavelet/internal/shared/response" @@ -33,8 +33,8 @@ func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.A tokenHash := model.HashToken(tokenStr) tokenRecord, err := GetCachedToken(ctx, tokenHash) if err != nil { - var dbToken model.AccessToken - if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&dbToken).Error; err != nil { + dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash) + if err != nil { return nil, nil, err } tokenRecord = &dbToken @@ -43,8 +43,8 @@ func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.A user, err := GetCachedUser(ctx, tokenRecord.UserID) if err != nil || !user.IsActive { - var dbUser model.User - if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil { + dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID) + if err != nil { return nil, nil, err } user = &dbUser @@ -87,11 +87,10 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) { user, err := GetCachedUser(ctx, userID) if err != nil || !user.IsActive { - var dbUser model.User // load user from db to make sure is active - tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&dbUser) - if tx.Error != nil { - return nil, tx.Error + dbUser, loadErr := repository.GetActiveUserByID(ctx, userID) + if loadErr != nil { + return nil, loadErr } user = &dbUser SetCachedUser(ctx, userID, user) diff --git a/internal/apps/oauth/oauth_userinfo.go b/internal/apps/oauth/oauth_userinfo.go index 3eb45ea5..c8b9c657 100644 --- a/internal/apps/oauth/oauth_userinfo.go +++ b/internal/apps/oauth/oauth_userinfo.go @@ -9,8 +9,8 @@ import ( "fmt" "strings" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/coreos/go-oidc/v3/oidc" "golang.org/x/oauth2" ) @@ -21,10 +21,8 @@ func uniqueUsername(ctx context.Context, base string) (string, error) { base = "user" } - var existingUsernames []string - if err := db.DB(ctx).Model(&model.User{}). - Where("username = ? OR username LIKE ?", base, base+"-%"). - Pluck("username", &existingUsernames).Error; err != nil { + existingUsernames, err := repository.ListUsernamesMatchingBase(ctx, base) + if err != nil { return "", err } diff --git a/internal/apps/openflare/agent/auth_cache.go b/internal/apps/openflare/agent/auth_cache.go index 15f6b681..42b8d487 100644 --- a/internal/apps/openflare/agent/auth_cache.go +++ b/internal/apps/openflare/agent/auth_cache.go @@ -10,6 +10,8 @@ import ( "sync" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "gorm.io/gorm" ) @@ -39,7 +41,7 @@ func newAccessTokenAuthCache() *accessTokenAuthCache { positive: make(map[string]cachedAgentNode), negative: make(map[string]time.Time), now: time.Now, - loadNodeByToken: model.GetOpenFlareNodeByAccessToken, + loadNodeByToken: repository.GetOpenFlareNodeByAccessToken, } } diff --git a/internal/apps/openflare/agent/config.go b/internal/apps/openflare/agent/config.go index a9656ffd..1955fb6c 100644 --- a/internal/apps/openflare/agent/config.go +++ b/internal/apps/openflare/agent/config.go @@ -9,14 +9,15 @@ import ( "errors" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/pages" - "github.com/Rain-kl/Wavelet/internal/model" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" "gorm.io/gorm" ) func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) { - version, err := model.GetActiveConfigVersion(ctx) + version, err := repository.GetActiveConfigVersion(ctx) if err != nil { return nil, err } @@ -27,7 +28,7 @@ func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) { } func getActiveConfigForAgent(ctx context.Context) (*ConfigResponse, error) { - version, err := model.GetActiveConfigVersion(ctx) + version, err := repository.GetActiveConfigVersion(ctx) if err != nil { return nil, err } diff --git a/internal/apps/openflare/agent/logics.go b/internal/apps/openflare/agent/logics.go index b982f61d..9b7d9474 100644 --- a/internal/apps/openflare/agent/logics.go +++ b/internal/apps/openflare/agent/logics.go @@ -9,11 +9,11 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip" "github.com/Rain-kl/Wavelet/internal/apps/openflare/node" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" - "gorm.io/gorm" ) // RegisterWithAccessToken registers an agent on a reserved node token. @@ -27,7 +27,7 @@ func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode, return nil, err } applyNodeRuntime(ctx, authNode, payload, true) - if err := model.SaveOpenFlareNode(ctx, authNode); err != nil { + if err := repository.SaveOpenFlareNode(ctx, authNode); err != nil { return nil, err } RefreshAccessTokenCache(ctx, authNode) @@ -71,7 +71,7 @@ func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*Registrat } applyNodeRuntime(ctx, record, payload, false) - if err = model.CreateOpenFlareNode(ctx, record); err != nil { + if err = repository.CreateOpenFlareNode(ctx, record); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errNodeIDConflict) } @@ -115,7 +115,7 @@ func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload N for field := range changes { fields = append(fields, field) } - if err := model.UpdateOpenFlareNodeFields(ctx, authNode, fields...); err != nil { + if err := repository.UpdateOpenFlareNodeFields(ctx, authNode, fields...); err != nil { return nil, err } } @@ -181,12 +181,12 @@ func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFl return nil, errors.New(errInvalidApplyResult) } - latest, err := model.GetLatestOpenFlareApplyLogByNodeID(ctx, payload.NodeID) + latest, err := repository.GetLatestOpenFlareApplyLogByNodeID(ctx, payload.NodeID) if err != nil { return nil, err } if model.IsRepeatSuccessApplyLog(latest, payload.Version, payload.Checksum, payload.Result) { - if err := updateNodeFromApplyLog(ctx, payload, now); err != nil { + if err := repository.UpdateOpenFlareNodeFromApplyResult(ctx, payload.NodeID, payload.Result, payload.Version, payload.Message, now); err != nil { return nil, err } return latest, nil @@ -204,49 +204,12 @@ func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFl CreatedAt: now, } - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New("database not initialized") - } - - err = conn.Transaction(func(tx *gorm.DB) error { - if err := tx.Create(log).Error; err != nil { - return err - } - return updateNodeFromApplyLogTx(tx, payload, now) - }) - if err != nil { + if err := repository.CreateOpenFlareApplyLogAndUpdateNode(ctx, log, payload.Result, payload.Version, payload.Message); err != nil { return nil, err } return log, nil } -func updateNodeFromApplyLog(ctx context.Context, payload ApplyLogPayload, now time.Time) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New("database not initialized") - } - return conn.Transaction(func(tx *gorm.DB) error { - return updateNodeFromApplyLogTx(tx, payload, now) - }) -} - -func updateNodeFromApplyLogTx(tx *gorm.DB, payload ApplyLogPayload, now time.Time) error { - record := &model.OpenFlareNode{} - if err := tx.Where("node_id = ?", payload.NodeID).First(record).Error; err != nil { - return err - } - record.Status = nodeStatusOnline - record.LastSeenAt = &now - if payload.Result == applyResultOK { - record.CurrentVersion = payload.Version - record.LastError = "" - } else { - record.LastError = payload.Message - } - return tx.Model(record).Select("status", "last_seen_at", "current_version", "last_error").Updates(record).Error -} - // ValidateDiscoveryToken delegates to the node package discovery token helper. func ValidateDiscoveryToken(ctx context.Context, token string) error { return node.ValidateDiscoveryToken(ctx, token) diff --git a/internal/apps/openflare/agent/observability.go b/internal/apps/openflare/agent/observability.go index 0fbcdf0e..0b81b633 100644 --- a/internal/apps/openflare/agent/observability.go +++ b/internal/apps/openflare/agent/observability.go @@ -10,23 +10,16 @@ import ( "strings" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "go.uber.org/zap" - "gorm.io/gorm" - "gorm.io/gorm/clause" ) const ( - healthEventStatusActive = "active" - healthEventStatusResolved = "resolved" - healthSeverityInfo = "info" - healthSeverityWarning = "warning" - healthSeverityCritical = "critical" accessLogPathMaxLength = 100 accessLogUserAgentMaxLength = 512 accessLogCacheStatusMaxLength = 32 - healthEventMessageMaxLength = 4096 ) // PersistHeartbeatObservability stores profile, host metrics, edge health, and access logs. @@ -43,28 +36,23 @@ func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload N return } - conn := db.DB(ctx) - if conn == nil { - return - } - accessLogRecords, err := buildNodeAccessLogRecords(nodeID, payload.AccessLogs, payload.Buffered, reportedAt) if err != nil { zap.L().Error("build heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err)) return } - if err := conn.Transaction(func(tx *gorm.DB) error { - if err := persistNodeSystemProfile(tx, nodeID, payload.Profile, reportedAt); err != nil { - return err - } - if payload.HealthEvents != nil { - if err := reconcileNodeHealthEvents(tx, nodeID, payload.HealthEvents, reportedAt); err != nil { - return err - } - } - return nil - }); err != nil { + profile := buildNodeSystemProfileModel(nodeID, payload.Profile, reportedAt) + healthEvents := healthEventInputs(payload.HealthEvents) + if err := repository.PersistOpenFlareNodePGObservability( + ctx, + profile, + nodeID, + healthEvents, + payload.HealthEvents != nil, + reportedAt, + nil, + ); err != nil { zap.L().Error("persist heartbeat observability failed", zap.String("node_id", nodeID), zap.Error(err)) return } @@ -107,7 +95,7 @@ func persistNodeEdgeHealth(ctx context.Context, nodeID string, health *NodeEdgeH if status == "" { status = openrestyStatusUnknown } - return model.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{ + return repository.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{ NodeID: nodeID, CapturedAt: timeFromUnix(health.CapturedAtUnix, reportedAt), Status: status, @@ -115,11 +103,11 @@ func persistNodeEdgeHealth(ctx context.Context, nodeID string, health *NodeEdgeH }) } -func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *NodeSystemProfile, reportedAt time.Time) error { +func buildNodeSystemProfileModel(nodeID string, profile *NodeSystemProfile, reportedAt time.Time) *model.OpenFlareNodeSystemProfile { if profile == nil { return nil } - record := &model.OpenFlareNodeSystemProfile{ + return &model.OpenFlareNodeSystemProfile{ NodeID: nodeID, Hostname: strings.TrimSpace(profile.Hostname), OSName: strings.TrimSpace(profile.OSName), @@ -133,23 +121,23 @@ func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *NodeSystemPro UptimeSeconds: profile.UptimeSeconds, ReportedAt: timeFromUnix(profile.ReportedAtUnix, reportedAt), } - return tx.Clauses(clause.OnConflict{ - Columns: []clause.Column{{Name: "node_id"}}, - DoUpdates: clause.AssignmentColumns([]string{ - "hostname", - "os_name", - "os_version", - "kernel_version", - "architecture", - "cpu_model", - "cpu_cores", - "total_memory_bytes", - "total_disk_bytes", - "uptime_seconds", - "reported_at", - "updated_at", - }), - }).Create(record).Error +} + +func healthEventInputs(events []NodeHealthEvent) []repository.OpenFlareHealthEventInput { + if events == nil { + return nil + } + out := make([]repository.OpenFlareHealthEventInput, 0, len(events)) + for _, event := range events { + out = append(out, repository.OpenFlareHealthEventInput{ + EventType: event.EventType, + Severity: event.Severity, + Message: event.Message, + TriggeredAtUnix: event.TriggeredAtUnix, + Metadata: event.Metadata, + }) + } + return out } func persistNodeMetricSnapshot(ctx context.Context, nodeID string, snapshot *NodeMetricSnapshot, reportedAt time.Time) error { @@ -168,7 +156,7 @@ func persistNodeMetricSnapshot(ctx context.Context, nodeID string, snapshot *Nod DiskWriteBytes: snapshot.DiskWriteBytes, // NetworkRx/Tx no longer collected from agents; CH columns remain 0. } - return model.InsertOpenFlareMetricSnapshot(ctx, record) + return repository.InsertOpenFlareMetricSnapshot(ctx, record) } func buildNodeAccessLogRecords(nodeID string, direct []NodeAccessLog, buffered []BufferedObservabilityRecord, reportedAt time.Time) ([]*model.OpenFlareAccessLog, error) { @@ -234,123 +222,12 @@ func persistNodeAccessLogs(ctx context.Context, _ string, records []*model.OpenF if len(records) == 0 { return nil } - return model.InsertOpenFlareAccessLogsBatch(ctx, records) -} - -func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time) error { - return ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil) + return repository.InsertOpenFlareAccessLogsBatch(ctx, records) } // ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types. -func ReconcileScopedNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error { - activeTypes := make(map[string]NodeHealthEvent, len(events)) - for _, event := range events { - eventType := normalizeHealthEventType(event.EventType) - if eventType == "" { - continue - } - if len(managedEventTypes) > 0 { - if _, ok := managedEventTypes[eventType]; !ok { - continue - } - } - event.EventType = eventType - event.Severity = normalizeHealthSeverity(event.Severity) - if event.TriggeredAtUnix <= 0 { - event.TriggeredAtUnix = reportedAt.Unix() - } - activeTypes[eventType] = event - } - - var activeEvents []*model.OpenFlareHealthEvent - query := tx.Where("node_id = ? AND status = ?", nodeID, healthEventStatusActive) - if len(managedEventTypes) > 0 { - scopedTypes := make([]string, 0, len(managedEventTypes)) - for eventType := range managedEventTypes { - eventType = normalizeHealthEventType(eventType) - if eventType != "" { - scopedTypes = append(scopedTypes, eventType) - } - } - if len(scopedTypes) == 0 { - return nil - } - query = query.Where("event_type IN ?", scopedTypes) - } - if err := query.Find(&activeEvents).Error; err != nil { - return err - } - - activeByType := make(map[string]*model.OpenFlareHealthEvent, len(activeEvents)) - for _, event := range activeEvents { - activeByType[event.EventType] = event - } - - for eventType, event := range activeTypes { - triggeredAt := timeFromUnix(event.TriggeredAtUnix, reportedAt) - if existing, ok := activeByType[eventType]; ok { - existing.Severity = event.Severity - existing.Message = normalizeHealthEventMessage(event.Message) - existing.LastTriggeredAt = triggeredAt - existing.ReportedAt = reportedAt - existing.MetadataJSON = marshalJSON(event.Metadata) - existing.ResolvedAt = nil - if err := tx.Save(existing).Error; err != nil { - return err - } - continue - } - record := &model.OpenFlareHealthEvent{ - NodeID: nodeID, - EventType: eventType, - Severity: event.Severity, - Status: healthEventStatusActive, - Message: normalizeHealthEventMessage(event.Message), - FirstTriggeredAt: triggeredAt, - LastTriggeredAt: triggeredAt, - ReportedAt: reportedAt, - MetadataJSON: marshalJSON(event.Metadata), - } - if err := tx.Create(record).Error; err != nil { - return err - } - } - - for _, existing := range activeEvents { - if _, ok := activeTypes[existing.EventType]; ok { - continue - } - resolvedAt := reportedAt - existing.Status = healthEventStatusResolved - existing.ReportedAt = reportedAt - existing.ResolvedAt = &resolvedAt - if err := tx.Save(existing).Error; err != nil { - return err - } - } - - return nil -} - -func normalizeHealthEventType(eventType string) string { - eventType = strings.TrimSpace(strings.ToLower(eventType)) - eventType = strings.ReplaceAll(eventType, " ", "_") - return eventType -} - -func normalizeHealthSeverity(severity string) string { - switch strings.ToLower(strings.TrimSpace(severity)) { - case healthSeverityCritical: - return healthSeverityCritical - case healthSeverityInfo: - return healthSeverityInfo - default: - return healthSeverityWarning - } -} - -func normalizeHealthEventMessage(message string) string { - return truncateForDatabase(message, healthEventMessageMaxLength) +func ReconcileScopedNodeHealthEvents(ctx context.Context, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error { + return repository.ReconcileOpenFlareHealthEvents(ctx, nodeID, healthEventInputs(events), reportedAt, managedEventTypes) } func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time { diff --git a/internal/apps/openflare/agent/waf_ip_group.go b/internal/apps/openflare/agent/waf_ip_group.go index 4b5cbe21..af0de03b 100644 --- a/internal/apps/openflare/agent/waf_ip_group.go +++ b/internal/apps/openflare/agent/waf_ip_group.go @@ -15,6 +15,8 @@ import ( "strconv" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/pkg/protocol" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" @@ -96,7 +98,7 @@ func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error return []WAFIPGroup{}, nil } sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) - groups, err := model.ListOpenFlareWAFIPGroupsByIDs(ctx, ids) + groups, err := repository.ListOpenFlareWAFIPGroupsByIDs(ctx, ids) if err != nil { return nil, err } @@ -158,7 +160,7 @@ func checksumAgentWAFIPGroup(group WAFIPGroup) string { } func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) { - version, err := model.GetActiveConfigVersion(ctx) + version, err := repository.GetActiveConfigVersion(ctx) if err != nil { if isActiveConfigNotFound(err) { return []uint{}, nil diff --git a/internal/apps/openflare/agent/waf_ip_group_test.go b/internal/apps/openflare/agent/waf_ip_group_test.go index 661055b4..327cabe3 100644 --- a/internal/apps/openflare/agent/waf_ip_group_test.go +++ b/internal/apps/openflare/agent/waf_ip_group_test.go @@ -10,6 +10,8 @@ import ( "strings" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/pkg/protocol" @@ -124,7 +126,7 @@ func TestChangedWAFIPGroupsForAgentDiscoversGraphReferences(t *testing.T) { Enabled: true, IPList: `["192.0.2.88"]`, } - require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) + require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) seedActiveConfigWithWAFGraphIPGroup(t, ctx, ipGroup.ID) groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil) @@ -163,7 +165,7 @@ func TestChangedWAFIPGroupsForAgentRejectsOversizedSnapshotBeforeChecksumDelta(t Enabled: true, IPList: `[]`, } - require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) + require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) agentGroup, err := buildAgentWAFIPGroup(ipGroup) require.NoError(t, err) @@ -185,7 +187,7 @@ func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) { Enabled: true, IPList: `["203.0.113.44"]`, } - require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) + require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID) groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil) @@ -201,7 +203,7 @@ func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) { assert.Empty(t, same) ipGroup.IPList = `["203.0.113.45"]` - require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup)) + require.NoError(t, repository.UpdateOpenFlareWAFIPGroup(ctx, ipGroup)) delta, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum}) require.NoError(t, err) @@ -222,7 +224,7 @@ func TestSyncWAFIPGroupsReturnsChangedGroups(t *testing.T) { Enabled: true, IPList: `["198.51.100.10"]`, } - require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) + require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID) result, err := SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{ @@ -255,9 +257,9 @@ func TestChangedWAFIPGroupsForAgentDisabledGroupClearsIPList(t *testing.T) { Enabled: true, IPList: `["203.0.113.10"]`, } - require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) + require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) ipGroup.Enabled = false - require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup)) + require.NoError(t, repository.UpdateOpenFlareWAFIPGroup(ctx, ipGroup)) seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID) groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil) diff --git a/internal/apps/openflare/agent/ws_status.go b/internal/apps/openflare/agent/ws_status.go index 8d299f50..81478fc5 100644 --- a/internal/apps/openflare/agent/ws_status.go +++ b/internal/apps/openflare/agent/ws_status.go @@ -8,8 +8,9 @@ import ( "encoding/json" "log/slog" + "github.com/Rain-kl/Wavelet/internal/repository" + ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" - "github.com/Rain-kl/Wavelet/internal/model" ) // HandleWSStatus processes an agent websocket status payload (replaces HTTP heartbeat in WS mode). @@ -20,7 +21,7 @@ func HandleWSStatus(ctx context.Context, nodeID, remoteAddr string, rawPayload j return } - authNode, err := model.GetOpenFlareNodeByNodeID(ctx, nodeID) + authNode, err := repository.GetOpenFlareNodeByNodeID(ctx, nodeID) if err != nil { slog.Debug("agent ws status reload node failed", "node_id", nodeID, "error", err) return diff --git a/internal/apps/openflare/apply_log/logics.go b/internal/apps/openflare/apply_log/logics.go index f61dbce7..a97507e0 100644 --- a/internal/apps/openflare/apply_log/logics.go +++ b/internal/apps/openflare/apply_log/logics.go @@ -9,6 +9,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" ) @@ -53,7 +55,7 @@ func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) { pageSize := normalizePageSize(input.PageSize) nodeID := strings.TrimSpace(input.NodeID) - rows, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ + rows, err := repository.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ NodeID: nodeID, PageNo: pageNo, PageSize: pageSize, @@ -62,7 +64,7 @@ func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) { return nil, err } - total, err := model.CountOpenFlareApplyLogs(ctx, nodeID) + total, err := repository.CountOpenFlareApplyLogs(ctx, nodeID) if err != nil { return nil, err } @@ -83,7 +85,7 @@ func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) { // Cleanup removes old apply logs or deletes all records. func Cleanup(ctx context.Context, input CleanupInput) (*CleanupResult, error) { if input.DeleteAll { - deleted, err := model.DeleteAllOpenFlareApplyLogs(ctx) + deleted, err := repository.DeleteAllOpenFlareApplyLogs(ctx) if err != nil { return nil, err } @@ -98,7 +100,7 @@ func Cleanup(ctx context.Context, input CleanupInput) (*CleanupResult, error) { } cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour) - deleted, err := model.DeleteOpenFlareApplyLogsBefore(ctx, cutoff) + deleted, err := repository.DeleteOpenFlareApplyLogsBefore(ctx, cutoff) if err != nil { return nil, err } diff --git a/internal/apps/openflare/apply_log/logics_test.go b/internal/apps/openflare/apply_log/logics_test.go index c8462bf8..eec5dc98 100644 --- a/internal/apps/openflare/apply_log/logics_test.go +++ b/internal/apps/openflare/apply_log/logics_test.go @@ -8,6 +8,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/glebarez/sqlite" @@ -67,7 +69,7 @@ func TestListPageAndCleanup(t *testing.T) { assert.Equal(t, int64(1), cleanupResult.DeletedCount) assert.NotNil(t, cleanupResult.Cutoff) - remaining, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ + remaining, err := repository.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ NodeID: "node-logs", PageNo: 1, PageSize: 10, @@ -80,7 +82,7 @@ func TestListPageAndCleanup(t *testing.T) { assert.Equal(t, int64(2), cleanupAll.DeletedCount) assert.True(t, cleanupAll.DeleteAll) - finalLogs, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ + finalLogs, err := repository.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ NodeID: "node-logs", PageNo: 1, PageSize: 10, diff --git a/internal/apps/openflare/async_tasks_test.go b/internal/apps/openflare/async_tasks_test.go index a8ecc9ae..9ecba285 100644 --- a/internal/apps/openflare/async_tasks_test.go +++ b/internal/apps/openflare/async_tasks_test.go @@ -42,8 +42,8 @@ func TestDatabaseAutoCleanupHandlerDeletesRowsWhenEnabled(t *testing.T) { require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{})) db.SetDB(sqliteDB) - resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) - resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) + resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore()) + resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore()) t.Cleanup(func() { resetObservabilityStore() resetAccessLogStore() @@ -52,7 +52,7 @@ func TestDatabaseAutoCleanupHandlerDeletesRowsWhenEnabled(t *testing.T) { ctx := context.Background() now := time.Now().UTC() - require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{{ + require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{{ NodeID: "node-a", LoggedAt: now.Add(-95 * 24 * time.Hour), RemoteAddr: "203.0.113.10", @@ -69,7 +69,7 @@ func TestDatabaseAutoCleanupHandlerDeletesRowsWhenEnabled(t *testing.T) { require.NotNil(t, result) assert.Contains(t, result.Message, "共删除") - rows, err := model.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10}) + rows, err := repository.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10}) require.NoError(t, err) assert.Empty(t, rows) } diff --git a/internal/apps/openflare/chwriter/live_ch_test.go b/internal/apps/openflare/chwriter/live_ch_test.go index 85e1f920..e418abb2 100644 --- a/internal/apps/openflare/chwriter/live_ch_test.go +++ b/internal/apps/openflare/chwriter/live_ch_test.go @@ -10,6 +10,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/chwriter" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" @@ -27,7 +29,7 @@ func TestLiveAppWritePath(t *testing.T) { now := time.Now().UTC() nodeID := "e2e-app-write-" + now.Format("150405") - if err := model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ + if err := repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ NodeID: nodeID, CapturedAt: now, CPUUsagePercent: 33.3, @@ -46,7 +48,7 @@ func TestLiveAppWritePath(t *testing.T) { deadline := time.Now().Add(45 * time.Second) var found bool for time.Now().Before(deadline) { - rows, err := model.ListOpenFlareMetricSnapshotsSince(ctx, nodeID, now.Add(-time.Minute), 10) + rows, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, nodeID, now.Add(-time.Minute), 10) if err != nil { t.Fatalf("ListOpenFlareMetricSnapshotsSince: %v", err) } @@ -61,7 +63,7 @@ func TestLiveAppWritePath(t *testing.T) { t.Fatal("metric snapshot not visible in ClickHouse after flush wait") } - latest, err := model.ListOpenFlareLatestMetricSnapshotsSince(ctx, "", now.Add(-time.Hour)) + latest, err := repository.ListOpenFlareLatestMetricSnapshotsSince(ctx, "", now.Add(-time.Hour)) if err != nil { t.Fatalf("ListOpenFlareLatestMetricSnapshotsSince: %v", err) } diff --git a/internal/apps/openflare/chwriter/writer.go b/internal/apps/openflare/chwriter/writer.go index 772d42f5..95751217 100644 --- a/internal/apps/openflare/chwriter/writer.go +++ b/internal/apps/openflare/chwriter/writer.go @@ -11,9 +11,10 @@ import ( "sync" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/persistence/batchwriter" - "github.com/Rain-kl/Wavelet/internal/model" analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics" "github.com/Rain-kl/Wavelet/internal/platform/lifecycle" analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" @@ -279,13 +280,13 @@ func withFlushRetries[T any](flush batchwriter.FlushFunc[T]) batchwriter.FlushFu } func wireModelInsertHooks() { - model.SetObservabilityInsertHooks(model.ObservabilityInsertHooks{ + repository.SetObservabilityInsertHooks(repository.ObservabilityInsertHooks{ QueueMetricSnapshot: QueueMetricSnapshot, QueueEdgeHealth: QueueEdgeHealth, QueueFrpsObservation: QueueFrpsObservation, QueueFrpcObservation: QueueFrpcObservation, }) - model.SetAccessLogInsertHooks(model.AccessLogInsertHooks{ + repository.SetAccessLogInsertHooks(repository.AccessLogInsertHooks{ QueueNodeAccessLogs: QueueNodeAccessLogs, }) } diff --git a/internal/apps/openflare/config_version/certificate_snapshot_test.go b/internal/apps/openflare/config_version/certificate_snapshot_test.go index 8a3b5647..de536c17 100644 --- a/internal/apps/openflare/config_version/certificate_snapshot_test.go +++ b/internal/apps/openflare/config_version/certificate_snapshot_test.go @@ -15,6 +15,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + oftls "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls" "github.com/Rain-kl/Wavelet/internal/infra/config" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" @@ -76,7 +78,7 @@ func TestBuildSnapshotReadsZoneDomainCertificates(t *testing.T) { require.NoError(t, err) route := &model.ProxyRoute{SiteName: "tls-site", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true, EnableHTTPS: true} - require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, route)) zone := &model.Zone{Domain: "example.com"} require.NoError(t, db.DB(ctx).Create(zone).Error) require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &route.ID, Domain: "one.example.com", CertID: &first.ID}).Error) diff --git a/internal/apps/openflare/config_version/helpers.go b/internal/apps/openflare/config_version/helpers.go index 9e39dcdd..49c8b3a2 100644 --- a/internal/apps/openflare/config_version/helpers.go +++ b/internal/apps/openflare/config_version/helpers.go @@ -11,6 +11,8 @@ import ( "strconv" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" ) @@ -208,7 +210,7 @@ func relayAgentAddress(node *model.OpenFlareNode) string { } func resolveTunnelOpenRestyUpstreamURL(ctx context.Context) string { - nodes, err := model.ListOpenFlareNodes(ctx) + nodes, err := repository.ListOpenFlareNodes(ctx) if err == nil { for index := range nodes { node := &nodes[index] @@ -230,7 +232,7 @@ func listWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*model.OpenFlareWA } groups := make([]*model.OpenFlareWAFIPGroup, 0, len(ids)) for _, id := range ids { - group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id) if err != nil { return nil, err } diff --git a/internal/apps/openflare/config_version/logics.go b/internal/apps/openflare/config_version/logics.go index b8146af2..88437be9 100644 --- a/internal/apps/openflare/config_version/logics.go +++ b/internal/apps/openflare/config_version/logics.go @@ -14,6 +14,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" "github.com/Rain-kl/Wavelet/internal/model" pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol" @@ -75,17 +77,17 @@ type CleanupResult struct { // ListConfigVersions returns all config version summaries. func ListConfigVersions(ctx context.Context) ([]*model.ConfigVersionSummary, error) { - return model.ListConfigVersionSummaries(ctx) + return repository.ListConfigVersionSummaries(ctx) } // GetConfigVersionDetail returns a config version by version. func GetConfigVersionDetail(ctx context.Context, version string) (*model.ConfigVersion, error) { - return model.GetConfigVersionByVersion(ctx, version) + return repository.GetConfigVersionByVersion(ctx, version) } // GetActiveConfigVersion returns the active config version. func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) { - return model.GetActiveConfigVersion(ctx) + return repository.GetActiveConfigVersion(ctx) } // PreviewConfigVersion renders the current draft configuration. @@ -123,7 +125,7 @@ func DiffConfigVersion(ctx context.Context) (*ConfigDiffResult, error) { ChangedOptionDetails: []ConfigOptionDiffItem{}, CurrentWebsiteCount: len(bundle.SnapshotRoutes), } - activeVersion, err := model.GetActiveConfigVersion(ctx) + activeVersion, err := repository.GetActiveConfigVersion(ctx) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { for _, route := range bundle.SnapshotRoutes { @@ -203,7 +205,7 @@ func PublishConfigVersion(ctx context.Context, createdBy string, force bool) (*m if len(bundle.Routes) == 0 { return nil, errors.New(errNoEnabledRoutes) } - activeVersion, err := model.GetActiveConfigVersion(ctx) + activeVersion, err := repository.GetActiveConfigVersion(ctx) if !force && err == nil && activeVersion.Checksum == bundle.Checksum { return nil, errors.New(errNoChangesToPublish) } @@ -228,7 +230,7 @@ func PublishConfigVersion(ctx context.Context, createdBy string, force bool) (*m IsActive: true, CreatedBy: createdBy, } - if err = model.PublishConfigVersionTx(ctx, record); err != nil { + if err = repository.PublishConfigVersionTx(ctx, record); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errVersionConflict) } @@ -243,11 +245,11 @@ func PublishConfigVersion(ctx context.Context, createdBy string, force bool) (*m // ActivateConfigVersion activates an existing config version. func ActivateConfigVersion(ctx context.Context, versionStr string) (*model.ConfigVersion, error) { - version, err := model.GetConfigVersionByVersion(ctx, versionStr) + version, err := repository.GetConfigVersionByVersion(ctx, versionStr) if err != nil { return nil, err } - if err = model.ActivateConfigVersionTx(ctx, versionStr); err != nil { + if err = repository.ActivateConfigVersionTx(ctx, versionStr); err != nil { return nil, err } version.IsActive = true @@ -263,7 +265,7 @@ func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult, if keepCount < minConfigVersionKeepCount { keepCount = minConfigVersionKeepCount } - versions, err := model.ListConfigVersionSummaries(ctx) + versions, err := repository.ListConfigVersionSummaries(ctx) if err != nil { return nil, err } @@ -283,7 +285,7 @@ func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult, if len(deleteVersions) == 0 { return &CleanupResult{DeletedCount: 0, Message: cleanupSuccessMessage}, nil } - deletedCount, err := model.DeleteConfigVersionsByVersions(ctx, deleteVersions) + deletedCount, err := repository.DeleteConfigVersionsByVersions(ctx, deleteVersions) if err != nil { return nil, err } @@ -292,7 +294,7 @@ func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult, func nextVersionNumber(ctx context.Context, now time.Time) (string, error) { prefix := now.Format("20060102") - latest, err := model.GetLatestConfigVersionByPrefix(ctx, prefix) + latest, err := repository.GetLatestConfigVersionByPrefix(ctx, prefix) if errors.Is(err, gorm.ErrRecordNotFound) { return fmt.Sprintf("%s-%03d", prefix, 1), nil } diff --git a/internal/apps/openflare/config_version/logics_test.go b/internal/apps/openflare/config_version/logics_test.go index bbdf2afe..3bafd3a5 100644 --- a/internal/apps/openflare/config_version/logics_test.go +++ b/internal/apps/openflare/config_version/logics_test.go @@ -10,6 +10,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/waf" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" @@ -101,7 +103,7 @@ func TestPublishConfigVersionCreatesVersion(t *testing.T) { Upstreams: `["http://origin.publish.example.com:8080"]`, Enabled: true, } - require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, route)) createSnapshotZoneDomains(t, ctx, route, "publish.example.com") version, err := PublishConfigVersion(ctx, "tester", false) @@ -145,15 +147,15 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) { Upstreams: `["http://origin.example.com:8080"]`, Enabled: true, } - require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, route)) createSnapshotZoneDomains(t, ctx, route, "example.com", "www.example.com") require.NoError(t, waf.EnsureDefaultRuleGroup(ctx)) - globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx) + globalGroup, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx) require.NoError(t, err) customGroup := createSnapshotRule(t, ctx, "pow-group", waf.DefaultRuleGraph()) - require.NoError(t, model.ReplaceOpenFlareWAFRuleGroupBindings(ctx, customGroup.ID, []uint{route.ID})) + require.NoError(t, repository.ReplaceOpenFlareWAFRuleGroupBindings(ctx, customGroup.ID, []uint{route.ID})) bundle, err := buildCurrentConfigBundle(ctx, true) require.NoError(t, err) @@ -203,11 +205,11 @@ func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testi Upstreams: `["http://origin.example.com:8080"]`, Enabled: true, } - require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, route)) createSnapshotZoneDomains(t, ctx, route, "pow-global.example.com") require.NoError(t, waf.EnsureDefaultRuleGroup(ctx)) - globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx) + globalGroup, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx) require.NoError(t, err) graphJSON, err := json.Marshal(snapshotPoWGraph()) require.NoError(t, err) diff --git a/internal/apps/openflare/config_version/pages_snapshot.go b/internal/apps/openflare/config_version/pages_snapshot.go index b7486f67..2f270193 100644 --- a/internal/apps/openflare/config_version/pages_snapshot.go +++ b/internal/apps/openflare/config_version/pages_snapshot.go @@ -10,6 +10,8 @@ import ( "path" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/pkg/pagesarchive" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" @@ -26,13 +28,13 @@ func buildPagesRouteSnapshot( if route == nil { return "", nil, nil, nil, errors.New("pages 路由配置无效") } - if !model.HasPagesProjectsTable(ctx) { + if !repository.HasPagesProjectsTable(ctx) { return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 模块不可用", route.SiteName) } if route.PagesProjectID == nil || *route.PagesProjectID == 0 { return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: 未绑定 Pages 项目", route.SiteName) } - project, err := model.GetPagesProjectByID(ctx, *route.PagesProjectID) + project, err := repository.GetPagesProjectByID(ctx, *route.PagesProjectID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目不存在", route.SiteName) @@ -45,7 +47,7 @@ func buildPagesRouteSnapshot( if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 { return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目没有激活部署", route.SiteName) } - activeDeployment, err := model.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) + activeDeployment, err := repository.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 激活部署不存在", route.SiteName) diff --git a/internal/apps/openflare/config_version/pages_snapshot_test.go b/internal/apps/openflare/config_version/pages_snapshot_test.go index fe99a138..7d3fa496 100644 --- a/internal/apps/openflare/config_version/pages_snapshot_test.go +++ b/internal/apps/openflare/config_version/pages_snapshot_test.go @@ -8,6 +8,8 @@ import ( "encoding/json" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" @@ -52,7 +54,7 @@ func TestBuildSnapshotRoutesPages(t *testing.T) { UpstreamType: "pages", PagesProjectID: &project.ID, } - require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, route)) createSnapshotZoneDomains(t, ctx, route, "speedtest.arctel.net") bundle, err := buildCurrentConfigBundle(ctx, true) diff --git a/internal/apps/openflare/config_version/snapshot.go b/internal/apps/openflare/config_version/snapshot.go index 5921466f..e6aa6d86 100644 --- a/internal/apps/openflare/config_version/snapshot.go +++ b/internal/apps/openflare/config_version/snapshot.go @@ -162,7 +162,7 @@ type configBundle struct { } func buildCurrentConfigBundle(ctx context.Context, requireRoutes bool) (*configBundle, error) { - routes, err := model.ListEnabledProxyRoutes(ctx) + routes, err := repository.ListEnabledProxyRoutes(ctx) if err != nil { return nil, err } @@ -224,7 +224,7 @@ func buildCurrentConfigBundle(ctx context.Context, requireRoutes bool) (*configB func buildSnapshotRoutes(ctx context.Context, routes []*model.ProxyRoute) ([]snapshotRoute, error) { items := make([]snapshotRoute, 0, len(routes)) for _, route := range routes { - zoneDomains, err := model.ListZoneDomainsByRouteID(ctx, route.ID) + zoneDomains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID) if err != nil { return nil, err } @@ -310,7 +310,7 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) ( if err := waf.EnsureDefaultRuleGroup(ctx); err != nil { return snapshotWAFDocument{}, err } - groups, err := model.ListOpenFlareWAFRuleGroups(ctx) + groups, err := repository.ListOpenFlareWAFRuleGroups(ctx) if err != nil { return snapshotWAFDocument{}, err } @@ -349,7 +349,7 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) ( if route == nil { continue } - domains, domainErr := model.ListZoneDomainsByRouteID(ctx, route.ID) + domains, domainErr := repository.ListZoneDomainsByRouteID(ctx, route.ID) if domainErr != nil { return snapshotWAFDocument{}, domainErr } @@ -358,7 +358,7 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) ( } enabledRouteSiteNames[route.ID] = route.SiteName } - rawBindings, err := model.ListOpenFlareWAFRuleGroupBindings(ctx) + rawBindings, err := repository.ListOpenFlareWAFRuleGroupBindings(ctx) if err != nil { return snapshotWAFDocument{}, err } @@ -456,7 +456,7 @@ func buildSnapshotWAFIPGroups(ctx context.Context, idSet map[uint]struct{}) ([]s } func snapshotWAFIPGroupExists(ctx context.Context, id uint) (bool, error) { - group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id) if errors.Is(err, gorm.ErrRecordNotFound) { return false, nil } @@ -593,7 +593,7 @@ func buildCertificateSupportFiles(ctx context.Context, routes []snapshotRoute) ( sort.Slice(certIDs, func(i, j int) bool { return certIDs[i] < certIDs[j] }) files := make([]SupportFile, 0, len(certIDs)*supportFilesPerCertificate) for _, certID := range certIDs { - certificate, err := model.GetTLSCertificateByID(ctx, certID) + certificate, err := repository.GetTLSCertificateByID(ctx, certID) if err != nil { return nil, err } diff --git a/internal/apps/openflare/config_version/waf_graph_snapshot_test.go b/internal/apps/openflare/config_version/waf_graph_snapshot_test.go index 572de5d8..c97e6b23 100644 --- a/internal/apps/openflare/config_version/waf_graph_snapshot_test.go +++ b/internal/apps/openflare/config_version/waf_graph_snapshot_test.go @@ -9,6 +9,8 @@ import ( "strings" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/waf" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" @@ -51,7 +53,7 @@ func TestWAFGraphSnapshotPreservesOrderAndGraphReferences(t *testing.T) { ctx := context.Background() route := &model.ProxyRoute{SiteName: "ordered.example.com", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true} - require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, route)) createSnapshotZoneDomains(t, ctx, route, route.SiteName) referenced := &model.OpenFlareWAFIPGroup{Name: "referenced", Type: "manual", Enabled: true, IPList: `["192.0.2.1"]`} @@ -61,7 +63,7 @@ func TestWAFGraphSnapshotPreservesOrderAndGraphReferences(t *testing.T) { customA := createSnapshotRule(t, ctx, "custom-a", waf.DefaultRuleGraph()) customB := createSnapshotRule(t, ctx, "custom-b", snapshotIPMatchGraph(referenced.ID)) - require.NoError(t, model.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, route.ID, []uint{customB.ID, customA.ID})) + require.NoError(t, repository.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, route.ID, []uint{customB.ID, customA.ID})) snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route}) require.NoError(t, err) @@ -91,7 +93,7 @@ func TestWAFGraphSnapshotEncodesEmptyBindingsAsArrays(t *testing.T) { ctx := context.Background() route := &model.ProxyRoute{SiteName: "empty-binding.example.com", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true} - require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, route)) createSnapshotZoneDomains(t, ctx, route, route.SiteName) snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route}) diff --git a/internal/apps/openflare/dashboard/logics.go b/internal/apps/openflare/dashboard/logics.go index bd2934ad..cadcba89 100644 --- a/internal/apps/openflare/dashboard/logics.go +++ b/internal/apps/openflare/dashboard/logics.go @@ -8,6 +8,8 @@ import ( "sort" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/observability" "github.com/Rain-kl/Wavelet/internal/model" ) @@ -113,24 +115,24 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) { now := time.Now() since := now.Add(-24 * time.Hour) - nodes, err := model.ListOpenFlareNodes(ctx) + nodes, err := repository.ListOpenFlareNodes(ctx) if err != nil { return nil, err } // Latest-per-node health: dedicated LIMIT 1 BY queries (not a global raw LIMIT). - latestSnapshotRows, err := model.ListOpenFlareLatestMetricSnapshotsSince(ctx, "", since) + latestSnapshotRows, err := repository.ListOpenFlareLatestMetricSnapshotsSince(ctx, "", since) if err != nil { return nil, err } - snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, "", since, dashboardOverviewSnapshotLimit) + snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, "", since, dashboardOverviewSnapshotLimit) if err != nil { return nil, err } - accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, "", since, dashboardDistributionLimit) + accessLogRegions, err := repository.ListOpenFlareAccessLogRegionCounts(ctx, "", since, dashboardDistributionLimit) if err != nil { return nil, err } - activeEvents, err := model.ListOpenFlareActiveHealthEvents(ctx) + activeEvents, err := repository.ListOpenFlareActiveHealthEvents(ctx) if err != nil { return nil, err } @@ -146,7 +148,7 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) { } // Global traffic summary uses true window uniqExact for UV (not sum of hourly uniques). - if summary, sumErr := model.TrafficSummaryOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ + if summary, sumErr := repository.TrafficSummaryOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ Since: since, Until: now, }); sumErr == nil { @@ -164,7 +166,7 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) { } nodeTraffic := map[string]model.OpenFlareAccessLogNodeAggregate{} - if aggregates, aggErr := model.NodeAggregatesOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ + if aggregates, aggErr := repository.NodeAggregatesOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ Since: since, Until: now, }); aggErr == nil { diff --git a/internal/apps/openflare/dashboard/logics_test.go b/internal/apps/openflare/dashboard/logics_test.go index 9fbc6328..d0973cb4 100644 --- a/internal/apps/openflare/dashboard/logics_test.go +++ b/internal/apps/openflare/dashboard/logics_test.go @@ -8,6 +8,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/glebarez/sqlite" @@ -24,8 +26,8 @@ func setupDashboardTestDB(t *testing.T) func() { require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareNode{})) db.SetDB(sqliteDB) - resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) - resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) + resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore()) + resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore()) return func() { resetObservabilityStore() resetAccessLogStore() @@ -59,14 +61,14 @@ func TestGetOverviewStructure(t *testing.T) { }).Error) // Seed older + newer snapshots per node; health must use latest-per-node, not a global raw limit. - require.NoError(t, model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ + require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ NodeID: "node-dashboard-1", CapturedAt: now.Add(-2 * time.Hour), CPUUsagePercent: 10, MemoryUsedBytes: 1, MemoryTotalBytes: 10, })) - require.NoError(t, model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ + require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ NodeID: "node-dashboard-1", CapturedAt: now.Add(-time.Minute), CPUUsagePercent: 55, @@ -101,7 +103,7 @@ func TestGetOverviewStructure(t *testing.T) { StatusCode: 502, BytesSent: 10, }) - require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, logs)) + require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, logs)) overview, err := GetOverview(ctx) require.NoError(t, err) diff --git a/internal/apps/openflare/flared/helpers.go b/internal/apps/openflare/flared/helpers.go index 9a51fc69..3843b211 100644 --- a/internal/apps/openflare/flared/helpers.go +++ b/internal/apps/openflare/flared/helpers.go @@ -10,6 +10,8 @@ import ( "strconv" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/relay" "github.com/Rain-kl/Wavelet/internal/model" ) @@ -51,7 +53,7 @@ func normalizeFlaredHeartbeatPayload(payload HeartbeatPayload) HeartbeatPayload } func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) { - version, err := model.GetActiveConfigVersion(ctx) + version, err := repository.GetActiveConfigVersion(ctx) if err != nil { return nil, err } @@ -62,7 +64,7 @@ func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) { } func listTunnelRelayNodes(ctx context.Context) ([]model.OpenFlareNode, error) { - nodes, err := model.ListOpenFlareNodes(ctx) + nodes, err := repository.ListOpenFlareNodes(ctx) if err != nil { return nil, err } diff --git a/internal/apps/openflare/flared/logics.go b/internal/apps/openflare/flared/logics.go index 19a4d011..c20ecf26 100644 --- a/internal/apps/openflare/flared/logics.go +++ b/internal/apps/openflare/flared/logics.go @@ -10,8 +10,9 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "gorm.io/gorm" ) @@ -74,7 +75,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat node.LastSeenAt = &lastSeen node.Status = nodeStatusOnline - if err := db.DB(ctx).Model(node).Updates(changes).Error; err != nil { + if err := repository.UpdateOpenFlareNodeColumns(ctx, node, changes); err != nil { return nil, fmt.Errorf("update flared heartbeat: %w", err) } agent.RefreshAccessTokenCache(ctx, node) @@ -102,7 +103,7 @@ func GetTunnelConfig(ctx context.Context, node *model.OpenFlareNode) (*TunnelCon return nil, fmt.Errorf("no active config version: %w", err) } - routes, err := model.ListProxyRoutes(ctx) + routes, err := repository.ListProxyRoutes(ctx) if err != nil { return nil, fmt.Errorf("get proxy routes: %w", err) } @@ -133,7 +134,7 @@ func GetTunnelConfig(ctx context.Context, node *model.OpenFlareNode) (*TunnelCon if !route.Enabled { continue } - zoneDomains, domainErr := model.ListZoneDomainsByRouteID(ctx, route.ID) + zoneDomains, domainErr := repository.ListZoneDomainsByRouteID(ctx, route.ID) if domainErr != nil || len(zoneDomains) == 0 { continue } @@ -177,12 +178,12 @@ func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFl return nil, errors.New("result 仅支持 success、warning 或 failed") } - latest, err := model.GetLatestOpenFlareApplyLogByNodeID(ctx, payload.NodeID) + latest, err := repository.GetLatestOpenFlareApplyLogByNodeID(ctx, payload.NodeID) if err != nil { return nil, err } if model.IsRepeatSuccessApplyLog(latest, payload.Version, payload.Checksum, payload.Result) { - if err := updateFlaredNodeFromApplyLog(ctx, payload, now); err != nil { + if err := repository.UpdateOpenFlareNodeFromApplyResult(ctx, payload.NodeID, payload.Result, payload.Version, payload.Message, now); err != nil { return nil, err } return latest, nil @@ -200,45 +201,12 @@ func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFl CreatedAt: now, } - err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Create(log).Error; err != nil { - return err - } - return updateFlaredNodeFromApplyLogTx(tx, payload, now) - }) - if err != nil { + if err := repository.CreateOpenFlareApplyLogAndUpdateNode(ctx, log, payload.Result, payload.Version, payload.Message); err != nil { return nil, err } return log, nil } -func updateFlaredNodeFromApplyLog(ctx context.Context, payload ApplyLogPayload, now time.Time) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New("database not initialized") - } - return conn.Transaction(func(tx *gorm.DB) error { - return updateFlaredNodeFromApplyLogTx(tx, payload, now) - }) -} - -func updateFlaredNodeFromApplyLogTx(tx *gorm.DB, payload ApplyLogPayload, now time.Time) error { - var node model.OpenFlareNode - if err := tx.Where("node_id = ?", payload.NodeID).First(&node).Error; err != nil { - return err - } - node.Status = nodeStatusOnline - lastSeen := now - node.LastSeenAt = &lastSeen - if payload.Result == applyResultOK { - node.CurrentVersion = payload.Version - node.LastError = "" - } else { - node.LastError = payload.Message - } - return tx.Model(&node).Select("status", "last_seen_at", "current_version", "last_error").Updates(&node).Error -} - func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload { payload.NodeID = strings.TrimSpace(payload.NodeID) payload.Version = strings.TrimSpace(payload.Version) diff --git a/internal/apps/openflare/flared/middleware.go b/internal/apps/openflare/flared/middleware.go index 89a8b53c..1e800efe 100644 --- a/internal/apps/openflare/flared/middleware.go +++ b/internal/apps/openflare/flared/middleware.go @@ -8,6 +8,8 @@ import ( "errors" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/gin-gonic/gin" @@ -38,7 +40,7 @@ func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlar if token == "" { return nil, errors.New("missing tunnel token") } - node, err := model.GetOpenFlareNodeByAccessToken(ctx, token) + node, err := repository.GetOpenFlareNodeByAccessToken(ctx, token) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, errors.New("invalid tunnel token") diff --git a/internal/apps/openflare/flared/middleware_test.go b/internal/apps/openflare/flared/middleware_test.go index a8122e78..4baeb020 100644 --- a/internal/apps/openflare/flared/middleware_test.go +++ b/internal/apps/openflare/flared/middleware_test.go @@ -9,6 +9,8 @@ import ( "net/http/httptest" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/shared/response" @@ -44,7 +46,7 @@ func seedFlaredNode(t *testing.T, nodeType, accessToken string) *model.OpenFlare NodeType: nodeType, AccessToken: accessToken, } - require.NoError(t, model.CreateOpenFlareNode(ctx, node)) + require.NoError(t, repository.CreateOpenFlareNode(ctx, node)) return node } diff --git a/internal/apps/openflare/flared/observability.go b/internal/apps/openflare/flared/observability.go index ea9b5bb3..0a5a28b6 100644 --- a/internal/apps/openflare/flared/observability.go +++ b/internal/apps/openflare/flared/observability.go @@ -10,7 +10,6 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "go.uber.org/zap" ) @@ -40,11 +39,7 @@ func persistFlaredObservability(ctx context.Context, nodeID string, payload Hear }, }) } - conn := db.DB(ctx) - if conn == nil { - return - } - if err := agent.ReconcileScopedNodeHealthEvents(conn, nodeID, events, reportedAt, managedTypes); err != nil { + if err := agent.ReconcileScopedNodeHealthEvents(ctx, nodeID, events, reportedAt, managedTypes); err != nil { zap.L().Error("persist flared health events failed", zap.String("node_id", nodeID), zap.Error(err)) } } diff --git a/internal/apps/openflare/flared/observability_test.go b/internal/apps/openflare/flared/observability_test.go index 9a78ee14..4bc6e7de 100644 --- a/internal/apps/openflare/flared/observability_test.go +++ b/internal/apps/openflare/flared/observability_test.go @@ -7,6 +7,8 @@ import ( "context" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" @@ -62,7 +64,7 @@ func TestHeartbeatFlaredEmitsHealthEventOnUnhealthy(t *testing.T) { }) require.NoError(t, err) - events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 20) + events, err := repository.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 20) require.NoError(t, err) require.Len(t, events, 1) assert.Equal(t, flaredRuntimeUnhealthyEventType, events[0].EventType) diff --git a/internal/apps/openflare/integration/agent_protocol_test.go b/internal/apps/openflare/integration/agent_protocol_test.go index 674c4f67..92f00abc 100644 --- a/internal/apps/openflare/integration/agent_protocol_test.go +++ b/internal/apps/openflare/integration/agent_protocol_test.go @@ -8,6 +8,8 @@ import ( "net/http" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" ofnode "github.com/Rain-kl/Wavelet/internal/apps/openflare/node" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" @@ -38,8 +40,8 @@ func setupProtocolTestEnv(t *testing.T) (*gin.Engine, func()) { db.SetDB(sqliteDB) agent.ResetAuthCacheForTest() - resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) - resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) + resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore()) + resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore()) engine := testhelper.NewTestGinEngine() mountOpenFlareTestRoutes(engine) @@ -110,7 +112,7 @@ func TestAgentRelayFlaredProtocol(t *testing.T) { assert.NotNil(t, heartbeatData.RelayConfig) assert.NotNil(t, heartbeatData.RelaySettings) - stored, err := model.GetOpenFlareNodeByNodeID(ctx, relayNode.NodeID) + stored, err := repository.GetOpenFlareNodeByNodeID(ctx, relayNode.NodeID) require.NoError(t, err) assert.Equal(t, "online", stored.Status) assert.Equal(t, "healthy", stored.RelayStatus) @@ -134,7 +136,7 @@ func TestAgentRelayFlaredProtocol(t *testing.T) { assert.Equal(t, http.StatusOK, rec.Code) requireAPIOK(t, rec) - stored, err := model.GetOpenFlareNodeByNodeID(ctx, clientNode.NodeID) + stored, err := repository.GetOpenFlareNodeByNodeID(ctx, clientNode.NodeID) require.NoError(t, err) assert.Equal(t, "online", stored.Status) assert.Equal(t, "v0.2.0", stored.Version) @@ -161,7 +163,7 @@ func TestAgentRelayFlaredProtocol(t *testing.T) { assert.NotEmpty(t, registration.AccessToken) assert.Equal(t, "discovered-edge", registration.Name) - stored, err := model.GetOpenFlareNodeByNodeID(ctx, registration.NodeID) + stored, err := repository.GetOpenFlareNodeByNodeID(ctx, registration.NodeID) require.NoError(t, err) assert.Equal(t, "online", stored.Status) assert.Equal(t, registration.AccessToken, stored.AccessToken) @@ -190,7 +192,7 @@ func TestAgentRelayFlaredProtocol(t *testing.T) { assert.Equal(t, "success", applyLog.Result) assert.Equal(t, "20260618-001", applyLog.Version) - stored, err := model.GetOpenFlareNodeByNodeID(ctx, edge.NodeID) + stored, err := repository.GetOpenFlareNodeByNodeID(ctx, edge.NodeID) require.NoError(t, err) assert.Equal(t, "online", stored.Status) assert.Equal(t, "20260618-001", stored.CurrentVersion) diff --git a/internal/apps/openflare/node/logics.go b/internal/apps/openflare/node/logics.go index 3553db35..0cdb28d9 100644 --- a/internal/apps/openflare/node/logics.go +++ b/internal/apps/openflare/node/logics.go @@ -134,7 +134,7 @@ type HealthEventCleanupResult = observability.HealthEventCleanupResult // ListNodes returns all node views with latest apply log metadata. func ListNodes(ctx context.Context) ([]*View, error) { - nodes, err := model.ListOpenFlareNodes(ctx) + nodes, err := repository.ListOpenFlareNodes(ctx) if err != nil { return nil, err } @@ -142,7 +142,7 @@ func ListNodes(ctx context.Context) ([]*View, error) { for _, node := range nodes { nodeIDs = append(nodeIDs, node.NodeID) } - latestLogs, err := model.GetLatestOpenFlareApplyLogsByNodeIDs(ctx, nodeIDs) + latestLogs, err := repository.GetLatestOpenFlareApplyLogsByNodeIDs(ctx, nodeIDs) if err != nil { return nil, err } @@ -209,7 +209,7 @@ func CreateNode(ctx context.Context, input Input) (*View, error) { node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL) node.RelayWebServerEnabled = input.RelayWebServerEnabled } - if err = model.CreateOpenFlareNode(ctx, node); err != nil { + if err = repository.CreateOpenFlareNode(ctx, node); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errNodeIDConflict) } @@ -227,7 +227,7 @@ func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) { if name == "" { return nil, errors.New(errNodeNameRequired) } - node, err := model.GetOpenFlareNodeByID(ctx, id) + node, err := repository.GetOpenFlareNodeByID(ctx, id) if err != nil { return nil, err } @@ -252,7 +252,7 @@ func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) { node.RelayVhostHTTPPort = input.RelayVhostHTTPPort } } - if err = model.SaveOpenFlareNode(ctx, node); err != nil { + if err = repository.SaveOpenFlareNode(ctx, node); err != nil { return nil, err } return buildNodeView(node), nil @@ -260,10 +260,10 @@ func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) { // DeleteNode removes a node by id. func DeleteNode(ctx context.Context, id uint) error { - if _, err := model.GetOpenFlareNodeByID(ctx, id); err != nil { + if _, err := repository.GetOpenFlareNodeByID(ctx, id); err != nil { return err } - return model.DeleteOpenFlareNode(ctx, id) + return repository.DeleteOpenFlareNode(ctx, id) } // GetBootstrapToken returns the global discovery token, creating one if missing. @@ -289,7 +289,7 @@ func RotateBootstrapToken(ctx context.Context) (*BootstrapView, error) { // GetAgentRelease checks the latest agent release for a node. func GetAgentRelease(ctx context.Context, id uint, channel string) (*AgentReleaseInfo, error) { - node, err := model.GetOpenFlareNodeByID(ctx, id) + node, err := repository.GetOpenFlareNodeByID(ctx, id) if err != nil { return nil, err } @@ -302,7 +302,7 @@ func GetAgentRelease(ctx context.Context, id uint, channel string) (*AgentReleas // RequestAgentUpdate marks a node for manual agent update. func RequestAgentUpdate(ctx context.Context, id uint, input AgentUpdateInput) (*View, error) { - node, err := model.GetOpenFlareNodeByID(ctx, id) + node, err := repository.GetOpenFlareNodeByID(ctx, id) if err != nil { return nil, err } @@ -323,7 +323,7 @@ func RequestAgentUpdate(ctx context.Context, id uint, input AgentUpdateInput) (* node.UpdateRequested = true node.UpdateChannel = channel.String() node.UpdateTag = tagName - if err = model.UpdateOpenFlareNodeFields(ctx, node, "update_requested", "update_channel", "update_tag"); err != nil { + if err = repository.UpdateOpenFlareNodeFields(ctx, node, "update_requested", "update_channel", "update_tag"); err != nil { return nil, err } return buildNodeView(node), nil @@ -331,12 +331,12 @@ func RequestAgentUpdate(ctx context.Context, id uint, input AgentUpdateInput) (* // RequestOpenrestyRestart marks a node for openresty restart. func RequestOpenrestyRestart(ctx context.Context, id uint) (*View, error) { - node, err := model.GetOpenFlareNodeByID(ctx, id) + node, err := repository.GetOpenFlareNodeByID(ctx, id) if err != nil { return nil, err } node.RestartOpenrestyRequested = true - if err = model.UpdateOpenFlareNodeFields(ctx, node, "restart_openresty_requested"); err != nil { + if err = repository.UpdateOpenFlareNodeFields(ctx, node, "restart_openresty_requested"); err != nil { return nil, err } return buildNodeView(node), nil @@ -344,11 +344,11 @@ func RequestOpenrestyRestart(ctx context.Context, id uint) (*View, error) { // RequestForceSync pushes force_sync_config to a connected agent websocket. func RequestForceSync(ctx context.Context, id uint) (*View, error) { - node, err := model.GetOpenFlareNodeByID(ctx, id) + node, err := repository.GetOpenFlareNodeByID(ctx, id) if err != nil { return nil, err } - activeConfig, err := model.GetActiveConfigVersion(ctx) + activeConfig, err := repository.GetActiveConfigVersion(ctx) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, fmt.Errorf("无法获取当前激活的配置版本:%s", errNoActiveConfigVersion) diff --git a/internal/apps/openflare/node/logics_test.go b/internal/apps/openflare/node/logics_test.go index 1147ef4d..b19ba4a0 100644 --- a/internal/apps/openflare/node/logics_test.go +++ b/internal/apps/openflare/node/logics_test.go @@ -40,8 +40,8 @@ func setupNodeTestDB(t *testing.T) func() { )) db.SetDB(sqliteDB) - resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) - resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) + resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore()) + resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore()) return func() { resetObservabilityStore() @@ -83,7 +83,7 @@ func TestCreateTunnelRelayNode(t *testing.T) { assert.Equal(t, 7000, view.RelayBindPort) assert.Equal(t, 8080, view.RelayVhostHTTPPort) - stored, err := model.GetOpenFlareNodeByID(ctx, view.ID) + stored, err := repository.GetOpenFlareNodeByID(ctx, view.ID) require.NoError(t, err) assert.NotEmpty(t, stored.RelayAuthToken) } @@ -139,7 +139,7 @@ func TestDeleteNode(t *testing.T) { require.NoError(t, err) require.NoError(t, DeleteNode(ctx, created.ID)) - _, err = model.GetOpenFlareNodeByID(ctx, created.ID) + _, err = repository.GetOpenFlareNodeByID(ctx, created.ID) require.Error(t, err) assert.ErrorIs(t, err, gorm.ErrRecordNotFound) } diff --git a/internal/apps/openflare/observability/access_log_logics.go b/internal/apps/openflare/observability/access_log_logics.go index 6b81a381..dc55b928 100644 --- a/internal/apps/openflare/observability/access_log_logics.go +++ b/internal/apps/openflare/observability/access_log_logics.go @@ -8,6 +8,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" "github.com/Rain-kl/Wavelet/pkg/logger" @@ -309,7 +311,7 @@ func GetAccessLogOverview(ctx context.Context, input AccessLogOverviewQuery) (*A Until: now, } - summaryRow, err := model.TrafficSummaryOpenFlareAccessLogs(ctx, query) + summaryRow, err := repository.TrafficSummaryOpenFlareAccessLogs(ctx, query) if err != nil { return nil, err } @@ -408,7 +410,7 @@ func valueCountDistribution( column string, limit int, ) []DistributionItem { - rows, err := model.ValueCountsOpenFlareAccessLogs(ctx, query, column, limit) + rows, err := repository.ValueCountsOpenFlareAccessLogs(ctx, query, column, limit) if err != nil { logger.ErrorF(ctx, "[AccessLog] ValueCountsOpenFlareAccessLogs failed for column %s: %v", column, err) return []DistributionItem{} @@ -496,7 +498,7 @@ func buildAccessLogOverviewTrends( bandwidth[index].BucketStartedAt = bucketAt } - buckets, err := model.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ + buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ NodeID: query.NodeID, Host: query.Host, Hosts: query.Hosts, @@ -532,11 +534,11 @@ func buildAccessLogOverviewTrends( func ListAccessLogs(ctx context.Context, input AccessLogQuery) (*AccessLogList, error) { normalized := normalizeAccessLogQuery(input) modelQuery := buildModelAccessLogQuery(normalized) - logs, err := model.ListOpenFlareAccessLogs(ctx, modelQuery) + logs, err := repository.ListOpenFlareAccessLogs(ctx, modelQuery) if err != nil { return nil, err } - totalRecords, totalIPs, _, err := model.CountOpenFlareAccessLogs(ctx, modelQuery) + totalRecords, totalIPs, _, err := repository.CountOpenFlareAccessLogs(ctx, modelQuery) if err != nil { return nil, err } @@ -597,15 +599,15 @@ func ListFoldedAccessLogs(ctx context.Context, input AccessLogQuery) (*FoldedAcc SortOrder: normalized.SortOrder, FoldMinutes: foldMinutes, } - items, err := model.ListOpenFlareAccessLogBuckets(ctx, bucketQuery) + items, err := repository.ListOpenFlareAccessLogBuckets(ctx, bucketQuery) if err != nil { return nil, err } - totalBuckets, err := model.CountOpenFlareAccessLogBuckets(ctx, bucketQuery) + totalBuckets, err := repository.CountOpenFlareAccessLogBuckets(ctx, bucketQuery) if err != nil { return nil, err } - totalRecords, totalIPs, _, err := model.CountOpenFlareAccessLogs(ctx, modelQuery) + totalRecords, totalIPs, _, err := repository.CountOpenFlareAccessLogs(ctx, modelQuery) if err != nil { return nil, err } @@ -654,11 +656,11 @@ func ListFoldedAccessLogIPs(ctx context.Context, input FoldedAccessLogIPQuery) ( SortBy: normalized.SortBy, SortOrder: normalized.SortOrder, } - items, err := model.ListOpenFlareAccessLogBucketIPs(ctx, modelQuery) + items, err := repository.ListOpenFlareAccessLogBucketIPs(ctx, modelQuery) if err != nil { return nil, err } - totalIP, err := model.CountOpenFlareAccessLogBucketIPs(ctx, modelQuery) + totalIP, err := repository.CountOpenFlareAccessLogBucketIPs(ctx, modelQuery) if err != nil { return nil, err } @@ -706,11 +708,11 @@ func ListAccessLogIPSummaries(ctx context.Context, input AccessLogIPSummaryQuery SortBy: normalized.SortBy, SortOrder: normalized.SortOrder, } - items, err := model.ListOpenFlareAccessLogIPSummaries(ctx, query, time.Time{}) + items, err := repository.ListOpenFlareAccessLogIPSummaries(ctx, query, time.Time{}) if err != nil { return nil, err } - totalIP, err := model.CountOpenFlareAccessLogIPSummaries(ctx, query) + totalIP, err := repository.CountOpenFlareAccessLogIPSummaries(ctx, query) if err != nil { return nil, err } @@ -751,7 +753,7 @@ func GetAccessLogIPTrend(ctx context.Context, input AccessLogIPTrendQuery) (*Acc if err != nil { return nil, err } - points, err := model.ListOpenFlareAccessLogIPTrend(ctx, model.OpenFlareAccessLogIPTrendQuery{ + points, err := repository.ListOpenFlareAccessLogIPTrend(ctx, model.OpenFlareAccessLogIPTrendQuery{ NodeID: strings.TrimSpace(normalized.NodeID), RemoteAddr: strings.TrimSpace(normalized.RemoteAddr), Host: strings.TrimSpace(normalized.Host), @@ -802,7 +804,7 @@ func GetAccessLogIPAnalysis(ctx context.Context, input AccessLogIPAnalysisQuery) Until: now, } - summaryRow, err := model.TrafficSummaryOpenFlareAccessLogs(ctx, query) + summaryRow, err := repository.TrafficSummaryOpenFlareAccessLogs(ctx, query) if err != nil { return nil, err } @@ -868,7 +870,7 @@ func CleanupAccessLogs(ctx context.Context, input AccessLogCleanupInput) (*Acces return nil, errors.New("retention_days 必须在 1 到 90 之间") } cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour) - deleted, err := model.DeleteOpenFlareAccessLogsBefore(ctx, cutoff) + deleted, err := repository.DeleteOpenFlareAccessLogsBefore(ctx, cutoff) if err != nil { return nil, err } @@ -913,7 +915,7 @@ func listNodeNameMap(ctx context.Context, logs []*model.OpenFlareAccessLog) (map if len(nodeIDs) == 0 { return map[string]string{}, nil } - nodes, err := model.ListOpenFlareNodesByNodeIDs(ctx, nodeIDs) + nodes, err := repository.ListOpenFlareNodesByNodeIDs(ctx, nodeIDs) if err != nil { return nil, err } diff --git a/internal/apps/openflare/observability/analytics.go b/internal/apps/openflare/observability/analytics.go index 0df2304b..c0c590ce 100644 --- a/internal/apps/openflare/observability/analytics.go +++ b/internal/apps/openflare/observability/analytics.go @@ -9,6 +9,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" ) @@ -135,7 +137,7 @@ func buildTrafficWindowSummaryFromAccessLogs( nodeID string, since, until time.Time, ) *TrafficWindowSummary { - row, err := model.TrafficSummaryOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ + row, err := repository.TrafficSummaryOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ NodeID: nodeID, Since: since, Until: until, @@ -203,7 +205,7 @@ func BuildTrafficDistributionsFromAccessLogs( topDomains := make(distributionAccumulator) query := model.OpenFlareAccessLogQuery{Since: since, Until: until} - if statusRows, err := model.ValueCountsOpenFlareAccessLogs(ctx, query, "status_code", limit); err == nil { + if statusRows, err := repository.ValueCountsOpenFlareAccessLogs(ctx, query, "status_code", limit); err == nil { for _, row := range statusRows { if strings.TrimSpace(row.Value) == "" || row.Count <= 0 { continue @@ -211,7 +213,7 @@ func BuildTrafficDistributionsFromAccessLogs( statusCodes[row.Value] = row.Count } } - if hostRows, err := model.ValueCountsOpenFlareAccessLogs(ctx, query, "host", limit); err == nil { + if hostRows, err := repository.ValueCountsOpenFlareAccessLogs(ctx, query, "host", limit); err == nil { for _, row := range hostRows { if strings.TrimSpace(row.Value) == "" || row.Count <= 0 { continue @@ -287,7 +289,7 @@ func BuildNodeTrends( applyAccessLogBytesToNetworkTrend(ctx, now, nodeID, trendSince, networkTrend) diskIOTrend := BuildDiskIOTrendPoints(now, snapshots) - metricHourly, metricErr := model.ListOpenFlareMetricHourlySince(ctx, nodeID, trendSince) + metricHourly, metricErr := repository.ListOpenFlareMetricHourlySince(ctx, nodeID, trendSince) if metricErr == nil && len(metricHourly) > 0 { capacityTrend = BuildCapacityTrendPointsFromHourly(now, metricHourly) diskIOTrend = BuildDiskIOTrendPointsFromHourly(now, metricHourly) @@ -311,7 +313,7 @@ func BuildTrafficTrendPointsFromAccessLogs(ctx context.Context, now time.Time, n points[index].BucketStartedAt = start.Add(time.Duration(index) * time.Hour) } - if hourly, err := model.ListOpenFlareTrafficHourlySince(ctx, nodeID, since); err == nil && len(hourly) > 0 { + if hourly, err := repository.ListOpenFlareTrafficHourlySince(ctx, nodeID, since); err == nil && len(hourly) > 0 { for _, row := range hourly { if row == nil { continue @@ -327,7 +329,7 @@ func BuildTrafficTrendPointsFromAccessLogs(ctx context.Context, now time.Time, n return points } - buckets, err := model.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ + buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ NodeID: nodeID, Since: since, Until: now, @@ -373,7 +375,7 @@ func applyAccessLogBytesToNetworkTrend(ctx context.Context, now time.Time, nodeI return } - buckets, err := model.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ + buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ NodeID: nodeID, Since: since, Until: now, @@ -406,7 +408,7 @@ type accessLogHourBytes struct { } func analyticsListAccessLogHourlyBytes(ctx context.Context, nodeID string, since time.Time) (map[int64]accessLogHourBytes, error) { - rows, err := model.ListOpenFlareAccessLogHourlySince(ctx, nodeID, since) + rows, err := repository.ListOpenFlareAccessLogHourlySince(ctx, nodeID, since) if err != nil { return nil, err } diff --git a/internal/apps/openflare/observability/node_logics.go b/internal/apps/openflare/observability/node_logics.go index a6322ed3..b586519c 100644 --- a/internal/apps/openflare/observability/node_logics.go +++ b/internal/apps/openflare/observability/node_logics.go @@ -10,6 +10,8 @@ import ( "sync" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "gorm.io/gorm" ) @@ -94,7 +96,7 @@ type HealthEventCleanupResult struct { // GetNodeObservability returns observability details for a node. func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeView, error) { now := time.Now() - node, err := model.GetOpenFlareNodeByID(ctx, id) + node, err := repository.GetOpenFlareNodeByID(ctx, id) if err != nil { return nil, err } @@ -105,7 +107,7 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV limit := normalizeObservabilityLimit(query.Limit) since := now.Add(-normalizeObservabilityWindow(query.Hours)) - profile, err := model.GetOpenFlareNodeSystemProfile(ctx, node.NodeID) + profile, err := repository.GetOpenFlareNodeSystemProfile(ctx, node.NodeID) if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { return nil, err } @@ -113,19 +115,19 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV profile = nil } - snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, since, limit) + snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, since, limit) if err != nil { return nil, err } - edgeHealth, err := model.ListOpenFlareEdgeHealth(ctx, node.NodeID, since, limit) + edgeHealth, err := repository.ListOpenFlareEdgeHealth(ctx, node.NodeID, since, limit) if err != nil { return nil, err } - accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, defaultTrafficDistributionLimit) + accessLogRegions, err := repository.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, defaultTrafficDistributionLimit) if err != nil { return nil, err } - events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, limit) + events, err := repository.ListOpenFlareHealthEvents(ctx, node.NodeID, false, limit) if err != nil { return nil, err } @@ -146,7 +148,7 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV Trends: BuildNodeTrends(ctx, now, node.NodeID, snapshots), } if node.NodeType == "tunnel_relay" { - frpsObs, frpsErr := model.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1) + frpsObs, frpsErr := repository.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1) if frpsErr != nil { return nil, frpsErr } @@ -187,11 +189,11 @@ func setCachedNodeObservability(nodeID string, view *NodeView) { // CleanupHealthEvents removes all health events for a node. func CleanupHealthEvents(ctx context.Context, id uint) (*HealthEventCleanupResult, error) { - node, err := model.GetOpenFlareNodeByID(ctx, id) + node, err := repository.GetOpenFlareNodeByID(ctx, id) if err != nil { return nil, err } - deletedCount, err := model.DeleteOpenFlareHealthEventsByNodeID(ctx, node.NodeID) + deletedCount, err := repository.DeleteOpenFlareHealthEventsByNodeID(ctx, node.NodeID) if err != nil { return nil, err } diff --git a/internal/apps/openflare/option/logics.go b/internal/apps/openflare/option/logics.go index 3693c29b..ff716cd3 100644 --- a/internal/apps/openflare/option/logics.go +++ b/internal/apps/openflare/option/logics.go @@ -149,7 +149,7 @@ func getStatus(ctx context.Context, baseAPIPath string) (*statusView, error) { } func publicAuthSources(ctx context.Context, baseAPIPath string) ([]publicAuthSourceView, error) { - sources, err := model.GetActiveAuthSources(ctx) + sources, err := repository.GetActiveAuthSources(ctx) if err != nil { return nil, err } diff --git a/internal/apps/openflare/option/logics_test.go b/internal/apps/openflare/option/logics_test.go index c19d34db..abc1c619 100644 --- a/internal/apps/openflare/option/logics_test.go +++ b/internal/apps/openflare/option/logics_test.go @@ -123,11 +123,11 @@ func TestCleanupDatabaseObservabilityDeletesRows(t *testing.T) { defer cleanup() ctx := context.Background() - resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) + resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore()) defer resetAccessLogStore() now := time.Now().UTC() - require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{ + require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{ { NodeID: "node-a", LoggedAt: now.Add(-10 * 24 * time.Hour), @@ -165,7 +165,7 @@ func TestCleanupDatabaseObservabilityDeletesRows(t *testing.T) { assert.True(t, result.DeleteAll) assert.Equal(t, "truncate", result.CleanupMode) - rows, err := model.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10}) + rows, err := repository.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10}) require.NoError(t, err) assert.Empty(t, rows) } diff --git a/internal/apps/openflare/origin/logics.go b/internal/apps/openflare/origin/logics.go index 6e676c61..d1c19ccd 100644 --- a/internal/apps/openflare/origin/logics.go +++ b/internal/apps/openflare/origin/logics.go @@ -11,8 +11,8 @@ import ( "sort" "strings" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "gorm.io/gorm" ) @@ -51,7 +51,7 @@ type DetailView struct { // ListOrigins 列出全部源站。 func ListOrigins(ctx context.Context) ([]View, error) { - origins, err := model.ListOrigins(ctx) + origins, err := repository.ListOrigins(ctx) if err != nil { return nil, err } @@ -60,7 +60,7 @@ func ListOrigins(ctx context.Context) ([]View, error) { // GetOriginDetail 获取源站详情。 func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) { - origin, err := model.GetOriginByID(ctx, id) + origin, err := repository.GetOriginByID(ctx, id) if err != nil { return nil, err } @@ -68,13 +68,13 @@ func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) { if err != nil { return nil, err } - routes, err := model.ListProxyRoutesByOriginID(ctx, id) + routes, err := repository.ListProxyRoutesByOriginID(ctx, id) if err != nil { return nil, err } items := make([]RouteSummary, 0, len(routes)) for _, route := range routes { - domains, err := model.ListZoneDomainsByRouteID(ctx, route.ID) + domains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID) if err != nil { return nil, err } @@ -105,7 +105,7 @@ func CreateOrigin(ctx context.Context, input Input) (*model.Origin, error) { if err != nil { return nil, err } - if err = model.CreateOriginRecord(ctx, origin); err != nil { + if err = repository.CreateOriginRecord(ctx, origin); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errOriginAddressExists) } @@ -116,7 +116,7 @@ func CreateOrigin(ctx context.Context, input Input) (*model.Origin, error) { // UpdateOrigin 更新源站。 func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, error) { - origin, err := model.GetOriginByID(ctx, id) + origin, err := repository.GetOriginByID(ctx, id) if err != nil { return nil, err } @@ -125,8 +125,8 @@ func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, err if err != nil { return nil, err } - err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Save(nextOrigin).Error; err != nil { + err = repository.WithOriginTx(ctx, func(tx *gorm.DB) error { + if err := repository.SaveOriginTx(tx, nextOrigin); err != nil { if isUniqueConstraintError(err) { return errors.New(errOriginAddressExists) } @@ -145,17 +145,17 @@ func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, err // DeleteOrigin 删除源站。 func DeleteOrigin(ctx context.Context, id uint) error { - count, err := model.CountProxyRoutesByOriginID(ctx, id) + count, err := repository.CountProxyRoutesByOriginID(ctx, id) if err != nil { return err } if count > 0 { return errors.New(errOriginDeleteReferenced) } - if _, err = model.GetOriginByID(ctx, id); err != nil { + if _, err = repository.GetOriginByID(ctx, id); err != nil { return err } - return model.DeleteOriginRecord(ctx, id) + return repository.DeleteOriginRecord(ctx, id) } func buildOrigin(existing *model.Origin, input Input) (*model.Origin, error) { @@ -173,7 +173,7 @@ func buildOrigin(existing *model.Origin, input Input) (*model.Origin, error) { } func buildOriginViews(ctx context.Context, origins []model.Origin) ([]View, error) { - countRows, err := model.ListOriginRouteCounts(ctx) + countRows, err := repository.ListOriginRouteCounts(ctx) if err != nil { return nil, err } @@ -197,11 +197,11 @@ func buildOriginViews(ctx context.Context, origins []model.Origin) ([]View, erro } func updateRoutesForOriginAddress(ctx context.Context, tx *gorm.DB, originID uint, address string) error { - if !model.HasProxyRoutesTable(ctx) { + if !repository.HasProxyRoutesTable(ctx) { return nil } - var routes []model.OriginProxyRoute - if err := tx.Where("origin_id = ?", originID).Order("id asc").Find(&routes).Error; err != nil { + routes, err := repository.ListProxyRoutesByOriginIDAscTx(tx, originID) + if err != nil { return fmt.Errorf("query routes for origin update failed: %w", err) } for _, route := range routes { @@ -224,12 +224,7 @@ func updateRoutesForOriginAddress(ctx context.Context, tx *gorm.DB, originID uin if err != nil { return fmt.Errorf("encode route %d upstreams failed: %w", route.ID, err) } - if err := tx.Model(&model.OriginProxyRoute{}). - Where("id = ?", route.ID). - Updates(map[string]any{ - "origin_url": rewrittenOriginURL, - "upstreams": string(upstreamsJSON), - }).Error; err != nil { + if err := repository.UpdateProxyRouteOriginAddressTx(tx, route.ID, rewrittenOriginURL, string(upstreamsJSON)); err != nil { return fmt.Errorf("update route %d origin address failed: %w", route.ID, err) } } diff --git a/internal/apps/openflare/pages/github_source.go b/internal/apps/openflare/pages/github_source.go index 05a6f3c3..cd929a19 100644 --- a/internal/apps/openflare/pages/github_source.go +++ b/internal/apps/openflare/pages/github_source.go @@ -17,11 +17,10 @@ import ( "unicode" "unicode/utf8" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" "gorm.io/gorm" - "gorm.io/gorm/clause" ) const ( @@ -213,8 +212,7 @@ func unsafeGitHubInputRune(character rune) bool { } func updateGitHubSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) (bool, error) { - var project model.PagesProject - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil { + if _, err := repository.LockPagesProjectByIDTx(tx, projectID); err != nil { return false, err } existing, hasExisting, err := loadProjectSourceForUpdate(tx, projectID) @@ -231,16 +229,15 @@ func updateGitHubSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) if !githubSourceConfigChanged(existing, config) { return false, nil } - var runtime model.PagesProjectSourceRuntime - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("source_id = ?", existing.ID).First(&runtime).Error; err != nil { + runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, existing.ID) + if err != nil { return false, err } identityChanged := existing.SourceIdentity != config.SourceIdentity - if err := tx.Model(existing).Updates(githubSourceUpdates(config, existing.ConfigVersion+1)).Error; err != nil { + if err := repository.UpdatePagesProjectSourceTx(tx, existing, githubSourceUpdates(config, existing.ConfigVersion+1)); err != nil { return false, err } - if err := resetRuntimeAfterGitHubUpdate(tx, &runtime, config, identityChanged); err != nil { + if err := resetRuntimeAfterGitHubUpdate(tx, runtime, config, identityChanged); err != nil { return false, err } return true, nil @@ -259,7 +256,7 @@ func createGitHubSourceTx(tx *gorm.DB, projectID uint, config githubSourceConfig ConfigVersion: 1, SourceIdentity: config.SourceIdentity, } - if err := tx.Create(source).Error; err != nil { + if err := repository.CreatePagesProjectSourceTx(tx, source); err != nil { return err } runtime := &model.PagesProjectSourceRuntime{SourceID: source.ID, SyncStatus: pagesSourceStatusIdle} @@ -267,7 +264,7 @@ func createGitHubSourceTx(tx *gorm.DB, projectID uint, config githubSourceConfig next := nextGitHubCheckAt(time.Now(), source.ID, config.CheckInterval) runtime.NextCheckAt = &next } - return tx.Create(runtime).Error + return repository.CreatePagesProjectSourceRuntimeTx(tx, runtime) } func githubSourceUpdates(config githubSourceConfig, version int) map[string]any { @@ -308,7 +305,7 @@ func resetRuntimeAfterGitHubUpdate( next := nextGitHubCheckAt(time.Now(), runtime.SourceID, config.CheckInterval) nextCheckAt = &next } - return tx.Model(runtime).Update("next_check_at", nextCheckAt).Error + return repository.UpdatePagesProjectSourceRuntimeFieldTx(tx, runtime, "next_check_at", nextCheckAt) } func nextGitHubCheckAt(now time.Time, sourceID uint, intervalMinutes int) time.Time { @@ -323,8 +320,8 @@ func markInitialCheckDispatchFailed(ctx context.Context, sourceID uint, configVe sourceRuntimeColumnSyncStatus: pagesSourceStatusFailed, sourceRuntimeColumnLastError: errPagesSourceInitialCheckWarning, } - var source model.PagesProjectSource - if err := db.DB(ctx).Where("id = ? AND config_version = ?", sourceID, configVersion).First(&source).Error; err != nil { + source, err := repository.GetPagesProjectSourceByIDAndConfigVersion(ctx, sourceID, configVersion) + if err != nil { if !errors.Is(err, gorm.ErrRecordNotFound) { logger.ErrorF(ctx, "[PagesSource] load initial check source snapshot failed: source_id=%d error=%v", sourceID, err) } @@ -335,12 +332,7 @@ func markInitialCheckDispatchFailed(ctx context.Context, sourceID uint, configVe updates["next_check_at"] = &next } now := time.Now() - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). - Where("source_id = ?", sourceID). - Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). - Where("EXISTS (SELECT 1 FROM of_pages_project_sources source WHERE source.id = ? AND source.config_version = ?)", sourceID, configVersion). - Updates(updates) - if result.Error != nil { - logger.ErrorF(ctx, "[PagesSource] mark initial check dispatch failure: source_id=%d error=%v", sourceID, result.Error) + if _, err := repository.MarkPagesSourceInitialCheckDispatchFailed(ctx, sourceID, configVersion, now, updates); err != nil { + logger.ErrorF(ctx, "[PagesSource] mark initial check dispatch failure: source_id=%d error=%v", sourceID, err) } } diff --git a/internal/apps/openflare/pages/github_source_action.go b/internal/apps/openflare/pages/github_source_action.go index 93e8d950..755019c9 100644 --- a/internal/apps/openflare/pages/github_source_action.go +++ b/internal/apps/openflare/pages/github_source_action.go @@ -13,14 +13,13 @@ import ( "strings" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/integration/githubrelease" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/pagesarchive" "gorm.io/gorm" - "gorm.io/gorm/clause" ) const githubSourceDetailProvider = "github" @@ -201,7 +200,7 @@ func finishGitHubCheckNotModified( ) (string, string, error) { var revision string var status string - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { runtime, now, err := lockOwnedSourceRuntime(tx, snapshot) if err != nil { return err @@ -211,7 +210,7 @@ func finishGitHubCheckNotModified( updates := githubCheckTerminalUpdates(snapshot, now, result.RetryAt) updates["etag"] = result.ETag updates[sourceRuntimeColumnSyncStatus] = status - return tx.Model(runtime).Updates(updates).Error + return repository.UpdatePagesProjectSourceRuntimeTx(tx, runtime, updates) }) return revision, status, err } @@ -223,7 +222,7 @@ func finishGitHubCheckTarget( target *githubSourceTarget, ) (string, error) { status := pagesSourceStatusIdle - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { runtime, now, err := lockOwnedSourceRuntime(tx, snapshot) if err != nil { return err @@ -234,7 +233,7 @@ func finishGitHubCheckTarget( updates["last_seen_revision"] = target.Revision updates["last_seen_detail"] = target.DetailJSON updates[sourceRuntimeColumnSyncStatus] = status - return tx.Model(runtime).Updates(updates).Error + return repository.UpdatePagesProjectSourceRuntimeTx(tx, runtime, updates) }) return status, err } @@ -273,9 +272,8 @@ func lockOwnedSourceRuntime( tx *gorm.DB, snapshot *sourceExecutionSnapshot, ) (*model.PagesProjectSourceRuntime, time.Time, error) { - var runtime model.PagesProjectSourceRuntime - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("source_id = ?", snapshot.SourceID).First(&runtime).Error; err != nil { + runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, snapshot.SourceID) + if err != nil { return nil, time.Time{}, err } now := time.Now() @@ -283,7 +281,7 @@ func lockOwnedSourceRuntime( !runtime.LeaseExpiresAt.After(now) { return nil, time.Time{}, errSourceFinalFence } - return &runtime, now, nil + return runtime, now, nil } func failGitHubCheckLease( @@ -309,13 +307,13 @@ func failGitHubCheckLease( } else { updates[sourceRuntimeColumnNextCheckAt] = nil } - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). - Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now). - Updates(updates) - if result.Error != nil { - return result.Error + rows, err := repository.UpdatePagesSourceRuntimeByActiveLease( + ctx, snapshot.SourceID, snapshot.LeaseToken, now, updates, + ) + if err != nil { + return err } - if result.RowsAffected != 1 { + if rows != 1 { return errSourceFinalFence } return nil @@ -338,11 +336,11 @@ func targetRuntimeStatus( } func preflightGitHubSyncConfirmation(ctx context.Context, sourceID uint, confirmedRevision string) error { - var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { + runtime, err := repository.GetPagesProjectSourceRuntimeBySourceID(ctx, sourceID) + if err != nil { return err } - replacement := sourceHasSameReleaseReplacement(&runtime) + replacement := sourceHasSameReleaseReplacement(runtime) if replacement && confirmedRevision == "" { return errors.New(errPagesSourceConfirmationNeeded) } @@ -673,13 +671,13 @@ func releaseGitHubSyncWithoutActivation( sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseExpiresAt: nil, } - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). - Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now). - Updates(updates) - if result.Error != nil { - return result.Error + rows, err := repository.UpdatePagesSourceRuntimeByActiveLease( + ctx, snapshot.SourceID, snapshot.LeaseToken, now, updates, + ) + if err != nil { + return err } - if result.RowsAffected != 1 { + if rows != 1 { return errSourceFinalFence } return nil diff --git a/internal/apps/openflare/pages/github_source_test.go b/internal/apps/openflare/pages/github_source_test.go index 58e4f75d..d773a2c9 100644 --- a/internal/apps/openflare/pages/github_source_test.go +++ b/internal/apps/openflare/pages/github_source_test.go @@ -16,7 +16,10 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/integration/githubrelease" "github.com/Rain-kl/Wavelet/internal/model" "github.com/hibiken/asynq" @@ -100,7 +103,7 @@ func TestGitHubSourceValidationNormalizationAndProviderSwitch(t *testing.T) { if result.CheckTask == nil || result.CheckTask.Action != sourceActionCheck || result.Warning != "" { t.Errorf("UpdateSourceAs(GitHub) result = %+v, want initial check receipt without warning", result) } - execution, err := model.GetTaskExecutionByTaskID(ctx, result.CheckTask.TaskID) + execution, err := repository.GetTaskExecutionByTaskID(ctx, result.CheckTask.TaskID) if err != nil { t.Fatalf("GetTaskExecutionByTaskID(%q) error = %v, want nil", result.CheckTask.TaskID, err) } @@ -171,6 +174,11 @@ func TestGitHubSourceValidationNormalizationAndProviderSwitch(t *testing.T) { func TestGitHubSourceSaveSurvivesInitialCheckDispatchFailure(t *testing.T) { ctx := setupPagesSourceTest(t) + // Isolate from other tests that may leave a global Asynq client registered. + previousClient := task.AsynqClient + task.AsynqClient = nil + t.Cleanup(func() { task.AsynqClient = previousClient }) + project := mustCreatePagesSourceProject(t, ctx, "github-dispatch-warning") result, err := UpdateSourceAs(ctx, project.ID, SourceUpdateInput{ SourceType: PagesSourceTypeGitHubRelease, @@ -679,7 +687,7 @@ func TestGitHubSyncActivatesWithMetadataRevisionAndPackageChecksum(t *testing.T) if outcome == nil || outcome.Deployment == nil || outcome.Stale { t.Fatalf("syncGitHubSource() = %+v, want active deployment", outcome) } - deployment, err := model.GetPagesDeploymentByID(ctx, outcome.Deployment.ID) + deployment, err := repository.GetPagesDeploymentByID(ctx, outcome.Deployment.ID) if err != nil { t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", outcome.Deployment.ID, err) } @@ -878,7 +886,7 @@ func TestGitHubSyncRejectsStaleConfirmationWithoutChangingActive(t *testing.T) { if err == nil || err.Error() != errPagesSourceConfirmationStale { t.Errorf("syncGitHubSource(stale confirmation) error = %v, want %q", err, errPagesSourceConfirmationStale) } - storedProject, loadErr := model.GetPagesProjectByID(ctx, project.ID) + storedProject, loadErr := repository.GetPagesProjectByID(ctx, project.ID) if loadErr != nil { t.Fatalf("GetPagesProjectByID() error = %v, want nil", loadErr) } diff --git a/internal/apps/openflare/pages/helpers.go b/internal/apps/openflare/pages/helpers.go index b6020a6f..8e017ea2 100644 --- a/internal/apps/openflare/pages/helpers.go +++ b/internal/apps/openflare/pages/helpers.go @@ -19,7 +19,6 @@ import ( "strings" "github.com/Rain-kl/Wavelet/internal/apps/upload" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" @@ -325,11 +324,10 @@ func removePagesUploadIfUnreferenced(ctx context.Context, projectID uint, upload if uploadID == 0 { return nil } - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { if projectID != 0 { - var project model.PagesProject - projectErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error - if projectErr != nil && !errors.Is(projectErr, gorm.ErrRecordNotFound) { + if _, projectErr := repository.LockPagesProjectByIDTx(tx, projectID); projectErr != nil && + !errors.Is(projectErr, gorm.ErrRecordNotFound) { return projectErr } } diff --git a/internal/apps/openflare/pages/logics.go b/internal/apps/openflare/pages/logics.go index 13b12eab..007aeeaa 100644 --- a/internal/apps/openflare/pages/logics.go +++ b/internal/apps/openflare/pages/logics.go @@ -17,8 +17,8 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/apps/upload" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/pagesarchive" "gorm.io/gorm" @@ -101,7 +101,7 @@ type View struct { // ListProjects 列出全部 Pages 项目。 func ListProjects(ctx context.Context) ([]View, error) { - projects, err := model.ListPagesProjects(ctx) + projects, err := repository.ListPagesProjects(ctx) if err != nil { return nil, err } @@ -118,7 +118,7 @@ func ListProjects(ctx context.Context) ([]View, error) { // GetProject 获取 Pages 项目详情。 func GetProject(ctx context.Context, id uint) (*View, error) { - project, err := model.GetPagesProjectByID(ctx, id) + project, err := repository.GetPagesProjectByID(ctx, id) if err != nil { return nil, err } @@ -131,7 +131,7 @@ func CreateProject(ctx context.Context, input Input) (*View, error) { if err != nil { return nil, err } - if err = model.CreatePagesProjectRecord(ctx, project); err != nil { + if err = repository.CreatePagesProjectRecord(ctx, project); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errPagesSlugExists) } @@ -143,7 +143,7 @@ func CreateProject(ctx context.Context, input Input) (*View, error) { // UpdateProject 更新 Pages 项目。 func UpdateProject(ctx context.Context, id uint, input Input) (*View, error) { var project *model.PagesProject - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { var existing model.PagesProject if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&existing, id).Error; err != nil { return err @@ -177,10 +177,7 @@ func UpdateProject(ctx context.Context, id uint, input Input) (*View, error) { } if contentConfigChanged { updates["content_config_version"] = existing.ContentConfigVersion + 1 - var source model.PagesProjectSource - sourceErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("project_id = ?", existing.ID). - First(&source).Error + source, sourceErr := repository.LockPagesProjectSourceByProjectIDTx(tx, existing.ID) if sourceErr != nil && !errors.Is(sourceErr, gorm.ErrRecordNotFound) { return sourceErr } @@ -220,12 +217,12 @@ func ensureDeploymentEntry(conn *gorm.DB, deploymentID uint, rootDir, entryFile // DeleteProject 删除 Pages 项目。 func DeleteProject(ctx context.Context, id uint) error { - project, err := model.GetPagesProjectByID(ctx, id) + project, err := repository.GetPagesProjectByID(ctx, id) if err != nil { return err } var deployments []model.PagesDeployment - err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err = repository.WithPagesTx(ctx, func(tx *gorm.DB) error { var lockedProject model.PagesProject if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedProject, project.ID).Error; err != nil { return err @@ -241,25 +238,19 @@ func DeleteProject(ctx context.Context, id uint) error { return errors.New(errPagesDeleteReferenced) } } - var source model.PagesProjectSource - sourceErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("project_id = ?", project.ID). - First(&source).Error + source, sourceErr := repository.LockPagesProjectSourceByProjectIDTx(tx, project.ID) if sourceErr != nil && !errors.Is(sourceErr, gorm.ErrRecordNotFound) { return sourceErr } if sourceErr == nil { - var runtime model.PagesProjectSourceRuntime - runtimeErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("source_id = ?", source.ID). - First(&runtime).Error - if runtimeErr != nil && !errors.Is(runtimeErr, gorm.ErrRecordNotFound) { + if _, runtimeErr := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID); runtimeErr != nil && + !errors.Is(runtimeErr, gorm.ErrRecordNotFound) { return runtimeErr } - if err := tx.Where("source_id = ?", source.ID).Delete(&model.PagesProjectSourceRuntime{}).Error; err != nil { + if err := repository.DeletePagesProjectSourceRuntimeBySourceIDTx(tx, source.ID); err != nil { return err } - if err := tx.Delete(&source).Error; err != nil { + if err := repository.DeletePagesProjectSourceTx(tx, source); err != nil { return err } } @@ -291,10 +282,10 @@ func DeleteProject(ctx context.Context, id uint) error { // ListProjectDeployments 列出项目的全部部署。 func ListProjectDeployments(ctx context.Context, projectID uint) ([]DeploymentView, error) { - if _, err := model.GetPagesProjectByID(ctx, projectID); err != nil { + if _, err := repository.GetPagesProjectByID(ctx, projectID); err != nil { return nil, err } - deployments, err := model.ListPagesDeployments(ctx, projectID) + deployments, err := repository.ListPagesDeployments(ctx, projectID) if err != nil { return nil, err } @@ -307,10 +298,10 @@ func ListProjectDeployments(ctx context.Context, projectID uint) ([]DeploymentVi // ListDeploymentFiles 列出部署文件清单。 func ListDeploymentFiles(ctx context.Context, deploymentID uint) ([]DeploymentFileView, error) { - if _, err := model.GetPagesDeploymentByID(ctx, deploymentID); err != nil { + if _, err := repository.GetPagesDeploymentByID(ctx, deploymentID); err != nil { return nil, err } - files, err := model.ListPagesDeploymentFiles(ctx, deploymentID) + files, err := repository.ListPagesDeploymentFiles(ctx, deploymentID) if err != nil { return nil, err } @@ -330,7 +321,7 @@ func ListDeploymentFiles(ctx context.Context, deploymentID uint) ([]DeploymentFi // UploadDeployment 上传 Pages 部署包(本地 multipart 文件)。 func UploadDeployment(ctx context.Context, projectID uint, fileHeader *multipart.FileHeader, createdBy string) (*DeploymentView, error) { - project, err := model.GetPagesProjectByID(ctx, projectID) + project, err := repository.GetPagesProjectByID(ctx, projectID) if err != nil { return nil, err } @@ -364,7 +355,7 @@ type UploadFromURLInput struct { // UploadDeploymentFromURL downloads a package from url and creates a deployment. func UploadDeploymentFromURL(ctx context.Context, projectID uint, rawURL string, createdBy string) (*DeploymentView, error) { - project, err := model.GetPagesProjectByID(ctx, projectID) + project, err := repository.GetPagesProjectByID(ctx, projectID) if err != nil { return nil, err } @@ -438,7 +429,7 @@ func createDeploymentFromTempPackage( } }() deployment := &model.PagesDeployment{} - err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err = repository.WithPagesTx(ctx, func(tx *gorm.DB) error { var lockedProject model.PagesProject if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedProject, project.ID).Error; err != nil { return err @@ -539,7 +530,7 @@ func pruneProjectDeploymentHistory(ctx context.Context, projectID uint, keepCoun // Returns the number of deployments deleted from the database. func pruneProjectDeploymentHistoryOnce(ctx context.Context, projectID uint, keepCount int, preserveCandidateID uint) (int, error) { var deletedDeployments []model.PagesDeployment - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { var project model.PagesProject if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil { return fmt.Errorf("load pages project: %w", err) @@ -705,7 +696,7 @@ type deploymentActivationSource struct { } func ensureActivationDeploymentUpload(ctx context.Context, projectID uint, deploymentID uint) error { - deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) + deployment, err := repository.GetPagesDeploymentByID(ctx, deploymentID) if err != nil { return err } @@ -725,7 +716,7 @@ func activateDeploymentTransaction( now time.Time, ) (deploymentActivationAudit, error) { audit := deploymentActivationAudit{} - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { var project model.PagesProject if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil { return err @@ -755,23 +746,18 @@ func activateDeploymentTransaction( } func lockDeploymentActivationSource(tx *gorm.DB, projectID uint) (*deploymentActivationSource, error) { - var source model.PagesProjectSource - err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("project_id = ?", projectID). - First(&source).Error + source, err := repository.LockPagesProjectSourceByProjectIDTx(tx, projectID) if errors.Is(err, gorm.ErrRecordNotFound) { return nil, nil } if err != nil { return nil, err } - var runtime model.PagesProjectSourceRuntime - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("source_id = ?", source.ID). - First(&runtime).Error; err != nil { + runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID) + if err != nil { return nil, err } - return &deploymentActivationSource{Source: &source, Runtime: &runtime}, nil + return &deploymentActivationSource{Source: source, Runtime: runtime}, nil } func loadDeploymentActivationTarget( @@ -821,10 +807,10 @@ func fenceDeploymentActivationSource( audit.SourceType = state.Source.SourceType audit.SourceIdentity = state.Source.SourceIdentity audit.AutoDisabled = state.Source.AutoUpdateEnabled - if err := tx.Model(state.Source).Updates(map[string]any{ + if err := repository.UpdatePagesProjectSourceTx(tx, state.Source, map[string]any{ sourceColumnConfigVersion: state.Source.ConfigVersion + 1, sourceColumnAutoUpdateEnabled: false, - }).Error; err != nil { + }); err != nil { return err } if deployment.SourceIdentity != nil && *deployment.SourceIdentity == state.Source.SourceIdentity && @@ -835,13 +821,13 @@ func fenceDeploymentActivationSource( state.Runtime.LastAppliedRevision = "" state.Runtime.LastAppliedDetail = "" } - return tx.Model(state.Runtime).Updates(map[string]any{ + return repository.UpdatePagesProjectSourceRuntimeTx(tx, state.Runtime, map[string]any{ "last_applied_revision": state.Runtime.LastAppliedRevision, "last_applied_detail": state.Runtime.LastAppliedDetail, sourceRuntimeColumnSyncStatus: normalizedSourceRuntimeStatus(state.Runtime), sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseExpiresAt: nil, - }).Error + }) } func switchActiveDeploymentTx( @@ -867,7 +853,7 @@ func switchActiveDeploymentTx( // GetDeploymentPackageHash returns the upload SHA-256 hash of the deployment package. // Prefer GetProjectLatestPackageHash for Agent latest-pointer pulls. func GetDeploymentPackageHash(ctx context.Context, deploymentID uint) (string, error) { - deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) + deployment, err := repository.GetPagesDeploymentByID(ctx, deploymentID) if err != nil { return "", err } @@ -893,7 +879,7 @@ func GetProjectLatestPackageHash(ctx context.Context, projectID uint) (uint, str // OpenDeploymentPackage opens the deployment artifact from the upload storage framework. func OpenDeploymentPackage(ctx context.Context, deploymentID uint) (DeploymentPackage, error) { - deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) + deployment, err := repository.GetPagesDeploymentByID(ctx, deploymentID) if err != nil { return DeploymentPackage{}, err } @@ -926,7 +912,7 @@ func resolveProjectActiveDeploymentForAgent(ctx context.Context, projectID uint) if projectID == 0 { return nil, errors.New(errPagesProjectNotFound) } - project, err := model.GetPagesProjectByID(ctx, projectID) + project, err := repository.GetPagesProjectByID(ctx, projectID) if err != nil { return nil, err } @@ -939,7 +925,7 @@ func resolveProjectActiveDeploymentForAgent(ctx context.Context, projectID uint) if err := ensureProjectInActiveConfig(ctx, project.ID); err != nil { return nil, err } - deployment, err := model.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) + deployment, err := repository.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) if err != nil { return nil, err } @@ -957,7 +943,7 @@ func deploymentPackageHash(ctx context.Context, deployment *model.PagesDeploymen if err := ensureDeploymentUploadRecord(ctx, deployment); err != nil { return "", err } - reloaded, err := model.GetPagesDeploymentByID(ctx, deployment.ID) + reloaded, err := repository.GetPagesDeploymentByID(ctx, deployment.ID) if err != nil { return "", err } @@ -1012,7 +998,7 @@ func ensureDeploymentUploadRecord(ctx context.Context, deployment *model.PagesDe if deployment.UploadID > 0 { return nil } - project, err := model.GetPagesProjectByID(ctx, deployment.ProjectID) + project, err := repository.GetPagesProjectByID(ctx, deployment.ProjectID) if err != nil { return err } @@ -1089,7 +1075,7 @@ func attachLegacyDeploymentUpload( uploadID uint64, ) (uint64, error) { winnerUploadID := uint64(0) - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { var err error winnerUploadID, err = attachLegacyDeploymentUploadTx(tx, projectID, deploymentID, uploadID) return err @@ -1146,11 +1132,11 @@ func attachLegacyDeploymentUploadTx( // it is the project's current active deployment and the project is used by the // active main config. func ensureDeploymentInActiveSnapshot(ctx context.Context, deploymentID uint) error { - deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) + deployment, err := repository.GetPagesDeploymentByID(ctx, deploymentID) if err != nil { return err } - project, err := model.GetPagesProjectByID(ctx, deployment.ProjectID) + project, err := repository.GetPagesProjectByID(ctx, deployment.ProjectID) if err != nil { return err } @@ -1163,7 +1149,7 @@ func ensureDeploymentInActiveSnapshot(ctx context.Context, deploymentID uint) er // ensureProjectInActiveConfig checks that the Pages project is referenced by at // least one pages route in the active main config snapshot. func ensureProjectInActiveConfig(ctx context.Context, projectID uint) error { - version, err := model.GetActiveConfigVersion(ctx) + version, err := repository.GetActiveConfigVersion(ctx) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return errors.New(errPagesPackageNotInActiveConfig) @@ -1192,7 +1178,7 @@ func ensureProjectInActiveConfig(ctx context.Context, projectID uint) error { } if route.PagesDeployment.DeploymentID != 0 { // Historical snapshot may only pin deployment_id. - snapDeployment, snapErr := model.GetPagesDeploymentByID(ctx, route.PagesDeployment.DeploymentID) + snapDeployment, snapErr := repository.GetPagesDeploymentByID(ctx, route.PagesDeployment.DeploymentID) if snapErr != nil { continue } @@ -1239,7 +1225,7 @@ func parseSnapshotRoutes(snapshotJSON string) ([]snapshotRouteRef, error) { // DeleteDeployment 删除 Pages 部署。 func DeleteDeployment(ctx context.Context, projectID uint, deploymentID uint) error { var removed model.PagesDeployment - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { var project model.PagesProject if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil { return err @@ -1354,13 +1340,13 @@ func buildProjectView(ctx context.Context, project *model.PagesProject) (*View, CreatedAt: project.CreatedAt, UpdatedAt: project.UpdatedAt, } - count, err := model.CountPagesDeploymentsByProjectID(ctx, project.ID) + count, err := repository.CountPagesDeploymentsByProjectID(ctx, project.ID) if err != nil { return nil, err } view.DeploymentCount = count if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID != 0 { - deployment, err := model.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) + deployment, err := repository.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) if err == nil { active := buildDeploymentView(deployment) view.ActiveDeployment = &active diff --git a/internal/apps/openflare/pages/logics_test.go b/internal/apps/openflare/pages/logics_test.go index 6300e939..6d282a5a 100644 --- a/internal/apps/openflare/pages/logics_test.go +++ b/internal/apps/openflare/pages/logics_test.go @@ -243,7 +243,7 @@ func TestUpdateProjectValidatesActiveDeploymentEntry(t *testing.T) { require.Error(t, err) assert.Contains(t, err.Error(), errPagesEntryFileMissing) - stored, err := model.GetPagesProjectByID(ctx, project.ID) + stored, err := repository.GetPagesProjectByID(ctx, project.ID) require.NoError(t, err) assert.Equal(t, "dist", stored.RootDir) assert.Equal(t, "index.html", stored.EntryFile) @@ -291,7 +291,7 @@ func TestUploadDeploymentStoresPackageInUploadFramework(t *testing.T) { require.NoError(t, err) assert.NotZero(t, deployment.UploadID) - storedDeployment, err := model.GetPagesDeploymentByID(ctx, deployment.ID) + storedDeployment, err := repository.GetPagesDeploymentByID(ctx, deployment.ID) require.NoError(t, err) assert.NotZero(t, storedDeployment.UploadID) assert.Empty(t, storedDeployment.ArtifactPath) @@ -371,7 +371,7 @@ func TestOpenDeploymentPackageHydratesLegacyArtifactPath(t *testing.T) { require.Len(t, reader.File, 1) assert.Equal(t, "index.html", reader.File[0].Name) - storedDeployment, err := model.GetPagesDeploymentByID(ctx, deployment.ID) + storedDeployment, err := repository.GetPagesDeploymentByID(ctx, deployment.ID) require.NoError(t, err) assert.NotZero(t, storedDeployment.UploadID) assert.Empty(t, storedDeployment.ArtifactPath) @@ -607,7 +607,7 @@ func TestPruneProjectDeploymentHistory(t *testing.T) { ids = append(ids, deployment.ID) } // After 3 uploads with keep=2 and no active: only 2 newest remain. - deployments, err := model.ListPagesDeployments(ctx, project.ID) + deployments, err := repository.ListPagesDeployments(ctx, project.ID) require.NoError(t, err) require.Len(t, deployments, 2) assert.Equal(t, ids[2], deployments[0].ID) @@ -622,11 +622,11 @@ func TestPruneProjectDeploymentHistory(t *testing.T) { })), "root") require.NoError(t, err) - deployments, err = model.ListPagesDeployments(ctx, project.ID) + deployments, err = repository.ListPagesDeployments(ctx, project.ID) require.NoError(t, err) require.Len(t, deployments, 2, "must be at most N=2, not active+N newest") - storedProject, err := model.GetPagesProjectByID(ctx, project.ID) + storedProject, err := repository.GetPagesProjectByID(ctx, project.ID) require.NoError(t, err) require.NotNil(t, storedProject.ActiveDeploymentID) assert.Equal(t, ids[1], *storedProject.ActiveDeploymentID) @@ -666,7 +666,7 @@ func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) { "index.html": "v2", })), "user:1") require.NoError(t, err) - deployments, err := model.ListPagesDeployments(ctx, project.ID) + deployments, err := repository.ListPagesDeployments(ctx, project.ID) require.NoError(t, err) require.Len(t, deployments, 2) @@ -674,7 +674,7 @@ func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) { "index.html": "v3", })), "user:1") require.NoError(t, err) - deployments, err = model.ListPagesDeployments(ctx, project.ID) + deployments, err = repository.ListPagesDeployments(ctx, project.ID) require.NoError(t, err) require.Len(t, deployments, 2) kept := map[uint]bool{} @@ -690,7 +690,7 @@ func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) { _, err = ActivateDeployment(ctx, project.ID, newCandidate.ID) require.NoError(t, err) - deployments, err = model.ListPagesDeployments(ctx, project.ID) + deployments, err = repository.ListPagesDeployments(ctx, project.ID) require.NoError(t, err) require.Len(t, deployments, 1) assert.Equal(t, newCandidate.ID, deployments[0].ID) @@ -723,7 +723,7 @@ func TestPruneUsesLockTimeNewestCandidateInsteadOfStaleCaller(t *testing.T) { deleted, err := pruneProjectDeploymentHistoryOnce(ctx, project.ID, 1, staleCandidate.ID) require.NoError(t, err) assert.Equal(t, 1, deleted) - deployments, err := model.ListPagesDeployments(ctx, project.ID) + deployments, err := repository.ListPagesDeployments(ctx, project.ID) require.NoError(t, err) require.Len(t, deployments, 2) kept := map[uint]bool{} @@ -762,7 +762,7 @@ func TestDeleteDeploymentAndProjectSoftDeleteUnreferencedArtifacts(t *testing.T) var firstUpload model.Upload require.NoError(t, db.DB(ctx).First(&firstUpload, first.UploadID).Error) assert.Equal(t, model.UploadStatusDeleted, firstUpload.Status) - _, err = model.GetPagesProjectByID(ctx, project.ID) + _, err = repository.GetPagesProjectByID(ctx, project.ID) assert.Error(t, err) } diff --git a/internal/apps/openflare/pages/package_metadata.go b/internal/apps/openflare/pages/package_metadata.go index 0e99c358..76caef92 100644 --- a/internal/apps/openflare/pages/package_metadata.go +++ b/internal/apps/openflare/pages/package_metadata.go @@ -9,8 +9,9 @@ import ( "fmt" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/upload" - "github.com/Rain-kl/Wavelet/internal/model" ) // ProjectLatestPackageMetadata describes the active package limits published to Agents. @@ -33,7 +34,7 @@ func GetProjectLatestPackageMetadata(ctx context.Context, projectID uint) (*Proj if err := ensureDeploymentUploadRecord(ctx, deployment); err != nil { return nil, err } - deployment, err = model.GetPagesDeploymentByID(ctx, deployment.ID) + deployment, err = repository.GetPagesDeploymentByID(ctx, deployment.ID) if err != nil { return nil, err } diff --git a/internal/apps/openflare/pages/rebind.go b/internal/apps/openflare/pages/rebind.go index d26d34d8..b55bb3d3 100644 --- a/internal/apps/openflare/pages/rebind.go +++ b/internal/apps/openflare/pages/rebind.go @@ -10,6 +10,8 @@ import ( "path" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" "gorm.io/gorm" @@ -156,7 +158,7 @@ func loadActivePagesProject(ctx context.Context, projectID uint, siteName string if siteName == "" { siteName = "pages" } - project, err := model.GetPagesProjectByID(ctx, projectID) + project, err := repository.GetPagesProjectByID(ctx, projectID) if err != nil { if errorsIsNotFound(err) { return nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目不存在", siteName) @@ -169,7 +171,7 @@ func loadActivePagesProject(ctx context.Context, projectID uint, siteName string if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 { return nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目没有激活部署", siteName) } - activeDeployment, err := model.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) + activeDeployment, err := repository.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) if err != nil { if errorsIsNotFound(err) { return nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 激活部署不存在", siteName) diff --git a/internal/apps/openflare/pages/routers_test.go b/internal/apps/openflare/pages/routers_test.go index c73a55c5..2ad9eda8 100644 --- a/internal/apps/openflare/pages/routers_test.go +++ b/internal/apps/openflare/pages/routers_test.go @@ -13,6 +13,8 @@ import ( "strconv" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/model" "github.com/gin-gonic/gin" @@ -48,7 +50,7 @@ func TestUploadDeploymentHandlerRecordsCurrentUserActor(t *testing.T) { UploadDeploymentHandler(c) assert.Equal(t, http.StatusOK, recorder.Code) - deployments, err := model.ListPagesDeployments(t.Context(), project.ID) + deployments, err := repository.ListPagesDeployments(t.Context(), project.ID) require.NoError(t, err) require.Len(t, deployments, 1) assert.Equal(t, "user:42", deployments[0].CreatedBy) @@ -83,7 +85,7 @@ func TestUploadDeploymentFromURLHandlerRecordsCurrentUserActor(t *testing.T) { UploadDeploymentFromURLHandler(c) assert.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) - deployments, err := model.ListPagesDeployments(t.Context(), project.ID) + deployments, err := repository.ListPagesDeployments(t.Context(), project.ID) require.NoError(t, err) require.Len(t, deployments, 1) assert.Equal(t, "user:77", deployments[0].CreatedBy) diff --git a/internal/apps/openflare/pages/source.go b/internal/apps/openflare/pages/source.go index 3b206821..af95202b 100644 --- a/internal/apps/openflare/pages/source.go +++ b/internal/apps/openflare/pages/source.go @@ -15,10 +15,9 @@ import ( "strings" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "gorm.io/gorm" - "gorm.io/gorm/clause" ) const ( @@ -120,7 +119,7 @@ type remoteSourceConfig struct { // GetSource returns the current persisted source or a manual discriminator. func GetSource(ctx context.Context, projectID uint) (*SourceView, error) { - if _, err := model.GetPagesProjectByID(ctx, projectID); err != nil { + if _, err := repository.GetPagesProjectByID(ctx, projectID); err != nil { return nil, err } source, runtime, err := loadSourceByProject(ctx, projectID) @@ -156,7 +155,7 @@ func UpdateSourceAs( changed := false var persistedSource model.PagesProjectSource - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { var err error switch strings.TrimSpace(input.SourceType) { case PagesSourceTypeRemoteURL: @@ -169,7 +168,12 @@ func UpdateSourceAs( if err != nil || !changed || strings.TrimSpace(input.SourceType) != PagesSourceTypeGitHubRelease { return err } - return tx.Where("project_id = ?", projectID).First(&persistedSource).Error + source, loadErr := repository.LockPagesProjectSourceByProjectIDTx(tx, projectID) + if loadErr != nil { + return loadErr + } + persistedSource = *source + return nil }) if err != nil { return nil, err @@ -193,8 +197,7 @@ func UpdateSourceAs( } func updateRemoteSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) (bool, error) { - var project model.PagesProject - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil { + if _, err := repository.LockPagesProjectByIDTx(tx, projectID); err != nil { return false, err } existing, hasExisting, err := loadProjectSourceForUpdate(tx, projectID) @@ -212,17 +215,14 @@ func updateRemoteSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) } func loadProjectSourceForUpdate(tx *gorm.DB, projectID uint) (*model.PagesProjectSource, bool, error) { - var source model.PagesProjectSource - err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("project_id = ?", projectID). - First(&source).Error + source, err := repository.LockPagesProjectSourceByProjectIDTx(tx, projectID) if errors.Is(err, gorm.ErrRecordNotFound) { - return &source, false, nil + return &model.PagesProjectSource{}, false, nil } if err != nil { return nil, false, err } - return &source, true, nil + return source, true, nil } func buildRemoteSourceConfig( @@ -256,13 +256,13 @@ func createRemoteSourceTx(tx *gorm.DB, projectID uint, config remoteSourceConfig ConfigVersion: 1, SourceIdentity: config.Identity, } - if err := tx.Create(source).Error; err != nil { + if err := repository.CreatePagesProjectSourceTx(tx, source); err != nil { return err } - return tx.Create(&model.PagesProjectSourceRuntime{ + return repository.CreatePagesProjectSourceRuntimeTx(tx, &model.PagesProjectSourceRuntime{ SourceID: source.ID, SyncStatus: pagesSourceStatusIdle, - }).Error + }) } func updateExistingRemoteSourceTx( @@ -273,14 +273,12 @@ func updateExistingRemoteSourceTx( if !remoteSourceConfigChanged(existing, config) { return false, nil } - var runtime model.PagesProjectSourceRuntime - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("source_id = ?", existing.ID). - First(&runtime).Error; err != nil { + runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, existing.ID) + if err != nil { return false, err } identityChanged := existing.SourceIdentity != config.Identity - if err := tx.Model(existing).Updates(map[string]any{ + if err := repository.UpdatePagesProjectSourceTx(tx, existing, map[string]any{ "source_type": PagesSourceTypeRemoteURL, "remote_url": config.URL, "allow_insecure": config.AllowInsecure, @@ -292,10 +290,10 @@ func updateExistingRemoteSourceTx( "check_interval_minutes": 0, sourceColumnConfigVersion: existing.ConfigVersion + 1, "source_identity": config.Identity, - }).Error; err != nil { + }); err != nil { return false, err } - return true, resetRuntimeAfterSourceUpdate(tx, &runtime, identityChanged) + return true, resetRuntimeAfterSourceUpdate(tx, runtime, identityChanged) } func remoteSourceConfigChanged(existing *model.PagesProjectSource, config remoteSourceConfig) bool { @@ -312,31 +310,25 @@ func remoteSourceConfigChanged(existing *model.PagesProjectSource, config remote // DeleteSource idempotently switches a project back to manual mode. func DeleteSource(ctx context.Context, projectID uint) (*SourceView, error) { - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { - var project model.PagesProject - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { + if _, err := repository.LockPagesProjectByIDTx(tx, projectID); err != nil { return err } - var source model.PagesProjectSource - err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("project_id = ?", projectID). - First(&source).Error + source, err := repository.LockPagesProjectSourceByProjectIDTx(tx, projectID) if errors.Is(err, gorm.ErrRecordNotFound) { return nil } if err != nil { return err } - var runtime model.PagesProjectSourceRuntime - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("source_id = ?", source.ID). - First(&runtime).Error; err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { + if _, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID); err != nil && + !errors.Is(err, gorm.ErrRecordNotFound) { return err } - if err := tx.Where("source_id = ?", source.ID).Delete(&model.PagesProjectSourceRuntime{}).Error; err != nil { + if err := repository.DeletePagesProjectSourceRuntimeBySourceIDTx(tx, source.ID); err != nil { return err } - return tx.Delete(&source).Error + return repository.DeletePagesProjectSourceTx(tx, source) }) if err != nil { return nil, err @@ -414,15 +406,7 @@ func remoteSourceIdentity(parsed *url.URL) string { } func loadSourceByProject(ctx context.Context, projectID uint) (*model.PagesProjectSource, *model.PagesProjectSourceRuntime, error) { - var source model.PagesProjectSource - if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { - return nil, nil, err - } - var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { - return nil, nil, err - } - return &source, &runtime, nil + return repository.GetPagesProjectSourceAndRuntimeByProjectID(ctx, projectID) } func buildSourceView(source *model.PagesProjectSource, runtime *model.PagesProjectSourceRuntime) (*SourceView, error) { @@ -515,7 +499,7 @@ func resetRuntimeAfterSourceUpdate(tx *gorm.DB, runtime *model.PagesProjectSourc } else { updates[sourceRuntimeColumnSyncStatus] = normalizedSourceRuntimeStatus(runtime) } - return tx.Model(runtime).Updates(updates).Error + return repository.UpdatePagesProjectSourceRuntimeTx(tx, runtime, updates) } func normalizedSourceRuntimeStatus(runtime *model.PagesProjectSourceRuntime) string { diff --git a/internal/apps/openflare/pages/source_orphan_cleanup.go b/internal/apps/openflare/pages/source_orphan_cleanup.go index 843d55ae..9de3bd9d 100644 --- a/internal/apps/openflare/pages/source_orphan_cleanup.go +++ b/internal/apps/openflare/pages/source_orphan_cleanup.go @@ -11,7 +11,6 @@ import ( "time" "github.com/Rain-kl/Wavelet/internal/apps/upload" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" @@ -62,7 +61,7 @@ func ReconcilePagesOrphanUploads( } cutoff := now.UTC().Add(-pagesOrphanUploadIsolation) systemUser := repository.GetSystemUser(ctx) - candidates, err := model.ListPagesOrphanUploadCandidates(ctx, model.PagesOrphanUploadCandidateQuery{ + candidates, err := repository.ListPagesOrphanUploadCandidates(ctx, model.PagesOrphanUploadCandidateQuery{ SystemUserID: systemUser.ID, UploadType: upload.ReservedPagesDeploymentType, Marker: pagesIngestMarkerV2, @@ -130,7 +129,7 @@ func reconcilePagesOrphanUploadCandidate( outcome := pagesOrphanCleanupSkipped uploadLocked := false - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { scopeOutcome, proceed, err := lockPagesOrphanCleanupScope(ctx, tx, candidate.ID, marker) if err != nil { return err @@ -179,14 +178,13 @@ func lockPagesOrphanCleanupScope( return pagesOrphanCleanupSkipped, true, nil } - var source model.PagesProjectSource - sourceExists, err := lockOptionalPagesCleanupRecord(tx, &source, "id = ?", *marker.SourceID) + source, err := repository.LockPagesProjectSourceByIDTx(tx, *marker.SourceID) + if errors.Is(err, gorm.ErrRecordNotFound) { + return pagesOrphanCleanupSkipped, true, nil + } if err != nil { return pagesOrphanCleanupSkipped, false, err } - if !sourceExists { - return pagesOrphanCleanupSkipped, true, nil - } if source.ProjectID != marker.ProjectID { logger.WarnF(ctx, "[PagesSource] orphan upload source ownership mismatch: upload_id=%d project_id=%d source_id=%d source_project_id=%d", @@ -198,15 +196,14 @@ func lockPagesOrphanCleanupScope( return pagesOrphanCleanupInvalidMarker, false, nil } - var runtime model.PagesProjectSourceRuntime - runtimeExists, err := lockOptionalPagesCleanupRecord(tx, &runtime, "source_id = ?", source.ID) - if err != nil { + runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID) + if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { return pagesOrphanCleanupSkipped, false, err } // Read the real clock only after obtaining the runtime row lock. The scanner // snapshot time is only an isolation cutoff and may be stale after lock wait. leaseCheckedAt := time.Now() - if runtimeExists && runtime.LeaseExpiresAt != nil && runtime.LeaseExpiresAt.After(leaseCheckedAt) { + if err == nil && runtime.LeaseExpiresAt != nil && runtime.LeaseExpiresAt.After(leaseCheckedAt) { return pagesOrphanCleanupLeaseBusy, false, nil } return pagesOrphanCleanupSkipped, true, nil diff --git a/internal/apps/openflare/pages/source_runtime.go b/internal/apps/openflare/pages/source_runtime.go index f3f1467b..4f2bf571 100644 --- a/internal/apps/openflare/pages/source_runtime.go +++ b/internal/apps/openflare/pages/source_runtime.go @@ -11,10 +11,9 @@ import ( "strings" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "gorm.io/gorm" - "gorm.io/gorm/clause" ) const ( @@ -89,24 +88,16 @@ func acquireSourceLease( now := time.Now() expiresAt := now.Add(leaseDuration) - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). - Where("source_id = ?", sourceID). - Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). - Where( - "EXISTS (SELECT 1 FROM of_pages_project_sources source WHERE source.id = ? AND source.config_version = ?)", - sourceID, - expectedConfigVersion, - ). - Updates(map[string]any{ - sourceRuntimeColumnLeaseToken: token, - sourceRuntimeColumnLeaseExpiresAt: expiresAt, - sourceRuntimeColumnSyncStatus: status, - sourceRuntimeColumnLastError: "", - }) - if result.Error != nil { - return nil, sourceLeaseStale, result.Error + rows, err := repository.TryAcquirePagesSourceRuntimeLease(ctx, sourceID, expectedConfigVersion, now, map[string]any{ + sourceRuntimeColumnLeaseToken: token, + sourceRuntimeColumnLeaseExpiresAt: expiresAt, + sourceRuntimeColumnSyncStatus: status, + sourceRuntimeColumnLastError: "", + }) + if err != nil { + return nil, sourceLeaseStale, err } - if result.RowsAffected == 0 { + if rows == 0 { outcome, inspectErr := inspectSourceLeaseMiss(ctx, sourceID, expectedConfigVersion, now) return nil, outcome, inspectErr } @@ -129,34 +120,30 @@ func loadSourceExecutionSnapshot( token string, ) (*sourceExecutionSnapshot, error) { var snapshot sourceExecutionSnapshot - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { - var source model.PagesProjectSource - if err := tx.Where("id = ?", sourceID).First(&source).Error; err != nil { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { + source, err := repository.GetPagesProjectSourceByIDTx(tx, sourceID) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return errSourceLeaseSnapshotStale } return err } - var project model.PagesProject - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - First(&project, source.ProjectID).Error; err != nil { + project, err := repository.LockPagesProjectByIDTx(tx, source.ProjectID) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return errSourceLeaseSnapshotStale } return err } - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("id = ?", sourceID). - First(&source).Error; err != nil { + source, err = repository.LockPagesProjectSourceByIDTx(tx, sourceID) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return errSourceLeaseSnapshotStale } return err } - var runtime model.PagesProjectSourceRuntime - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("source_id = ?", source.ID). - First(&runtime).Error; err != nil { + runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return errSourceLeaseSnapshotStale } @@ -207,8 +194,8 @@ func inspectSourceLeaseMiss( expectedConfigVersion int, now time.Time, ) (sourceLeaseOutcome, error) { - var source model.PagesProjectSource - if err := db.DB(ctx).Where("id = ?", sourceID).First(&source).Error; err != nil { + source, err := repository.GetPagesProjectSourceByID(ctx, sourceID) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return sourceLeaseStale, nil } @@ -217,8 +204,8 @@ func inspectSourceLeaseMiss( if source.ConfigVersion != expectedConfigVersion { return sourceLeaseStale, nil } - var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { + runtime, err := repository.GetPagesProjectSourceRuntimeBySourceID(ctx, sourceID) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return sourceLeaseStale, nil } @@ -255,13 +242,11 @@ func renewSourceLease(ctx context.Context, snapshot *sourceExecutionSnapshot, du } now := time.Now() expiresAt := now.Add(duration) - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). - Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now). - Updates(map[string]any{sourceRuntimeColumnLeaseExpiresAt: expiresAt}) - if result.Error != nil { - return false, result.Error + rows, err := repository.RenewPagesSourceRuntimeLease(ctx, snapshot.SourceID, snapshot.LeaseToken, now, expiresAt) + if err != nil { + return false, err } - if result.RowsAffected == 0 { + if rows == 0 { return false, nil } snapshot.LeaseExpiresAt = expiresAt @@ -274,14 +259,19 @@ func failSourceLease(ctx context.Context, snapshot *sourceExecutionSnapshot, mes } message = safeSourceRuntimeError(message) now := time.Now() - return db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). - Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now). - Updates(map[string]any{ + _, err := repository.UpdatePagesSourceRuntimeByActiveLease( + ctx, + snapshot.SourceID, + snapshot.LeaseToken, + now, + map[string]any{ sourceRuntimeColumnSyncStatus: pagesSourceStatusFailed, sourceRuntimeColumnLastError: message, sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseExpiresAt: nil, - }).Error + }, + ) + return err } func safeSourceRuntimeError(message string) string { @@ -296,8 +286,8 @@ func safeSourceRuntimeError(message string) string { } func sourceLeaseIsBusy(ctx context.Context, sourceID uint) (bool, error) { - var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { + runtime, err := repository.GetPagesProjectSourceRuntimeBySourceID(ctx, sourceID) + if err != nil { return false, err } return runtime.LeaseExpiresAt != nil && runtime.LeaseExpiresAt.After(time.Now()), nil @@ -326,35 +316,30 @@ func recoverExpiredSourceLease( sourceRuntimeColumnLeaseExpiresAt: nil, sourceRuntimeColumnNextCheckAt: nextCheckAt, } - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). - Where("source_id = ?", sourceID). - Where("lease_token = ?", token). - Where("lease_expires_at = ? AND lease_expires_at <= ?", expiresAt, now). - Where("sync_status = ?", status). - Updates(updates) - if result.Error != nil { - return false, result.Error + rows, err := repository.RecoverExpiredPagesSourceRuntimeLease( + ctx, sourceID, token, expiresAt, status, now, updates, + ) + if err != nil { + return false, err } - return result.RowsAffected == 1, nil + return rows == 1, nil } // fenceAndNormalizeRuntime invalidates in-flight work while preserving safe // seen/applied cursors. The caller must already hold the source row lock. func fenceAndNormalizeRuntime(tx *gorm.DB, sourceID uint) error { - var runtime model.PagesProjectSourceRuntime - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("source_id = ?", sourceID). - First(&runtime).Error; err != nil { + runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, sourceID) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil } return err } - return tx.Model(&runtime).Updates(map[string]any{ + return repository.UpdatePagesProjectSourceRuntimeTx(tx, runtime, map[string]any{ sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseExpiresAt: nil, - sourceRuntimeColumnSyncStatus: normalizedSourceRuntimeStatus(&runtime), - }).Error + sourceRuntimeColumnSyncStatus: normalizedSourceRuntimeStatus(runtime), + }) } func sourceHasSameReleaseReplacement(runtime *model.PagesProjectSourceRuntime) bool { diff --git a/internal/apps/openflare/pages/source_runtime_test.go b/internal/apps/openflare/pages/source_runtime_test.go index c1f44d29..88e51918 100644 --- a/internal/apps/openflare/pages/source_runtime_test.go +++ b/internal/apps/openflare/pages/source_runtime_test.go @@ -10,6 +10,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" ) @@ -286,7 +288,7 @@ func TestSourceConfigAndProjectContentChangesFenceLease(t *testing.T) { if renewed { t.Error("renewSourceLease(after content update) = true, want false") } - storedProject, err := model.GetPagesProjectByID(ctx, project.ID) + storedProject, err := repository.GetPagesProjectByID(ctx, project.ID) if err != nil { t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) } diff --git a/internal/apps/openflare/pages/source_scanner.go b/internal/apps/openflare/pages/source_scanner.go index c6389b65..56722fef 100644 --- a/internal/apps/openflare/pages/source_scanner.go +++ b/internal/apps/openflare/pages/source_scanner.go @@ -11,11 +11,10 @@ import ( "fmt" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" - "gorm.io/gorm" ) const ( @@ -65,20 +64,6 @@ type pagesSourceProviderBackoff struct { RetryAt string `json:"retry_at"` } -type expiredSourceLeaseCandidate struct { - SourceID uint - LeaseToken string - LeaseExpiresAt time.Time - SyncStatus string - SourceType string - ReleaseSelector string -} - -type dueGitHubSourceCandidate struct { - SourceID uint - ConfigVersion int -} - var ( pagesSourceScanNow = time.Now reconcilePagesSourceOrphans = ReconcilePagesOrphanUploads @@ -173,17 +158,11 @@ func recoverExpiredPagesSourceLeases( now time.Time, summary *pagesSourceScanSummary, ) error { - var candidates []expiredSourceLeaseCandidate - err := db.DB(ctx). - Table("of_pages_project_source_runtime AS runtime"). - Select(`runtime.source_id, runtime.lease_token, runtime.lease_expires_at, - runtime.sync_status, source.source_type, source.release_selector`). - Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id"). - Where("runtime.lease_token <> ''"). - Where("runtime.lease_expires_at IS NOT NULL AND runtime.lease_expires_at <= ?", now). - Where("runtime.sync_status IN ?", []string{pagesSourceStatusChecking, pagesSourceStatusSyncing}). - Order("runtime.source_id ASC"). - Scan(&candidates).Error + candidates, err := repository.ListExpiredPagesSourceLeaseCandidates( + ctx, + now, + []string{pagesSourceStatusChecking, pagesSourceStatusSyncing}, + ) if err != nil { return err } @@ -227,27 +206,18 @@ func scanDueGitHubSources( now time.Time, summary *pagesSourceScanSummary, ) error { - dueQuery := func() *gorm.DB { - return db.DB(ctx). - Table("of_pages_project_source_runtime AS runtime"). - Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id"). - Where("source.source_type = ?", PagesSourceTypeGitHubRelease). - Where("source.release_selector = ?", githubReleaseSelectorLatest). - Where("runtime.next_check_at IS NOT NULL AND runtime.next_check_at <= ?", now) - } - var dueCount int64 - if err := dueQuery().Count(&dueCount).Error; err != nil { + dueCount, err := repository.CountDueGitHubPagesSourceChecks( + ctx, now, PagesSourceTypeGitHubRelease, githubReleaseSelectorLatest, + ) + if err != nil { return err } summary.DueSources = int(dueCount) - var candidates []dueGitHubSourceCandidate - if err := dueQuery(). - Select("source.id AS source_id, source.config_version"). - Order("runtime.next_check_at ASC"). - Order("source.id ASC"). - Limit(pagesSourceScanBatchSize). - Scan(&candidates).Error; err != nil { + candidates, err := repository.ListDueGitHubPagesSourceChecks( + ctx, now, PagesSourceTypeGitHubRelease, githubReleaseSelectorLatest, pagesSourceScanBatchSize, + ) + if err != nil { return err } summary.SelectedSources = len(candidates) @@ -261,8 +231,10 @@ func scanDueGitHubSources( for _, candidate := range candidates { scanOneDueGitHubSource(ctx, candidate, summary) } - var remainingDue int64 - if err := dueQuery().Count(&remainingDue).Error; err != nil { + remainingDue, err := repository.CountDueGitHubPagesSourceChecks( + ctx, now, PagesSourceTypeGitHubRelease, githubReleaseSelectorLatest, + ) + if err != nil { return err } summary.Backlog = int(remainingDue) @@ -272,7 +244,7 @@ func scanDueGitHubSources( func scanOneDueGitHubSource( ctx context.Context, - candidate dueGitHubSourceCandidate, + candidate model.PagesDueGitHubSourceCandidate, summary *pagesSourceScanSummary, ) { snapshot, outcome, err := acquireSourceLease( @@ -343,11 +315,8 @@ func recordPagesSourceProviderBackoff( } retryAt := domainError.retryAt - var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx). - Select("next_check_at"). - Where("source_id = ?", sourceID). - First(&runtime).Error; err != nil { + runtime, err := repository.GetPagesProjectSourceRuntimeBySourceID(ctx, sourceID) + if err != nil { logger.WarnF(ctx, "[PagesSourceScan] load provider backoff deadline failed: source_id=%d error=%v", sourceID, err) } else if runtime.NextCheckAt != nil { retryAt = runtime.NextCheckAt @@ -448,29 +417,23 @@ func recordPagesSourceAutoDispatchFailure( if retryAt != nil && retryAt.After(next) { next = retryAt.In(now.Location()) } - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). - Where("source_id = ?", snapshot.SourceID). - Where("sync_status = ? AND last_seen_revision = ?", pagesSourceStatusUpdateAvailable, revision). - Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). - Where(`EXISTS ( - SELECT 1 FROM of_pages_project_sources AS source - WHERE source.id = ? AND source.config_version = ? - AND source.source_type = ? AND source.release_selector = ? - AND source.auto_update_enabled = ? - )`, - snapshot.SourceID, - snapshot.SourceConfigVersion, - PagesSourceTypeGitHubRelease, - githubReleaseSelectorLatest, - true, - ). - Updates(map[string]any{ + rows, err := repository.RecordPagesSourceAutoDispatchFailure( + ctx, + snapshot.SourceID, + snapshot.SourceConfigVersion, + PagesSourceTypeGitHubRelease, + githubReleaseSelectorLatest, + revision, + pagesSourceStatusUpdateAvailable, + now, + map[string]any{ sourceRuntimeColumnSyncStatus: pagesSourceStatusUpdateAvailable, sourceRuntimeColumnLastError: errPagesSourceTaskDispatchFailed, sourceRuntimeColumnNextCheckAt: &next, - }) - if result.Error != nil { - return false, result.Error + }, + ) + if err != nil { + return false, err } - return result.RowsAffected == 1, nil + return rows == 1, nil } diff --git a/internal/apps/openflare/pages/source_scanner_test.go b/internal/apps/openflare/pages/source_scanner_test.go index 161be65c..13565544 100644 --- a/internal/apps/openflare/pages/source_scanner_test.go +++ b/internal/apps/openflare/pages/source_scanner_test.go @@ -16,6 +16,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/integration/githubrelease" "github.com/Rain-kl/Wavelet/internal/model" @@ -241,7 +243,7 @@ func TestScheduledAutoSyncPersistsExplicitDeploymentTrigger(t *testing.T) { if err != nil || synced == nil || synced.Deployment == nil || synced.Stale { t.Fatalf("syncGitHubSourceWithTrigger() = %+v, %v; want active deployment", synced, err) } - deployment, err := model.GetPagesDeploymentByID(ctx, synced.Deployment.ID) + deployment, err := repository.GetPagesDeploymentByID(ctx, synced.Deployment.ID) if err != nil { t.Fatalf("GetPagesDeploymentByID(%d) error = %v", synced.Deployment.ID, err) } diff --git a/internal/apps/openflare/pages/source_sync.go b/internal/apps/openflare/pages/source_sync.go index 19b00e25..68293777 100644 --- a/internal/apps/openflare/pages/source_sync.go +++ b/internal/apps/openflare/pages/source_sync.go @@ -16,9 +16,9 @@ import ( "unicode/utf8" "github.com/Rain-kl/Wavelet/internal/apps/upload" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" "gorm.io/gorm" "gorm.io/gorm/clause" @@ -375,14 +375,7 @@ func findSourceDeployment( sourceIdentity string, revision string, ) (*model.PagesDeployment, error) { - var deployment model.PagesDeployment - err := db.DB(ctx). - Where("project_id = ? AND source_identity = ? AND source_revision = ?", projectID, sourceIdentity, revision). - First(&deployment).Error - if err != nil { - return nil, err - } - return &deployment, nil + return repository.GetPagesDeploymentBySourceRevision(ctx, projectID, sourceIdentity, revision) } func commitSourceDeploymentWithTrigger( @@ -408,7 +401,7 @@ func commitSourceDeploymentWithTrigger( var committed model.PagesDeployment reused := false ingestReferenced := false - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { state, err := lockSourceCommitState(tx, snapshot) if err != nil { return err @@ -445,18 +438,15 @@ func commitSourceDeploymentWithTrigger( func lockSourceCommitState(tx *gorm.DB, snapshot *sourceExecutionSnapshot) (*sourceCommitState, error) { state := &sourceCommitState{} - var project model.PagesProject - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - First(&project, snapshot.ProjectID).Error; err != nil { + project, err := repository.LockPagesProjectByIDTx(tx, snapshot.ProjectID) + if err != nil { return nil, sourceFenceRecordError(err) } if project.ContentConfigVersion != snapshot.ContentConfigVersion { return nil, errSourceFinalFence } - var source model.PagesProjectSource - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("id = ? AND project_id = ?", snapshot.SourceID, snapshot.ProjectID). - First(&source).Error; err != nil { + source, err := repository.LockPagesProjectSourceByIDAndProjectIDTx(tx, snapshot.SourceID, snapshot.ProjectID) + if err != nil { return nil, sourceFenceRecordError(err) } if source.ConfigVersion != snapshot.SourceConfigVersion || @@ -464,15 +454,13 @@ func lockSourceCommitState(tx *gorm.DB, snapshot *sourceExecutionSnapshot) (*sou source.SourceType != snapshot.SourceType { return nil, errSourceFinalFence } - var runtime model.PagesProjectSourceRuntime - if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). - Where("source_id = ?", source.ID). - First(&runtime).Error; err != nil { + runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID) + if err != nil { return nil, sourceFenceRecordError(err) } - state.Project = &project - state.Source = &source - state.Runtime = &runtime + state.Project = project + state.Source = source + state.Runtime = runtime if err := refreshSourceCommitLease(state, snapshot); err != nil { return nil, err } @@ -675,13 +663,12 @@ func activateSourceDeploymentTx( } nextCheckAt = &next } - result := tx.Model(&model.PagesProjectSourceRuntime{}). - Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", - state.Runtime.SourceID, - state.Runtime.LeaseToken, - finishedAt, - ). - Updates(map[string]any{ + rows, err := repository.UpdatePagesSourceRuntimeByActiveLeaseTx( + tx, + state.Runtime.SourceID, + state.Runtime.LeaseToken, + finishedAt, + map[string]any{ "last_seen_revision": revision, "last_seen_detail": detailJSON, "last_applied_revision": revision, @@ -693,11 +680,12 @@ func activateSourceDeploymentTx( "next_check_at": nextCheckAt, sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseExpiresAt: nil, - }) - if result.Error != nil { - return result.Error + }, + ) + if err != nil { + return err } - if result.RowsAffected != 1 { + if rows != 1 { return errSourceFinalFence } return nil diff --git a/internal/apps/openflare/pages/source_sync_test.go b/internal/apps/openflare/pages/source_sync_test.go index 8cf2b6e0..982c5243 100644 --- a/internal/apps/openflare/pages/source_sync_test.go +++ b/internal/apps/openflare/pages/source_sync_test.go @@ -17,6 +17,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/upload" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" @@ -70,7 +72,7 @@ func mustCreateActiveManualDeployment( if _, err := ActivateDeploymentAs(ctx, projectID, view.ID, "user:1"); err != nil { t.Fatalf("ActivateDeploymentAs(project=%d, deployment=%d) error = %v, want nil", projectID, view.ID, err) } - deployment, err := model.GetPagesDeploymentByID(ctx, view.ID) + deployment, err := repository.GetPagesDeploymentByID(ctx, view.ID) if err != nil { t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", view.ID, err) } @@ -120,14 +122,14 @@ func TestSyncRemoteSourceAtomicallyActivatesAndReusesChecksum(t *testing.T) { if first == nil || first.Stale || first.Reused || first.Deployment == nil { t.Fatalf("syncRemoteSource(first) = %+v, want new active deployment", first) } - storedProject, err := model.GetPagesProjectByID(ctx, project.ID) + storedProject, err := repository.GetPagesProjectByID(ctx, project.ID) if err != nil { t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) } if storedProject.ActiveDeploymentID == nil || *storedProject.ActiveDeploymentID != first.Deployment.ID { t.Fatalf("project ActiveDeploymentID = %v, want %d", storedProject.ActiveDeploymentID, first.Deployment.ID) } - deployment, err := model.GetPagesDeploymentByID(ctx, first.Deployment.ID) + deployment, err := repository.GetPagesDeploymentByID(ctx, first.Deployment.ID) if err != nil { t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", first.Deployment.ID, err) } @@ -293,7 +295,7 @@ func TestSyncRemoteSourceFinalFenceCompensatesIngest(t *testing.T) { if outcome == nil || !outcome.Stale || outcome.Deployment != nil { t.Fatalf("syncRemoteSource(final fence) = %+v, want stale outcome without deployment", outcome) } - storedProject, err := model.GetPagesProjectByID(ctx, project.ID) + storedProject, err := repository.GetPagesProjectByID(ctx, project.ID) if err != nil { t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) } @@ -351,7 +353,7 @@ func TestCommitSourceDeploymentRechecksLeaseAfterUploadLocks(t *testing.T) { if err != nil || first == nil || first.Deployment == nil { t.Fatalf("syncRemoteSource(seed) = (%+v, %v), want deployment", first, err) } - deployment, err := model.GetPagesDeploymentByID(ctx, first.Deployment.ID) + deployment, err := repository.GetPagesDeploymentByID(ctx, first.Deployment.ID) if err != nil { t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", first.Deployment.ID, err) } @@ -403,14 +405,14 @@ func TestCommitSourceDeploymentRechecksLeaseAfterUploadLocks(t *testing.T) { if nowCalls != 2 { t.Fatalf("sourceCommitNow calls = %d, want runtime-lock and post-upload checks", nowCalls) } - storedProject, err := model.GetPagesProjectByID(ctx, project.ID) + storedProject, err := repository.GetPagesProjectByID(ctx, project.ID) if err != nil { t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) } if storedProject.ActiveDeploymentID != nil { t.Errorf("ActiveDeploymentID after expiry recheck = %v, want nil", storedProject.ActiveDeploymentID) } - storedDeployment, err := model.GetPagesDeploymentByID(ctx, deployment.ID) + storedDeployment, err := repository.GetPagesDeploymentByID(ctx, deployment.ID) if err != nil { t.Fatalf("GetPagesDeploymentByID(%d) after expiry error = %v, want nil", deployment.ID, err) } @@ -479,7 +481,7 @@ func assertPagesSyncFailureState( wantDeploymentCount int64, ) { t.Helper() - project, err := model.GetPagesProjectByID(ctx, projectID) + project, err := repository.GetPagesProjectByID(ctx, projectID) if err != nil { t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", projectID, err) } @@ -589,7 +591,7 @@ func TestCommitSourceDeploymentRejectsDeletedTargetUpload(t *testing.T) { if !errors.Is(err, errSourceFinalFence) { t.Errorf("commitSourceDeployment(deleted upload) error = %v, want %v", err, errSourceFinalFence) } - storedProject, err := model.GetPagesProjectByID(ctx, project.ID) + storedProject, err := repository.GetPagesProjectByID(ctx, project.ID) if err != nil { t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) } diff --git a/internal/apps/openflare/pages/source_tasks.go b/internal/apps/openflare/pages/source_tasks.go index 7ddb8a6d..a8248328 100644 --- a/internal/apps/openflare/pages/source_tasks.go +++ b/internal/apps/openflare/pages/source_tasks.go @@ -13,9 +13,9 @@ import ( "strconv" "strings" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" "gorm.io/gorm" ) @@ -135,8 +135,8 @@ func (h *SourceActionHandler) Execute(ctx context.Context, payload []byte) (*tas return nil, task.PermanentError(errPagesSourceActionInvalid) } - var source model.PagesProjectSource - if err := db.DB(ctx).Where("id = ?", input.SourceID).First(&source).Error; err != nil { + source, err := repository.GetPagesProjectSourceByID(ctx, input.SourceID) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { task.AppendLog(ctx, "[resolve] 来源已不存在,本次任务跳过") return &task.TaskResult{Message: errPagesSourceActionStale}, nil @@ -179,7 +179,7 @@ func (h *SourceActionHandler) Execute(ctx context.Context, payload []byte) (*tas if input.Action == sourceActionCheck { return executeGitHubCheckAction(ctx, snapshot) } - return executeSourceSyncAction(ctx, &source, snapshot, input) + return executeSourceSyncAction(ctx, source, snapshot, input) } func executeGitHubCheckAction(ctx context.Context, snapshot *sourceExecutionSnapshot) (*task.TaskResult, error) { @@ -313,14 +313,14 @@ func dispatchSourceActionByProject( return nil, errors.New(errPagesSourceActionInvalid) } - var source model.PagesProjectSource - if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { + source, err := repository.GetPagesProjectSourceByProjectID(ctx, projectID) + if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, errors.New(errPagesSourceNotFound) } return nil, err } - if err := validateSourceActionPreflight(ctx, &source, action, targetRevision, confirmedRevision); err != nil { + if err := validateSourceActionPreflight(ctx, source, action, targetRevision, confirmedRevision); err != nil { return nil, err } busy, err := sourceLeaseIsBusy(ctx, source.ID) @@ -330,7 +330,7 @@ func dispatchSourceActionByProject( if busy { return nil, errors.New(errPagesSourceActionBusy) } - return dispatchSourceActionSnapshot(ctx, source, action, actor, targetRevision, confirmedRevision, "manual") + return dispatchSourceActionSnapshot(ctx, *source, action, actor, targetRevision, confirmedRevision, "manual") } func validateSourceActionPreflight( @@ -416,7 +416,7 @@ func dispatchSourceActionSnapshotWithTrigger( logger.ErrorF(ctx, "[PagesSource] dispatch action failed: project_id=%d source_id=%d action=%s error=%v", source.ProjectID, source.ID, action, err) return nil, errors.New(errPagesSourceTaskDispatchFailed) } - execution, err := model.GetTaskExecutionByTaskID(ctx, taskID) + execution, err := repository.GetTaskExecutionByTaskID(ctx, taskID) if err != nil { logger.ErrorF(ctx, "[PagesSource] load dispatched execution failed: source_id=%d task_id=%s error=%v", source.ID, taskID, err) return nil, errors.New(errPagesSourceTaskDispatchFailed) diff --git a/internal/apps/openflare/pages/source_test.go b/internal/apps/openflare/pages/source_test.go index e888c83a..30d302ea 100644 --- a/internal/apps/openflare/pages/source_test.go +++ b/internal/apps/openflare/pages/source_test.go @@ -10,6 +10,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" ) @@ -67,7 +69,7 @@ func mustCreatePagesSourceProject(t *testing.T, ctx context.Context, slug string if err != nil { t.Fatalf("CreateProject(%q) error = %v, want nil", slug, err) } - project, err := model.GetPagesProjectByID(ctx, view.ID) + project, err := repository.GetPagesProjectByID(ctx, view.ID) if err != nil { t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", view.ID, err) } @@ -356,7 +358,7 @@ func TestDeleteSourceIsIdempotentAndKeepsDeploymentState(t *testing.T) { if sourceCount != 0 || runtimeCount != 0 || deploymentCount != 1 { t.Errorf("DeleteSource counts = source:%d runtime:%d deployment:%d, want 0, 0, 1", sourceCount, runtimeCount, deploymentCount) } - storedProject, err := model.GetPagesProjectByID(ctx, project.ID) + storedProject, err := repository.GetPagesProjectByID(ctx, project.ID) if err != nil { t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) } diff --git a/internal/apps/openflare/proxy_route/build_helpers.go b/internal/apps/openflare/proxy_route/build_helpers.go index 10136e65..b2f457d8 100644 --- a/internal/apps/openflare/proxy_route/build_helpers.go +++ b/internal/apps/openflare/proxy_route/build_helpers.go @@ -11,7 +11,6 @@ import ( "strings" "github.com/Rain-kl/Wavelet/internal/model" - "gorm.io/gorm" ) type proxyRouteJSONFields struct { @@ -141,38 +140,3 @@ func applyProxyRouteUpstreamType(ctx context.Context, route *model.ProxyRoute, u } return nil } - -func updateProxyRouteRecord(tx *gorm.DB, route *model.ProxyRoute) error { - return tx.Model(&model.ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{ - "site_name": route.SiteName, "origin_id": route.OriginID, "origin_url": route.OriginURL, - "origin_host": route.OriginHost, "upstreams": route.Upstreams, "enabled": route.Enabled, - "enable_https": route.EnableHTTPS, "redirect_http": route.RedirectHTTP, - "limit_conn_per_server": route.LimitConnPerServer, "limit_conn_per_ip": route.LimitConnPerIP, - "limit_rate": route.LimitRate, "cache_enabled": route.CacheEnabled, "cache_policy": route.CachePolicy, - "cache_rules": route.CacheRules, "custom_headers": route.CustomHeaders, - "basic_auth_enabled": route.BasicAuthEnabled, "basic_auth_username": route.BasicAuthUsername, - "basic_auth_password": route.BasicAuthPassword, - "upstream_type": route.UpstreamType, "tunnel_node_id": route.TunnelNodeID, - "tunnel_target_addr": route.TunnelTargetAddr, "tunnel_target_protocol": route.TunnelTargetProtocol, - "pages_project_id": route.PagesProjectID, - }).Error -} - -func replaceZoneDomainRouteBindings(tx *gorm.DB, routeID uint, domainIDs []uint) error { - var requested []model.ZoneDomain - if err := tx.Where("id IN ?", domainIDs).Find(&requested).Error; err != nil { - return err - } - if len(requested) != len(domainIDs) { - return errors.New(errProxyRouteZoneDomainNotFound) - } - for _, domain := range requested { - if domain.ProxyRouteID != nil && *domain.ProxyRouteID != routeID { - return errors.New(errProxyRouteZoneDomainBound) - } - } - if err := tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ? AND id NOT IN ?", routeID, domainIDs).Update("proxy_route_id", nil).Error; err != nil { - return err - } - return tx.Model(&model.ZoneDomain{}).Where("id IN ?", domainIDs).Update("proxy_route_id", routeID).Error -} diff --git a/internal/apps/openflare/proxy_route/helpers.go b/internal/apps/openflare/proxy_route/helpers.go index 9e631042..e5f07a4a 100644 --- a/internal/apps/openflare/proxy_route/helpers.go +++ b/internal/apps/openflare/proxy_route/helpers.go @@ -17,6 +17,8 @@ import ( "strings" "unicode" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "gorm.io/gorm" ) @@ -195,7 +197,7 @@ func getOrCreateOriginByAddress(ctx context.Context, address string) (*model.Ori if err := validateOriginAddress(normalizedAddress); err != nil { return nil, err } - existing, err := model.GetOriginByAddress(ctx, normalizedAddress) + existing, err := repository.GetOriginByAddress(ctx, normalizedAddress) if err == nil { return existing, nil } @@ -207,9 +209,9 @@ func getOrCreateOriginByAddress(ctx context.Context, address string) (*model.Ori Address: normalizedAddress, Remark: "", } - if err := model.CreateOriginRecord(ctx, origin); err != nil { + if err := repository.CreateOriginRecord(ctx, origin); err != nil { if isUniqueConstraintError(err) { - return model.GetOriginByAddress(ctx, normalizedAddress) + return repository.GetOriginByAddress(ctx, normalizedAddress) } return nil, err } @@ -217,15 +219,15 @@ func getOrCreateOriginByAddress(ctx context.Context, address string) (*model.Ori } func lookupTLSCertificateByID(ctx context.Context, id uint) (*model.TLSCertificate, error) { - return model.GetTLSCertificateByID(ctx, id) + return repository.GetTLSCertificateByID(ctx, id) } func lookupTunnelNodeByID(ctx context.Context, id uint) (*model.OpenFlareNode, error) { - return model.GetOpenFlareNodeByID(ctx, id) + return repository.GetOpenFlareNodeByID(ctx, id) } func lookupPagesProjectByID(ctx context.Context, id uint) (*model.PagesProject, error) { - return model.GetPagesProjectByID(ctx, id) + return repository.GetPagesProjectByID(ctx, id) } func parseLeafCertificate(certPEM string) (*x509.Certificate, error) { @@ -270,7 +272,7 @@ func loadProxyRouteZoneDomains(ctx context.Context, ids []uint) ([]model.ZoneDom } seen[id] = struct{}{} } - domains, err := model.ListZoneDomainsByIDs(ctx, ids) + domains, err := repository.ListZoneDomainsByIDs(ctx, ids) if err != nil { return nil, errors.New(errProxyRouteZoneDomainNotFound) } @@ -285,7 +287,7 @@ func validateProxyRouteSiteName(siteName string) error { } func validateProxyRouteSiteNameUniqueness(ctx context.Context, route *model.ProxyRoute, siteName string) error { - routes, err := model.ListProxyRoutes(ctx) + routes, err := repository.ListProxyRoutes(ctx) if err != nil { return err } diff --git a/internal/apps/openflare/proxy_route/logics.go b/internal/apps/openflare/proxy_route/logics.go index 7c9ddb18..f4890e3b 100644 --- a/internal/apps/openflare/proxy_route/logics.go +++ b/internal/apps/openflare/proxy_route/logics.go @@ -10,10 +10,9 @@ import ( "strings" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "gorm.io/gorm" - "gorm.io/gorm/clause" ) // CustomHeaderInput 自定义响应头。 @@ -103,7 +102,7 @@ type ZoneDomainView struct { // ListProxyRoutes 列出全部代理规则。 func ListProxyRoutes(ctx context.Context) ([]*View, error) { - routes, err := model.ListProxyRoutes(ctx) + routes, err := repository.ListProxyRoutes(ctx) if err != nil { return nil, err } @@ -112,7 +111,7 @@ func ListProxyRoutes(ctx context.Context) ([]*View, error) { // GetProxyRoute 获取代理规则详情。 func GetProxyRoute(ctx context.Context, id uint) (*View, error) { - route, err := model.GetProxyRouteByID(ctx, id) + route, err := repository.GetProxyRouteByID(ctx, id) if err != nil { return nil, err } @@ -125,17 +124,17 @@ func CreateProxyRoute(ctx context.Context, input Input) (*View, error) { if err != nil { return nil, err } - if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err = repository.WithProxyRouteTx(ctx, func(tx *gorm.DB) error { if err := lockPagesProjectsForRouteMutation(tx, 0, route); err != nil { return err } - if err := tx.Create(route).Error; err != nil { + if err := repository.CreateProxyRouteRecordTx(tx, route); err != nil { return err } - return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs) + return repository.ReplaceZoneDomainRouteBindingsTx(tx, route.ID, input.ZoneDomainIDs) }); err != nil { - if isUniqueConstraintError(err) { - return nil, errors.New(errProxyRouteIdentityExists) + if mapped := mapProxyRoutePersistError(err); mapped != nil { + return nil, mapped } return nil, err } @@ -144,7 +143,7 @@ func CreateProxyRoute(ctx context.Context, input Input) (*View, error) { // UpdateProxyRoute 更新代理规则。 func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error) { - route, err := model.GetProxyRouteByID(ctx, id) + route, err := repository.GetProxyRouteByID(ctx, id) if err != nil { return nil, err } @@ -153,23 +152,39 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error) if err != nil { return nil, err } - if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err = repository.WithProxyRouteTx(ctx, func(tx *gorm.DB) error { if err := lockPagesProjectsForRouteMutation(tx, previousPagesProjectID, route); err != nil { return err } - if err := updateProxyRouteRecord(tx, route); err != nil { + if err := repository.UpdateProxyRouteRecordTx(tx, route); err != nil { return err } - return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs) + return repository.ReplaceZoneDomainRouteBindingsTx(tx, route.ID, input.ZoneDomainIDs) }); err != nil { - if isUniqueConstraintError(err) { - return nil, errors.New(errProxyRouteIdentityExists) + if mapped := mapProxyRoutePersistError(err); mapped != nil { + return nil, mapped } return nil, err } return buildProxyRouteView(ctx, route) } +func mapProxyRoutePersistError(err error) error { + if err == nil { + return nil + } + if isUniqueConstraintError(err) { + return errors.New(errProxyRouteIdentityExists) + } + if errors.Is(err, repository.ErrZoneDomainBoundToAnotherRoute) { + return errors.New(errProxyRouteZoneDomainBound) + } + if errors.Is(err, repository.ErrZoneDomainNotFound) { + return errors.New(errProxyRouteZoneDomainNotFound) + } + return nil +} + func pagesProjectIDForRoute(route *model.ProxyRoute) uint { if route == nil || route.UpstreamType != proxyRouteUpstreamTypePages || route.PagesProjectID == nil { return 0 @@ -189,8 +204,7 @@ func lockPagesProjectsForRouteMutation(tx *gorm.DB, previousProjectID uint, rout sort.Slice(projectIDs, func(i int, j int) bool { return projectIDs[i] < projectIDs[j] }) for _, projectID := range projectIDs { - var project model.PagesProject - err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&project, projectID).Error + project, err := repository.LockPagesProjectByIDTx(tx, projectID) if errors.Is(err, gorm.ErrRecordNotFound) && projectID != nextProjectID { continue } @@ -201,7 +215,7 @@ func lockPagesProjectsForRouteMutation(tx *gorm.DB, previousProjectID uint, rout return err } if projectID == nextProjectID { - if err := validateLockedPagesRouteProject(&project); err != nil { + if err := validateLockedPagesRouteProject(project); err != nil { return err } } @@ -224,15 +238,10 @@ func validateLockedPagesRouteProject(project *model.PagesProject) error { // DeleteProxyRoute 删除代理规则。 func DeleteProxyRoute(ctx context.Context, id uint) error { - if _, err := model.GetProxyRouteByID(ctx, id); err != nil { + if _, err := repository.GetProxyRouteByID(ctx, id); err != nil { return err } - return db.DB(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ?", id).Update("proxy_route_id", nil).Error; err != nil { - return err - } - return tx.Delete(&model.ProxyRoute{}, id).Error - }) + return repository.DeleteProxyRouteAndUnbind(ctx, id) } func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, []model.ZoneDomain, error) { @@ -338,7 +347,7 @@ func buildProxyRouteView(ctx context.Context, route *model.ProxyRoute) (*View, e if route == nil { return nil, errors.New("proxy route is nil") } - domains, err := model.ListZoneDomainsByRouteID(ctx, route.ID) + domains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID) if err != nil { return nil, err } diff --git a/internal/apps/openflare/proxy_route/logics_test.go b/internal/apps/openflare/proxy_route/logics_test.go index f63a2f01..2b3bb203 100644 --- a/internal/apps/openflare/proxy_route/logics_test.go +++ b/internal/apps/openflare/proxy_route/logics_test.go @@ -9,6 +9,7 @@ import ( db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -115,7 +116,7 @@ func TestPagesRouteLocksAndRevalidatesTargetProject(t *testing.T) { require.NoError(t, db.DB(ctx).Delete(&model.PagesProject{}, project.ID).Error) route := &model.ProxyRoute{UpstreamType: proxyRouteUpstreamTypePages, PagesProjectID: &project.ID} - err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err = repository.WithProxyRouteTx(ctx, func(tx *gorm.DB) error { return lockPagesProjectsForRouteMutation(tx, 0, route) }) require.EqualError(t, err, errProxyRoutePagesNotFound) @@ -128,7 +129,7 @@ func TestRouteCanMoveAwayFromAlreadyMissingPagesProject(t *testing.T) { missingProjectID := uint(404) route := &model.ProxyRoute{UpstreamType: "direct"} - err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err := repository.WithProxyRouteTx(ctx, func(tx *gorm.DB) error { return lockPagesProjectsForRouteMutation(tx, missingProjectID, route) }) require.NoError(t, err) diff --git a/internal/apps/openflare/proxy_route/origin_helpers.go b/internal/apps/openflare/proxy_route/origin_helpers.go index c0899ac6..0bc50260 100644 --- a/internal/apps/openflare/proxy_route/origin_helpers.go +++ b/internal/apps/openflare/proxy_route/origin_helpers.go @@ -8,7 +8,8 @@ import ( "errors" "strings" - "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "gorm.io/gorm" ) @@ -32,7 +33,7 @@ func resolveStructuredOriginInput(ctx context.Context, input Input) (string, *ui } func resolveOriginByID(ctx context.Context, scheme, port, uri string, originID uint) (string, *uint, error) { - origin, err := model.GetOriginByID(ctx, originID) + origin, err := repository.GetOriginByID(ctx, originID) if err != nil { return "", nil, errors.New(errProxyRouteOriginNotFound) } @@ -67,7 +68,7 @@ func resolveLegacyOriginInput(ctx context.Context, originURL string) (string, *u if err != nil { return "", nil, err } - origin, findErr := model.GetOriginByAddress(ctx, address) + origin, findErr := repository.GetOriginByAddress(ctx, address) if findErr == nil { return originURL, &origin.ID, nil } diff --git a/internal/apps/openflare/relay/logics.go b/internal/apps/openflare/relay/logics.go index 5b196652..5339c26f 100644 --- a/internal/apps/openflare/relay/logics.go +++ b/internal/apps/openflare/relay/logics.go @@ -11,8 +11,8 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" ) const nodeStatusOnline = "online" @@ -88,7 +88,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat node.LastSeenAt = &lastSeen node.Status = nodeStatusOnline - if err := db.DB(ctx).Model(node).Updates(changes).Error; err != nil { + if err := repository.UpdateOpenFlareNodeColumns(ctx, node, changes); err != nil { return nil, fmt.Errorf("update relay heartbeat: %w", err) } if err := reconcileRelayHealthEvents(ctx, node.NodeID, payload.RelayStatus, now); err != nil { diff --git a/internal/apps/openflare/relay/logics_test.go b/internal/apps/openflare/relay/logics_test.go index b93acb79..1a85b5f8 100644 --- a/internal/apps/openflare/relay/logics_test.go +++ b/internal/apps/openflare/relay/logics_test.go @@ -9,6 +9,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" @@ -36,7 +38,7 @@ func setupRelayTestDB(t *testing.T) func() { db.SetDB(sqliteDB) agent.ResetAuthCacheForTest() - resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) + resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore()) return func() { resetObservabilityStore() @@ -107,17 +109,17 @@ func TestHeartbeatPayloadBindingAndFrpsObservationInsert(t *testing.T) { assert.Equal(t, "v0.1.0", stored.Version) assert.Equal(t, "0.61.0", stored.ExtVersion) - profile, err := model.GetOpenFlareNodeSystemProfile(ctx, node.NodeID) + profile, err := repository.GetOpenFlareNodeSystemProfile(ctx, node.NodeID) require.NoError(t, err) assert.Equal(t, "relay-runtime", profile.Hostname) assert.Equal(t, "Ubuntu", profile.OSName) - snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-time.Minute), 10) + snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, now.Add(-time.Minute), 10) require.NoError(t, err) require.Len(t, snapshots, 1) assert.Equal(t, 12.5, snapshots[0].CPUUsagePercent) - frpsObs, err := model.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1) + frpsObs, err := repository.ListOpenFlareNodeObservationFrps(ctx, node.NodeID, time.Time{}, 1) require.NoError(t, err) require.Len(t, frpsObs, 1) assert.Equal(t, 7, frpsObs[0].FrpsConnections) @@ -153,7 +155,7 @@ func TestHeartbeatRelayReconcilesFrpsUnhealthyEvent(t *testing.T) { }) require.NoError(t, err) - events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, true, 10) + events, err := repository.ListOpenFlareHealthEvents(ctx, node.NodeID, true, 10) require.NoError(t, err) require.Len(t, events, 1) assert.Equal(t, relayFrpsUnhealthyEventType, events[0].EventType) @@ -166,7 +168,7 @@ func TestHeartbeatRelayReconcilesFrpsUnhealthyEvent(t *testing.T) { }) require.NoError(t, err) - events, err = model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 10) + events, err = repository.ListOpenFlareHealthEvents(ctx, node.NodeID, false, 10) require.NoError(t, err) require.Len(t, events, 1) assert.Equal(t, "resolved", events[0].Status) diff --git a/internal/apps/openflare/relay/middleware.go b/internal/apps/openflare/relay/middleware.go index efbcf13f..b0e3d10e 100644 --- a/internal/apps/openflare/relay/middleware.go +++ b/internal/apps/openflare/relay/middleware.go @@ -8,6 +8,8 @@ import ( "errors" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/gin-gonic/gin" @@ -38,7 +40,7 @@ func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlar if token == "" { return nil, errors.New("missing agent token") } - node, err := model.GetOpenFlareNodeByAccessToken(ctx, token) + node, err := repository.GetOpenFlareNodeByAccessToken(ctx, token) if err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { return nil, errors.New("invalid agent token") diff --git a/internal/apps/openflare/relay/middleware_test.go b/internal/apps/openflare/relay/middleware_test.go index 840ad6e3..14520894 100644 --- a/internal/apps/openflare/relay/middleware_test.go +++ b/internal/apps/openflare/relay/middleware_test.go @@ -9,6 +9,8 @@ import ( "net/http/httptest" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/shared/response" @@ -44,7 +46,7 @@ func seedRelayNode(t *testing.T, nodeType, accessToken string) *model.OpenFlareN NodeType: nodeType, AccessToken: accessToken, } - require.NoError(t, model.CreateOpenFlareNode(ctx, node)) + require.NoError(t, repository.CreateOpenFlareNode(ctx, node)) return node } diff --git a/internal/apps/openflare/relay/observability.go b/internal/apps/openflare/relay/observability.go index ab281bc7..c148e2a1 100644 --- a/internal/apps/openflare/relay/observability.go +++ b/internal/apps/openflare/relay/observability.go @@ -7,11 +7,11 @@ import ( "context" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "go.uber.org/zap" - "gorm.io/gorm" ) const relayFrpsUnhealthyEventType = "frps_unhealthy" @@ -35,13 +35,7 @@ func reconcileRelayHealthEvents(ctx context.Context, nodeID string, relayStatus }, }) } - conn := db.DB(ctx) - if conn == nil { - return nil - } - return conn.Transaction(func(tx *gorm.DB) error { - return agent.ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, managedTypes) - }) + return agent.ReconcileScopedNodeHealthEvents(ctx, nodeID, events, reportedAt, managedTypes) } func persistRelayHeartbeatObservability(ctx context.Context, nodeID string, payload HeartbeatPayload, reportedAt time.Time) { @@ -59,7 +53,7 @@ func persistRelayHeartbeatObservability(ctx context.Context, nodeID string, payl FrpsClientCount: payload.FrpsClientCount, FrpsProxies: agent.MarshalJSON(payload.FrpsProxies), } - if err := model.InsertOpenFlareNodeObservationFrps(ctx, frpsObs); err != nil { + if err := repository.InsertOpenFlareNodeObservationFrps(ctx, frpsObs); err != nil { zap.L().Error("persist relay frps observation failed", zap.String("node_id", nodeID), zap.Error(err)) } } diff --git a/internal/apps/openflare/tasks/database_cleanup.go b/internal/apps/openflare/tasks/database_cleanup.go index ec918b96..382dd599 100644 --- a/internal/apps/openflare/tasks/database_cleanup.go +++ b/internal/apps/openflare/tasks/database_cleanup.go @@ -193,15 +193,15 @@ func deleteAllObservabilityRows(ctx context.Context, target string) (int64, stri ) switch target { case DatabaseCleanupTargetAccessLogs: - deleted, err = model.DeleteAllOpenFlareAccessLogs(ctx) + deleted, err = repository.DeleteAllOpenFlareAccessLogs(ctx) case DatabaseCleanupTargetMetricSnapshots: - deleted, err = model.DeleteAllOpenFlareMetricSnapshots(ctx) + deleted, err = repository.DeleteAllOpenFlareMetricSnapshots(ctx) case DatabaseCleanupTargetEdgeHealth: - deleted, err = model.DeleteAllOpenFlareEdgeHealth(ctx) + deleted, err = repository.DeleteAllOpenFlareEdgeHealth(ctx) case DatabaseCleanupTargetObsFrps: - deleted, err = model.DeleteAllOpenFlareNodeObservationFrps(ctx) + deleted, err = repository.DeleteAllOpenFlareNodeObservationFrps(ctx) case DatabaseCleanupTargetObsFrpc: - deleted, err = model.DeleteAllOpenFlareNodeObservationFrpc(ctx) + deleted, err = repository.DeleteAllOpenFlareNodeObservationFrpc(ctx) default: return 0, "", errors.New("unsupported cleanup target") } @@ -226,15 +226,15 @@ func materializeObservabilityTableTTL(ctx context.Context, target string) (int64 ) switch target { case DatabaseCleanupTargetAccessLogs: - eligible, err = model.DeleteOpenFlareAccessLogsBefore(ctx, cutoff) + eligible, err = repository.DeleteOpenFlareAccessLogsBefore(ctx, cutoff) case DatabaseCleanupTargetMetricSnapshots: - eligible, err = model.DeleteOpenFlareMetricSnapshotsBefore(ctx, cutoff) + eligible, err = repository.DeleteOpenFlareMetricSnapshotsBefore(ctx, cutoff) case DatabaseCleanupTargetEdgeHealth: - eligible, err = model.DeleteOpenFlareEdgeHealthBefore(ctx, cutoff) + eligible, err = repository.DeleteOpenFlareEdgeHealthBefore(ctx, cutoff) case DatabaseCleanupTargetObsFrps: - eligible, err = model.DeleteOpenFlareNodeObservationFrpsBefore(ctx, cutoff) + eligible, err = repository.DeleteOpenFlareNodeObservationFrpsBefore(ctx, cutoff) case DatabaseCleanupTargetObsFrpc: - eligible, err = model.DeleteOpenFlareNodeObservationFrpcBefore(ctx, cutoff) + eligible, err = repository.DeleteOpenFlareNodeObservationFrpcBefore(ctx, cutoff) default: return 0, "", errors.New("unsupported cleanup target") } diff --git a/internal/apps/openflare/tasks/database_cleanup_test.go b/internal/apps/openflare/tasks/database_cleanup_test.go index 19c02be6..65f3c1f7 100644 --- a/internal/apps/openflare/tasks/database_cleanup_test.go +++ b/internal/apps/openflare/tasks/database_cleanup_test.go @@ -27,8 +27,8 @@ func setupDatabaseCleanupTestDB(t *testing.T) context.Context { require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{})) db.SetDB(sqliteDB) - resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) - resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) + resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore()) + resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore()) t.Cleanup(func() { resetObservabilityStore() resetAccessLogStore() @@ -69,12 +69,12 @@ func TestCleanupDatabaseObservabilityMaterializeDoesNotClaimHardDelete(t *testin now := time.Now().UTC() // One row past metric table TTL (30d), one still inside the window. - require.NoError(t, model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ + require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ NodeID: "node-a", CapturedAt: now.Add(-40 * 24 * time.Hour), CPUUsagePercent: 10, })) - require.NoError(t, model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ + require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ NodeID: "node-a", CapturedAt: now.Add(-12 * time.Hour), CPUUsagePercent: 20, @@ -96,7 +96,7 @@ func TestCleanupDatabaseObservabilityMaterializeDoesNotClaimHardDelete(t *testin assert.True(t, result.Cutoff.Before(now.Add(-29*24*time.Hour))) // Memory store applies the table-TTL cutoff for tests; only the recent row remains. - rows, err := model.ListOpenFlareMetricSnapshotsSince(ctx, "", time.Time{}, 0) + rows, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, "", time.Time{}, 0) require.NoError(t, err) require.Len(t, rows, 1) assert.Equal(t, float64(20), rows[0].CPUUsagePercent) @@ -106,7 +106,7 @@ func TestCleanupDatabaseObservabilityDeletesAllRowsWhenRetentionMissing(t *testi ctx := setupDatabaseCleanupTestDB(t) now := time.Now().UTC() - require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{ + require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{ { NodeID: "node-a", LoggedAt: now.Add(-3 * time.Hour), @@ -134,7 +134,7 @@ func TestCleanupDatabaseObservabilityDeletesAllRowsWhenRetentionMissing(t *testi assert.Equal(t, int64(2), result.DeletedCount) assert.Equal(t, int64(2), result.EligibleCount) - rows, err := model.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10}) + rows, err := repository.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10}) require.NoError(t, err) assert.Empty(t, rows) } @@ -144,7 +144,7 @@ func TestRunDatabaseAutoCleanupOnceClampsRetentionToTableTTL(t *testing.T) { now := time.Now().UTC() // Access logs TTL=90d, metrics TTL=30d. Config retention=1 must clamp, not reject. - require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{{ + require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{{ NodeID: "node-a", LoggedAt: now.Add(-100 * 24 * time.Hour), RemoteAddr: "203.0.113.10", @@ -152,12 +152,12 @@ func TestRunDatabaseAutoCleanupOnceClampsRetentionToTableTTL(t *testing.T) { Path: "/access", StatusCode: 200, }})) - require.NoError(t, model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ + require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ NodeID: "node-a", CapturedAt: now.Add(-40 * 24 * time.Hour), CPUUsagePercent: 10, })) - require.NoError(t, model.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{ + require.NoError(t, repository.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{ NodeID: "node-a", CapturedAt: now.Add(-40 * 24 * time.Hour), Status: "healthy", @@ -181,15 +181,15 @@ func TestRunDatabaseAutoCleanupOnceClampsRetentionToTableTTL(t *testing.T) { assert.GreaterOrEqual(t, *result.RetentionDays, result.TableTTLDays) } - accessLogs, err := model.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10}) + accessLogs, err := repository.ListOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Page: 0, PageSize: 10}) require.NoError(t, err) assert.Empty(t, accessLogs) - metricSnapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, "", time.Time{}, 0) + metricSnapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, "", time.Time{}, 0) require.NoError(t, err) assert.Empty(t, metricSnapshots) - edgeHealth, err := model.ListOpenFlareEdgeHealth(ctx, "", time.Time{}, 0) + edgeHealth, err := repository.ListOpenFlareEdgeHealth(ctx, "", time.Time{}, 0) require.NoError(t, err) assert.Empty(t, edgeHealth) } diff --git a/internal/apps/openflare/tasks/ssl_renew.go b/internal/apps/openflare/tasks/ssl_renew.go index a7012a3a..ec70ded2 100644 --- a/internal/apps/openflare/tasks/ssl_renew.go +++ b/internal/apps/openflare/tasks/ssl_renew.go @@ -7,8 +7,9 @@ import ( "context" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls" - "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/pkg/logger" ) @@ -16,7 +17,7 @@ import ( func RunSSLRenewJob(ctx context.Context) error { logger.InfoF(ctx, "[OpenFlareTasks] SSL renew job started") - certificates, err := model.ListTLSCertificates(ctx) + certificates, err := repository.ListTLSCertificates(ctx) if err != nil { logger.ErrorF(ctx, "[OpenFlareTasks] list certificates failed: %v", err) return err diff --git a/internal/apps/openflare/tasks/ssl_renew_test.go b/internal/apps/openflare/tasks/ssl_renew_test.go index 822ac879..f79b8d7f 100644 --- a/internal/apps/openflare/tasks/ssl_renew_test.go +++ b/internal/apps/openflare/tasks/ssl_renew_test.go @@ -8,6 +8,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls" "github.com/Rain-kl/Wavelet/internal/infra/config" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" @@ -70,16 +72,16 @@ func TestRunSSLRenewJobTriggersDueCertificates(t *testing.T) { KeyPEM: " ", NotAfter: now.Add(30 * 24 * time.Hour), } - require.NoError(t, model.CreateTLSCertificateRecord(ctx, due)) - require.NoError(t, model.CreateTLSCertificateRecord(ctx, fresh)) + require.NoError(t, repository.CreateTLSCertificateRecord(ctx, due)) + require.NoError(t, repository.CreateTLSCertificateRecord(ctx, fresh)) require.NoError(t, RunSSLRenewJob(ctx)) - renewed, err := model.GetTLSCertificateByID(ctx, due.ID) + renewed, err := repository.GetTLSCertificateByID(ctx, due.ID) require.NoError(t, err) assert.Equal(t, "applying", renewed.ApplyStatus) - unchanged, err := model.GetTLSCertificateByID(ctx, fresh.ID) + unchanged, err := repository.GetTLSCertificateByID(ctx, fresh.ID) require.NoError(t, err) assert.Equal(t, "ready", unchanged.ApplyStatus) } diff --git a/internal/apps/openflare/tls/acme_obtain_test.go b/internal/apps/openflare/tls/acme_obtain_test.go index 93bd7b85..d291ecce 100644 --- a/internal/apps/openflare/tls/acme_obtain_test.go +++ b/internal/apps/openflare/tls/acme_obtain_test.go @@ -9,6 +9,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -81,7 +83,7 @@ func TestRenewCertificateSetsApplying(t *testing.T) { CertPEM: " ", KeyPEM: " ", } - require.NoError(t, model.CreateTLSCertificateRecord(ctx, cert)) + require.NoError(t, repository.CreateTLSCertificateRecord(ctx, cert)) restore := SetObtainCertificateFuncForTest(func(ctx context.Context, c *model.TLSCertificate) error { return nil @@ -106,19 +108,19 @@ func TestConvertCertificateToACMEPreservesUploadOnFailure(t *testing.T) { }) require.NoError(t, err) - stored, err := model.GetTLSCertificateByID(ctx, cert.ID) + stored, err := repository.GetTLSCertificateByID(ctx, cert.ID) require.NoError(t, err) originalStoredCertPEM := stored.CertPEM originalStoredKeyPEM := stored.KeyPEM stored.ApplyStatus = "applying" stored.PrimaryDomain = "manual.example.com" - require.NoError(t, model.SaveTLSCertificate(ctx, stored)) + require.NoError(t, repository.SaveTLSCertificate(ctx, stored)) err = updateCertError(ctx, stored, "dns challenge failed") require.Error(t, err) - finalCert, err := model.GetTLSCertificateByID(ctx, cert.ID) + finalCert, err := repository.GetTLSCertificateByID(ctx, cert.ID) require.NoError(t, err) assert.Equal(t, "upload", finalCert.Provider) assert.Equal(t, "error", finalCert.ApplyStatus) @@ -141,14 +143,14 @@ func TestConvertCertificateToACMERejectsInvalidStates(t *testing.T) { require.NoError(t, err) cert.Provider = "acme" - require.NoError(t, model.SaveTLSCertificate(ctx, cert)) + require.NoError(t, repository.SaveTLSCertificate(ctx, cert)) _, err = ConvertCertificateToACME(ctx, cert.ID, ApplyInput{Name: "manual-cert"}) require.Error(t, err) assert.Contains(t, err.Error(), "only uploaded") cert.Provider = "upload" cert.ApplyStatus = "applying" - require.NoError(t, model.SaveTLSCertificate(ctx, cert)) + require.NoError(t, repository.SaveTLSCertificate(ctx, cert)) _, err = ConvertCertificateToACME(ctx, cert.ID, ApplyInput{Name: "manual-cert"}) require.Error(t, err) assert.Contains(t, err.Error(), "already applying") diff --git a/internal/apps/openflare/tls/logics.go b/internal/apps/openflare/tls/logics.go index 40987326..9fc5c7c2 100644 --- a/internal/apps/openflare/tls/logics.go +++ b/internal/apps/openflare/tls/logics.go @@ -12,6 +12,8 @@ import ( "mime/multipart" "strings" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" ) @@ -71,17 +73,17 @@ type DNSAccountInput struct { // ListCertificates 列出全部证书(不含 PEM)。 func ListCertificates(ctx context.Context) ([]model.TLSCertificate, error) { - return model.ListTLSCertificates(ctx) + return repository.ListTLSCertificates(ctx) } // GetCertificate 获取证书详情(不含 PEM)。 func GetCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) { - return model.GetTLSCertificateByID(ctx, id) + return repository.GetTLSCertificateByID(ctx, id) } // GetCertificateContent 获取证书 PEM 内容。 func GetCertificateContent(ctx context.Context, id uint) (*CertificateContent, error) { - certificate, err := model.GetTLSCertificateByID(ctx, id) + certificate, err := repository.GetTLSCertificateByID(ctx, id) if err != nil { return nil, err } @@ -117,7 +119,7 @@ func CreateCertificate(ctx context.Context, input CertificateInput) (*model.TLSC if err != nil { return nil, err } - if err = model.CreateTLSCertificateRecord(ctx, certificate); err != nil { + if err = repository.CreateTLSCertificateRecord(ctx, certificate); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errCertificateNameExists) } @@ -149,7 +151,7 @@ func CreateCertificateFromFiles(ctx context.Context, name string, certFile *mult // UpdateCertificate 更新上传证书。 func UpdateCertificate(ctx context.Context, id uint, input CertificateInput) (*model.TLSCertificate, error) { - existing, err := model.GetTLSCertificateByID(ctx, id) + existing, err := repository.GetTLSCertificateByID(ctx, id) if err != nil { return nil, err } @@ -157,7 +159,7 @@ func UpdateCertificate(ctx context.Context, id uint, input CertificateInput) (*m if err != nil { return nil, err } - if err = model.SaveTLSCertificate(ctx, certificate); err != nil { + if err = repository.SaveTLSCertificate(ctx, certificate); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errCertificateNameExists) } @@ -171,10 +173,10 @@ func DeleteCertificate(ctx context.Context, id uint) error { if err := ensureCertificateNotReferenced(ctx, id); err != nil { return err } - if _, err := model.GetTLSCertificateByID(ctx, id); err != nil { + if _, err := repository.GetTLSCertificateByID(ctx, id); err != nil { return err } - return model.DeleteTLSCertificateRecord(ctx, id) + return repository.DeleteTLSCertificateRecord(ctx, id) } // ApplyCertificate 申请 ACME 证书。 @@ -188,7 +190,7 @@ func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertific if cert.Name == "" { return nil, errors.New(errCertificateNameRequired) } - if err := model.CreateTLSCertificateRecord(ctx, cert); err != nil { + if err := repository.CreateTLSCertificateRecord(ctx, cert); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errCertificateNameExists) } @@ -205,7 +207,7 @@ func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertific // UpdateACMECertificate 更新 ACME 证书配置。 func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) { - cert, err := model.GetTLSCertificateByID(ctx, id) + cert, err := repository.GetTLSCertificateByID(ctx, id) if err != nil { return nil, err } @@ -216,7 +218,7 @@ func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*mod if cert.Name == "" { return nil, errors.New(errCertificateNameRequired) } - if err := model.SaveTLSCertificate(ctx, cert); err != nil { + if err := repository.SaveTLSCertificate(ctx, cert); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errCertificateNameExists) } @@ -233,7 +235,7 @@ func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*mod // ConvertCertificateToACME 将上传证书转为 ACME 管理。 func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) { - cert, err := model.GetTLSCertificateByID(ctx, id) + cert, err := repository.GetTLSCertificateByID(ctx, id) if err != nil { return nil, err } @@ -248,7 +250,7 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (* return nil, errors.New(errCertificateNameRequired) } cert.ApplyMessage = "" - if err := model.SaveTLSCertificate(ctx, cert); err != nil { + if err := repository.SaveTLSCertificate(ctx, cert); err != nil { if isUniqueConstraintError(err) { return nil, errors.New(errCertificateNameExists) } @@ -260,14 +262,14 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (* if err := obtainTLSCertificate(asyncCtx, c); err != nil { return } - latest, err := model.GetTLSCertificateByID(asyncCtx, c.ID) + latest, err := repository.GetTLSCertificateByID(asyncCtx, c.ID) if err != nil { return } latest.Provider = tlsProviderACME latest.ApplyStatus = tlsApplyStatusReady latest.ApplyMessage = "" - _ = model.SaveTLSCertificate(asyncCtx, latest) + _ = repository.SaveTLSCertificate(asyncCtx, latest) }(cert) return sanitizeCertificateForResponse(cert), nil @@ -275,7 +277,7 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (* // RenewCertificate 续期 ACME 证书。 func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) { - cert, err := model.GetTLSCertificateByID(ctx, id) + cert, err := repository.GetTLSCertificateByID(ctx, id) if err != nil { return nil, err } @@ -295,7 +297,7 @@ func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, erro cert.ApplyStatus = tlsApplyStatusApplying cert.ApplyMessage = "" - if err := model.SaveTLSCertificate(ctx, cert); err != nil { + if err := repository.SaveTLSCertificate(ctx, cert); err != nil { return nil, err } return sanitizeCertificateForResponse(cert), nil @@ -303,7 +305,7 @@ func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, erro // ListDNSAccounts 列出 DNS 账号。 func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) { - return model.ListDNSAccounts(ctx) + return repository.ListDNSAccounts(ctx) } // CreateDNSAccount 创建 DNS 账号。 @@ -320,7 +322,7 @@ func CreateDNSAccount(ctx context.Context, input DNSAccountInput) (*model.DNSAcc if account.Name == "" || account.Type == "" || authorization == "" { return nil, errors.New("DNS 账号参数不完整") } - if err := model.CreateDNSAccountRecord(ctx, account); err != nil { + if err := repository.CreateDNSAccountRecord(ctx, account); err != nil { return nil, err } return sanitizeDNSAccountForResponse(account), nil @@ -328,7 +330,7 @@ func CreateDNSAccount(ctx context.Context, input DNSAccountInput) (*model.DNSAcc // UpdateDNSAccount 更新 DNS 账号。 func UpdateDNSAccount(ctx context.Context, id uint, input DNSAccountInput) (*model.DNSAccount, error) { - account, err := model.GetDNSAccountByID(ctx, id) + account, err := repository.GetDNSAccountByID(ctx, id) if err != nil { return nil, err } @@ -342,7 +344,7 @@ func UpdateDNSAccount(ctx context.Context, id uint, input DNSAccountInput) (*mod if account.Name == "" || account.Type == "" || authorization == "" { return nil, errors.New("DNS 账号参数不完整") } - if err := model.SaveDNSAccount(ctx, account); err != nil { + if err := repository.SaveDNSAccount(ctx, account); err != nil { return nil, err } return sanitizeDNSAccountForResponse(account), nil @@ -350,22 +352,22 @@ func UpdateDNSAccount(ctx context.Context, id uint, input DNSAccountInput) (*mod // DeleteDNSAccount 删除 DNS 账号。 func DeleteDNSAccount(ctx context.Context, id uint) error { - if _, err := model.GetDNSAccountByID(ctx, id); err != nil { + if _, err := repository.GetDNSAccountByID(ctx, id); err != nil { return err } - count, err := model.CountTLSCertificatesByDNSAccountID(ctx, id) + count, err := repository.CountTLSCertificatesByDNSAccountID(ctx, id) if err != nil { return err } if count > 0 { return errors.New(errDNSAccountInUse) } - return model.DeleteDNSAccountRecord(ctx, id) + return repository.DeleteDNSAccountRecord(ctx, id) } // GetDefaultAcmeAccount 获取默认 ACME 账号。 func GetDefaultAcmeAccount(ctx context.Context) (*model.AcmeAccount, error) { - account, err := model.GetDefaultAcmeAccount(ctx) + account, err := repository.GetDefaultAcmeAccount(ctx) if err != nil { return nil, err } @@ -430,7 +432,7 @@ func fillAcmeCertificateFields(cert *model.TLSCertificate, input ApplyInput) { } func ensureCertificateNotReferenced(ctx context.Context, id uint) error { - count, err := model.CountZoneDomainsByCertificateID(ctx, id) + count, err := repository.CountZoneDomainsByCertificateID(ctx, id) if err != nil { return err } diff --git a/internal/apps/openflare/tls/logics_test.go b/internal/apps/openflare/tls/logics_test.go index 855632f9..31cdf5b7 100644 --- a/internal/apps/openflare/tls/logics_test.go +++ b/internal/apps/openflare/tls/logics_test.go @@ -16,6 +16,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/infra/config" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" @@ -123,7 +125,7 @@ func TestCreateCertificateEncryptsPrivateKey(t *testing.T) { }) require.NoError(t, err) - stored, err := model.GetTLSCertificateByID(ctx, certificate.ID) + stored, err := repository.GetTLSCertificateByID(ctx, certificate.ID) require.NoError(t, err) assert.NotEqual(t, keyPEM, stored.KeyPEM) assert.Contains(t, stored.KeyPEM, sensitiveValuePrefix) diff --git a/internal/apps/openflare/tls/obtain.go b/internal/apps/openflare/tls/obtain.go index 4012b920..266c72dc 100644 --- a/internal/apps/openflare/tls/obtain.go +++ b/internal/apps/openflare/tls/obtain.go @@ -9,6 +9,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls/acme" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" @@ -35,7 +37,7 @@ func SetObtainCertificateFuncForTest(fn func(context.Context, *model.TLSCertific func obtainCertificate(ctx context.Context, cert *model.TLSCertificate) error { task.AppendLog(ctx, "【续签任务】开始续签,设置申请状态为 applying...") cert.ApplyStatus = tlsApplyStatusApplying - if err := model.SaveTLSCertificate(ctx, cert); err != nil { + if err := repository.SaveTLSCertificate(ctx, cert); err != nil { return err } @@ -46,7 +48,7 @@ func obtainCertificate(ctx context.Context, cert *model.TLSCertificate) error { } task.AppendLog(ctx, "【续签任务】正在解析 DNS 账户信息 (ID=%d)...", cert.DNSAccountID) - dnsAccount, err := model.GetDNSAccountByID(ctx, cert.DNSAccountID) + dnsAccount, err := repository.GetDNSAccountByID(ctx, cert.DNSAccountID) if err != nil { return updateCertError(ctx, cert, fmt.Sprintf("Failed to get DNS account: %v", err)) } @@ -100,7 +102,7 @@ func obtainCertificate(ctx context.Context, cert *model.TLSCertificate) error { func updateCertError(ctx context.Context, cert *model.TLSCertificate, message string) error { cert.ApplyStatus = "error" cert.ApplyMessage = message - if err := model.SaveTLSCertificate(ctx, cert); err != nil { + if err := repository.SaveTLSCertificate(ctx, cert); err != nil { return err } return fmt.Errorf("%s", message) diff --git a/internal/apps/openflare/tls/obtain_helpers.go b/internal/apps/openflare/tls/obtain_helpers.go index 904df2ee..d135fa98 100644 --- a/internal/apps/openflare/tls/obtain_helpers.go +++ b/internal/apps/openflare/tls/obtain_helpers.go @@ -7,21 +7,23 @@ import ( "context" "fmt" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls/acme" "github.com/Rain-kl/Wavelet/internal/model" ) func resolveAcmeAccount(ctx context.Context, cert *model.TLSCertificate) (*model.AcmeAccount, error) { - acmeAccount, err := model.GetAcmeAccountByID(ctx, cert.AcmeAccountID) + acmeAccount, err := repository.GetAcmeAccountByID(ctx, cert.AcmeAccountID) if err == nil { return acmeAccount, nil } - acmeAccount, err = model.GetDefaultAcmeAccount(ctx) + acmeAccount, err = repository.GetDefaultAcmeAccount(ctx) if err != nil { return nil, fmt.Errorf("failed to get ACME account: %w", err) } cert.AcmeAccountID = acmeAccount.ID - if err := model.SaveTLSCertificate(ctx, cert); err != nil { + if err := repository.SaveTLSCertificate(ctx, cert); err != nil { return nil, err } return acmeAccount, nil @@ -49,14 +51,14 @@ func persistAcmeAccountUpdates( acmeAccount.URL = newAccountURL } if acmeAccount.ID == 0 { - if dbErr := model.CreateAcmeAccountRecord(ctx, acmeAccount); dbErr != nil { + if dbErr := repository.CreateAcmeAccountRecord(ctx, acmeAccount); dbErr != nil { return fmt.Errorf("failed to create ACME account: %w", dbErr) } - } else if dbErr := model.SaveAcmeAccount(ctx, acmeAccount); dbErr != nil { + } else if dbErr := repository.SaveAcmeAccount(ctx, acmeAccount); dbErr != nil { return fmt.Errorf("failed to save ACME account: %w", dbErr) } cert.AcmeAccountID = acmeAccount.ID - return model.SaveTLSCertificate(ctx, cert) + return repository.SaveTLSCertificate(ctx, cert) } func saveObtainedCertificate(ctx context.Context, cert *model.TLSCertificate, result *acme.CertificateResult) error { @@ -70,5 +72,5 @@ func saveObtainedCertificate(ctx context.Context, cert *model.TLSCertificate, re cert.NotAfter = result.NotAfter cert.ApplyStatus = tlsApplyStatusReady cert.ApplyMessage = "" - return model.SaveTLSCertificate(ctx, cert) + return repository.SaveTLSCertificate(ctx, cert) } diff --git a/internal/apps/openflare/tls/tasks.go b/internal/apps/openflare/tls/tasks.go index 7d9eb2b6..3dceff00 100644 --- a/internal/apps/openflare/tls/tasks.go +++ b/internal/apps/openflare/tls/tasks.go @@ -9,8 +9,9 @@ import ( "errors" "fmt" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" ) const ( @@ -77,7 +78,7 @@ func (h *SSLSingleRenewHandler) Execute(ctx context.Context, payload []byte) (*t task.AppendLog(ctx, "开始续期证书,ID: %d", req.ID) - cert, err := model.GetTLSCertificateByID(ctx, req.ID) + cert, err := repository.GetTLSCertificateByID(ctx, req.ID) if err != nil { task.AppendLog(ctx, "获取证书记录失败 ID=%d: %v", req.ID, err) return nil, fmt.Errorf("获取证书记录失败: %w", err) diff --git a/internal/apps/openflare/tls/tasks_test.go b/internal/apps/openflare/tls/tasks_test.go index 77021ec0..27267008 100644 --- a/internal/apps/openflare/tls/tasks_test.go +++ b/internal/apps/openflare/tls/tasks_test.go @@ -9,6 +9,8 @@ import ( "errors" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -91,7 +93,7 @@ func TestSSLSingleRenewHandler_Execute(t *testing.T) { Provider: "custom", PrimaryDomain: "example.com", } - err := model.CreateTLSCertificateRecord(ctx, cert) + err := repository.CreateTLSCertificateRecord(ctx, cert) require.NoError(t, err) payload, err := json.Marshal(SSLSingleRenewPayload{ID: cert.ID}) @@ -108,13 +110,13 @@ func TestSSLSingleRenewHandler_Execute(t *testing.T) { Provider: tlsProviderACME, PrimaryDomain: "success.example.com", } - err := model.CreateTLSCertificateRecord(ctx, cert) + err := repository.CreateTLSCertificateRecord(ctx, cert) require.NoError(t, err) // Mock obtainCertificate to succeed restore := SetObtainCertificateFuncForTest(func(ctx context.Context, c *model.TLSCertificate) error { c.ApplyStatus = tlsApplyStatusReady - return model.SaveTLSCertificate(ctx, c) + return repository.SaveTLSCertificate(ctx, c) }) defer restore() @@ -125,7 +127,7 @@ func TestSSLSingleRenewHandler_Execute(t *testing.T) { require.NoError(t, err) assert.Contains(t, res.Message, "续签成功") - updated, err := model.GetTLSCertificateByID(ctx, cert.ID) + updated, err := repository.GetTLSCertificateByID(ctx, cert.ID) require.NoError(t, err) assert.Equal(t, tlsApplyStatusReady, updated.ApplyStatus) }) @@ -136,7 +138,7 @@ func TestSSLSingleRenewHandler_Execute(t *testing.T) { Provider: tlsProviderACME, PrimaryDomain: "fail.example.com", } - err := model.CreateTLSCertificateRecord(ctx, cert) + err := repository.CreateTLSCertificateRecord(ctx, cert) require.NoError(t, err) // Mock obtainCertificate to fail diff --git a/internal/apps/openflare/uptimekuma/sync.go b/internal/apps/openflare/uptimekuma/sync.go index ca2a6a50..062b1fc2 100644 --- a/internal/apps/openflare/uptimekuma/sync.go +++ b/internal/apps/openflare/uptimekuma/sync.go @@ -100,7 +100,7 @@ func SyncToUptimeKuma(ctx context.Context) error { "scope", config.MonitorScope, ) - allRoutes, err := model.ListProxyRoutes(ctx) + allRoutes, err := repository.ListProxyRoutes(ctx) if err != nil { return fmt.Errorf("failed to list local proxy routes: %w", err) } @@ -227,7 +227,7 @@ func routeMonitorURL(ctx context.Context, route *model.ProxyRoute) (string, erro if route == nil { return "", fmt.Errorf("proxy route is nil") } - domains, err := model.ListZoneDomainsByRouteID(ctx, route.ID) + domains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID) if err != nil { return "", err } diff --git a/internal/apps/openflare/uptimekuma/sync_test.go b/internal/apps/openflare/uptimekuma/sync_test.go index c06d1b0c..eaebfdb9 100644 --- a/internal/apps/openflare/uptimekuma/sync_test.go +++ b/internal/apps/openflare/uptimekuma/sync_test.go @@ -217,9 +217,9 @@ func TestSyncToUptimeKumaSuccess(t *testing.T) { EnableHTTPS: false, } - require.NoError(t, model.CreateProxyRouteRecord(ctx, routeA)) - require.NoError(t, model.CreateProxyRouteRecord(ctx, routeB)) - require.NoError(t, model.CreateProxyRouteRecord(ctx, routeC)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeA)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeB)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeC)) createRouteZoneDomain(t, ctx, routeA, "site-a.com") createRouteZoneDomain(t, ctx, routeB, "site-b.com") createRouteZoneDomain(t, ctx, routeC, "site-c.com") @@ -320,8 +320,8 @@ func TestSyncToUptimeKumaSelectedScope(t *testing.T) { EnableHTTPS: false, } - require.NoError(t, model.CreateProxyRouteRecord(ctx, routeA)) - require.NoError(t, model.CreateProxyRouteRecord(ctx, routeB)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeA)) + require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeB)) createRouteZoneDomain(t, ctx, routeA, "site-a.com") createRouteZoneDomain(t, ctx, routeB, "site-b.com") diff --git a/internal/apps/openflare/waf/ip_group_sync.go b/internal/apps/openflare/waf/ip_group_sync.go index e49342c5..1016b556 100644 --- a/internal/apps/openflare/waf/ip_group_sync.go +++ b/internal/apps/openflare/waf/ip_group_sync.go @@ -17,6 +17,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" "github.com/Rain-kl/Wavelet/internal/model" @@ -132,7 +134,7 @@ type ipGroupAutoAccumulator struct { // SyncDueWAFIPGroups syncs all enabled automatic/subscription IP groups that are due. func SyncDueWAFIPGroups(ctx context.Context) error { now := time.Now().UTC() - groups, err := model.ListDueOpenFlareWAFIPGroups(ctx, now) + groups, err := repository.ListDueOpenFlareWAFIPGroups(ctx, now) if err != nil { return err } @@ -176,7 +178,7 @@ func syncIPGroupSubscription(ctx context.Context, group *model.OpenFlareWAFIPGro group.NextSyncAt = &nextSyncAt group.LastSyncStatus = "success" group.LastSyncMessage = fmt.Sprintf("同步成功,共 %d 条 IP/IP 段", len(ips)) - if err := model.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group); err != nil { + if err := repository.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group); err != nil { return nil, err } broadcastIPGroupToAgents(ctx, group.ID) @@ -259,7 +261,7 @@ func syncIPGroupAutomatic(ctx context.Context, group *model.OpenFlareWAFIPGroup, group.NextSyncAt = &nextSyncAt group.LastSyncStatus = "success" group.LastSyncMessage = fmt.Sprintf("自动规则执行成功,共命中 %d 个 IP,当前生效 %d 个 IP", len(ips), len(finalIPs)) - if err := model.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group); err != nil { + if err := repository.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group); err != nil { return nil, err } broadcastIPGroupToAgents(ctx, group.ID) @@ -283,7 +285,7 @@ func recordIPGroupSyncFailure(ctx context.Context, group *model.OpenFlareWAFIPGr group.NextSyncAt = &nextSyncAt group.LastSyncStatus = "failed" group.LastSyncMessage = syncErr.Error() - _ = model.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group) + _ = repository.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group) } func evaluateParsedIPGroupAutoConfig(ctx context.Context, config ipGroupAutoConfig, now time.Time) ([]string, error) { @@ -302,7 +304,7 @@ func evaluateParsedIPGroupAutoConfig(ctx context.Context, config ipGroupAutoConf if lookback <= 0 { lookback = defaultWAFIPGroupAutoLookbackDur } - aggregates, err := model.ListOpenFlareAccessLogWAFIPAggregates(ctx, model.OpenFlareAccessLogQuery{ + aggregates, err := repository.ListOpenFlareAccessLogWAFIPAggregates(ctx, model.OpenFlareAccessLogQuery{ Since: now.Add(-lookback), Until: now, }) diff --git a/internal/apps/openflare/waf/ip_group_sync_test.go b/internal/apps/openflare/waf/ip_group_sync_test.go index da67ae02..9fb5b0f7 100644 --- a/internal/apps/openflare/waf/ip_group_sync_test.go +++ b/internal/apps/openflare/waf/ip_group_sync_test.go @@ -11,6 +11,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/glebarez/sqlite" @@ -32,7 +34,7 @@ func setupIPGroupSyncTestDB(t *testing.T) func() { )) db.SetDB(sqliteDB) - resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) + resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore()) return func() { resetAccessLogStore() db.SetDB(nil) @@ -161,28 +163,28 @@ func TestListDueOpenFlareWAFIPGroups(t *testing.T) { Name: "due auto", Type: wafIPGroupTypeAutomatic, Enabled: true, IPList: "[]", AutoConfig: "{}", ExtIPs: "[]", NextSyncAt: &past, } - require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, dueAuto)) + require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, dueAuto)) futureAuto := &model.OpenFlareWAFIPGroup{ Name: "future auto", Type: wafIPGroupTypeAutomatic, Enabled: true, IPList: "[]", AutoConfig: "{}", ExtIPs: "[]", NextSyncAt: &future, } - require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, futureAuto)) + require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, futureAuto)) dueSub := &model.OpenFlareWAFIPGroup{ Name: "due sub", Type: wafIPGroupTypeSubscription, Enabled: true, IPList: "[]", AutoConfig: "{}", ExtIPs: "[]", SubscriptionURL: "https://example.com/list", NextSyncAt: &past, } - require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, dueSub)) + require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, dueSub)) manual := &model.OpenFlareWAFIPGroup{ Name: "manual", Type: wafIPGroupTypeManual, Enabled: true, IPList: "[]", AutoConfig: "{}", ExtIPs: "[]", NextSyncAt: &past, } - require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, manual)) + require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, manual)) - groups, err := model.ListDueOpenFlareWAFIPGroups(ctx, time.Now().UTC()) + groups, err := repository.ListDueOpenFlareWAFIPGroups(ctx, time.Now().UTC()) require.NoError(t, err) require.Len(t, groups, 2) ids := []uint{groups[0].ID, groups[1].ID} @@ -207,7 +209,7 @@ func seedWAFAccessLogs(t *testing.T, ctx context.Context, loggedAt time.Time, re StatusCode: statusCode, }) } - require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, records)) + require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, records)) } func TestParseIPGroupAutoConfigLookback(t *testing.T) { diff --git a/internal/apps/openflare/waf/logics.go b/internal/apps/openflare/waf/logics.go index 6a869178..499e182a 100644 --- a/internal/apps/openflare/waf/logics.go +++ b/internal/apps/openflare/waf/logics.go @@ -14,6 +14,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" exprlang "github.com/expr-lang/expr" @@ -141,7 +143,7 @@ type ipGroupExtIP struct { // GetSiteRuleGroups returns WAF rule groups for a proxy route. func GetSiteRuleGroups(ctx context.Context, routeID uint) (*SiteRuleGroupsView, error) { - if _, err := model.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil { + if _, err := repository.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil { return nil, err } groups, err := ListRules(ctx) @@ -182,14 +184,14 @@ func GetSiteRuleGroups(ctx context.Context, routeID uint) (*SiteRuleGroupsView, // ReplaceSiteRuleGroups replaces rule group bindings for a proxy route. func ReplaceSiteRuleGroups(ctx context.Context, routeID uint, groupIDs []uint) (*SiteRuleGroupsView, error) { - if _, err := model.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil { + if _, err := repository.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil { return nil, err } normalized, err := normalizeRuleGroupIDs(ctx, groupIDs) if err != nil { return nil, &RuleValidationError{Err: err} } - if err = model.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, routeID, normalized); err != nil { + if err = repository.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, routeID, normalized); err != nil { return nil, err } return GetSiteRuleGroups(ctx, routeID) @@ -197,7 +199,7 @@ func ReplaceSiteRuleGroups(ctx context.Context, routeID uint, groupIDs []uint) ( // ListSiteRuleGroupIDs returns rule group ids bound to a proxy route. func ListSiteRuleGroupIDs(ctx context.Context, routeID uint) ([]uint, error) { - bindings, err := model.ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx, routeID) + bindings, err := repository.ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx, routeID) if err != nil { return nil, err } @@ -210,7 +212,7 @@ func ListSiteRuleGroupIDs(ctx context.Context, routeID uint) ([]uint, error) { // EnsureDefaultRuleGroup ensures the global WAF rule group exists. func EnsureDefaultRuleGroup(ctx context.Context) error { - _, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx) + _, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx) if err == nil { return nil } @@ -224,12 +226,12 @@ func EnsureDefaultRuleGroup(ctx context.Context) error { group := &model.OpenFlareWAFRuleGroup{ Name: "全局规则组", Enabled: true, IsGlobal: true, Graph: string(graph), Revision: 1, } - return model.CreateOpenFlareWAFRuleGroup(ctx, group) + return repository.CreateOpenFlareWAFRuleGroup(ctx, group) } // ListIPGroups returns all WAF IP groups. func ListIPGroups(ctx context.Context) ([]IPGroupView, error) { - groups, err := model.ListOpenFlareWAFIPGroups(ctx) + groups, err := repository.ListOpenFlareWAFIPGroups(ctx) if err != nil { return nil, err } @@ -250,7 +252,7 @@ func ListIPGroups(ctx context.Context) ([]IPGroupView, error) { // GetIPGroup returns a WAF IP group by id. func GetIPGroup(ctx context.Context, id uint) (*IPGroupView, error) { - group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id) if err != nil { return nil, err } @@ -271,7 +273,7 @@ func CreateIPGroup(ctx context.Context, input IPGroupInput) (*IPGroupView, error if err != nil { return nil, &RuleValidationError{Err: err} } - if err = model.CreateOpenFlareWAFIPGroup(ctx, group); err != nil { + if err = repository.CreateOpenFlareWAFIPGroup(ctx, group); err != nil { return nil, err } broadcastIPGroupToAgents(ctx, group.ID) @@ -280,7 +282,7 @@ func CreateIPGroup(ctx context.Context, input IPGroupInput) (*IPGroupView, error // UpdateIPGroup updates a WAF IP group. func UpdateIPGroup(ctx context.Context, id uint, input IPGroupInput) (*IPGroupView, error) { - group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id) if err != nil { return nil, err } @@ -288,7 +290,7 @@ func UpdateIPGroup(ctx context.Context, id uint, input IPGroupInput) (*IPGroupVi if err != nil { return nil, &RuleValidationError{Err: err} } - if err = model.UpdateOpenFlareWAFIPGroup(ctx, group); err != nil { + if err = repository.UpdateOpenFlareWAFIPGroup(ctx, group); err != nil { return nil, err } broadcastIPGroupToAgents(ctx, group.ID) @@ -297,7 +299,7 @@ func UpdateIPGroup(ctx context.Context, id uint, input IPGroupInput) (*IPGroupVi // DeleteIPGroup deletes a WAF IP group when not referenced. func DeleteIPGroup(ctx context.Context, id uint) error { - group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id) if err != nil { return err } @@ -308,12 +310,12 @@ func DeleteIPGroup(ctx context.Context, id uint) error { if counts[group.ID] > 0 { return &RuleValidationError{Err: errors.New("IP 组已被 WAF 规则引用,请先移除引用")} } - return model.DeleteOpenFlareWAFIPGroup(ctx, group.ID) + return repository.DeleteOpenFlareWAFIPGroup(ctx, group.ID) } // SyncIPGroup synchronizes a subscription or automatic WAF IP group. func SyncIPGroup(ctx context.Context, id uint) (*IPGroupSyncResult, error) { - group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id) if err != nil { return nil, err } @@ -341,7 +343,7 @@ func TestIPGroupAutoConfig(ctx context.Context, input IPGroupAutoTestInput) (*IP } func loadRuleGroupBindings(ctx context.Context) (map[uint][]uint, error) { - bindings, err := model.ListOpenFlareWAFRuleGroupBindings(ctx) + bindings, err := repository.ListOpenFlareWAFRuleGroupBindings(ctx) if err != nil { return nil, err } @@ -353,7 +355,7 @@ func loadRuleGroupBindings(ctx context.Context) (map[uint][]uint, error) { } func loadIPGroupReferenceCounts(ctx context.Context) (map[uint]int, error) { - groups, err := model.ListOpenFlareWAFRuleGroups(ctx) + groups, err := repository.ListOpenFlareWAFRuleGroups(ctx) if err != nil { return nil, err } @@ -442,7 +444,7 @@ func decodeStringList(raw string) ([]string, error) { func normalizeRuleGroupIDs(ctx context.Context, groupIDs []uint) ([]uint, error) { normalized := uniqueUintIDsInOrder(groupIDs) for _, groupID := range normalized { - group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, groupID) + group, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, groupID) if err != nil { return nil, fmt.Errorf("WAF 规则组 %d 不存在", groupID) } diff --git a/internal/apps/openflare/waf/logics_test.go b/internal/apps/openflare/waf/logics_test.go index fae99c64..e96d6c54 100644 --- a/internal/apps/openflare/waf/logics_test.go +++ b/internal/apps/openflare/waf/logics_test.go @@ -7,6 +7,8 @@ import ( "context" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/glebarez/sqlite" @@ -57,11 +59,11 @@ func TestUpdateIPGroupPrunesAutomaticExtIPs(t *testing.T) { }) require.NoError(t, err) - group, err := model.GetOpenFlareWAFIPGroupByID(ctx, created.ID) + group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, created.ID) require.NoError(t, err) group.IPList = `["203.0.113.10","203.0.113.11"]` group.ExtIPs = `[{"ip":"203.0.113.10","captured_at":"2026-06-18T10:00:00Z"},{"ip":"203.0.113.11","captured_at":"2026-06-18T11:00:00Z"}]` - require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, group)) + require.NoError(t, repository.UpdateOpenFlareWAFIPGroup(ctx, group)) updated, err := UpdateIPGroup(ctx, created.ID, IPGroupInput{ Name: created.Name, diff --git a/internal/apps/openflare/waf/rule_logics.go b/internal/apps/openflare/waf/rule_logics.go index deae54fa..c128f293 100644 --- a/internal/apps/openflare/waf/rule_logics.go +++ b/internal/apps/openflare/waf/rule_logics.go @@ -11,6 +11,8 @@ import ( "strings" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/model" "gorm.io/gorm" ) @@ -57,7 +59,7 @@ func ListRules(ctx context.Context) ([]RuleView, error) { if err := EnsureDefaultRuleGroup(ctx); err != nil { return nil, err } - groups, err := model.ListOpenFlareWAFRuleGroups(ctx) + groups, err := repository.ListOpenFlareWAFRuleGroups(ctx) if err != nil { return nil, err } @@ -78,7 +80,7 @@ func ListRules(ctx context.Context) ([]RuleView, error) { // GetRule returns one orchestrated WAF rule. func GetRule(ctx context.Context, id uint) (*RuleView, error) { - group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id) + group, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, id) if err != nil { return nil, err } @@ -101,13 +103,13 @@ func CreateRule(ctx context.Context, input CreateRuleInput) (*RuleView, error) { return nil, err } group := &model.OpenFlareWAFRuleGroup{Name: name, Enabled: false, IsGlobal: false, Graph: string(raw), Revision: 1} - if err = model.CreateOpenFlareWAFRuleGroup(ctx, group); err != nil { + if err = repository.CreateOpenFlareWAFRuleGroup(ctx, group); err != nil { return nil, err } // GORM applies the model's database default to a false bool on Create, so // explicitly persist the safe disabled state after the row has an ID. group.Enabled = false - if err = model.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil { + if err = repository.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil { return nil, err } return GetRule(ctx, group.ID) @@ -115,7 +117,7 @@ func CreateRule(ctx context.Context, input CreateRuleInput) (*RuleView, error) { // UpdateRuleMeta updates rule metadata without touching its graph revision. func UpdateRuleMeta(ctx context.Context, id uint, input UpdateRuleMetaInput) (*RuleView, error) { - group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id) + group, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, id) if err != nil { return nil, err } @@ -124,7 +126,7 @@ func UpdateRuleMeta(ctx context.Context, id uint, input UpdateRuleMetaInput) (*R return nil, &RuleValidationError{Err: errors.New("WAF 规则名称不能为空")} } group.Name, group.Enabled = name, input.Enabled - if err = model.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil { + if err = repository.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil { return nil, err } return GetRule(ctx, id) @@ -132,19 +134,19 @@ func UpdateRuleMeta(ctx context.Context, id uint, input UpdateRuleMetaInput) (*R // DeleteRuleGroup deletes a non-global orchestrated WAF rule. func DeleteRuleGroup(ctx context.Context, id uint) error { - group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id) + group, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, id) if err != nil { return err } if group.IsGlobal { return &RuleValidationError{Err: errors.New("全局 WAF 规则不能删除")} } - return model.DeleteOpenFlareWAFRuleGroupWithBindings(ctx, id) + return repository.DeleteOpenFlareWAFRuleGroupWithBindings(ctx, id) } // SaveRuleGraph validates and atomically replaces a rule graph. func SaveRuleGraph(ctx context.Context, id uint, input SaveRuleGraphInput) (*RuleView, error) { - if _, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id); err != nil { + if _, err := repository.GetOpenFlareWAFRuleGroupByID(ctx, id); err != nil { return nil, err } if err := ValidateRuleGraph(ctx, input.Graph, ruleIPGroupExists); err != nil { @@ -154,14 +156,14 @@ func SaveRuleGraph(ctx context.Context, id uint, input SaveRuleGraphInput) (*Rul if err != nil { return nil, err } - if _, err = model.UpdateOpenFlareWAFRuleGraph(ctx, id, input.Revision, string(raw)); err != nil { + if _, err = repository.UpdateOpenFlareWAFRuleGraph(ctx, id, input.Revision, string(raw)); err != nil { return nil, err } return GetRule(ctx, id) } func ruleIPGroupExists(ctx context.Context, id uint) (bool, error) { - _, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) + _, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id) if errors.Is(err, gorm.ErrRecordNotFound) { return false, nil } diff --git a/internal/apps/openflare/waf/rule_logics_test.go b/internal/apps/openflare/waf/rule_logics_test.go index 27807607..71546a2d 100644 --- a/internal/apps/openflare/waf/rule_logics_test.go +++ b/internal/apps/openflare/waf/rule_logics_test.go @@ -13,6 +13,8 @@ import ( "strconv" "testing" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/shared/response" @@ -165,7 +167,7 @@ func TestReplaceSiteRuleGroupsPreservesOrderAndRejectsGlobal(t *testing.T) { assert.Equal(t, []uint{third.ID, first.ID, second.ID}, view.AppliedIDs) require.NoError(t, EnsureDefaultRuleGroup(ctx)) - global, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx) + global, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx) require.NoError(t, err) _, err = ReplaceSiteRuleGroups(ctx, 7, []uint{global.ID, second.ID}) require.Error(t, err) diff --git a/internal/apps/openflare/zone/logics.go b/internal/apps/openflare/zone/logics.go index 146acc3a..df518b67 100644 --- a/internal/apps/openflare/zone/logics.go +++ b/internal/apps/openflare/zone/logics.go @@ -10,8 +10,8 @@ import ( "strings" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "golang.org/x/net/publicsuffix" "gorm.io/gorm" ) @@ -74,7 +74,7 @@ func Create(ctx context.Context, input Input) (*model.Zone, error) { return nil, errors.New(errZoneRootInvalid) } zone := &model.Zone{Domain: domain} - if err := db.DB(ctx).Create(zone).Error; err != nil { + if err := repository.CreateZone(ctx, zone); err != nil { if isUnique(err) { return nil, errors.New(errDomainExists) } @@ -85,8 +85,8 @@ func Create(ctx context.Context, input Input) (*model.Zone, error) { // Update replaces a Zone's mutable fields. func Update(ctx context.Context, id uint, input Input) (*model.Zone, error) { - var zone model.Zone - if err := db.DB(ctx).First(&zone, id).Error; err != nil { + zone, err := repository.GetZoneByID(ctx, id) + if err != nil { return nil, err } domain, err := normalizeDomain(input.Domain) @@ -98,31 +98,24 @@ func Update(ctx context.Context, id uint, input Input) (*model.Zone, error) { return nil, errors.New(errZoneRootInvalid) } zone.Domain = domain - if err := db.DB(ctx).Save(&zone).Error; err != nil { + if err := repository.SaveZone(ctx, zone); err != nil { if isUnique(err) { return nil, errors.New(errDomainExists) } return nil, err } - return &zone, nil + return zone, nil } // List returns all Zones in stable domain order, with domain counts for list cards. func List(ctx context.Context) ([]ListItem, error) { - var zones []model.Zone - if err := db.DB(ctx).Order("domain asc").Find(&zones).Error; err != nil { + zones, err := repository.ListZones(ctx) + if err != nil { return nil, err } - type countRow struct { - ZoneID uint `gorm:"column:zone_id"` - Count int64 `gorm:"column:count"` - } - var rows []countRow - if err := db.DB(ctx).Model(&model.ZoneDomain{}). - Select("zone_id, count(*) as count"). - Group("zone_id"). - Scan(&rows).Error; err != nil { + rows, err := repository.ListZoneDomainCounts(ctx) + if err != nil { return nil, err } counts := make(map[uint]int64, len(rows)) @@ -145,21 +138,21 @@ func List(ctx context.Context) ([]ListItem, error) { // GetOverview returns a Zone and its domains. func GetOverview(ctx context.Context, id uint) (*Overview, error) { - var zone model.Zone - if err := db.DB(ctx).First(&zone, id).Error; err != nil { + zone, err := repository.GetZoneByID(ctx, id) + if err != nil { return nil, err } - var domains []model.ZoneDomain - if err := db.DB(ctx).Where("zone_id = ?", id).Order("domain asc").Find(&domains).Error; err != nil { + domains, err := repository.ListZoneDomainsByZoneID(ctx, id) + if err != nil { return nil, err } - return &Overview{Zone: zone, Domains: domains}, nil + return &Overview{Zone: *zone, Domains: domains}, nil } // CreateDomain adds a validated exact hostname to a Zone. func CreateDomain(ctx context.Context, zoneID uint, input DomainInput) (*model.ZoneDomain, error) { - var zone model.Zone - if err := db.DB(ctx).First(&zone, zoneID).Error; err != nil { + zone, err := repository.GetZoneByID(ctx, zoneID) + if err != nil { return nil, err } domain, err := normalizeDomain(input.Domain) @@ -171,12 +164,12 @@ func CreateDomain(ctx context.Context, zoneID uint, input DomainInput) (*model.Z return nil, errors.New(errDomainOutsideZone) } if input.CertID != nil { - if _, err := model.GetTLSCertificateByID(ctx, *input.CertID); err != nil { + if _, err := repository.GetTLSCertificateByID(ctx, *input.CertID); err != nil { return nil, errors.New(errCertificateNotFound) } } item := &model.ZoneDomain{ZoneID: zoneID, Domain: domain, CertID: input.CertID} - if err := db.DB(ctx).Create(item).Error; err != nil { + if err := repository.CreateZoneDomain(ctx, item); err != nil { if isUnique(err) { return nil, errors.New(errDomainExists) } @@ -187,16 +180,16 @@ func CreateDomain(ctx context.Context, zoneID uint, input DomainInput) (*model.Z // UpdateDomain replaces a Zone-domain's mutable fields. func UpdateDomain(ctx context.Context, zoneID, id uint, input DomainInput) (*model.ZoneDomain, error) { - var item model.ZoneDomain - if err := db.DB(ctx).Where("id = ? AND zone_id = ?", id, zoneID).First(&item).Error; err != nil { + item, err := repository.GetZoneDomainByZoneAndID(ctx, zoneID, id) + if err != nil { return nil, err } domain, err := normalizeDomain(input.Domain) if err != nil { return nil, err } - var zone model.Zone - if err = db.DB(ctx).First(&zone, zoneID).Error; err != nil { + zone, err := repository.GetZoneByID(ctx, zoneID) + if err != nil { return nil, err } root, err := zoneRoot(domain) @@ -204,46 +197,45 @@ func UpdateDomain(ctx context.Context, zoneID, id uint, input DomainInput) (*mod return nil, errors.New(errDomainOutsideZone) } if input.CertID != nil { - if _, err = model.GetTLSCertificateByID(ctx, *input.CertID); err != nil { + if _, err = repository.GetTLSCertificateByID(ctx, *input.CertID); err != nil { return nil, errors.New(errCertificateNotFound) } } item.Domain, item.CertID = domain, input.CertID - if err = db.DB(ctx).Save(&item).Error; err != nil { + if err = repository.SaveZoneDomain(ctx, item); err != nil { if isUnique(err) { return nil, errors.New(errDomainExists) } return nil, err } - return &item, nil + return item, nil } // DeleteDomain removes a Zone domain that is not bound to a proxy route. func DeleteDomain(ctx context.Context, zoneID, id uint) error { - var item model.ZoneDomain - if err := db.DB(ctx).Where("id = ? AND zone_id = ?", id, zoneID).First(&item).Error; err != nil { + item, err := repository.GetZoneDomainByZoneAndID(ctx, zoneID, id) + if err != nil { return err } if item.ProxyRouteID != nil { return errors.New(errDomainBoundToRoute) } - return db.DB(ctx).Delete(&item).Error + return repository.DeleteZoneDomain(ctx, item) } // Delete removes a Zone that has no remaining domains. func Delete(ctx context.Context, id uint) error { - var zone model.Zone - if err := db.DB(ctx).First(&zone, id).Error; err != nil { + if _, err := repository.GetZoneByID(ctx, id); err != nil { return err } - var count int64 - if err := db.DB(ctx).Model(&model.ZoneDomain{}).Where("zone_id = ?", id).Count(&count).Error; err != nil { + count, err := repository.CountZoneDomainsByZoneID(ctx, id) + if err != nil { return err } if count > 0 { return errors.New(errZoneHasDomains) } - return db.DB(ctx).Delete(&zone).Error + return repository.DeleteZone(ctx, id) } func isUnique(err error) bool { diff --git a/internal/apps/openflare/zone/logics_test.go b/internal/apps/openflare/zone/logics_test.go index d9a36aee..da735fd7 100644 --- a/internal/apps/openflare/zone/logics_test.go +++ b/internal/apps/openflare/zone/logics_test.go @@ -8,6 +8,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/glebarez/sqlite" @@ -41,13 +43,13 @@ func TestDeleteDomainRejectsBoundRoute(t *testing.T) { require.NoError(t, err) routeID := uint(9) item.ProxyRouteID = &routeID - require.NoError(t, db.DB(ctx).Save(item).Error) + require.NoError(t, repository.SaveZoneDomain(ctx, item)) err = DeleteDomain(ctx, zone.ID, item.ID) require.EqualError(t, err, errDomainBoundToRoute) item.ProxyRouteID = nil - require.NoError(t, db.DB(ctx).Save(item).Error) + require.NoError(t, repository.SaveZoneDomain(ctx, item)) require.NoError(t, DeleteDomain(ctx, zone.ID, item.ID)) } @@ -59,7 +61,7 @@ func TestLegacyImportUsesEffectiveTLDPlusOne(t *testing.T) { func TestGetStatsAggregatesZoneHosts(t *testing.T) { ctx := setupZoneDB(t) - reset := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) + reset := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore()) t.Cleanup(reset) zone, err := Create(ctx, Input{Domain: "example.com"}) @@ -70,7 +72,7 @@ func TestGetStatsAggregatesZoneHosts(t *testing.T) { require.NoError(t, err) now := time.Now().UTC() - require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{ + require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{ {NodeID: "n1", LoggedAt: now.Add(-1 * time.Hour), RemoteAddr: "1.1.1.1", Host: "api.example.com", Path: "/", StatusCode: 200, BytesSent: 1000}, {NodeID: "n1", LoggedAt: now.Add(-2 * time.Hour), RemoteAddr: "1.1.1.1", Host: "www.example.com", Path: "/", StatusCode: 200, BytesSent: 500}, {NodeID: "n1", LoggedAt: now.Add(-3 * time.Hour), RemoteAddr: "2.2.2.2", Host: "api.example.com", Path: "/x", StatusCode: 404, BytesSent: 200}, diff --git a/internal/apps/openflare/zone/stats.go b/internal/apps/openflare/zone/stats.go index 23f3207a..6a5dcfc9 100644 --- a/internal/apps/openflare/zone/stats.go +++ b/internal/apps/openflare/zone/stats.go @@ -9,8 +9,8 @@ import ( "strings" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "gorm.io/gorm" ) @@ -80,13 +80,12 @@ func GetStats(ctx context.Context, id uint, rangeRaw string) (*Stats, error) { return nil, err } - var zone model.Zone - if err := db.DB(ctx).First(&zone, id).Error; err != nil { + if _, err := repository.GetZoneByID(ctx, id); err != nil { return nil, err } - var domains []model.ZoneDomain - if err := db.DB(ctx).Where("zone_id = ?", id).Order("domain asc").Find(&domains).Error; err != nil { + domains, err := repository.ListZoneDomainsByZoneID(ctx, id) + if err != nil { return nil, err } @@ -120,7 +119,7 @@ func GetStats(ctx context.Context, id uint, rangeRaw string) (*Stats, error) { return result, nil } - requestCount, uniqueVisitors, totalBytesSent, err := model.CountOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ + requestCount, uniqueVisitors, totalBytesSent, err := repository.CountOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ Hosts: hosts, Since: since, Until: now, @@ -136,7 +135,7 @@ func GetStats(ctx context.Context, id uint, rangeRaw string) (*Stats, error) { result.UniqueVisitors = uniqueVisitors result.BytesSent = totalBytesSent - buckets, err := model.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ + buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ Hosts: hosts, Since: since, Until: now, diff --git a/internal/apps/upload/cache/meta_cache.go b/internal/apps/upload/cache/meta_cache.go index ff2b8691..f445d006 100644 --- a/internal/apps/upload/cache/meta_cache.go +++ b/internal/apps/upload/cache/meta_cache.go @@ -11,6 +11,7 @@ import ( db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/cache/ram" ) @@ -99,10 +100,8 @@ func GetUploadByID(ctx context.Context, id uint64) (model.Upload, error) { } } - var upload model.Upload - if err := db.DB(ctx). - Where("id = ? AND status IN (?, ?)", id, model.UploadStatusPending, model.UploadStatusUsed). - First(&upload).Error; err != nil { + upload, err := repository.GetCacheableUploadByID(ctx, id) + if err != nil { return model.Upload{}, err } diff --git a/internal/apps/upload/ingest/helpers.go b/internal/apps/upload/ingest/helpers.go index 3000cd8c..5905bbf4 100644 --- a/internal/apps/upload/ingest/helpers.go +++ b/internal/apps/upload/ingest/helpers.go @@ -11,17 +11,17 @@ import ( "strings" "time" + "gorm.io/gorm" + 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/infra/objectstore" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" - "gorm.io/gorm" ) func normalizeRequest(req *Request) { @@ -122,7 +122,8 @@ func cleanupUnpersistedObject(ctx context.Context, objectKey string) { } func createUploadWithStats(ctx context.Context, upload *model.Upload) error { - return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + // Multi-step: create upload row + apply incremental stats in one transaction. + return repository.RunInTransaction(ctx, func(tx *gorm.DB) error { if err := repository.CreateUploadTx(tx, upload); err != nil { return err } diff --git a/internal/apps/upload/ingest/remove.go b/internal/apps/upload/ingest/remove.go index 4a20100e..8ebb6a7c 100644 --- a/internal/apps/upload/ingest/remove.go +++ b/internal/apps/upload/ingest/remove.go @@ -6,14 +6,13 @@ package ingest import ( "context" + "gorm.io/gorm" + 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" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository" - "gorm.io/gorm" - "gorm.io/gorm/clause" ) // Remove soft-deletes an ordinary upload and decrements incremental stats once. @@ -36,20 +35,23 @@ func RemoveOwned(ctx context.Context, userID, uploadID uint64) (model.Upload, er func remove(ctx context.Context, userID, uploadID uint64, owned bool) (model.Upload, error) { var upload model.Upload - if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). - Where("id = ?", uploadID). - First(&upload).Error; err != nil { + // Multi-step: FOR UPDATE lock + ownership/type checks + soft-delete + stats delta. + if err := repository.RunInTransaction(ctx, func(tx *gorm.DB) error { + locked, err := repository.GetUploadByIDForUpdateTx(tx, uploadID) + if err != nil { return err } - if owned && upload.UserID != userID { + if owned && locked.UserID != userID { return ErrForbidden } - if upload.Type == shared.ReservedPagesDeploymentType { + if locked.Type == shared.ReservedPagesDeploymentType { return ErrReservedUploadType } - _, err := RemoveLockedTx(tx, &upload) - return err + if _, err := RemoveLockedTx(tx, &locked); err != nil { + return err + } + upload = locked + return nil }); err != nil { return model.Upload{}, err } diff --git a/internal/apps/upload/stats/stats_counter.go b/internal/apps/upload/stats/stats_counter.go index b43f631d..95342593 100644 --- a/internal/apps/upload/stats/stats_counter.go +++ b/internal/apps/upload/stats/stats_counter.go @@ -5,13 +5,12 @@ package stats import ( "context" - "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/pkg/logger" "gorm.io/gorm" - "gorm.io/gorm/clause" + + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/pkg/logger" ) // ApplyUploadStatsAdd increments incremental stats for a newly active upload record. @@ -26,22 +25,8 @@ func ApplyUploadStatsRemove(ctx context.Context, upload *model.Upload) error { // RebuildUploadStats rebuilds all incremental stats from current upload records. func RebuildUploadStats(ctx context.Context) error { - return db.DB(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Where("1 = 1").Delete(&model.UploadStat{}).Error; err != nil { - return err - } - - var uploads []model.Upload - if err := tx.Where("status != ?", model.UploadStatusDeleted).Find(&uploads).Error; err != nil { - return err - } - - for i := range uploads { - if err := ApplyUploadStatsDeltaTx(tx, &uploads[i], 1); err != nil { - return err - } - } - return nil + return repository.RebuildUploadStats(ctx, func(tx *gorm.DB, upload *model.Upload) error { + return ApplyUploadStatsDeltaTx(tx, upload, 1) }) } @@ -49,7 +34,7 @@ func applyUploadStatsDelta(ctx context.Context, upload *model.Upload, sign int64 if upload == nil || !isActiveUploadStatus(upload.Status) { return nil } - return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + return repository.RunInTransaction(ctx, func(tx *gorm.DB) error { return ApplyUploadStatsDeltaTx(tx, upload, sign) }) } @@ -78,40 +63,13 @@ func ApplyUploadStatsDeltaTx(tx *gorm.DB, upload *model.Upload, sign int64) erro } for _, entry := range entries { - if err := upsertUploadStatDelta(tx, entry.dimension, entry.key, countDelta, sizeDelta); err != nil { + if err := repository.UpsertUploadStatDeltaTx(tx, entry.dimension, entry.key, countDelta, sizeDelta); err != nil { return err } } return nil } -func upsertUploadStatDelta(tx *gorm.DB, dimension, key string, countDelta, sizeDelta int64) error { - return tx.Clauses(clause.OnConflict{ - Columns: []clause.Column{ - {Name: "dimension"}, - {Name: "stat_key"}, - }, - DoUpdates: clause.Assignments(map[string]any{ - "file_count": gorm.Expr( - "CASE WHEN w_upload_stats.file_count + ? < 0 THEN 0 ELSE w_upload_stats.file_count + ? END", - countDelta, - countDelta, - ), - "file_size": gorm.Expr( - "CASE WHEN w_upload_stats.file_size + ? < 0 THEN 0 ELSE w_upload_stats.file_size + ? END", - sizeDelta, - sizeDelta, - ), - "updated_at": time.Now(), - }), - }).Create(&model.UploadStat{ - Dimension: dimension, - StatKey: key, - FileCount: countDelta, - FileSize: sizeDelta, - }).Error -} - // RecordUploadStatsAdd logs and applies upload stats increment. func RecordUploadStatsAdd(ctx context.Context, upload *model.Upload) { if err := ApplyUploadStatsAdd(ctx, upload); err != nil { diff --git a/internal/apps/upload/stats/stats_counter_test.go b/internal/apps/upload/stats/stats_counter_test.go index 5e197716..850ff83f 100644 --- a/internal/apps/upload/stats/stats_counter_test.go +++ b/internal/apps/upload/stats/stats_counter_test.go @@ -8,10 +8,11 @@ import ( "testing" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/model" - "github.com/Rain-kl/Wavelet/internal/testhelper" "gorm.io/gorm" + + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/testhelper" ) func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) { @@ -29,7 +30,7 @@ func TestApplyUploadStatsDeltaTxWithinTransaction(t *testing.T) { CreatedAt: time.Now(), } - if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := repository.RunInTransaction(ctx, func(tx *gorm.DB) error { return ApplyUploadStatsDeltaTx(tx, upload, 1) }); err != nil { t.Fatalf("ApplyUploadStatsDeltaTx returned error: %v", err) @@ -89,8 +90,8 @@ type uploadStatsSnapshot struct { } func loadUploadStats(ctx context.Context) (uploadStatsSnapshot, error) { - var rows []model.UploadStat - if err := db.DB(ctx).Where("dimension = ?", model.UploadStatDimensionTotal).Find(&rows).Error; err != nil { + rows, err := repository.ListUploadStatsByDimension(ctx, model.UploadStatDimensionTotal) + if err != nil { return uploadStatsSnapshot{}, err } if len(rows) == 0 { diff --git a/internal/apps/upload/storage/migration.go b/internal/apps/upload/storage/migration.go index 59f8c46d..3fa6015b 100644 --- a/internal/apps/upload/storage/migration.go +++ b/internal/apps/upload/storage/migration.go @@ -11,9 +11,8 @@ import ( "strings" "github.com/Rain-kl/Wavelet/internal/infra/objectstore" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" - "gorm.io/gorm" + "github.com/Rain-kl/Wavelet/internal/repository" ) // StorageMigrationTask is the Asynq task name for storage migration. @@ -21,18 +20,7 @@ const StorageMigrationTask = "storage:migrate" // LatestMigrationExecution returns the most recent storage migration task execution. func LatestMigrationExecution(ctx context.Context) (*model.TaskExecution, bool, error) { - var execution model.TaskExecution - err := db.DB(ctx). - Where("task_type = ?", StorageMigrationTask). - Order("id DESC"). - First(&execution).Error - if err == nil { - return &execution, true, nil - } - if !errors.Is(err, gorm.ErrRecordNotFound) { - return nil, false, err - } - return nil, false, nil + return repository.GetLatestTaskExecutionByTaskType(ctx, StorageMigrationTask) } // ParseMigrationTargetConfig parses and validates a storage migration target payload. diff --git a/internal/apps/upload/task/cleanup.go b/internal/apps/upload/task/cleanup.go index f8220c8f..6f522f14 100644 --- a/internal/apps/upload/task/cleanup.go +++ b/internal/apps/upload/task/cleanup.go @@ -10,15 +10,15 @@ import ( "fmt" "time" + "gorm.io/gorm" + "github.com/Rain-kl/Wavelet/internal/apps/upload/ingest" "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/pkg/logger" - "gorm.io/gorm" - "gorm.io/gorm/clause" ) const ( @@ -58,12 +58,8 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas task.AppendLog(ctx, "开始扫描未使用上传文件,阈值: %s", oneHourAgo.Format(time.RFC3339)) for { - var unusedUploads []model.Upload - if err := db.DB(ctx). - Where("id > ? AND status = ? AND created_at < ?", lastID, model.UploadStatusPending, oneHourAgo). - Order("id ASC"). - Limit(batchSize). - Find(&unusedUploads).Error; err != nil { + unusedUploads, err := repository.ListPendingUploadsOlderThan(ctx, lastID, oneHourAgo, batchSize) + if err != nil { task.AppendLog(ctx, "查询未使用的上传文件失败: %v", err) return nil, fmt.Errorf(shared.ErrQueryUnusedUploadsFailed, err) } @@ -78,19 +74,18 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas totalProcessed++ transitioned := false - if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { - var locked model.Upload - if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). - Where("id = ?", u.ID). - First(&locked).Error; err != nil { + // Multi-step: row lock + ownership-safe soft-delete + stats delta stay orchestrated here. + if err := repository.RunInTransaction(ctx, func(tx *gorm.DB) error { + locked, err := repository.GetUploadByIDForUpdateTx(tx, u.ID) + if err != nil { return err } if locked.Status != model.UploadStatusPending || !locked.CreatedAt.Before(oneHourAgo) { return nil } - var err error - transitioned, err = ingest.RemoveLockedTx(tx, &locked) - return err + var removeErr error + transitioned, removeErr = ingest.RemoveLockedTx(tx, &locked) + return removeErr }); err != nil { task.AppendLog(ctx, "清理上传文件失败 [ID:%d]: %v", u.ID, err) lastID = u.ID @@ -107,21 +102,22 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas task.AppendLog(ctx, "开始清理历史推送审计日志,只保留最近7天数据...") cutoff := time.Now().AddDate(0, 0, -7) - var pushHistoryCount int64 - if err := db.DB(ctx).Model(&model.PushHistory{}).Where("created_at < ?", cutoff).Count(&pushHistoryCount).Error; err != nil { + pushHistoryCount, err := repository.CountPushHistoriesCreatedBefore(ctx, cutoff) + switch { + case err != nil: task.AppendLog(ctx, "统计待清理的历史推送记录失败: %v", err) - } else if pushHistoryCount > 0 { - if err := db.DB(ctx).Where("created_at < ?", cutoff).Delete(&model.PushHistory{}).Error; err != nil { - task.AppendLog(ctx, "删除历史推送记录失败: %v", err) + case pushHistoryCount == 0: + task.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05")) + default: + if _, delErr := repository.DeletePushHistoriesCreatedBefore(ctx, cutoff); delErr != nil { + task.AppendLog(ctx, "删除历史推送记录失败: %v", delErr) } else { task.AppendLog(ctx, "成功删除 %d 条历史推送记录 (截止时间: %s)", pushHistoryCount, cutoff.Format("2006-01-02 15:04:05")) } - } else { - task.AppendLog(ctx, "没有需要清理的历史推送记录 (截止时间: %s)", cutoff.Format("2006-01-02 15:04:05")) } task.AppendLog(ctx, "开始清理任务执行日志:高频任务保留最近3天,低频任务保留最近30天...") - taskLogStats, err := model.CleanupTaskExecutionLogs(ctx, time.Now()) + taskLogStats, err := repository.CleanupTaskExecutionLogs(ctx, time.Now()) if err != nil { task.AppendLog(ctx, "清理任务执行日志失败: %v", err) logger.ErrorF(ctx, "清理任务执行日志失败: %v", err) diff --git a/internal/apps/upload/task/rebuild_stats.go b/internal/apps/upload/task/rebuild_stats.go index e5a3a163..2d0b59cf 100644 --- a/internal/apps/upload/task/rebuild_stats.go +++ b/internal/apps/upload/task/rebuild_stats.go @@ -8,9 +8,8 @@ import ( "fmt" uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" ) const ( @@ -37,11 +36,8 @@ type RebuildUploadStatsHandler struct{} // Execute scans active uploads and rebuilds all upload stat dimensions. func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*task.TaskResult, error) { - var activeCount int64 - if err := db.DB(ctx). - Model(&model.Upload{}). - Where("status != ?", model.UploadStatusDeleted). - Count(&activeCount).Error; err != nil { + activeCount, err := repository.CountActiveUploads(ctx) + if err != nil { task.AppendLog(ctx, "统计活跃上传记录失败: %v", err) return nil, fmt.Errorf("count active uploads: %w", err) } @@ -53,10 +49,8 @@ func (h *RebuildUploadStatsHandler) Execute(ctx context.Context, _ []byte) (*tas return nil, fmt.Errorf("rebuild upload stats: %w", err) } - var totalStat model.UploadStat - if err := db.DB(ctx). - Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, ""). - First(&totalStat).Error; err != nil { + totalStat, err := repository.GetTotalUploadStat(ctx) + if err != nil { task.AppendLog(ctx, "读取总量统计失败: %v", err) return nil, fmt.Errorf("load total upload stats: %w", err) } diff --git a/internal/apps/upload/task/storage_migration.go b/internal/apps/upload/task/storage_migration.go index bcf79f40..616d2d56 100644 --- a/internal/apps/upload/task/storage_migration.go +++ b/internal/apps/upload/task/storage_migration.go @@ -15,13 +15,15 @@ import ( "sync/atomic" "time" + "golang.org/x/sync/errgroup" + 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/infra/objectstore" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/model" - "golang.org/x/sync/errgroup" + "github.com/Rain-kl/Wavelet/internal/repository" ) const ( @@ -171,12 +173,7 @@ func (h *MigrationHandler) Execute(ctx context.Context, payload []byte) (*task.T } func countStorageObjects(ctx context.Context) (int64, error) { - var count int64 - err := db.DB(ctx).Model(&model.Upload{}). - Where("status != ?", model.UploadStatusDeleted). - Distinct("file_path"). - Count(&count).Error - return count, err + return repository.CountDistinctActiveFilePaths(ctx) } func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) { @@ -187,13 +184,6 @@ func hasUnresolvedMigrationTask(ctx context.Context) (bool, error) { return execution.Status == model.TaskExecutionStatusPending || execution.Status == model.TaskExecutionStatusRunning, nil } -type migrationObject struct { - FilePath string `gorm:"column:file_path"` - FileSize int64 `gorm:"column:file_size"` - MimeType string `gorm:"column:mime_type"` - Hash string `gorm:"column:hash"` -} - func migrateObjects( ctx context.Context, sourceBackend objectstore.Backend, @@ -212,17 +202,8 @@ func migrateObjects( task.AppendLog(ctx, "正在查询待迁移对象批次,当前已完成迁移: %d/%d", atomic.LoadInt64(&migrated), total) - var objects []migrationObject - query := db.DB(ctx).Model(&model.Upload{}). - Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type, MAX(hash) AS hash"). - Where("status != ?", model.UploadStatusDeleted) - if lastFilePath != "" { - query = query.Where("file_path > ?", lastFilePath) - } - if err := query.Group("file_path"). - Order("file_path ASC"). - Limit(batchSize). - Scan(&objects).Error; err != nil { + objects, err := repository.ListDistinctActiveStorageObjects(ctx, lastFilePath, batchSize) + if err != nil { return atomic.LoadInt64(&migrated), fmt.Errorf("query source objects: %w", err) } if len(objects) == 0 { @@ -260,7 +241,7 @@ func migrateSingleObject( ctx context.Context, sourceBackend objectstore.Backend, targetBackend objectstore.Backend, - obj migrationObject, + obj repository.UploadStorageObject, sha256HexLength int, ) error { if shouldSkipMigration(ctx, targetBackend, obj) { @@ -310,9 +291,7 @@ func migrateSingleObject( if targetResult.Key != obj.FilePath { task.AppendLog(ctx, "[更新数据库] 正在更新文件路径: %s -> %s", obj.FilePath, targetResult.Key) - if err := db.DB(ctx).Model(&model.Upload{}). - Where("file_path = ? AND status != ?", obj.FilePath, model.UploadStatusDeleted). - Update("file_path", targetResult.Key).Error; err != nil { + if err := repository.UpdateActiveUploadsFilePath(ctx, obj.FilePath, targetResult.Key); err != nil { return fmt.Errorf("update migrated object %q: %w", obj.FilePath, err) } } @@ -323,7 +302,7 @@ func migrateSingleObject( func shouldSkipMigration( ctx context.Context, targetBackend objectstore.Backend, - obj migrationObject, + obj repository.UploadStorageObject, ) bool { targetObj, err := targetBackend.Get(ctx, obj.FilePath) if err != nil || targetObj == nil || targetObj.Body == nil { @@ -343,16 +322,9 @@ func markMissingMigrationObjectDeleted( ) error { task.AppendLog(ctx, "警告: 源存储中物理文件不存在,标记为已删除并跳过: %s (错误: %v)", filePath, sourceErr) - var affectedUploads []model.Upload - if err := db.DB(ctx). - Where("file_path = ? AND status != ?", filePath, model.UploadStatusDeleted). - Find(&affectedUploads).Error; err != nil { - return fmt.Errorf("load missing object uploads %q: %w", filePath, err) - } - if err := db.DB(ctx).Model(&model.Upload{}). - Where("file_path = ?", filePath). - Update("status", model.UploadStatusDeleted).Error; err != nil { - return fmt.Errorf("update missing object %q: %w", filePath, err) + affectedUploads, err := repository.MarkActiveUploadsDeletedByFilePath(ctx, filePath) + if err != nil { + return fmt.Errorf("mark missing object deleted %q: %w", filePath, err) } for i := range affectedUploads { uploadstats.RecordUploadStatsRemove(ctx, &affectedUploads[i]) diff --git a/internal/apps/upload/task/tasks.go b/internal/apps/upload/task/tasks.go index a36a6a96..f21056ba 100644 --- a/internal/apps/upload/task/tasks.go +++ b/internal/apps/upload/task/tasks.go @@ -14,9 +14,8 @@ import ( "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/repository" ) const ( @@ -113,17 +112,8 @@ func (h *WarmImageCacheHandler) Execute(ctx context.Context, payload []byte) (*t return nil, fmt.Errorf("image cache warmup canceled: %w", err) } - var uploads []model.Upload - if err := db.DB(ctx). - Where("id > ? AND status != ? AND (LOWER(mime_type) LIKE ? OR LOWER(extension) IN ?)", - lastID, - model.UploadStatusDeleted, - "image/%", - []string{"jpg", "jpeg", "png", "webp", "gif"}, - ). - Order("id ASC"). - Limit(batchSize). - Find(&uploads).Error; err != nil { + uploads, err := repository.ListActiveImageUploadsAfterID(ctx, lastID, batchSize) + if err != nil { task.AppendLog(ctx, "查询图片上传记录失败: %v", err) return nil, fmt.Errorf(shared.ErrQueryImagesForCacheWarmup, err) } diff --git a/internal/apps/upload/task/tasks_test.go b/internal/apps/upload/task/tasks_test.go index adab09a0..08ca1d1d 100644 --- a/internal/apps/upload/task/tasks_test.go +++ b/internal/apps/upload/task/tasks_test.go @@ -17,6 +17,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv" "github.com/Rain-kl/Wavelet/internal/apps/upload/shared" uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats" @@ -129,7 +131,7 @@ func TestSystemCleanupHandler_Execute(t *testing.T) { UpdatedAt: now.AddDate(0, 0, -31), TriggeredBy: "system", } - err = model.CreateTaskExecution(ctx, oldTaskLog) + err = repository.CreateTaskExecution(ctx, oldTaskLog) require.NoError(t, err) // 执行 handler diff --git a/internal/apps/user/logics.go b/internal/apps/user/logics.go index bd97a277..4fc79faa 100644 --- a/internal/apps/user/logics.go +++ b/internal/apps/user/logics.go @@ -218,8 +218,8 @@ func sendRegisterEmailCode(ctx context.Context, email string) error { return errors.New(errEmailRequired) } - var count int64 - if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", email).Count(&count).Error; err != nil { + count, err := repository.CountUsersByEmail(ctx, email) + if err != nil { return err } if count > 0 { @@ -249,8 +249,8 @@ func validateRegisterEmailVerification(ctx context.Context, email, code string) } func updateUserProfile(ctx context.Context, userID uint64, input updateProfileInput) (*model.User, error) { - var dbUser model.User - if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil { + dbUser, err := repository.GetUserByID(ctx, userID) + if err != nil { return nil, errors.New(errUserNotFound) } @@ -260,8 +260,8 @@ func updateUserProfile(ctx context.Context, userID uint64, input updateProfileIn return nil, errors.New(errEmailFormatInvalid) } - var count int64 - if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", input.Email, dbUser.ID).Count(&count).Error; err != nil { + count, err := repository.CountUsersByEmailExceptID(ctx, input.Email, dbUser.ID) + if err != nil { return nil, err } if count > 0 { @@ -281,26 +281,26 @@ func updateUserProfile(ctx context.Context, userID uint64, input updateProfileIn dbUser.Website = strings.TrimSpace(input.Website) dbUser.Location = strings.TrimSpace(input.Location) - if err := db.DB(ctx).Save(&dbUser).Error; err != nil { + if err := repository.UpdateUser(ctx, &dbUser); err != nil { return nil, err } return &dbUser, nil } func getUserByUsernameOrEmail(ctx context.Context, input string) (*model.User, error) { - var user model.User - if err := db.DB(ctx).Where("username = ? OR email = ?", input, input).First(&user).Error; err != nil { + user, err := repository.GetUserByUsernameOrEmail(ctx, input) + if err != nil { return nil, err } return &user, nil } func updateLastLogin(ctx context.Context, user *model.User) error { - return db.DB(ctx).Model(user).Update("last_login_at", user.LastLoginAt).Error + return repository.UpdateUserLastLoginAt(ctx, user.ID, user.LastLoginAt) } func registerUserLogic(ctx context.Context, u *model.User) error { - if err := u.RegisterUser(ctx, db.DB(ctx)); err != nil { + if err := repository.RegisterUserWithChecks(ctx, u); err != nil { if strings.Contains(err.Error(), "duplicate key") || strings.Contains(err.Error(), "UNIQUE") { return errors.New("用户名或邮箱已被占用") } @@ -310,8 +310,8 @@ func registerUserLogic(ctx context.Context, u *model.User) error { } func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass string) error { - var dbUser model.User - if err := db.DB(ctx).Where("id = ?", userID).First(&dbUser).Error; err != nil { + dbUser, err := repository.GetUserByID(ctx, userID) + if err != nil { return errors.New(errUserNotFound) } @@ -323,18 +323,17 @@ func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass st return errors.New(errPasswordEncryptFailed) } - if err := db.DB(ctx).Model(&dbUser).Update("password", dbUser.Password).Error; err != nil { + if err := repository.UpdateUserPassword(ctx, dbUser.ID, dbUser.Password); err != nil { return errors.New("更新密码失败,请稍后再试") } // 吊销该用户所有的 Access Token - var tokens []model.AccessToken - if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Find(&tokens).Error; err == nil { + if tokens, err := repository.ListAccessTokensByUserID(ctx, dbUser.ID); err == nil { for _, token := range tokens { oauth.InvalidateCachedToken(ctx, token.TokenHash) } } - if err := db.DB(ctx).Where("user_id = ?", dbUser.ID).Delete(&model.AccessToken{}).Error; err != nil { + if err := repository.DeleteAccessTokensByUserID(ctx, dbUser.ID); err != nil { return errors.New("吊销 Access Token 失败,请稍后再试") } @@ -343,23 +342,23 @@ func changePasswordLogic(ctx context.Context, userID uint64, oldPass, newPass st } func listAccessTokensLogic(ctx context.Context, userID uint64) ([]model.AccessToken, error) { - var tokens []model.AccessToken - if err := db.DB(ctx).Where("user_id = ?", userID).Order("created_at desc").Find(&tokens).Error; err != nil { + tokens, err := repository.ListAccessTokensByUserID(ctx, userID) + if err != nil { return nil, errors.New("获取令牌列表失败,请稍后再试") } return tokens, nil } func countAccessTokensLogic(ctx context.Context, userID uint64) (int64, error) { - var count int64 - if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", userID).Count(&count).Error; err != nil { + count, err := repository.CountAccessTokensByUserID(ctx, userID) + if err != nil { return 0, errors.New("查询令牌数量失败,请稍后再试") } return count, nil } func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) error { - if err := db.DB(ctx).Create(record).Error; err != nil { + if err := repository.CreateAccessToken(ctx, record); err != nil { return errors.New("创建令牌失败,请稍后再试") } oauth.SetCachedToken(ctx, record.TokenHash, record) @@ -367,25 +366,25 @@ func createAccessTokenLogic(ctx context.Context, record *model.AccessToken) erro } func deleteAccessTokenLogic(ctx context.Context, id, userID uint64) error { - var tokenRecord model.AccessToken - if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil { + tokenRecord, err := repository.GetAccessTokenByIDAndUserID(ctx, id, userID) + if err != nil { return errors.New(errTokenNotFoundOrForbidden) } oauth.InvalidateCachedToken(ctx, tokenRecord.TokenHash) - tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.AccessToken{}) - if tx.Error != nil { + rows, err := repository.DeleteAccessTokenForUser(ctx, id, userID) + if err != nil { return errors.New("删除令牌失败,请稍后再试") } - if tx.RowsAffected == 0 { + if rows == 0 { return errors.New(errTokenNotFoundOrForbidden) } return nil } func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *model.AccessToken, error) { - var tokenRecord model.AccessToken - if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&tokenRecord).Error; err != nil { + tokenRecord, err := repository.GetAccessTokenByIDAndUserID(ctx, id, userID) + if err != nil { return "", nil, errors.New(errTokenNotFoundOrForbidden) } @@ -402,7 +401,7 @@ func rotateAccessTokenLogic(ctx context.Context, id, userID uint64) (string, *mo tokenRecord.TokenHash = newTokenHash tokenRecord.MaskedToken = newMaskedToken - if err := db.DB(ctx).Save(&tokenRecord).Error; err != nil { + if err := repository.SaveAccessToken(ctx, &tokenRecord); err != nil { return "", nil, errors.New("轮换令牌失败,请稍后再试") } diff --git a/internal/infra/task/executor.go b/internal/infra/task/executor.go index e7cb979b..c8dfac41 100644 --- a/internal/infra/task/executor.go +++ b/internal/infra/task/executor.go @@ -11,6 +11,8 @@ import ( "fmt" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/pkg/logger" @@ -110,7 +112,7 @@ func AppendLog(ctx context.Context, format string, args ...interface{}) { } logLine := fmt.Sprintf(format, args...) - if err := model.AppendTaskExecutionLog(ctx, taskID, logLine); err != nil { + if err := repository.AppendTaskExecutionLog(ctx, taskID, logLine); err != nil { logger.ErrorF(ctx, "[TaskExecutor] 追加任务日志失败 taskID=%s: %v", taskID, err) } } @@ -138,7 +140,7 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere TriggeredBy: triggeredBy, } - if err := model.CreateTaskExecution(ctx, execution); err != nil { + if err := repository.CreateTaskExecution(ctx, execution); err != nil { return "", fmt.Errorf(errCreateTaskExecutionFailed, err) } @@ -156,11 +158,11 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere now := time.Now() execution.StartedAt = &now execution.FinishedAt = &now - _ = model.UpdateTaskExecution(ctx, execution) + _ = repository.UpdateTaskExecution(ctx, execution) return "", fmt.Errorf(errTaskEnqueueFailed, err) } - if err := model.AppendTaskExecutionLog(ctx, taskID, fmt.Sprintf("[系统] 任务已成功入队,等待调度执行 (队列: %s, 最大重试次数: %d)", meta.Queue, meta.MaxRetry)); err != nil { + if err := repository.AppendTaskExecutionLog(ctx, taskID, fmt.Sprintf("[系统] 任务已成功入队,等待调度执行 (队列: %s, 最大重试次数: %d)", meta.Queue, meta.MaxRetry)); err != nil { logger.ErrorF(ctx, "[TaskExecutor] 追加入队日志失败 taskID=%s: %v", taskID, err) } @@ -169,7 +171,7 @@ func DispatchTask(ctx context.Context, taskType string, payload []byte, triggere // RetryTask 重试失败的任务 func RetryTask(ctx context.Context, id uint64) (string, error) { - execution, err := model.GetTaskExecutionByID(ctx, id) + execution, err := repository.GetTaskExecutionByID(ctx, id) if err != nil { return "", fmt.Errorf(errTaskExecutionNotFound, err) } @@ -198,7 +200,7 @@ func RetryTask(ctx context.Context, id uint64) (string, error) { TriggeredBy: "retry", } - if err := model.CreateTaskExecution(ctx, newExecution); err != nil { + if err := repository.CreateTaskExecution(ctx, newExecution); err != nil { return "", fmt.Errorf(errCreateRetryExecutionFailed, err) } @@ -221,11 +223,11 @@ func RetryTask(ctx context.Context, id uint64) (string, error) { now := time.Now() newExecution.StartedAt = &now newExecution.FinishedAt = &now - _ = model.UpdateTaskExecution(ctx, newExecution) + _ = repository.UpdateTaskExecution(ctx, newExecution) return "", fmt.Errorf(errRetryTaskEnqueueFailed, err) } - if err := model.AppendTaskExecutionLog(ctx, newTaskID, fmt.Sprintf("[系统] 手动触发重试,已重新创建任务并入队 (原任务ID: %s, 重试次数: %d/%d)", execution.TaskID, execution.RetryCount+1, execution.MaxRetry)); err != nil { + if err := repository.AppendTaskExecutionLog(ctx, newTaskID, fmt.Sprintf("[系统] 手动触发重试,已重新创建任务并入队 (原任务ID: %s, 重试次数: %d/%d)", execution.TaskID, execution.RetryCount+1, execution.MaxRetry)); err != nil { logger.ErrorF(ctx, "[TaskExecutor] 追加重试日志失败 taskID=%s: %v", newTaskID, err) } @@ -350,7 +352,7 @@ func updateExecutionOnStart(ctx context.Context, execution *model.TaskExecution, dirty = true } if dirty { - if updateErr := model.UpdateTaskExecution(ctx, execution); updateErr != nil { + if updateErr := repository.UpdateTaskExecution(ctx, execution); updateErr != nil { logger.ErrorF(ctx, "[TaskExecutor] 更新执行状态失败 taskID=%s: %v", execution.TaskID, updateErr) } } @@ -358,7 +360,7 @@ func updateExecutionOnStart(ctx context.Context, execution *model.TaskExecution, // getOrCreateTaskExecution 获取已有的任务执行记录,如果不存在则针对已知任务类型动态创建记录 func getOrCreateTaskExecution(ctx context.Context, taskID string, t *asynq.Task, payload []byte, now time.Time) (*model.TaskExecution, error) { - execution, err := model.GetTaskExecutionByTaskID(ctx, taskID) + execution, err := repository.GetTaskExecutionByTaskID(ctx, taskID) if err == nil { return execution, nil } @@ -381,7 +383,7 @@ func getOrCreateTaskExecution(ctx context.Context, taskID string, t *asynq.Task, StartedAt: &now, } - if createErr := model.CreateTaskExecution(ctx, execution); createErr != nil { + if createErr := repository.CreateTaskExecution(ctx, execution); createErr != nil { logger.ErrorF(ctx, "[TaskExecutor] 动态创建执行记录失败 taskID=%s: %v", taskID, createErr) return nil, createErr } @@ -404,11 +406,11 @@ func completeTaskExecution(ctx context.Context, execution *model.TaskExecution, handleSuccessfulTask(ctx, execution, t, duration, result) } - if err := model.UpdateTaskExecution(ctx, execution); err != nil { + if err := repository.UpdateTaskExecution(ctx, execution); err != nil { logger.ErrorF(ctx, "[TaskExecutor] 更新执行记录失败 taskID=%s: %v", execution.TaskID, err) } if shouldFlushTaskExecutionLog(ctx, execErr) { - if err := model.FlushTaskExecutionLog(ctx, execution.TaskID); err != nil { + if err := repository.FlushTaskExecutionLog(ctx, execution.TaskID); err != nil { logger.ErrorF(ctx, "[TaskExecutor] 持久化任务日志失败 taskID=%s: %v", execution.TaskID, err) } } diff --git a/internal/infra/task/executor_test.go b/internal/infra/task/executor_test.go index c65f2d7e..a58adc2b 100644 --- a/internal/infra/task/executor_test.go +++ b/internal/infra/task/executor_test.go @@ -11,6 +11,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/testhelper" @@ -122,7 +124,7 @@ func TestAppendLogWithTaskID(t *testing.T) { Status: model.TaskExecutionStatusRunning, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) // 注入 taskID 并追加日志 @@ -131,7 +133,7 @@ func TestAppendLogWithTaskID(t *testing.T) { AppendLog(ctx, "处理了 %d 条数据", 50) // 验证日志 - found, err := model.GetTaskExecutionByTaskID(ctx, "log_test_001") + found, err := repository.GetTaskExecutionByTaskID(ctx, "log_test_001") require.NoError(t, err) assert.Contains(t, found.Log, "第一条日志") assert.Contains(t, found.Log, "处理了 50 条数据") @@ -190,7 +192,7 @@ func TestProcessTaskSuccess(t *testing.T) { MaxRetry: 3, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) // 通过 asynq 的 Task 不能直接设置 taskID,ProcessTask 通过 t.ResultWriter().TaskID() 获取 @@ -208,7 +210,7 @@ func TestProcessTaskSuccess(t *testing.T) { assert.Equal(t, "处理完成,共 100 条", result.Message) // 验证日志被追加 - found, err := model.GetTaskExecutionByTaskID(ctx, "process_success_001") + found, err := repository.GetTaskExecutionByTaskID(ctx, "process_success_001") require.NoError(t, err) assert.Contains(t, found.Log, "执行成功,处理了 100 条数据") } @@ -231,7 +233,7 @@ func TestProcessTaskFailure(t *testing.T) { MaxRetry: 3, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) // 直接调用 handler @@ -244,7 +246,7 @@ func TestProcessTaskFailure(t *testing.T) { assert.Contains(t, err.Error(), "模拟执行失败") // 验证日志 - found, err := model.GetTaskExecutionByTaskID(ctx, "process_fail_001") + found, err := repository.GetTaskExecutionByTaskID(ctx, "process_fail_001") require.NoError(t, err) assert.Contains(t, found.Log, "开始执行任务") } @@ -261,7 +263,7 @@ func TestCompleteTaskExecutionFlushesLog(t *testing.T) { Status: model.TaskExecutionStatusRunning, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) ctx = withTaskID(ctx, execution.TaskID) @@ -279,7 +281,7 @@ func TestCompleteTaskExecutionFlushesLog(t *testing.T) { trace.SpanFromContext(ctx), ) - found, err := model.GetTaskExecutionByTaskID(ctx, execution.TaskID) + found, err := repository.GetTaskExecutionByTaskID(ctx, execution.TaskID) require.NoError(t, err) assert.Equal(t, model.TaskExecutionStatusSucceeded, found.Status) assert.Contains(t, found.Log, "任务执行中的日志") @@ -300,7 +302,7 @@ func TestCompleteTaskExecutionFlushesPermanentFailureLog(t *testing.T) { MaxRetry: 3, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) ctx = withTaskID(ctx, execution.TaskID) @@ -319,7 +321,7 @@ func TestCompleteTaskExecutionFlushesPermanentFailureLog(t *testing.T) { trace.SpanFromContext(ctx), ) - found, err := model.GetTaskExecutionByTaskID(ctx, execution.TaskID) + found, err := repository.GetTaskExecutionByTaskID(ctx, execution.TaskID) require.NoError(t, err) assert.Equal(t, model.TaskExecutionStatusFailed, found.Status) assert.Equal(t, "来源配置无效", found.ErrorMessage) @@ -356,7 +358,7 @@ func TestRetryTask(t *testing.T) { Duration: 100, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) // 重试 @@ -366,7 +368,7 @@ func TestRetryTask(t *testing.T) { assert.Contains(t, newTaskID, "retry_1_") // 验证新记录 - newExecution, err := model.GetTaskExecutionByTaskID(ctx, newTaskID) + newExecution, err := repository.GetTaskExecutionByTaskID(ctx, newTaskID) require.NoError(t, err) assert.Equal(t, model.TaskExecutionStatusPending, newExecution.Status) assert.Equal(t, 1, newExecution.RetryCount) @@ -375,7 +377,7 @@ func TestRetryTask(t *testing.T) { assert.True(t, newExecution.Retryable) // 原记录不变 - original, err := model.GetTaskExecutionByID(ctx, execution.ID) + original, err := repository.GetTaskExecutionByID(ctx, execution.ID) require.NoError(t, err) assert.Equal(t, model.TaskExecutionStatusFailed, original.Status) assert.Equal(t, 0, original.RetryCount) @@ -396,7 +398,7 @@ func TestRetryTaskNotFailed(t *testing.T) { MaxRetry: 3, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) // 尝试重试成功的任务 @@ -419,7 +421,7 @@ func TestRetryTaskNotRetryable(t *testing.T) { MaxRetry: 0, TriggeredBy: "manual", } - err := model.CreateTaskExecution(ctx, execution) + err := repository.CreateTaskExecution(ctx, execution) require.NoError(t, err) _, err = RetryTask(ctx, execution.ID) diff --git a/internal/infra/task/scheduler/scheduler.go b/internal/infra/task/scheduler/scheduler.go index f45fcb64..5c8e17c1 100644 --- a/internal/infra/task/scheduler/scheduler.go +++ b/internal/infra/task/scheduler/scheduler.go @@ -11,8 +11,9 @@ import ( "syscall" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/infra/task" - "github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/platform/bootstrap" "github.com/Rain-kl/Wavelet/pkg/logger" @@ -84,7 +85,7 @@ func ReloadScheduler() error { } // 2. 从数据库载入启用的定时任务配置 - schedules, err := model.ListActiveSchedules(context.Background()) + schedules, err := repository.ListActiveSchedules(context.Background()) if err != nil { return fmt.Errorf("load schedules from db failed: %w", err) } diff --git a/internal/model/auth_source.go b/internal/model/auth_source.go index 51bbb8d7..00bc4144 100644 --- a/internal/model/auth_source.go +++ b/internal/model/auth_source.go @@ -4,14 +4,10 @@ package model import ( - "context" "errors" "regexp" "strings" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "gorm.io/gorm" ) // 认证源类型 @@ -124,203 +120,3 @@ func (source *AuthSource) Sanitize() { source.ClientSecretConfigured = source.ClientSecret != "" source.ClientSecret = "" } - -// GetAuthSources 获取所有认证源(已脱敏) -func GetAuthSources(ctx context.Context) ([]AuthSource, error) { - var sources []AuthSource - if err := db.DB(ctx).Order("id asc").Find(&sources).Error; err != nil { - return nil, err - } - for i := range sources { - sources[i].Sanitize() - } - return sources, nil -} - -// GetActiveAuthSources 获取所有已启用的认证源(已脱敏) -func GetActiveAuthSources(ctx context.Context) ([]AuthSource, error) { - var sources []AuthSource - if err := db.DB(ctx).Where("is_active = ?", true).Order("id asc").Find(&sources).Error; err != nil { - return nil, err - } - for i := range sources { - sources[i].Sanitize() - } - return sources, nil -} - -// GetAuthSourceByID 根据 ID 获取认证源 -func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) { - if id == 0 { - return nil, errors.New(errAuthSourceIDRequired) - } - var source AuthSource - if err := db.DB(ctx).First(&source, "id = ?", id).Error; err != nil { - return nil, err - } - source.ClientSecretConfigured = source.ClientSecret != "" - return &source, nil -} - -// GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写) -func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) { - name = strings.TrimSpace(name) - if name == "" { - return nil, errors.New(errAuthSourceNameRequired) - } - var source AuthSource - if err := db.DB(ctx).First(&source, "LOWER(name) = LOWER(?)", name).Error; err != nil { - return nil, err - } - source.ClientSecretConfigured = source.ClientSecret != "" - return &source, nil -} - -// CreateAuthSource 创建认证源 -func CreateAuthSource(ctx context.Context, source *AuthSource) error { - if err := source.Validate(); err != nil { - return err - } - return db.DB(ctx).Create(source).Error -} - -// UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥 -func UpdateAuthSource(ctx context.Context, source *AuthSource, keepSecret bool) error { - if source.ID == 0 { - return errors.New(errAuthSourceIDRequired) - } - var current AuthSource - if err := db.DB(ctx).First(¤t, "id = ?", source.ID).Error; err != nil { - return err - } - if keepSecret { - source.ClientSecret = current.ClientSecret - } - if err := source.Validate(); err != nil { - return err - } - return db.DB(ctx).Model(¤t).Updates(map[string]any{ - colName: source.Name, - "type": source.Type, - "display_name": source.DisplayName, - "is_active": source.IsActive, - "client_id": source.ClientID, - "client_secret": source.ClientSecret, - "openid_discovery_url": source.OpenIDDiscoveryURL, - "scopes": source.Scopes, - "icon_url": source.IconURL, - }).Error -} - -// ToggleAuthSource 切换认证源启用状态 -func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error { - source, err := GetAuthSourceByID(ctx, id) - if err != nil { - return err - } - source.IsActive = isActive - if err := source.Validate(); err != nil { - return err - } - return db.DB(ctx).Model(&AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error -} - -// DeleteAuthSource 删除认证源及其关联的外部帐号绑定 -func DeleteAuthSource(ctx context.Context, id uint64) error { - if id == 0 { - return errors.New(errAuthSourceIDRequired) - } - return db.DB(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Where("auth_source_id = ?", id).Delete(&ExternalAccount{}).Error; err != nil { - return err - } - return tx.Delete(&AuthSource{}, "id = ?", id).Error - }) -} - -// FindExternalAccount 查找外部帐号绑定记录 -func FindExternalAccount(ctx context.Context, sourceID uint64, externalID string) (*ExternalAccount, error) { - var account ExternalAccount - if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error; err != nil { - return nil, err - } - return &account, nil -} - -// BindExternalAccount 绑定外部帐号(已存在时更新用户名和邮箱) -func BindExternalAccount(ctx context.Context, account *ExternalAccount) error { - if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" { - return errors.New(errExternalAccountBindingIncomplete) - } - account.ExternalID = strings.TrimSpace(account.ExternalID) - account.ExternalUsername = strings.TrimSpace(account.ExternalUsername) - account.Email = strings.TrimSpace(account.Email) - - return db.DB(ctx).Transaction(func(tx *gorm.DB) error { - var current ExternalAccount - err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(¤t).Error - if err == nil { - if current.UserID != account.UserID { - return errors.New(errExternalAccountAlreadyBoundToAnother) - } - return tx.Model(¤t).Updates(map[string]any{ - "external_username": account.ExternalUsername, - "email": account.Email, - }).Error - } - if !errors.Is(err, gorm.ErrRecordNotFound) { - return err - } - return tx.Create(account).Error - }) -} - -// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图 -func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccountView, error) { - if userID == 0 { - return nil, errors.New(errUserIDRequired) - } - var accounts []ExternalAccount - if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil { - return nil, err - } - views := make([]ExternalAccountView, 0, len(accounts)) - for _, account := range accounts { - var name, sourceType, label string - if account.AuthSourceID == 0 { - name = "default" - sourceType = "oidc" - label = "历史认证源" - } else { - source, err := GetAuthSourceByID(ctx, account.AuthSourceID) - if err != nil { - continue - } - name = source.Name - sourceType = source.Type - label = source.DisplayName - if label == "" { - label = source.Name - } - } - views = append(views, ExternalAccountView{ - ID: account.ID, - AuthSourceID: account.AuthSourceID, - AuthSourceName: name, - AuthSourceType: sourceType, - AuthSourceLabel: label, - ExternalUsername: account.ExternalUsername, - Email: account.Email, - CreatedAt: account.CreatedAt, - }) - } - return views, nil -} - -// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定 -func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error { - if id == 0 || userID == 0 { - return errors.New(errExternalAccountBindingIDRequired) - } - return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error -} diff --git a/internal/model/errs.go b/internal/model/errs.go index e8723ad6..51da0c48 100644 --- a/internal/model/errs.go +++ b/internal/model/errs.go @@ -3,29 +3,15 @@ package model +// Domain validation messages used by model.Validate and other no-IO rules. +// Persistence / data-access messages belong in internal/repository (do not import repository). const ( - errRegistrationDisabled = "注册已关闭" - errDatabaseNotInitialized = "database not initialized" - errClickHouseNotInitialized = "clickhouse not initialized" - errUsernameExists = "用户名已存在" - errEmailAlreadyBound = "该邮箱已被其他账号绑定" - errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w" - errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w" - errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w" - errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w" - errTemplateKeyRequired = "模板标识符不能为空" - errTemplateNameRequired = "模板名称不能为空" - errTemplateContentRequired = "模板内容不能为空" - errTemplateUnavailable = "模板 %s 不存在或不可用: %w" - errTemplateRenderFailed = "模板 %s 渲染失败: %w" - errAuthSourceNameRequired = "认证源名称不能为空" - errAuthSourceNameInvalid = "认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头" - errAuthSourceTypeUnsupported = "认证源类型仅支持 oidc" - errAuthSourceDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL" - errAuthSourceClientCredentialsRequired = "启用认证源前必须配置 Client ID 和 Client Secret" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errAuthSourceIDRequired = "认证源 ID 不能为空" - errExternalAccountBindingIncomplete = "外部账号绑定信息不完整" - errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户" - errUserIDRequired = "用户 ID 不能为空" - errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空" + errTemplateKeyRequired = "模板标识符不能为空" + errTemplateNameRequired = "模板名称不能为空" + errTemplateContentRequired = "模板内容不能为空" + errAuthSourceNameRequired = "认证源名称不能为空" + errAuthSourceNameInvalid = "认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头" + errAuthSourceTypeUnsupported = "认证源类型仅支持 oidc" + errAuthSourceDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL" + errAuthSourceClientCredentialsRequired = "启用认证源前必须配置 Client ID 和 Client Secret" //nolint:gosec // false positive: this is an error message, not hardcoded credentials ) diff --git a/internal/model/openflare_access_log.go b/internal/model/openflare_access_log.go index 363f9a40..0f1c65db 100644 --- a/internal/model/openflare_access_log.go +++ b/internal/model/openflare_access_log.go @@ -3,472 +3,26 @@ package model -import ( - "context" - "math" - "sort" - "strings" - "time" - - analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics" -) - -const ( - columnRemoteAddr = "remote_addr" - columnHost = "host" - sortOrderAsc = "asc" - secondsPerMinute = 60 -) - -type openFlareAccessLogBucketAggregateRow = analyticsmodel.NodeAccessLogBucketAggregate -type openFlareAccessLogBucketDimensionRow = analyticsmodel.NodeAccessLogBucketDimension -type openFlareAccessLogIPAggregateRow = analyticsmodel.NodeAccessLogIPAggregate -type openFlareAccessLogIPSummaryRow = analyticsmodel.NodeAccessLogIPSummary -type openFlareAccessLogIPTrendRow = analyticsmodel.NodeAccessLogIPTrend -type openFlareAccessLogWAFIPAggregateRow = analyticsmodel.NodeAccessLogWAFIPAggregate - -// ListOpenFlareAccessLogWAFIPAggregates returns per-IP aggregates for WAF automatic rules. -func ListOpenFlareAccessLogWAFIPAggregates(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLogWAFIPAggregate, error) { - rows, err := currentAccessLogStore().WAFIPAggregates(ctx, query) - if err != nil { - return nil, err - } - result := make([]*OpenFlareAccessLogWAFIPAggregate, 0, len(rows)) - for _, row := range rows { - remoteAddr := strings.TrimSpace(row.RemoteAddr) - if remoteAddr == "" { - continue - } - statusCounts := make(map[int]int, len(row.StatusCounts)) - for code, count := range row.StatusCounts { - statusCounts[code] = int(count) - } - result = append(result, &OpenFlareAccessLogWAFIPAggregate{ - RemoteAddr: remoteAddr, - RequestCount: int(row.RequestCount), - Status404Count: int(row.Status404Count), - ClientErrorCount: int(row.ClientErrorCount), - ServerErrorCount: int(row.ServerErrorCount), - IPHostCount: int(row.IPHostCount), - LastSeenEpoch: row.LastSeenEpoch, - StatusCounts: statusCounts, - }) - } - return result, nil +// OpenFlareAccessLogTrafficSummary is a window-level traffic summary from access logs. +type OpenFlareAccessLogTrafficSummary struct { + RequestCount int64 + ErrorCount int64 + UniqueIPCount int64 + BytesSent int64 + RequestLength int64 + NodeCount int64 } -// InsertOpenFlareAccessLogsBatch inserts access log rows into ClickHouse. -func InsertOpenFlareAccessLogsBatch(ctx context.Context, records []*OpenFlareAccessLog) error { - return currentAccessLogStore().InsertBatch(ctx, records) +// OpenFlareAccessLogValueCount is a dimension value count. +type OpenFlareAccessLogValueCount struct { + Value string + Count int64 } -// ListOpenFlareAccessLogs lists access logs matching the query. -func ListOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) { - return currentAccessLogStore().List(ctx, query) -} - -// CountOpenFlareAccessLogs counts access logs, distinct IPs, and total bytes sent matching the query. -func CountOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) (int64, int64, int64, error) { - return currentAccessLogStore().Count(ctx, query) -} - -// TrafficSummaryOpenFlareAccessLogs returns window-level request/error/UV/bytes summary. -func TrafficSummaryOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) (OpenFlareAccessLogTrafficSummary, error) { - return currentAccessLogStore().TrafficSummary(ctx, query) -} - -// ValueCountsOpenFlareAccessLogs groups logs by status_code, host, path, remote_addr, or user_agent. -func ValueCountsOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery, column string, limit int) ([]OpenFlareAccessLogValueCount, error) { - return currentAccessLogStore().ValueCounts(ctx, query, column, limit) -} - -// NodeAggregatesOpenFlareAccessLogs returns per-node request/error/UV for the window. -func NodeAggregatesOpenFlareAccessLogs(ctx context.Context, query OpenFlareAccessLogQuery) ([]OpenFlareAccessLogNodeAggregate, error) { - return currentAccessLogStore().NodeAggregates(ctx, query) -} - -// ListOpenFlareAccessLogRegionCounts returns region counts for access logs. -func ListOpenFlareAccessLogRegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error) { - return currentAccessLogStore().RegionCounts(ctx, nodeID, since, limit) -} - -// ListOpenFlareAccessLogBuckets lists folded access log buckets. -func ListOpenFlareAccessLogBuckets(ctx context.Context, query OpenFlareAccessLogBucketQuery) ([]*OpenFlareAccessLogBucketRow, error) { - return buildOpenFlareAccessLogBucketRows(ctx, query) -} - -// CountOpenFlareAccessLogBuckets counts folded access log buckets. -func CountOpenFlareAccessLogBuckets(ctx context.Context, query OpenFlareAccessLogBucketQuery) (int64, error) { - filter := openFlareAccessLogQueryFromBucket(query) - bucketSeconds := int64(query.FoldMinutes * secondsPerMinute) - if bucketSeconds <= 0 { - bucketSeconds = 180 - } - return currentAccessLogStore().CountBuckets(ctx, filter, bucketSeconds) -} - -// ListOpenFlareAccessLogBucketIPs lists folded IP rows for a bucket window. -func ListOpenFlareAccessLogBucketIPs(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) ([]*OpenFlareAccessLogBucketIPRow, error) { - rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query) - if err != nil { - return nil, err - } - start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize) - if start >= len(rows) { - return []*OpenFlareAccessLogBucketIPRow{}, nil - } - return rows[start:end], nil -} - -// CountOpenFlareAccessLogBucketIPs counts folded IP rows for a bucket window. -func CountOpenFlareAccessLogBucketIPs(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) (int64, error) { - rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query) - if err != nil { - return 0, err - } - return int64(len(rows)), nil -} - -// ListOpenFlareAccessLogIPSummaries lists IP summaries. -func ListOpenFlareAccessLogIPSummaries(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*OpenFlareAccessLogIPSummaryRow, error) { - return buildOpenFlareAccessLogIPSummaryRows(ctx, query, recentSince) -} - -// CountOpenFlareAccessLogIPSummaries counts IP summaries. -func CountOpenFlareAccessLogIPSummaries(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery) (int64, error) { - filter := openFlareAccessLogQueryFromIPSummary(query) - return currentAccessLogStore().CountIPSummaries(ctx, filter) -} - -// ListOpenFlareAccessLogIPTrend lists IP trend points. -func ListOpenFlareAccessLogIPTrend(ctx context.Context, query OpenFlareAccessLogIPTrendQuery) ([]*OpenFlareAccessLogIPTrendRow, error) { - remoteAddr := strings.TrimSpace(query.RemoteAddr) - if remoteAddr == "" { - return []*OpenFlareAccessLogIPTrendRow{}, nil - } - filter := OpenFlareAccessLogQuery{ - NodeID: query.NodeID, - RemoteAddr: remoteAddr, - Host: query.Host, - Since: query.Since, - } - bucketSeconds := int64(query.BucketMinutes * secondsPerMinute) - if bucketSeconds <= 0 { - bucketSeconds = 1800 - } - rows, err := currentAccessLogStore().IPTrend(ctx, filter, bucketSeconds) - if err != nil { - return nil, err - } - result := make([]*OpenFlareAccessLogIPTrendRow, len(rows)) - for index, row := range rows { - result[index] = &OpenFlareAccessLogIPTrendRow{ - BucketEpoch: row.BucketEpoch, - RequestCount: row.RequestCount, - } - } - return result, nil -} - -// DeleteAllOpenFlareAccessLogs deletes all access logs. -func DeleteAllOpenFlareAccessLogs(ctx context.Context) (int64, error) { - return currentAccessLogStore().DeleteAll(ctx) -} - -// DeleteOpenFlareAccessLogsBefore deletes access logs older than cutoff. -func DeleteOpenFlareAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) { - return currentAccessLogStore().DeleteBefore(ctx, cutoff) -} - -// DeleteOpenFlareAccessLogsByNodeBefore deletes access logs for a node older than cutoff. -func DeleteOpenFlareAccessLogsByNodeBefore(ctx context.Context, nodeID string, cutoff time.Time) (int64, error) { - return currentAccessLogStore().DeleteByNodeBefore(ctx, nodeID, cutoff) -} - -func buildOpenFlareAccessLogBucketRows(ctx context.Context, query OpenFlareAccessLogBucketQuery) ([]*OpenFlareAccessLogBucketRow, error) { - filter := openFlareAccessLogQueryFromBucket(query) - bucketSeconds := int64(query.FoldMinutes * secondsPerMinute) - if bucketSeconds <= 0 { - bucketSeconds = 180 - } - - partials, err := currentAccessLogStore().BucketAggregates(ctx, filter, bucketSeconds) - if err != nil { - return nil, err - } - rows := make([]*OpenFlareAccessLogBucketRow, 0, len(partials)) - for _, partial := range partials { - rows = append(rows, &OpenFlareAccessLogBucketRow{ - BucketEpoch: partial.BucketEpoch, - RequestCount: partial.RequestCount, - UniqueIPCount: partial.UniqueIPCount, - UniqueHostCount: partial.UniqueHostCount, - SuccessCount: partial.SuccessCount, - ClientErrorCount: partial.ClientErrorCount, - ServerErrorCount: partial.ServerErrorCount, - BytesSent: partial.BytesSent, - RequestLength: partial.RequestLength, - }) - } - return rows, nil -} - -func buildOpenFlareAccessLogBucketIPRows(ctx context.Context, query OpenFlareAccessLogBucketIPQuery) ([]*OpenFlareAccessLogBucketIPRow, error) { - if query.BucketStartedAt.IsZero() { - return []*OpenFlareAccessLogBucketIPRow{}, nil - } - foldMinutes := query.FoldMinutes - if foldMinutes <= 0 { - foldMinutes = 3 - } - bucketStartedAt := query.BucketStartedAt.UTC() - filter := OpenFlareAccessLogQuery{ - NodeID: query.NodeID, - RemoteAddr: query.RemoteAddr, - Host: query.Host, - Path: query.Path, - Since: bucketStartedAt, - Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute), - } - rows, err := queryOpenFlareAccessLogIPAggregateRows(ctx, filter, false) - if err != nil { - return nil, err - } - sortOpenFlareAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder) - return rows, nil -} - -func buildOpenFlareAccessLogIPSummaryRows(ctx context.Context, query OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*OpenFlareAccessLogIPSummaryRow, error) { - filter := openFlareAccessLogQueryFromIPSummary(query) - partials, err := currentAccessLogStore().IPSummaries(ctx, filter, recentSince) - if err != nil { - return nil, err - } - rows := make([]*OpenFlareAccessLogIPSummaryRow, 0, len(partials)) - for _, partial := range partials { - remoteAddr := strings.TrimSpace(partial.RemoteAddr) - if remoteAddr == "" { - continue - } - rows = append(rows, &OpenFlareAccessLogIPSummaryRow{ - RemoteAddr: remoteAddr, - Region: strings.TrimSpace(partial.Region), - TotalRequests: partial.TotalRequests, - Success2xxCount: partial.Success2xxCount, - SuccessRatio: partial.SuccessRatio, - BytesReceived: partial.BytesReceived, - BytesSent: partial.BytesSent, - RecentRequests: 0, - LastSeenEpoch: partial.LastSeenEpoch, - }) - } - return rows, nil -} - -func queryOpenFlareAccessLogIPAggregateRows(ctx context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]*OpenFlareAccessLogBucketIPRow, error) { - partials, err := currentAccessLogStore().IPAggregates(ctx, filter, exactRemoteAddr) - if err != nil { - return nil, err - } - rows := make([]*OpenFlareAccessLogBucketIPRow, 0, len(partials)) - for _, partial := range partials { - remoteAddr := strings.TrimSpace(partial.RemoteAddr) - if remoteAddr == "" { - continue - } - rows = append(rows, &OpenFlareAccessLogBucketIPRow{ - RemoteAddr: remoteAddr, - RequestCount: partial.RequestCount, - SuccessCount: partial.SuccessCount, - ClientErrorCount: partial.ClientErrorCount, - ServerErrorCount: partial.ServerErrorCount, - LastSeenEpoch: partial.LastSeenEpoch, - }) - } - return rows, nil -} - -func openFlareAccessLogQueryFromBucket(query OpenFlareAccessLogBucketQuery) OpenFlareAccessLogQuery { - return OpenFlareAccessLogQuery{ - NodeID: query.NodeID, - RemoteAddr: query.RemoteAddr, - Host: query.Host, - Hosts: query.Hosts, - Path: query.Path, - Since: query.Since, - Until: query.Until, - Page: query.Page, - PageSize: query.PageSize, - SortBy: query.SortBy, - SortOrder: query.SortOrder, - } -} - -func openFlareAccessLogQueryFromIPSummary(query OpenFlareAccessLogIPSummaryQuery) OpenFlareAccessLogQuery { - return OpenFlareAccessLogQuery{ - NodeID: query.NodeID, - RemoteAddr: query.RemoteAddr, - Host: query.Host, - Since: query.Since, - Until: query.Until, - Page: query.Page, - PageSize: query.PageSize, - SortBy: query.SortBy, - SortOrder: query.SortOrder, - } -} - -func sortOpenFlareAccessLogBucketIPRows(items []*OpenFlareAccessLogBucketIPRow, sortBy string, sortOrder string) { - desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc - sort.Slice(items, func(i, j int) bool { - left := items[i] - right := items[j] - if left == nil || right == nil { - return left != nil - } - var compare int - switch strings.TrimSpace(sortBy) { - case "last_seen_at": - compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) - case "remote_addr": - compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) - default: - compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount) - } - if compare == 0 { - compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) - } - if compare == 0 { - compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) - } - if desc { - return compare > 0 - } - return compare < 0 - }) -} - -func sortOpenFlareAccessLogBucketRows(items []*OpenFlareAccessLogBucketRow, sortBy string, sortOrder string) { - desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc - sort.Slice(items, func(i, j int) bool { - left := items[i] - right := items[j] - if left == nil || right == nil { - return left != nil - } - var compare int - switch strings.TrimSpace(sortBy) { - case "request_count": - compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount) - default: - compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch) - } - if compare == 0 { - compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch) - } - if desc { - return compare > 0 - } - return compare < 0 - }) -} - -func sortOpenFlareAccessLogIPSummaryRows(items []*OpenFlareAccessLogIPSummaryRow, sortBy string, sortOrder string) { - desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc - sort.Slice(items, func(i, j int) bool { - left := items[i] - right := items[j] - if left == nil || right == nil { - return left != nil - } - var compare int - switch strings.TrimSpace(sortBy) { - case "request_length", "bytes_received": - compare = openFlareAccessLogCompareInt64(left.BytesReceived, right.BytesReceived) - case "bytes_sent": - compare = openFlareAccessLogCompareInt64(left.BytesSent, right.BytesSent) - case "success_ratio": - compare = openFlareAccessLogCompareFloat64(left.SuccessRatio, right.SuccessRatio) - case "last_seen_at": - compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) - case "remote_addr": - compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) - default: - compare = openFlareAccessLogCompareInt64(left.TotalRequests, right.TotalRequests) - } - if compare == 0 { - compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) - } - if compare == 0 { - compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) - } - if desc { - return compare > 0 - } - return compare < 0 - }) -} - -func openFlareAccessLogCompareFloat64(left, right float64) int { - if left < right { - return -1 - } - if left > right { - return 1 - } - return 0 -} - -func openFlareAccessLogPaginateBounds(total int, page int, pageSize int) (int, int) { - if page < 0 { - page = 0 - } - if pageSize <= 0 { - return 0, total - } - start := page * pageSize - if start > total { - start = total - } - end := start + pageSize - if end > total { - end = total - } - return start, end -} - -func openFlareAccessLogNormalizeSortOrder(sortOrder string) string { - if strings.EqualFold(strings.TrimSpace(sortOrder), sortOrderAsc) { - return sortOrderAsc - } - return "desc" -} - -func openFlareAccessLogStatusCodeToInt32(code int) int32 { - switch { - case code > math.MaxInt32: - return math.MaxInt32 - case code < math.MinInt32: - return math.MinInt32 - default: - return int32(code) - } -} - -func openFlareAccessLogUintToInt64(value uint64) int64 { - if value > math.MaxInt64 { - return math.MaxInt64 - } - return int64(value) -} - -func openFlareAccessLogCompareInt64(left int64, right int64) int { - switch { - case left > right: - return 1 - case left < right: - return -1 - default: - return 0 - } +// OpenFlareAccessLogNodeAggregate is per-node traffic over a window. +type OpenFlareAccessLogNodeAggregate struct { + NodeID string + RequestCount int64 + ErrorCount int64 + UniqueIPCount int64 } diff --git a/internal/model/openflare_acme_account.go b/internal/model/openflare_acme_account.go index e3d8c78e..e1acdfb6 100644 --- a/internal/model/openflare_acme_account.go +++ b/internal/model/openflare_acme_account.go @@ -4,12 +4,7 @@ package model import ( - "context" - "errors" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "gorm.io/gorm" ) // AcmeAccount OpenFlare ACME 账号实体。 @@ -26,57 +21,3 @@ type AcmeAccount struct { func (AcmeAccount) TableName() string { return "of_acme_accounts" } - -// GetAcmeAccountByID 按 ID 查询 ACME 账号。 -func GetAcmeAccountByID(ctx context.Context, id uint) (*AcmeAccount, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var account AcmeAccount - if err := conn.First(&account, id).Error; err != nil { - return nil, err - } - return &account, nil -} - -// CreateAcmeAccountRecord 创建 ACME 账号。 -func CreateAcmeAccountRecord(ctx context.Context, account *AcmeAccount) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Create(account).Error -} - -// SaveAcmeAccount 保存 ACME 账号。 -func SaveAcmeAccount(ctx context.Context, account *AcmeAccount) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Save(account).Error -} - -// GetDefaultAcmeAccount 获取默认 ACME 账号,不存在时创建占位记录。 -func GetDefaultAcmeAccount(ctx context.Context) (*AcmeAccount, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var account AcmeAccount - err := conn.Order("id asc").First(&account).Error - if err == nil { - return &account, nil - } - if !errors.Is(err, gorm.ErrRecordNotFound) { - return nil, err - } - account = AcmeAccount{ - Email: "admin@openflare.dev", - } - if err = conn.Create(&account).Error; err != nil { - return nil, err - } - return &account, nil -} diff --git a/internal/model/openflare_apply_log.go b/internal/model/openflare_apply_log.go index a9b9a3f4..64d13e8e 100644 --- a/internal/model/openflare_apply_log.go +++ b/internal/model/openflare_apply_log.go @@ -4,13 +4,8 @@ package model import ( - "context" - "errors" "strings" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "gorm.io/gorm" ) // OpenFlareApplyLogQuery filters apply logs for list queries. @@ -39,74 +34,6 @@ func (OpenFlareApplyLog) TableName() string { return "of_apply_logs" } -// ListOpenFlareApplyLogs returns apply logs ordered by id desc with optional pagination. -func ListOpenFlareApplyLogs(ctx context.Context, query OpenFlareApplyLogQuery) ([]*OpenFlareApplyLog, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - - dbQuery := conn.Model(&OpenFlareApplyLog{}).Order("id desc") - if query.NodeID != "" { - dbQuery = dbQuery.Where("node_id = ?", query.NodeID) - } - if query.PageSize > 0 { - offset := 0 - if query.PageNo > 1 { - offset = (query.PageNo - 1) * query.PageSize - } - dbQuery = dbQuery.Limit(query.PageSize).Offset(offset) - } - - var logs []*OpenFlareApplyLog - if err := dbQuery.Find(&logs).Error; err != nil { - return nil, err - } - return logs, nil -} - -// CountOpenFlareApplyLogs returns total apply logs, optionally filtered by node_id. -func CountOpenFlareApplyLogs(ctx context.Context, nodeID string) (int64, error) { - conn := db.DB(ctx) - if conn == nil { - return 0, errors.New(errDatabaseNotInitialized) - } - - query := conn.Model(&OpenFlareApplyLog{}) - if nodeID != "" { - query = query.Where("node_id = ?", nodeID) - } - - var total int64 - if err := query.Count(&total).Error; err != nil { - return 0, err - } - return total, nil -} - -// GetLatestOpenFlareApplyLogByNodeID returns the most recent apply log for a node. -func GetLatestOpenFlareApplyLogByNodeID(ctx context.Context, nodeID string) (*OpenFlareApplyLog, error) { - nodeID = strings.TrimSpace(nodeID) - if nodeID == "" { - return nil, errors.New("node_id is required") - } - - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - - var log OpenFlareApplyLog - err := conn.Where("node_id = ?", nodeID).Order("id desc").First(&log).Error - if errors.Is(err, gorm.ErrRecordNotFound) { - return nil, nil - } - if err != nil { - return nil, err - } - return &log, nil -} - // IsRepeatSuccessApplyLog reports whether the payload repeats an already-recorded success entry. func IsRepeatSuccessApplyLog(latest *OpenFlareApplyLog, version, checksum, result string) bool { if latest == nil || result != "success" { @@ -116,51 +43,3 @@ func IsRepeatSuccessApplyLog(latest *OpenFlareApplyLog, version, checksum, resul strings.TrimSpace(latest.Version) == strings.TrimSpace(version) && strings.TrimSpace(latest.Checksum) == strings.TrimSpace(checksum) } - -// GetLatestOpenFlareApplyLogsByNodeIDs returns the latest apply log per node id. -func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string) (map[string]*OpenFlareApplyLog, error) { - result := make(map[string]*OpenFlareApplyLog) - if len(nodeIDs) == 0 { - return result, nil - } - - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - - var logs []*OpenFlareApplyLog - subQuery := conn.Model(&OpenFlareApplyLog{}). - Select("MAX(id) AS id"). - Where("node_id IN ?", nodeIDs). - Group("node_id") - if err := conn.Where("id IN (?)", subQuery).Find(&logs).Error; err != nil { - return nil, err - } - for _, log := range logs { - result[log.NodeID] = log - } - return result, nil -} - -// DeleteAllOpenFlareApplyLogs removes every apply log record. -func DeleteAllOpenFlareApplyLogs(ctx context.Context) (int64, error) { - conn := db.DB(ctx) - if conn == nil { - return 0, errors.New(errDatabaseNotInitialized) - } - - result := conn.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&OpenFlareApplyLog{}) - return result.RowsAffected, result.Error -} - -// DeleteOpenFlareApplyLogsBefore removes apply logs created before the cutoff time. -func DeleteOpenFlareApplyLogsBefore(ctx context.Context, before time.Time) (int64, error) { - conn := db.DB(ctx) - if conn == nil { - return 0, errors.New(errDatabaseNotInitialized) - } - - result := conn.Where("created_at < ?", before).Delete(&OpenFlareApplyLog{}) - return result.RowsAffected, result.Error -} diff --git a/internal/model/openflare_config_version.go b/internal/model/openflare_config_version.go index 9274acfa..d1f33fd1 100644 --- a/internal/model/openflare_config_version.go +++ b/internal/model/openflare_config_version.go @@ -4,11 +4,8 @@ package model import ( - "context" - "errors" "time" - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "gorm.io/gorm" ) @@ -58,124 +55,3 @@ func (cv *ConfigVersion) AfterCreate(_ *gorm.DB) (err error) { func (ConfigVersion) TableName() string { return "of_config_versions" } - -// ListConfigVersionSummaries returns config version summaries ordered by created_at desc. -func ListConfigVersionSummaries(ctx context.Context) ([]*ConfigVersionSummary, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var versions []*ConfigVersionSummary - err := conn.Model(&ConfigVersion{}). - Select("version", "checksum", "is_active", "created_by", "created_at"). - Order("created_at desc, version desc"). - Find(&versions).Error - return versions, err -} - -// GetConfigVersionByVersion returns a config version by version string. -func GetConfigVersionByVersion(ctx context.Context, version string) (*ConfigVersion, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var cv ConfigVersion - if err := conn.First(&cv, "version = ?", version).Error; err != nil { - return nil, err - } - return &cv, nil -} - -// GetActiveConfigVersion returns the currently active config version. -func GetActiveConfigVersion(ctx context.Context) (*ConfigVersion, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var version ConfigVersion - if err := conn.Where("is_active = ?", true).Order("version desc").First(&version).Error; err != nil { - return nil, err - } - return &version, nil -} - -// GetLatestConfigVersionByPrefix returns the latest version string matching a date prefix. -func GetLatestConfigVersionByPrefix(ctx context.Context, prefix string) (string, error) { - conn := db.DB(ctx) - if conn == nil { - return "", errors.New(errDatabaseNotInitialized) - } - var version ConfigVersion - err := conn.Model(&ConfigVersion{}). - Select("version"). - Where("version LIKE ?", prefix+"-%"). - Order("version desc"). - First(&version).Error - if err != nil { - return "", err - } - return version.Version, nil -} - -// CreateConfigVersion inserts a new config version record. -func CreateConfigVersion(ctx context.Context, version *ConfigVersion) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Create(version).Error -} - -// PublishConfigVersionTx deactivates all versions and creates a new active version. -func PublishConfigVersionTx(ctx context.Context, version *ConfigVersion) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Transaction(func(tx *gorm.DB) error { - if err := tx.Model(&ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil { - return err - } - return tx.Create(version).Error - }) -} - -// ActivateConfigVersionTx marks the given version active and deactivates others. -func ActivateConfigVersionTx(ctx context.Context, version string) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Transaction(func(tx *gorm.DB) error { - if err := tx.Model(&ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil { - return err - } - return tx.Model(&ConfigVersion{}).Where("version = ?", version).Update("is_active", true).Error - }) -} - -// DeleteConfigVersionsByVersions removes config versions by versions. -func DeleteConfigVersionsByVersions(ctx context.Context, versions []string) (int64, error) { - if len(versions) == 0 { - return 0, nil - } - conn := db.DB(ctx) - if conn == nil { - return 0, errors.New(errDatabaseNotInitialized) - } - result := conn.Where("version IN ?", versions).Delete(&ConfigVersion{}) - return result.RowsAffected, result.Error -} - -// ListEnabledProxyRoutes returns enabled proxy routes ordered by id asc. -func ListEnabledProxyRoutes(ctx context.Context) ([]*ProxyRoute, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var routes []*ProxyRoute - if err := conn.Where("enabled = ?", true).Order("id asc").Find(&routes).Error; err != nil { - return nil, err - } - return routes, nil -} diff --git a/internal/model/openflare_dns_account.go b/internal/model/openflare_dns_account.go index fe30ce2b..59b40936 100644 --- a/internal/model/openflare_dns_account.go +++ b/internal/model/openflare_dns_account.go @@ -4,11 +4,7 @@ package model import ( - "context" - "errors" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" ) // DNSAccount OpenFlare DNS 账号实体。 @@ -25,56 +21,3 @@ type DNSAccount struct { func (DNSAccount) TableName() string { return "of_dns_accounts" } - -// ListDNSAccounts 列出全部 DNS 账号(授权信息不通过 JSON 暴露)。 -func ListDNSAccounts(ctx context.Context) ([]DNSAccount, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var accounts []DNSAccount - if err := conn.Order("id desc").Find(&accounts).Error; err != nil { - return nil, err - } - return accounts, nil -} - -// GetDNSAccountByID 按 ID 查询 DNS 账号。 -func GetDNSAccountByID(ctx context.Context, id uint) (*DNSAccount, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var account DNSAccount - if err := conn.First(&account, id).Error; err != nil { - return nil, err - } - return &account, nil -} - -// CreateDNSAccountRecord 创建 DNS 账号。 -func CreateDNSAccountRecord(ctx context.Context, account *DNSAccount) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Create(account).Error -} - -// SaveDNSAccount 保存 DNS 账号。 -func SaveDNSAccount(ctx context.Context, account *DNSAccount) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Save(account).Error -} - -// DeleteDNSAccountRecord 删除 DNS 账号。 -func DeleteDNSAccountRecord(ctx context.Context, id uint) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Delete(&DNSAccount{}, id).Error -} diff --git a/internal/model/openflare_node.go b/internal/model/openflare_node.go index 9bab7932..786f48f0 100644 --- a/internal/model/openflare_node.go +++ b/internal/model/openflare_node.go @@ -4,11 +4,7 @@ package model import ( - "context" - "errors" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" ) // OpenFlareNode stores an edge, relay, or tunnel client node. @@ -54,110 +50,3 @@ type OpenFlareNode struct { func (OpenFlareNode) TableName() string { return "of_nodes" } - -// ListOpenFlareNodes returns all nodes ordered by id desc. -func ListOpenFlareNodes(ctx context.Context) ([]OpenFlareNode, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var nodes []OpenFlareNode - if err := conn.Order("id desc").Find(&nodes).Error; err != nil { - return nil, err - } - return nodes, nil -} - -// ListOpenFlareNodesByNodeIDs returns nodes matching the given node ids. -func ListOpenFlareNodesByNodeIDs(ctx context.Context, nodeIDs []string) ([]OpenFlareNode, error) { - if len(nodeIDs) == 0 { - return []OpenFlareNode{}, nil - } - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var nodes []OpenFlareNode - if err := conn.Where("node_id IN ?", nodeIDs).Find(&nodes).Error; err != nil { - return nil, err - } - return nodes, nil -} - -// GetOpenFlareNodeByID returns a node by primary key. -func GetOpenFlareNodeByID(ctx context.Context, id uint) (*OpenFlareNode, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var node OpenFlareNode - if err := conn.First(&node, id).Error; err != nil { - return nil, err - } - return &node, nil -} - -// GetOpenFlareNodeByNodeID returns a node by node_id. -func GetOpenFlareNodeByNodeID(ctx context.Context, nodeID string) (*OpenFlareNode, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var node OpenFlareNode - if err := conn.Where("node_id = ?", nodeID).First(&node).Error; err != nil { - return nil, err - } - return &node, nil -} - -// GetOpenFlareNodeByAccessToken returns a node by access token. -func GetOpenFlareNodeByAccessToken(ctx context.Context, token string) (*OpenFlareNode, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var node OpenFlareNode - if err := conn.Where("access_token = ?", token).First(&node).Error; err != nil { - return nil, err - } - return &node, nil -} - -// CreateOpenFlareNode inserts a new node. -func CreateOpenFlareNode(ctx context.Context, node *OpenFlareNode) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Create(node).Error -} - -// SaveOpenFlareNode persists node changes. -func SaveOpenFlareNode(ctx context.Context, node *OpenFlareNode) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Save(node).Error -} - -// UpdateOpenFlareNodeFields updates selected columns for a node. -func UpdateOpenFlareNodeFields(ctx context.Context, node *OpenFlareNode, fields ...string) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - if len(fields) == 0 { - return conn.Save(node).Error - } - return conn.Model(node).Select(fields).Updates(node).Error -} - -// DeleteOpenFlareNode removes a node by primary key. -func DeleteOpenFlareNode(ctx context.Context, id uint) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Delete(&OpenFlareNode{}, id).Error -} diff --git a/internal/model/openflare_observability.go b/internal/model/openflare_observability.go index ecf98035..3c446497 100644 --- a/internal/model/openflare_observability.go +++ b/internal/model/openflare_observability.go @@ -4,14 +4,7 @@ package model import ( - "context" - "errors" - "strings" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" - "gorm.io/gorm" ) // OpenFlareMetricSnapshot stores a node capacity snapshot in ClickHouse (database: openflare, table: of_node_metric_snapshots). @@ -285,80 +278,6 @@ type OpenFlareAccessLogWAFIPAggregate struct { StatusCounts map[int]int } -func isMissingTableError(err error) bool { - if err == nil { - return false - } - if errors.Is(err, gorm.ErrRecordNotFound) { - return false - } - msg := strings.ToLower(err.Error()) - return strings.Contains(msg, "no such table") || - strings.Contains(msg, "doesn't exist") || - strings.Contains(msg, "does not exist") -} - -// InsertOpenFlareMetricSnapshot inserts a metric snapshot into ClickHouse. -func InsertOpenFlareMetricSnapshot(ctx context.Context, record *OpenFlareMetricSnapshot) error { - return currentObservabilityStore().InsertMetricSnapshot(ctx, record) -} - -// InsertOpenFlareEdgeHealth inserts an L2 edge health snapshot into ClickHouse. -func InsertOpenFlareEdgeHealth(ctx context.Context, record *OpenFlareEdgeHealth) error { - return currentObservabilityStore().InsertEdgeHealth(ctx, record) -} - -// InsertOpenFlareNodeObservationFrps inserts an FRPS observation into ClickHouse. -func InsertOpenFlareNodeObservationFrps(ctx context.Context, record *OpenFlareNodeObservationFrps) error { - return currentObservabilityStore().InsertNodeObservationFrps(ctx, record) -} - -// InsertOpenFlareNodeObservationFrpc inserts an FRPC observation into ClickHouse. -func InsertOpenFlareNodeObservationFrpc(ctx context.Context, record *OpenFlareNodeObservationFrpc) error { - return currentObservabilityStore().InsertNodeObservationFrpc(ctx, record) -} - -// ListOpenFlareMetricSnapshotsSince returns metric snapshots since the given time. -func ListOpenFlareMetricSnapshotsSince(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error) { - return currentObservabilityStore().ListMetricSnapshots(ctx, nodeID, since, limit) -} - -// ListOpenFlareLatestMetricSnapshotsSince returns the latest metric snapshot per node. -// Prefer ClickHouse LIMIT 1 BY; on CH unavailability fall back to store list + reduce. -func ListOpenFlareLatestMetricSnapshotsSince(ctx context.Context, nodeID string, since time.Time) ([]*OpenFlareMetricSnapshot, error) { - rows, err := analyticsrepo.ListLatestNodeMetricSnapshots(ctx, analyticsrepo.NodeObservabilityFilter{ - NodeID: nodeID, - Since: since, - }) - if err == nil { - return fromAnalyticsNodeMetricSnapshots(rows), nil - } - // Fallback for unit tests (memory store) and environments without ClickHouse. - all, listErr := ListOpenFlareMetricSnapshotsSince(ctx, nodeID, since, 0) - if listErr != nil { - return nil, err - } - return openFlareLatestMetricSnapshots(all), nil -} - -func openFlareLatestMetricSnapshots(snapshots []*OpenFlareMetricSnapshot) []*OpenFlareMetricSnapshot { - latestByNode := make(map[string]*OpenFlareMetricSnapshot, len(snapshots)) - for _, snapshot := range snapshots { - if snapshot == nil || snapshot.NodeID == "" { - continue - } - if existing, ok := latestByNode[snapshot.NodeID]; ok && !snapshot.CapturedAt.After(existing.CapturedAt) { - continue - } - latestByNode[snapshot.NodeID] = snapshot - } - result := make([]*OpenFlareMetricSnapshot, 0, len(latestByNode)) - for _, snapshot := range latestByNode { - result = append(result, snapshot) - } - return result -} - // OpenFlareTrafficHourly is an hourly traffic rollup row. type OpenFlareTrafficHourly struct { NodeID string `json:"node_id"` @@ -368,29 +287,6 @@ type OpenFlareTrafficHourly struct { UniqueVisitorCount int64 `json:"unique_visitor_count"` } -// ListOpenFlareTrafficHourlySince returns hourly traffic rollup rows since the given time. -// Source: of_access_log_hourly (M5). -func ListOpenFlareTrafficHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*OpenFlareTrafficHourly, error) { - rows, err := analyticsrepo.ListNodeTrafficHourly(ctx, analyticsrepo.NodeObservabilityFilter{ - NodeID: nodeID, - Since: since, - }) - if err != nil { - return nil, err - } - result := make([]*OpenFlareTrafficHourly, len(rows)) - for index, row := range rows { - result[index] = &OpenFlareTrafficHourly{ - NodeID: row.NodeID, - Hour: row.Hour, - RequestCount: row.RequestCount, - ErrorCount: row.ErrorCount, - UniqueVisitorCount: row.UniqueVisitorCount, - } - } - return result, nil -} - // OpenFlareAccessLogHourly is a per-node/host hourly access log rollup. type OpenFlareAccessLogHourly struct { NodeID string `json:"node_id"` @@ -402,30 +298,6 @@ type OpenFlareAccessLogHourly struct { RequestLength int64 `json:"request_length"` } -// ListOpenFlareAccessLogHourlySince returns of_access_log_hourly rows since the given time. -func ListOpenFlareAccessLogHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*OpenFlareAccessLogHourly, error) { - rows, err := analyticsrepo.ListAccessLogHourly(ctx, analyticsrepo.NodeObservabilityFilter{ - NodeID: nodeID, - Since: since, - }) - if err != nil { - return nil, err - } - result := make([]*OpenFlareAccessLogHourly, len(rows)) - for index, row := range rows { - result[index] = &OpenFlareAccessLogHourly{ - NodeID: row.NodeID, - Hour: row.Hour, - Host: row.Host, - RequestCount: row.RequestCount, - ErrorCount: row.ErrorCount, - BytesSent: row.BytesSent, - RequestLength: row.RequestLength, - } - } - return result, nil -} - // OpenFlareMetricHourly is an hourly metric snapshot aggregation row. type OpenFlareMetricHourly struct { Hour time.Time `json:"hour"` @@ -437,154 +309,3 @@ type OpenFlareMetricHourly struct { DiskWriteBytes int64 `json:"disk_write_bytes"` ReportedNodes int `json:"reported_nodes"` } - -// ListOpenFlareMetricHourlySince returns hourly metric aggregates since the given time. -func ListOpenFlareMetricHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*OpenFlareMetricHourly, error) { - rows, err := analyticsrepo.ListNodeMetricHourly(ctx, analyticsrepo.NodeObservabilityFilter{ - NodeID: nodeID, - Since: since, - }) - if err != nil { - return nil, err - } - result := make([]*OpenFlareMetricHourly, len(rows)) - for index, row := range rows { - result[index] = &OpenFlareMetricHourly{ - Hour: row.Hour, - AverageCPUUsagePercent: row.AverageCPUUsagePercent, - AverageMemoryUsagePercent: row.AverageMemoryUsagePercent, - NetworkRxBytes: row.NetworkRxBytes, - NetworkTxBytes: row.NetworkTxBytes, - DiskReadBytes: row.DiskReadBytes, - DiskWriteBytes: row.DiskWriteBytes, - ReportedNodes: row.ReportedNodes, - } - } - return result, nil -} - -// ListOpenFlareActiveHealthEvents returns active health events across all nodes. -func ListOpenFlareActiveHealthEvents(ctx context.Context) ([]*OpenFlareHealthEvent, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var rows []*OpenFlareHealthEvent - if err := conn.Where("status = ?", "active").Order("last_triggered_at desc").Find(&rows).Error; err != nil { - if isMissingTableError(err) { - return []*OpenFlareHealthEvent{}, nil - } - return nil, err - } - return rows, nil -} - -// ListOpenFlareHealthEvents returns health events for a node. -func ListOpenFlareHealthEvents(ctx context.Context, nodeID string, activeOnly bool, limit int) ([]*OpenFlareHealthEvent, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - query := conn.Model(&OpenFlareHealthEvent{}).Where("node_id = ?", nodeID).Order("last_triggered_at desc") - if activeOnly { - query = query.Where("status = ?", "active") - } - if limit > 0 { - query = query.Limit(limit) - } - var rows []*OpenFlareHealthEvent - if err := query.Find(&rows).Error; err != nil { - if isMissingTableError(err) { - return []*OpenFlareHealthEvent{}, nil - } - return nil, err - } - return rows, nil -} - -// DeleteOpenFlareMetricSnapshotsBefore deletes metric snapshots captured before cutoff. -func DeleteOpenFlareMetricSnapshotsBefore(ctx context.Context, cutoff time.Time) (int64, error) { - return currentObservabilityStore().DeleteMetricSnapshotsBefore(ctx, cutoff) -} - -// DeleteAllOpenFlareMetricSnapshots deletes all metric snapshots. -func DeleteAllOpenFlareMetricSnapshots(ctx context.Context) (int64, error) { - return currentObservabilityStore().DeleteAllMetricSnapshots(ctx) -} - -// DeleteOpenFlareEdgeHealthBefore deletes edge health rows captured before cutoff. -func DeleteOpenFlareEdgeHealthBefore(ctx context.Context, cutoff time.Time) (int64, error) { - return currentObservabilityStore().DeleteEdgeHealthBefore(ctx, cutoff) -} - -// DeleteAllOpenFlareEdgeHealth deletes all edge health snapshots. -func DeleteAllOpenFlareEdgeHealth(ctx context.Context) (int64, error) { - return currentObservabilityStore().DeleteAllEdgeHealth(ctx) -} - -// DeleteOpenFlareNodeObservationFrpsBefore deletes FRPS observations captured before cutoff. -func DeleteOpenFlareNodeObservationFrpsBefore(ctx context.Context, cutoff time.Time) (int64, error) { - return currentObservabilityStore().DeleteNodeObservationFrpsBefore(ctx, cutoff) -} - -// DeleteAllOpenFlareNodeObservationFrps deletes all FRPS observations. -func DeleteAllOpenFlareNodeObservationFrps(ctx context.Context) (int64, error) { - return currentObservabilityStore().DeleteAllNodeObservationFrps(ctx) -} - -// DeleteOpenFlareNodeObservationFrpcBefore deletes FRPC observations captured before cutoff. -func DeleteOpenFlareNodeObservationFrpcBefore(ctx context.Context, cutoff time.Time) (int64, error) { - return currentObservabilityStore().DeleteNodeObservationFrpcBefore(ctx, cutoff) -} - -// DeleteAllOpenFlareNodeObservationFrpc deletes all FRPC observations. -func DeleteAllOpenFlareNodeObservationFrpc(ctx context.Context) (int64, error) { - return currentObservabilityStore().DeleteAllNodeObservationFrpc(ctx) -} - -// DeleteOpenFlareHealthEventsByNodeID deletes all health events for a node. -func DeleteOpenFlareHealthEventsByNodeID(ctx context.Context, nodeID string) (int64, error) { - conn := db.DB(ctx) - if conn == nil { - return 0, errors.New(errDatabaseNotInitialized) - } - result := conn.Where("node_id = ?", nodeID).Delete(&OpenFlareHealthEvent{}) - if result.Error != nil { - if isMissingTableError(result.Error) { - return 0, nil - } - return 0, result.Error - } - return result.RowsAffected, nil -} - -// GetOpenFlareNodeSystemProfile returns the system profile for a node. -func GetOpenFlareNodeSystemProfile(ctx context.Context, nodeID string) (*OpenFlareNodeSystemProfile, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var profile OpenFlareNodeSystemProfile - if err := conn.Where("node_id = ?", nodeID).First(&profile).Error; err != nil { - if errors.Is(err, gorm.ErrRecordNotFound) || isMissingTableError(err) { - return nil, gorm.ErrRecordNotFound - } - return nil, err - } - return &profile, nil -} - -// ListOpenFlareEdgeHealth returns L2 edge health snapshots. -func ListOpenFlareEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareEdgeHealth, error) { - return currentObservabilityStore().ListEdgeHealth(ctx, nodeID, since, limit) -} - -// ListOpenFlareNodeObservationFrpc returns frpc observations. -func ListOpenFlareNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error) { - return currentObservabilityStore().ListNodeObservationFrpc(ctx, nodeID, since, limit) -} - -// ListOpenFlareNodeObservationFrps returns frps observations. -func ListOpenFlareNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error) { - return currentObservabilityStore().ListNodeObservationFrps(ctx, nodeID, since, limit) -} diff --git a/internal/model/openflare_origin.go b/internal/model/openflare_origin.go index fb39bfcd..6bcca868 100644 --- a/internal/model/openflare_origin.go +++ b/internal/model/openflare_origin.go @@ -4,10 +4,7 @@ package model import ( - "context" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" ) // Origin OpenFlare 源站实体。 @@ -46,88 +43,3 @@ type OriginProxyRoute struct { func (OriginProxyRoute) TableName() string { return tableOfProxyRoutes } - -// HasProxyRoutesTable 判断代理规则表是否已迁移。 -func HasProxyRoutesTable(ctx context.Context) bool { - return db.DB(ctx).Migrator().HasTable(&OriginProxyRoute{}) -} - -// ListOrigins 列出全部源站。 -func ListOrigins(ctx context.Context) ([]Origin, error) { - var origins []Origin - if err := db.DB(ctx).Order("id desc").Find(&origins).Error; err != nil { - return nil, err - } - return origins, nil -} - -// GetOriginByID 按 ID 查询源站。 -func GetOriginByID(ctx context.Context, id uint) (*Origin, error) { - var origin Origin - if err := db.DB(ctx).First(&origin, id).Error; err != nil { - return nil, err - } - return &origin, nil -} - -// GetOriginByAddress 按地址查询源站。 -func GetOriginByAddress(ctx context.Context, address string) (*Origin, error) { - var origin Origin - if err := db.DB(ctx).Where("address = ?", address).First(&origin).Error; err != nil { - return nil, err - } - return &origin, nil -} - -// CreateOriginRecord 创建源站。 -func CreateOriginRecord(ctx context.Context, origin *Origin) error { - return db.DB(ctx).Create(origin).Error -} - -// SaveOrigin 保存源站。 -func SaveOrigin(ctx context.Context, origin *Origin) error { - return db.DB(ctx).Save(origin).Error -} - -// DeleteOriginRecord 删除源站。 -func DeleteOriginRecord(ctx context.Context, id uint) error { - return db.DB(ctx).Delete(&Origin{}, id).Error -} - -// ListOriginRouteCounts 统计各源站关联的代理规则数量。 -func ListOriginRouteCounts(ctx context.Context) ([]OriginRouteCount, error) { - if !HasProxyRoutesTable(ctx) { - return nil, nil - } - result := make([]OriginRouteCount, 0) - err := db.DB(ctx).Model(&OriginProxyRoute{}). - Select("origin_id, COUNT(*) AS route_count"). - Where("origin_id IS NOT NULL"). - Group("origin_id"). - Scan(&result).Error - return result, err -} - -// ListProxyRoutesByOriginID 列出源站关联的代理规则。 -func ListProxyRoutesByOriginID(ctx context.Context, originID uint) ([]OriginProxyRoute, error) { - if !HasProxyRoutesTable(ctx) { - return nil, nil - } - var routes []OriginProxyRoute - if err := db.DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil { - return nil, err - } - return routes, nil -} - -// CountProxyRoutesByOriginID 统计源站关联的代理规则数量。 -func CountProxyRoutesByOriginID(ctx context.Context, originID uint) (int64, error) { - if !HasProxyRoutesTable(ctx) { - return 0, nil - } - var count int64 - if err := db.DB(ctx).Model(&OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil { - return 0, err - } - return count, nil -} diff --git a/internal/model/openflare_pages.go b/internal/model/openflare_pages.go index 5c5b39ed..c6f5c507 100644 --- a/internal/model/openflare_pages.go +++ b/internal/model/openflare_pages.go @@ -4,10 +4,7 @@ package model import ( - "context" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" ) // Pages deployment status constants. @@ -83,88 +80,3 @@ type PagesDeploymentFile struct { func (PagesDeploymentFile) TableName() string { return "of_pages_deployment_files" } - -// HasPagesProjectsTable 判断 Pages 项目表是否已迁移。 -func HasPagesProjectsTable(ctx context.Context) bool { - return db.DB(ctx).Migrator().HasTable(&PagesProject{}) -} - -// ListPagesProjects 列出全部 Pages 项目。 -func ListPagesProjects(ctx context.Context) ([]PagesProject, error) { - var projects []PagesProject - if err := db.DB(ctx).Order("id desc").Find(&projects).Error; err != nil { - return nil, err - } - return projects, nil -} - -// GetPagesProjectByID 按 ID 查询 Pages 项目。 -func GetPagesProjectByID(ctx context.Context, id uint) (*PagesProject, error) { - var project PagesProject - if err := db.DB(ctx).First(&project, id).Error; err != nil { - return nil, err - } - return &project, nil -} - -// GetPagesProjectBySlug 按 slug 查询 Pages 项目。 -func GetPagesProjectBySlug(ctx context.Context, slug string) (*PagesProject, error) { - var project PagesProject - if err := db.DB(ctx).Where("slug = ?", slug).First(&project).Error; err != nil { - return nil, err - } - return &project, nil -} - -// CreatePagesProjectRecord 创建 Pages 项目。 -func CreatePagesProjectRecord(ctx context.Context, project *PagesProject) error { - return db.DB(ctx).Create(project).Error -} - -// ListPagesDeployments 列出项目的全部部署。 -func ListPagesDeployments(ctx context.Context, projectID uint) ([]PagesDeployment, error) { - var deployments []PagesDeployment - if err := db.DB(ctx).Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil { - return nil, err - } - return deployments, nil -} - -// GetPagesDeploymentByID 按 ID 查询 Pages 部署。 -func GetPagesDeploymentByID(ctx context.Context, id uint) (*PagesDeployment, error) { - var deployment PagesDeployment - if err := db.DB(ctx).First(&deployment, id).Error; err != nil { - return nil, err - } - return &deployment, nil -} - -// ListPagesDeploymentFiles 列出部署文件清单。 -func ListPagesDeploymentFiles(ctx context.Context, deploymentID uint) ([]PagesDeploymentFile, error) { - var files []PagesDeploymentFile - if err := db.DB(ctx).Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error; err != nil { - return nil, err - } - return files, nil -} - -// CountPagesDeploymentsByProjectID 统计项目部署数量。 -func CountPagesDeploymentsByProjectID(ctx context.Context, projectID uint) (int64, error) { - var count int64 - if err := db.DB(ctx).Model(&PagesDeployment{}).Where("project_id = ?", projectID).Count(&count).Error; err != nil { - return 0, err - } - return count, nil -} - -// CountProxyRoutesByPagesProjectID 统计引用 Pages 项目的代理规则数量。 -func CountProxyRoutesByPagesProjectID(ctx context.Context, projectID uint) (int64, error) { - if !HasProxyRoutesTable(ctx) { - return 0, nil - } - var count int64 - if err := db.DB(ctx).Model(&ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil { - return 0, err - } - return count, nil -} diff --git a/internal/model/openflare_pages_cleanup.go b/internal/model/openflare_pages_cleanup.go index 55cbf825..adff8e87 100644 --- a/internal/model/openflare_pages_cleanup.go +++ b/internal/model/openflare_pages_cleanup.go @@ -4,19 +4,16 @@ package model import ( - "context" - "errors" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" ) const ( // PagesOrphanUploadCandidateLimit bounds one delayed Pages upload cleanup pass. PagesOrphanUploadCandidateLimit = 100 - - pagesOrphanMarkerPredicatePostgres = "w_uploads.metadata #>> '{extra,pages_ingest_marker}' = ?" - pagesOrphanMarkerPredicateSQLite = "CASE WHEN json_valid(w_uploads.metadata) THEN json_extract(w_uploads.metadata, '$.extra.pages_ingest_marker') ELSE NULL END = ?" + // PagesOrphanMarkerPredicatePostgres is the Postgres JSON marker match SQL fragment. + PagesOrphanMarkerPredicatePostgres = "w_uploads.metadata #>> '{extra,pages_ingest_marker}' = ?" + // PagesOrphanMarkerPredicateSQLite is the SQLite JSON marker match SQL fragment. + PagesOrphanMarkerPredicateSQLite = "CASE WHEN json_valid(w_uploads.metadata) THEN json_extract(w_uploads.metadata, '$.extra.pages_ingest_marker') ELSE NULL END = ?" ) // PagesOrphanUploadCandidateQuery describes the fail-closed SQL candidate set @@ -27,49 +24,3 @@ type PagesOrphanUploadCandidateQuery struct { Marker string CreatedBefore time.Time } - -// ListPagesOrphanUploadCandidates returns at most 100 unreferenced, isolated -// Pages V2 upload records. Callers must still lock and recheck every condition -// before deleting a candidate. -func ListPagesOrphanUploadCandidates( - ctx context.Context, - input PagesOrphanUploadCandidateQuery, -) ([]Upload, error) { - if input.SystemUserID == 0 || input.UploadType == "" || input.Marker == "" || input.CreatedBefore.IsZero() { - return nil, errors.New("invalid pages orphan upload candidate query") - } - markerPredicate, err := pagesOrphanMarkerPredicate(db.DB(ctx).Name()) - if err != nil { - return nil, err - } - - deploymentTable := (PagesDeployment{}).TableName() - uploadTable := (Upload{}).TableName() - var candidates []Upload - err = db.DB(ctx). - Model(&Upload{}). - Where(uploadTable+".status = ?", UploadStatusUsed). - Where(uploadTable+".user_id = ?", input.SystemUserID). - Where(uploadTable+".type = ?", input.UploadType). - Where(uploadTable+".created_at < ?", input.CreatedBefore). - Where(markerPredicate, input.Marker). - Where("NOT EXISTS (SELECT 1 FROM " + deploymentTable + " WHERE " + deploymentTable + ".upload_id = " + uploadTable + ".id)"). - Order(uploadTable + ".id ASC"). - Limit(PagesOrphanUploadCandidateLimit). - Find(&candidates).Error - if err != nil { - return nil, err - } - return candidates, nil -} - -func pagesOrphanMarkerPredicate(dialect string) (string, error) { - switch dialect { - case "postgres": - return pagesOrphanMarkerPredicatePostgres, nil - case "sqlite": - return pagesOrphanMarkerPredicateSQLite, nil - default: - return "", errors.New("unsupported database dialect for Pages orphan cleanup") - } -} diff --git a/internal/model/openflare_pages_source.go b/internal/model/openflare_pages_source.go index 63c0ca89..62481ce8 100644 --- a/internal/model/openflare_pages_source.go +++ b/internal/model/openflare_pages_source.go @@ -56,3 +56,19 @@ type PagesProjectSourceRuntime struct { func (PagesProjectSourceRuntime) TableName() string { return "of_pages_project_source_runtime" } + +// PagesExpiredSourceLeaseCandidate is a scanner query DTO for expired runtime leases. +type PagesExpiredSourceLeaseCandidate struct { + SourceID uint + LeaseToken string + LeaseExpiresAt time.Time + SyncStatus string + SourceType string + ReleaseSelector string +} + +// PagesDueGitHubSourceCandidate is a scanner query DTO for due GitHub latest checks. +type PagesDueGitHubSourceCandidate struct { + SourceID uint + ConfigVersion int +} diff --git a/internal/model/openflare_proxy_route.go b/internal/model/openflare_proxy_route.go index 5ae7802c..4e95b00e 100644 --- a/internal/model/openflare_proxy_route.go +++ b/internal/model/openflare_proxy_route.go @@ -4,10 +4,7 @@ package model import ( - "context" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" ) // ProxyRoute OpenFlare 代理规则实体。 @@ -47,61 +44,3 @@ type ProxyRoute struct { func (ProxyRoute) TableName() string { return tableOfProxyRoutes } - -// ListProxyRoutes 列出全部代理规则。 -func ListProxyRoutes(ctx context.Context) ([]*ProxyRoute, error) { - var routes []*ProxyRoute - if err := db.DB(ctx).Order("id desc").Find(&routes).Error; err != nil { - return nil, err - } - return routes, nil -} - -// GetProxyRouteByID 按 ID 查询代理规则。 -func GetProxyRouteByID(ctx context.Context, id uint) (*ProxyRoute, error) { - var route ProxyRoute - if err := db.DB(ctx).First(&route, id).Error; err != nil { - return nil, err - } - return &route, nil -} - -// CreateProxyRouteRecord 创建代理规则。 -func CreateProxyRouteRecord(ctx context.Context, route *ProxyRoute) error { - return db.DB(ctx).Create(route).Error -} - -// UpdateProxyRouteRecord 更新代理规则。 -func UpdateProxyRouteRecord(ctx context.Context, route *ProxyRoute) error { - return db.DB(ctx).Model(&ProxyRoute{}).Where("id = ?", route.ID).Updates(map[string]any{ - "site_name": route.SiteName, - "origin_id": route.OriginID, - "origin_url": route.OriginURL, - "origin_host": route.OriginHost, - "upstreams": route.Upstreams, - colEnabled: route.Enabled, - "enable_https": route.EnableHTTPS, - "redirect_http": route.RedirectHTTP, - "limit_conn_per_server": route.LimitConnPerServer, - "limit_conn_per_ip": route.LimitConnPerIP, - "limit_rate": route.LimitRate, - "limit_req_per_ip": route.LimitReqPerIP, - "cache_enabled": route.CacheEnabled, - "cache_policy": route.CachePolicy, - "cache_rules": route.CacheRules, - "custom_headers": route.CustomHeaders, - "basic_auth_enabled": route.BasicAuthEnabled, - "basic_auth_username": route.BasicAuthUsername, - "basic_auth_password": route.BasicAuthPassword, - "upstream_type": route.UpstreamType, - "tunnel_node_id": route.TunnelNodeID, - "tunnel_target_addr": route.TunnelTargetAddr, - "tunnel_target_protocol": route.TunnelTargetProtocol, - "pages_project_id": route.PagesProjectID, - }).Error -} - -// DeleteProxyRouteRecord 删除代理规则。 -func DeleteProxyRouteRecord(ctx context.Context, id uint) error { - return db.DB(ctx).Delete(&ProxyRoute{}, id).Error -} diff --git a/internal/model/openflare_tls.go b/internal/model/openflare_tls.go index a8b01f68..71d64573 100644 --- a/internal/model/openflare_tls.go +++ b/internal/model/openflare_tls.go @@ -4,11 +4,7 @@ package model import ( - "context" - "errors" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" ) // TLSCertificate OpenFlare TLS 证书实体。 @@ -54,86 +50,3 @@ type TLSProxyRouteRef struct { func (TLSProxyRouteRef) TableName() string { return tableOfProxyRoutes } - -// HasTLSProxyRoutesTable 判断代理规则表是否已迁移。 -func HasTLSProxyRoutesTable(ctx context.Context) bool { - return db.DB(ctx).Migrator().HasTable(&TLSProxyRouteRef{}) -} - -// ListTLSCertificates 列出全部证书(不含 PEM 敏感字段的 JSON 暴露由 struct tag 控制)。 -func ListTLSCertificates(ctx context.Context) ([]TLSCertificate, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var certificates []TLSCertificate - if err := conn.Order("id desc").Find(&certificates).Error; err != nil { - return nil, err - } - return certificates, nil -} - -// GetTLSCertificateByID 按 ID 查询证书。 -func GetTLSCertificateByID(ctx context.Context, id uint) (*TLSCertificate, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - var certificate TLSCertificate - if err := conn.First(&certificate, id).Error; err != nil { - return nil, err - } - return &certificate, nil -} - -// CreateTLSCertificateRecord 创建证书记录。 -func CreateTLSCertificateRecord(ctx context.Context, certificate *TLSCertificate) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Create(certificate).Error -} - -// SaveTLSCertificate 保存证书记录。 -func SaveTLSCertificate(ctx context.Context, certificate *TLSCertificate) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Save(certificate).Error -} - -// DeleteTLSCertificateRecord 删除证书记录。 -func DeleteTLSCertificateRecord(ctx context.Context, id uint) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New(errDatabaseNotInitialized) - } - return conn.Delete(&TLSCertificate{}, id).Error -} - -// CountTLSCertificatesByDNSAccountID 统计引用指定 DNS 账号的证书数量。 -func CountTLSCertificatesByDNSAccountID(ctx context.Context, dnsAccountID uint) (int64, error) { - conn := db.DB(ctx) - if conn == nil { - return 0, errors.New(errDatabaseNotInitialized) - } - var count int64 - if err := conn.Model(&TLSCertificate{}).Where("dns_account_id = ?", dnsAccountID).Count(&count).Error; err != nil { - return 0, err - } - return count, nil -} - -// ListTLSProxyRouteRefs 列出代理规则证书引用字段。 -func ListTLSProxyRouteRefs(ctx context.Context) ([]TLSProxyRouteRef, error) { - if !HasTLSProxyRoutesTable(ctx) { - return nil, nil - } - var routes []TLSProxyRouteRef - if err := db.DB(ctx).Order("id asc").Find(&routes).Error; err != nil { - return nil, err - } - return routes, nil -} diff --git a/internal/model/openflare_waf.go b/internal/model/openflare_waf.go index 9e858306..b0e0bcee 100644 --- a/internal/model/openflare_waf.go +++ b/internal/model/openflare_waf.go @@ -4,12 +4,8 @@ package model import ( - "context" "errors" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "gorm.io/gorm" ) // OpenFlareWAFRuleGroup stores a WAF rule group. @@ -71,345 +67,3 @@ var ErrWAFRuleRevisionConflict = errors.New("waf rule revision conflict") func (OpenFlareWAFRuleGroupBinding) TableName() string { return "of_waf_rule_group_bindings" } - -func wafDB(ctx context.Context) (*gorm.DB, error) { - conn := db.DB(ctx) - if conn == nil { - return nil, errors.New(errDatabaseNotInitialized) - } - return conn, nil -} - -// ListOpenFlareWAFRuleGroups returns all rule groups. -func ListOpenFlareWAFRuleGroups(ctx context.Context) ([]*OpenFlareWAFRuleGroup, error) { - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var groups []*OpenFlareWAFRuleGroup - if err = conn.Order("is_global desc").Order("id asc").Find(&groups).Error; err != nil { - return nil, err - } - return groups, nil -} - -// GetOpenFlareWAFRuleGroupByID returns a rule group by id. -func GetOpenFlareWAFRuleGroupByID(ctx context.Context, id uint) (*OpenFlareWAFRuleGroup, error) { - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var group OpenFlareWAFRuleGroup - if err = conn.First(&group, id).Error; err != nil { - return nil, err - } - return &group, nil -} - -// GetGlobalOpenFlareWAFRuleGroup returns the global rule group if present. -func GetGlobalOpenFlareWAFRuleGroup(ctx context.Context) (*OpenFlareWAFRuleGroup, error) { - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var group OpenFlareWAFRuleGroup - if err = conn.Where("is_global = ?", true).Order("id asc").First(&group).Error; err != nil { - return nil, err - } - return &group, nil -} - -// CreateOpenFlareWAFRuleGroup inserts a rule group. -func CreateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGroup) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Create(group).Error -} - -// UpdateOpenFlareWAFRuleGroup updates mutable rule group fields. -func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGroup) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Model(&OpenFlareWAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ - "name": group.Name, - colEnabled: group.Enabled, - "is_global": group.IsGlobal, - }).Error -} - -// UpdateOpenFlareWAFRuleGraph atomically replaces a graph when revision is current. -func UpdateOpenFlareWAFRuleGraph(ctx context.Context, id uint, revision uint64, graph string) (uint64, error) { - conn, err := wafDB(ctx) - if err != nil { - return 0, err - } - result := conn.Model(&OpenFlareWAFRuleGroup{}). - Where("id = ? AND revision = ?", id, revision). - Updates(map[string]any{"graph": graph, "revision": gorm.Expr("revision + 1")}) - if result.Error != nil { - return 0, result.Error - } - if result.RowsAffected != 1 { - return 0, ErrWAFRuleRevisionConflict - } - return revision + 1, nil -} - -// DeleteOpenFlareWAFRuleGroup removes a rule group. -func DeleteOpenFlareWAFRuleGroup(ctx context.Context, id uint) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Delete(&OpenFlareWAFRuleGroup{}, id).Error -} - -// ListOpenFlareWAFIPGroups returns all IP groups. -func ListOpenFlareWAFIPGroups(ctx context.Context) ([]*OpenFlareWAFIPGroup, error) { - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var groups []*OpenFlareWAFIPGroup - if err = conn.Order("type asc").Order("id asc").Find(&groups).Error; err != nil { - return nil, err - } - return groups, nil -} - -// ListOpenFlareWAFIPGroupsByIDs returns IP groups for the given ids. -func ListOpenFlareWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*OpenFlareWAFIPGroup, error) { - if len(ids) == 0 { - return []*OpenFlareWAFIPGroup{}, nil - } - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var groups []*OpenFlareWAFIPGroup - if err = conn.Where("id IN ?", ids).Find(&groups).Error; err != nil { - return nil, err - } - return groups, nil -} - -// GetOpenFlareWAFIPGroupByID returns an IP group by id. -func GetOpenFlareWAFIPGroupByID(ctx context.Context, id uint) (*OpenFlareWAFIPGroup, error) { - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var group OpenFlareWAFIPGroup - if err = conn.First(&group, id).Error; err != nil { - return nil, err - } - return &group, nil -} - -// CreateOpenFlareWAFIPGroup inserts an IP group. -func CreateOpenFlareWAFIPGroup(ctx context.Context, group *OpenFlareWAFIPGroup) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Create(group).Error -} - -// UpdateOpenFlareWAFIPGroup updates mutable IP group fields. -func UpdateOpenFlareWAFIPGroup(ctx context.Context, group *OpenFlareWAFIPGroup) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Model(&OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ - "name": group.Name, - "type": group.Type, - "enabled": group.Enabled, - "ip_list": group.IPList, - "auto_config": group.AutoConfig, - "ext_ips": group.ExtIPs, - "subscription_url": group.SubscriptionURL, - "subscription_format": group.SubscriptionFormat, - "subscription_mapping_rule": group.SubscriptionMappingRule, - "sync_interval_minutes": group.SyncIntervalMinutes, - "next_sync_at": group.NextSyncAt, - "last_sync_status": group.LastSyncStatus, - "last_sync_message": group.LastSyncMessage, - }).Error -} - -// ListDueOpenFlareWAFIPGroups returns enabled automatic/subscription groups due for sync. -func ListDueOpenFlareWAFIPGroups(ctx context.Context, now time.Time) ([]*OpenFlareWAFIPGroup, error) { - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var groups []*OpenFlareWAFIPGroup - err = conn.Where( - "enabled = ? AND (type = ? OR (type = ? AND subscription_url <> '')) AND (next_sync_at IS NULL OR next_sync_at <= ?)", - true, "automatic", "subscription", now, - ).Order("id asc").Find(&groups).Error - return groups, err -} - -// UpdateOpenFlareWAFIPGroupSyncResult persists IP group sync outcome fields. -func UpdateOpenFlareWAFIPGroupSyncResult(ctx context.Context, group *OpenFlareWAFIPGroup) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Model(&OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ - "ip_list": group.IPList, - "ext_ips": group.ExtIPs, - "last_synced_at": group.LastSyncedAt, - "next_sync_at": group.NextSyncAt, - "last_sync_status": group.LastSyncStatus, - "last_sync_message": group.LastSyncMessage, - "subscription_format": group.SubscriptionFormat, - }).Error -} - -// DeleteOpenFlareWAFIPGroup removes an IP group. -func DeleteOpenFlareWAFIPGroup(ctx context.Context, id uint) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Delete(&OpenFlareWAFIPGroup{}, id).Error -} - -// ListOpenFlareWAFRuleGroupBindings returns all bindings. -func ListOpenFlareWAFRuleGroupBindings(ctx context.Context) ([]OpenFlareWAFRuleGroupBinding, error) { - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var bindings []OpenFlareWAFRuleGroupBinding - if err = conn.Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil { - return nil, err - } - return bindings, nil -} - -// ListOpenFlareWAFRuleGroupBindingsByRouteID returns bindings for a proxy route. -func ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx context.Context, routeID uint) ([]OpenFlareWAFRuleGroupBinding, error) { - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var bindings []OpenFlareWAFRuleGroupBinding - if err = conn.Where("proxy_route_id = ?", routeID).Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil { - return nil, err - } - return bindings, nil -} - -func syncWAFBindingIDSequence(tx *gorm.DB) error { - if tx == nil || tx.Dialector.Name() != "postgres" { //nolint:staticcheck // QF1008: keep explicit Dialector field access - return nil - } - return tx.Exec(` - SELECT setval( - pg_get_serial_sequence('of_waf_rule_group_bindings', 'id'), - GREATEST(COALESCE((SELECT MAX(id) FROM of_waf_rule_group_bindings), 0), 1), - COALESCE((SELECT MAX(id) FROM of_waf_rule_group_bindings), 0) > 0 - ) - `).Error -} - -func insertOpenFlareWAFRuleGroupBindings(tx *gorm.DB, bindings []OpenFlareWAFRuleGroupBinding) error { - if len(bindings) == 0 { - return nil - } - if err := syncWAFBindingIDSequence(tx); err != nil { - return err - } - return tx.Create(&bindings).Error -} - -// ReplaceOpenFlareWAFRuleGroupBindings replaces bindings for a rule group. -func ReplaceOpenFlareWAFRuleGroupBindings(ctx context.Context, groupID uint, routeIDs []uint) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Transaction(func(tx *gorm.DB) error { - if err = tx.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil { - return err - } - bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(routeIDs)) - for index, routeID := range routeIDs { - bindings = append(bindings, OpenFlareWAFRuleGroupBinding{ - RuleGroupID: groupID, - ProxyRouteID: routeID, - Sequence: index, - }) - } - return insertOpenFlareWAFRuleGroupBindings(tx, bindings) - }) -} - -// ReplaceOpenFlareWAFSiteRuleGroupBindings replaces bindings for a proxy route. -func ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx context.Context, routeID uint, groupIDs []uint) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Transaction(func(tx *gorm.DB) error { - if err = tx.Where("proxy_route_id = ?", routeID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil { - return err - } - bindings := make([]OpenFlareWAFRuleGroupBinding, 0, len(groupIDs)) - for index, groupID := range groupIDs { - bindings = append(bindings, OpenFlareWAFRuleGroupBinding{ - RuleGroupID: groupID, - ProxyRouteID: routeID, - Sequence: index, - }) - } - return insertOpenFlareWAFRuleGroupBindings(tx, bindings) - }) -} - -// DeleteOpenFlareWAFRuleGroupBindingsByGroupID removes bindings for a rule group. -func DeleteOpenFlareWAFRuleGroupBindingsByGroupID(ctx context.Context, groupID uint) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error -} - -// DeleteOpenFlareWAFRuleGroupWithBindings removes a rule group and its bindings. -func DeleteOpenFlareWAFRuleGroupWithBindings(ctx context.Context, groupID uint) error { - conn, err := wafDB(ctx) - if err != nil { - return err - } - return conn.Transaction(func(tx *gorm.DB) error { - if err = tx.Where("rule_group_id = ?", groupID).Delete(&OpenFlareWAFRuleGroupBinding{}).Error; err != nil { - return err - } - return tx.Delete(&OpenFlareWAFRuleGroup{}, groupID).Error - }) -} - -// GetOpenFlareProxyRouteByID returns a proxy route by id when the table exists. -func GetOpenFlareProxyRouteByID(ctx context.Context, id uint) (*OriginProxyRoute, error) { - if !HasProxyRoutesTable(ctx) { - return nil, gorm.ErrRecordNotFound - } - conn, err := wafDB(ctx) - if err != nil { - return nil, err - } - var route OriginProxyRoute - if err = conn.First(&route, id).Error; err != nil { - return nil, err - } - return &route, nil -} diff --git a/internal/model/openflare_zone.go b/internal/model/openflare_zone.go index 90100224..aff02dce 100644 --- a/internal/model/openflare_zone.go +++ b/internal/model/openflare_zone.go @@ -4,14 +4,7 @@ package model import ( - "context" - "errors" - "fmt" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "gorm.io/gorm" - "gorm.io/gorm/clause" ) const ( @@ -19,8 +12,6 @@ const ( tableOfZoneDomains = "of_zone_domains" ) -var errZoneDomainBoundToAnotherRoute = errors.New("zone domain is already bound to another proxy route") - // Zone OpenFlare 注册根域实体。 type Zone struct { ID uint `json:"id" gorm:"primaryKey;autoIncrement"` @@ -50,95 +41,8 @@ func (ZoneDomain) TableName() string { return tableOfZoneDomains } -// ListZoneDomainsByRouteID returns the domains bound to a proxy route. -func ListZoneDomainsByRouteID(ctx context.Context, routeID uint) ([]ZoneDomain, error) { - var domains []ZoneDomain - if err := db.DB(ctx).Where("proxy_route_id = ?", routeID).Order("id asc").Find(&domains).Error; err != nil { - return nil, err - } - return domains, nil -} - -// ListZoneDomainsByIDs returns explicit domains in the requested ID order. -func ListZoneDomainsByIDs(ctx context.Context, domainIDs []uint) ([]ZoneDomain, error) { - if len(domainIDs) == 0 { - return []ZoneDomain{}, nil - } - var domains []ZoneDomain - if err := db.DB(ctx).Where("id IN ?", domainIDs).Find(&domains).Error; err != nil { - return nil, err - } - byID := make(map[uint]ZoneDomain, len(domains)) - for _, domain := range domains { - byID[domain.ID] = domain - } - ordered := make([]ZoneDomain, 0, len(domainIDs)) - for _, id := range domainIDs { - domain, ok := byID[id] - if !ok { - return nil, fmt.Errorf("one or more zone domains do not exist") - } - ordered = append(ordered, domain) - } - return ordered, nil -} - -// CountZoneDomainsByCertificateID reports whether a certificate is assigned to a Zone domain. -func CountZoneDomainsByCertificateID(ctx context.Context, certificateID uint) (int64, error) { - var count int64 - err := db.DB(ctx).Model(&ZoneDomain{}).Where("cert_id = ?", certificateID).Count(&count).Error - return count, err -} - -// ReplaceZoneDomainRouteBindings replaces every ZoneDomain binding for a proxy route. -func ReplaceZoneDomainRouteBindings(ctx context.Context, routeID uint, domainIDs []uint) error { - conn := db.DB(ctx) - if conn == nil { - return errors.New("database is not initialized") - } - - return conn.Transaction(func(tx *gorm.DB) error { - var requested []ZoneDomain - if len(domainIDs) > 0 { - if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). - Where("id IN ?", domainIDs). - Find(&requested).Error; err != nil { - return err - } - if len(requested) != len(uniqueZoneDomainIDs(domainIDs)) { - return fmt.Errorf("one or more zone domains do not exist") - } - for _, domain := range requested { - if domain.ProxyRouteID != nil && *domain.ProxyRouteID != routeID { - return errZoneDomainBoundToAnotherRoute - } - } - } - - current := tx.Model(&ZoneDomain{}).Where("proxy_route_id = ?", routeID) - if len(domainIDs) > 0 { - current = current.Where("id NOT IN ?", domainIDs) - } - if err := current.Update("proxy_route_id", nil).Error; err != nil { - return err - } - - if len(domainIDs) == 0 { - return nil - } - return tx.Model(&ZoneDomain{}).Where("id IN ?", domainIDs).Update("proxy_route_id", routeID).Error - }) -} - -func uniqueZoneDomainIDs(domainIDs []uint) []uint { - seen := make(map[uint]struct{}, len(domainIDs)) - ids := make([]uint, 0, len(domainIDs)) - for _, id := range domainIDs { - if _, ok := seen[id]; ok { - continue - } - seen[id] = struct{}{} - ids = append(ids, id) - } - return ids +// ZoneDomainCount is the per-zone explicit domain count for list queries. +type ZoneDomainCount struct { + ZoneID uint `json:"zone_id" gorm:"column:zone_id"` + Count int64 `json:"count" gorm:"column:count"` } diff --git a/internal/model/schedule.go b/internal/model/schedule.go index b6337456..5d577f6b 100644 --- a/internal/model/schedule.go +++ b/internal/model/schedule.go @@ -4,10 +4,7 @@ package model import ( - "context" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" ) // Schedule 定时任务配置表 @@ -26,45 +23,3 @@ type Schedule struct { func (Schedule) TableName() string { return "w_schedules" } - -// CreateSchedule 创建定时任务 -func CreateSchedule(ctx context.Context, schedule *Schedule) error { - return db.DB(ctx).Create(schedule).Error -} - -// UpdateSchedule 更新定时任务 -func UpdateSchedule(ctx context.Context, schedule *Schedule) error { - return db.DB(ctx).Save(schedule).Error -} - -// DeleteSchedule 删除定时任务 -func DeleteSchedule(ctx context.Context, id uint64) error { - return db.DB(ctx).Delete(&Schedule{}, id).Error -} - -// GetScheduleByID 根据 ID 获取定时任务 -func GetScheduleByID(ctx context.Context, id uint64) (*Schedule, error) { - var schedule Schedule - if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil { - return nil, err - } - return &schedule, nil -} - -// ListSchedules 获取所有定时任务 -func ListSchedules(ctx context.Context) ([]Schedule, error) { - var schedules []Schedule - if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil { - return nil, err - } - return schedules, nil -} - -// ListActiveSchedules 获取所有启用的定时任务 -func ListActiveSchedules(ctx context.Context) ([]Schedule, error) { - var schedules []Schedule - if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil { - return nil, err - } - return schedules, nil -} diff --git a/internal/model/task_execution.go b/internal/model/task_execution.go index bd4b659e..2db07cdd 100644 --- a/internal/model/task_execution.go +++ b/internal/model/task_execution.go @@ -5,15 +5,7 @@ package model import ( - "context" - "errors" - "fmt" - "strings" "time" - - db "github.com/Rain-kl/Wavelet/internal/infra/persistence" - "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" - "github.com/redis/go-redis/v9" ) // TaskExecutionStatus 任务执行状态 @@ -25,10 +17,6 @@ const ( TaskExecutionStatusRunning TaskExecutionStatus = "running" TaskExecutionStatusSucceeded TaskExecutionStatus = "succeeded" TaskExecutionStatusFailed TaskExecutionStatus = "failed" - - taskExecutionLogRedisKeyPrefix = "task:execution:log:" - taskExecutionLogExpiration = 24 * time.Hour - taskExecutionLogMaxLines = 1000 ) // TaskExecution 任务执行记录 @@ -64,95 +52,6 @@ func (TaskExecution) TableName() string { return "w_task_executions" } -// CreateTaskExecution 创建任务执行记录 -func CreateTaskExecution(ctx context.Context, execution *TaskExecution) error { - execution.ID = idgen.NextUint64ID() - return db.DB(ctx).Create(execution).Error -} - -// UpdateTaskExecution 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。 -func UpdateTaskExecution(ctx context.Context, execution *TaskExecution) error { - return db.DB(ctx).Omit("log").Save(execution).Error -} - -// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录 -func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecution, error) { - var execution TaskExecution - if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil { - return nil, err - } - if err := loadTaskExecutionLog(ctx, &execution); err != nil { - return nil, err - } - return &execution, nil -} - -// GetTaskExecutionByID 根据 ID 获取执行记录 -func GetTaskExecutionByID(ctx context.Context, id uint64) (*TaskExecution, error) { - var execution TaskExecution - if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil { - return nil, err - } - if err := loadTaskExecutionLog(ctx, &execution); err != nil { - return nil, err - } - return &execution, nil -} - -// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。 -func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error { - if db.Redis == nil { - return errors.New("redis client is not initialized") - } - - now := time.Now().Format("15:04:05") - line := fmt.Sprintf("[%s] %s\n", now, logLine) - key := taskExecutionLogRedisKey(taskID) - - _, err := db.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error { - pipe.RPush(ctx, key, line) - pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1) - pipe.Expire(ctx, key, taskExecutionLogExpiration) - return nil - }) - if err != nil { - return fmt.Errorf("append task execution log to redis: %w", err) - } - return nil -} - -// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。 -func FlushTaskExecutionLog(ctx context.Context, taskID string) error { - if db.Redis == nil { - return errors.New("redis client is not initialized") - } - - key := taskExecutionLogRedisKey(taskID) - logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result() - if err != nil { - return fmt.Errorf("get task execution log from redis: %w", err) - } - if len(logLines) == 0 { - return nil - } - logText := strings.Join(logLines, "") - - result := db.DB(ctx).Model(&TaskExecution{}). - Where("task_id = ?", taskID). - Update("log", logText) - if result.Error != nil { - return fmt.Errorf("persist task execution log: %w", result.Error) - } - if result.RowsAffected == 0 { - return fmt.Errorf("persist task execution log: task %q not found", taskID) - } - - if err := db.Redis.Del(ctx, key).Err(); err != nil { - return fmt.Errorf("delete persisted task execution log from redis: %w", err) - } - return nil -} - // ListTaskExecutionsRequest 查询任务执行记录列表请求 type ListTaskExecutionsRequest struct { Status string `form:"status"` @@ -160,137 +59,3 @@ type ListTaskExecutionsRequest struct { Page int `form:"page"` PageSize int `form:"page_size"` } - -// ListTaskExecutions 分页查询任务执行记录 -func ListTaskExecutions(ctx context.Context, req ListTaskExecutionsRequest) ([]TaskExecution, int64, error) { - if req.Page <= 0 { - req.Page = 1 - } - if req.PageSize <= 0 { - req.PageSize = 20 - } - - query := db.DB(ctx).Model(&TaskExecution{}) - - if req.Status != "" { - query = query.Where("status = ?", req.Status) - } - if req.TaskType != "" { - query = query.Where("task_type = ?", req.TaskType) - } - - var total int64 - if err := query.Count(&total).Error; err != nil { - return nil, 0, err - } - - var executions []TaskExecution - offset := (req.Page - 1) * req.PageSize - if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil { - return nil, 0, err - } - if err := loadTaskExecutionLogs(ctx, executions); err != nil { - return nil, 0, err - } - - return executions, total, nil -} - -// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention. -func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (TaskExecutionCleanupStats, error) { - const ( - frequencyWindowDays = 30 - highFrequencyThreshold = frequencyWindowDays - ) - - frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays) - highFrequencyCutoff := now.AddDate(0, 0, -3) - lowFrequencyCutoff := now.AddDate(0, 0, -30) - terminalStatuses := []TaskExecutionStatus{TaskExecutionStatusSucceeded, TaskExecutionStatusFailed} - - var highFrequencyTaskTypes []string - if err := db.DB(ctx). - Model(&TaskExecution{}). - Select("task_type"). - Where("created_at >= ?", frequencyWindowStart). - Group("task_type"). - Having("COUNT(*) > ?", highFrequencyThreshold). - Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil { - return TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err) - } - - var highFrequencyDeleted int64 - if len(highFrequencyTaskTypes) > 0 { - highFrequencyResult := db.DB(ctx). - Where("status IN ?", terminalStatuses). - Where("created_at < ?", highFrequencyCutoff). - Where("task_type IN ?", highFrequencyTaskTypes). - Delete(&TaskExecution{}) - if highFrequencyResult.Error != nil { - return TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error) - } - highFrequencyDeleted = highFrequencyResult.RowsAffected - } - - lowFrequencyQuery := db.DB(ctx). - Where("status IN ?", terminalStatuses). - Where("created_at < ?", lowFrequencyCutoff) - if len(highFrequencyTaskTypes) > 0 { - lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes) - } - lowFrequencyResult := lowFrequencyQuery.Delete(&TaskExecution{}) - if lowFrequencyResult.Error != nil { - return TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error) - } - - return TaskExecutionCleanupStats{ - HighFrequencyDeleted: highFrequencyDeleted, - LowFrequencyDeleted: lowFrequencyResult.RowsAffected, - }, nil -} - -func taskExecutionLogRedisKey(taskID string) string { - return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID) -} - -func loadTaskExecutionLog(ctx context.Context, execution *TaskExecution) error { - if db.Redis == nil { - return nil - } - - logLines, err := db.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result() - if err != nil { - return fmt.Errorf("get task execution log from redis: %w", err) - } - if len(logLines) == 0 { - return nil - } - - execution.Log = strings.Join(logLines, "") - return nil -} - -func loadTaskExecutionLogs(ctx context.Context, executions []TaskExecution) error { - if db.Redis == nil || len(executions) == 0 { - return nil - } - - commands := make([]*redis.StringSliceCmd, len(executions)) - _, err := db.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error { - for i := range executions { - commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1) - } - return nil - }) - if err != nil { - return fmt.Errorf("get task execution logs from redis: %w", err) - } - - for i := range executions { - logLines := commands[i].Val() - if len(logLines) > 0 { - executions[i].Log = strings.Join(logLines, "") - } - } - return nil -} diff --git a/internal/model/users.go b/internal/model/users.go index a8495910..2a717f06 100644 --- a/internal/model/users.go +++ b/internal/model/users.go @@ -5,16 +5,13 @@ package model import ( - "context" "errors" "strconv" "strings" "time" - "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" "github.com/Rain-kl/Wavelet/internal/shared" "github.com/Rain-kl/Wavelet/pkg/util" - "gorm.io/gorm" ) // OAuthUserInfo 用户信息结构(同时支持 OIDC ID Token claims 和 UserEndpoint 响应) @@ -104,14 +101,6 @@ func (u *User) CheckPassword(password string) bool { return u.Password == password } -// GetByID 根据 ID 查询用户 -func (u *User) GetByID(tx *gorm.DB, id uint64) error { - if err := tx.Where("id = ?", id).First(u).Error; err != nil { - return err - } - return nil -} - // UpdateFromOAuthInfo 根据 OAuth 信息更新用户数据 func (u *User) UpdateFromOAuthInfo(oauthInfo *OAuthUserInfo) { u.Username = oauthInfo.Username @@ -129,67 +118,3 @@ func (u *User) CheckActive() error { } return nil } - -func (u *User) assignIDIfMissing() error { - if u.ID != 0 { - return nil - } - u.ID = idgen.NextUint64ID() - return nil -} - -// CreateUser 创建新用户(用于 OAuth/OIDC 自动注册,含底层权限校验) -func (u *User) CreateUser(_ context.Context, tx *gorm.DB, oauthInfo *OAuthUserInfo) error { - now := time.Now() - userID := oauthInfo.GetID() - newUser := User{ - ID: userID, - Username: oauthInfo.Username, - Nickname: oauthInfo.Name, - Email: oauthInfo.Email, - AvatarURL: oauthInfo.AvatarURL, - IsActive: oauthInfo.Active, - LastLoginAt: now, - IsAdmin: false, - } - if err := newUser.assignIDIfMissing(); err != nil { - return err - } - if err := tx.Create(&newUser).Error; err != nil { - return err - } - - *u = newUser - return nil -} - -// RegisterUser 创建新用户并注册(用于本地密码注册,含全局开关和唯一性多重底层校验) -func (u *User) RegisterUser(_ context.Context, tx *gorm.DB) error { - // 检查用户名冲突 - var count int64 - if err := tx.Model(&User{}).Where("username = ?", u.Username).Count(&count).Error; err != nil { - return err - } - if count > 0 { - return errors.New(errUsernameExists) - } - - // 检查邮箱冲突 - if u.Email != "" { - var emailCount int64 - if err := tx.Model(&User{}).Where("email = ?", u.Email).Count(&emailCount).Error; err != nil { - return err - } - if emailCount > 0 { - return errors.New(errEmailAlreadyBound) - } - } - - if err := u.assignIDIfMissing(); err != nil { - return err - } - if err := tx.Create(u).Error; err != nil { - return err - } - return nil -} diff --git a/internal/repository/access_token.go b/internal/repository/access_token.go new file mode 100644 index 00000000..ac5527dd --- /dev/null +++ b/internal/repository/access_token.go @@ -0,0 +1,69 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// ListAccessTokensByUserID returns all access tokens for a user ordered by created_at desc. +func ListAccessTokensByUserID(ctx context.Context, userID uint64) ([]model.AccessToken, error) { + var tokens []model.AccessToken + if err := db.DB(ctx).Where("user_id = ?", userID).Order("created_at desc").Find(&tokens).Error; err != nil { + return nil, err + } + return tokens, nil +} + +// CountAccessTokensByUserID returns how many access tokens a user owns. +func CountAccessTokensByUserID(ctx context.Context, userID uint64) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.AccessToken{}).Where("user_id = ?", userID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// CreateAccessToken inserts a new access token record. +func CreateAccessToken(ctx context.Context, record *model.AccessToken) error { + return db.DB(ctx).Create(record).Error +} + +// GetAccessTokenByIDAndUserID loads a token owned by the given user. +func GetAccessTokenByIDAndUserID(ctx context.Context, id, userID uint64) (model.AccessToken, error) { + var token model.AccessToken + if err := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).First(&token).Error; err != nil { + return model.AccessToken{}, err + } + return token, nil +} + +// DeleteAccessTokenForUser deletes a token if it belongs to the user. +// Returns the number of rows affected. +func DeleteAccessTokenForUser(ctx context.Context, id, userID uint64) (int64, error) { + tx := db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.AccessToken{}) + return tx.RowsAffected, tx.Error +} + +// GetAccessTokenByHash loads an access token by its token hash. +func GetAccessTokenByHash(ctx context.Context, tokenHash string) (model.AccessToken, error) { + var token model.AccessToken + if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&token).Error; err != nil { + return model.AccessToken{}, err + } + return token, nil +} + +// SaveAccessToken persists all fields of an existing access token. +func SaveAccessToken(ctx context.Context, record *model.AccessToken) error { + return db.DB(ctx).Save(record).Error +} + +// DeleteAccessTokensByUserID deletes all access tokens for a user. +func DeleteAccessTokensByUserID(ctx context.Context, userID uint64) error { + return db.DB(ctx).Where("user_id = ?", userID).Delete(&model.AccessToken{}).Error +} diff --git a/internal/repository/auth_source.go b/internal/repository/auth_source.go new file mode 100644 index 00000000..f45cb146 --- /dev/null +++ b/internal/repository/auth_source.go @@ -0,0 +1,215 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + "strings" + + "gorm.io/gorm" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// GetAuthSources 获取所有认证源(已脱敏) +func GetAuthSources(ctx context.Context) ([]model.AuthSource, error) { + var sources []model.AuthSource + if err := db.DB(ctx).Order("id asc").Find(&sources).Error; err != nil { + return nil, err + } + for i := range sources { + sources[i].Sanitize() + } + return sources, nil +} + +// GetActiveAuthSources 获取所有已启用的认证源(已脱敏) +func GetActiveAuthSources(ctx context.Context) ([]model.AuthSource, error) { + var sources []model.AuthSource + if err := db.DB(ctx).Where("is_active = ?", true).Order("id asc").Find(&sources).Error; err != nil { + return nil, err + } + for i := range sources { + sources[i].Sanitize() + } + return sources, nil +} + +// GetAuthSourceByID 根据 ID 获取认证源 +func GetAuthSourceByID(ctx context.Context, id uint64) (*model.AuthSource, error) { + if id == 0 { + return nil, errors.New(errAuthSourceIDRequired) + } + var source model.AuthSource + if err := db.DB(ctx).First(&source, "id = ?", id).Error; err != nil { + return nil, err + } + source.ClientSecretConfigured = source.ClientSecret != "" + return &source, nil +} + +// GetAuthSourceByName 根据名称获取认证源(名称比较不区分大小写) +func GetAuthSourceByName(ctx context.Context, name string) (*model.AuthSource, error) { + name = strings.TrimSpace(name) + if name == "" { + return nil, errors.New(errAuthSourceNameRequired) + } + var source model.AuthSource + if err := db.DB(ctx).First(&source, "LOWER(name) = LOWER(?)", name).Error; err != nil { + return nil, err + } + source.ClientSecretConfigured = source.ClientSecret != "" + return &source, nil +} + +// CreateAuthSource 创建认证源 +func CreateAuthSource(ctx context.Context, source *model.AuthSource) error { + if err := source.Validate(); err != nil { + return err + } + return db.DB(ctx).Create(source).Error +} + +// UpdateAuthSource 更新认证源,keepSecret 为 true 时保留原密钥 +func UpdateAuthSource(ctx context.Context, source *model.AuthSource, keepSecret bool) error { + if source.ID == 0 { + return errors.New(errAuthSourceIDRequired) + } + var current model.AuthSource + if err := db.DB(ctx).First(¤t, "id = ?", source.ID).Error; err != nil { + return err + } + if keepSecret { + source.ClientSecret = current.ClientSecret + } + if err := source.Validate(); err != nil { + return err + } + return db.DB(ctx).Model(¤t).Updates(map[string]any{ + colName: source.Name, + "type": source.Type, + "display_name": source.DisplayName, + "is_active": source.IsActive, + "client_id": source.ClientID, + "client_secret": source.ClientSecret, + "openid_discovery_url": source.OpenIDDiscoveryURL, + "scopes": source.Scopes, + "icon_url": source.IconURL, + }).Error +} + +// ToggleAuthSource 切换认证源启用状态 +func ToggleAuthSource(ctx context.Context, id uint64, isActive bool) error { + source, err := GetAuthSourceByID(ctx, id) + if err != nil { + return err + } + source.IsActive = isActive + if err := source.Validate(); err != nil { + return err + } + return db.DB(ctx).Model(&model.AuthSource{}).Where("id = ?", id).Update("is_active", isActive).Error +} + +// DeleteAuthSource 删除认证源及其关联的外部帐号绑定 +func DeleteAuthSource(ctx context.Context, id uint64) error { + if id == 0 { + return errors.New(errAuthSourceIDRequired) + } + return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx.Where("auth_source_id = ?", id).Delete(&model.ExternalAccount{}).Error; err != nil { + return err + } + return tx.Delete(&model.AuthSource{}, "id = ?", id).Error + }) +} + +// FindExternalAccount 查找外部帐号绑定记录 +func FindExternalAccount(ctx context.Context, sourceID uint64, externalID string) (*model.ExternalAccount, error) { + var account model.ExternalAccount + if err := db.DB(ctx).Where("auth_source_id = ? AND external_id = ?", sourceID, externalID).First(&account).Error; err != nil { + return nil, err + } + return &account, nil +} + +// BindExternalAccount 绑定外部帐号(已存在时更新用户名和邮箱) +func BindExternalAccount(ctx context.Context, account *model.ExternalAccount) error { + if account.UserID == 0 || strings.TrimSpace(account.ExternalID) == "" { + return errors.New(errExternalAccountBindingIncomplete) + } + account.ExternalID = strings.TrimSpace(account.ExternalID) + account.ExternalUsername = strings.TrimSpace(account.ExternalUsername) + account.Email = strings.TrimSpace(account.Email) + + return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + var current model.ExternalAccount + err := tx.Where("auth_source_id = ? AND external_id = ?", account.AuthSourceID, account.ExternalID).First(¤t).Error + if err == nil { + if current.UserID != account.UserID { + return errors.New(errExternalAccountAlreadyBoundToAnother) + } + return tx.Model(¤t).Updates(map[string]any{ + "external_username": account.ExternalUsername, + "email": account.Email, + }).Error + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + return tx.Create(account).Error + }) +} + +// ListExternalAccountsByUserID 获取指定用户的所有外部帐号绑定视图 +func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]model.ExternalAccountView, error) { + if userID == 0 { + return nil, errors.New(errUserIDRequired) + } + var accounts []model.ExternalAccount + if err := db.DB(ctx).Where("user_id = ?", userID).Order("id asc").Find(&accounts).Error; err != nil { + return nil, err + } + views := make([]model.ExternalAccountView, 0, len(accounts)) + for _, account := range accounts { + var name, sourceType, label string + if account.AuthSourceID == 0 { + name = "default" + sourceType = "oidc" + label = "历史认证源" + } else { + source, err := GetAuthSourceByID(ctx, account.AuthSourceID) + if err != nil { + continue + } + name = source.Name + sourceType = source.Type + label = source.DisplayName + if label == "" { + label = source.Name + } + } + views = append(views, model.ExternalAccountView{ + ID: account.ID, + AuthSourceID: account.AuthSourceID, + AuthSourceName: name, + AuthSourceType: sourceType, + AuthSourceLabel: label, + ExternalUsername: account.ExternalUsername, + Email: account.Email, + CreatedAt: account.CreatedAt, + }) + } + return views, nil +} + +// DeleteExternalAccountForUser 删除指定用户的外部帐号绑定 +func DeleteExternalAccountForUser(ctx context.Context, id uint64, userID uint64) error { + if id == 0 || userID == 0 { + return errors.New(errExternalAccountBindingIDRequired) + } + return db.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&model.ExternalAccount{}).Error +} diff --git a/internal/repository/auth_source_cache.go b/internal/repository/auth_source_cache.go index b78ca614..2dd4aa7d 100644 --- a/internal/repository/auth_source_cache.go +++ b/internal/repository/auth_source_cache.go @@ -178,7 +178,7 @@ func GetActiveAuthSourcesCached(ctx context.Context) ([]model.AuthSource, error) } } - sources, err := model.GetActiveAuthSources(ctx) + sources, err := GetActiveAuthSources(ctx) if err != nil { return nil, err } @@ -192,7 +192,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*model.AuthSou normalized := normalizeAuthSourceName(name) if normalized == "" { - return model.GetAuthSourceByName(ctx, name) + return GetAuthSourceByName(ctx, name) } if source, ok := authSourceByNameRAM.GetIfPresent(normalized); ok { @@ -210,7 +210,7 @@ func GetAuthSourceByNameCached(ctx context.Context, name string) (*model.AuthSou } } - source, err := model.GetAuthSourceByName(ctx, name) + source, err := GetAuthSourceByName(ctx, name) if err != nil { return nil, err } diff --git a/internal/repository/auth_source_cache_test.go b/internal/repository/auth_source_cache_test.go index 009ea3bc..dcf92afd 100644 --- a/internal/repository/auth_source_cache_test.go +++ b/internal/repository/auth_source_cache_test.go @@ -74,7 +74,7 @@ func TestGetActiveAuthSourcesCached_LoadsFromRedisBeforeDB(t *testing.T) { ClientSecret: "client-secret", OpenIDDiscoveryURL: "https://issuer.example.com", } - if err := model.CreateAuthSource(ctx, &source); err != nil { + if err := CreateAuthSource(ctx, &source); err != nil { t.Fatalf("CreateAuthSource() error = %v", err) } @@ -119,7 +119,7 @@ func TestGetAuthSourceByNameCached_LoadsFromRedisBeforeDB(t *testing.T) { ClientSecret: "client-secret", OpenIDDiscoveryURL: "https://issuer.example.com", } - if err := model.CreateAuthSource(ctx, &source); err != nil { + if err := CreateAuthSource(ctx, &source); err != nil { t.Fatalf("CreateAuthSource() error = %v", err) } @@ -164,7 +164,7 @@ func TestInvalidateAuthSourceCache_ClearsRedisKeys(t *testing.T) { ClientSecret: "client-secret", OpenIDDiscoveryURL: "https://issuer.example.com", } - if err := model.CreateAuthSource(ctx, &source); err != nil { + if err := CreateAuthSource(ctx, &source); err != nil { t.Fatalf("CreateAuthSource() error = %v", err) } if _, err := GetActiveAuthSourcesCached(ctx); err != nil { @@ -209,7 +209,7 @@ func TestAuthSourceInvalidationPubSubClearsPeerRAM(t *testing.T) { ClientSecret: "client-secret", OpenIDDiscoveryURL: "https://issuer.example.com", } - if err := model.CreateAuthSource(ctx, &source); err != nil { + if err := CreateAuthSource(ctx, &source); err != nil { t.Fatalf("CreateAuthSource() error = %v", err) } diff --git a/internal/repository/errs.go b/internal/repository/errs.go new file mode 100644 index 00000000..92f5e2b4 --- /dev/null +++ b/internal/repository/errs.go @@ -0,0 +1,27 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +// Persistence and repository-layer parameter messages live here (unexported). +// Domain field validation used by model.Validate stays in internal/model/errs.go; +// repository may call model.Validate and return those errors as-is. +// Keep wording aligned with model where the same user-facing phrase applies, +// but do not import or re-export model unexported consts (would require exporting). +const ( + errDatabaseNotInitialized = "database not initialized" + errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w" + errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w" + errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w" + errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w" + errAuthSourceNameRequired = "认证源名称不能为空" + errAuthSourceIDRequired = "认证源 ID 不能为空" + errExternalAccountBindingIncomplete = "外部账号绑定信息不完整" + errExternalAccountAlreadyBoundToAnother = "该外部账号已绑定到其他用户" + errUserIDRequired = "用户 ID 不能为空" + errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空" +) + +const colName = "name" + +const colEnabled = "enabled" diff --git a/internal/repository/openflare_access_log.go b/internal/repository/openflare_access_log.go new file mode 100644 index 00000000..e74062c5 --- /dev/null +++ b/internal/repository/openflare_access_log.go @@ -0,0 +1,476 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "math" + "sort" + "strings" + "time" + + analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics" + + "github.com/Rain-kl/Wavelet/internal/model" +) + +type openFlareAccessLogBucketAggregateRow = analyticsmodel.NodeAccessLogBucketAggregate +type openFlareAccessLogBucketDimensionRow = analyticsmodel.NodeAccessLogBucketDimension +type openFlareAccessLogIPAggregateRow = analyticsmodel.NodeAccessLogIPAggregate +type openFlareAccessLogIPSummaryRow = analyticsmodel.NodeAccessLogIPSummary +type openFlareAccessLogIPTrendRow = analyticsmodel.NodeAccessLogIPTrend +type openFlareAccessLogWAFIPAggregateRow = analyticsmodel.NodeAccessLogWAFIPAggregate + +const ( + sortOrderAsc = "asc" + columnRemoteAddr = "remote_addr" + columnHost = "host" + secondsPerMinute = 60 +) + +// ListOpenFlareAccessLogWAFIPAggregates returns per-IP aggregates for WAF automatic rules. +func ListOpenFlareAccessLogWAFIPAggregates(ctx context.Context, query model.OpenFlareAccessLogQuery) ([]*model.OpenFlareAccessLogWAFIPAggregate, error) { + rows, err := currentAccessLogStore().WAFIPAggregates(ctx, query) + if err != nil { + return nil, err + } + result := make([]*model.OpenFlareAccessLogWAFIPAggregate, 0, len(rows)) + for _, row := range rows { + remoteAddr := strings.TrimSpace(row.RemoteAddr) + if remoteAddr == "" { + continue + } + statusCounts := make(map[int]int, len(row.StatusCounts)) + for code, count := range row.StatusCounts { + statusCounts[code] = int(count) + } + result = append(result, &model.OpenFlareAccessLogWAFIPAggregate{ + RemoteAddr: remoteAddr, + RequestCount: int(row.RequestCount), + Status404Count: int(row.Status404Count), + ClientErrorCount: int(row.ClientErrorCount), + ServerErrorCount: int(row.ServerErrorCount), + IPHostCount: int(row.IPHostCount), + LastSeenEpoch: row.LastSeenEpoch, + StatusCounts: statusCounts, + }) + } + return result, nil +} + +// InsertOpenFlareAccessLogsBatch inserts access log rows into ClickHouse. +func InsertOpenFlareAccessLogsBatch(ctx context.Context, records []*model.OpenFlareAccessLog) error { + return currentAccessLogStore().InsertBatch(ctx, records) +} + +// ListOpenFlareAccessLogs lists access logs matching the query. +func ListOpenFlareAccessLogs(ctx context.Context, query model.OpenFlareAccessLogQuery) ([]*model.OpenFlareAccessLog, error) { + return currentAccessLogStore().List(ctx, query) +} + +// CountOpenFlareAccessLogs counts access logs, distinct IPs, and total bytes sent matching the query. +func CountOpenFlareAccessLogs(ctx context.Context, query model.OpenFlareAccessLogQuery) (int64, int64, int64, error) { + return currentAccessLogStore().Count(ctx, query) +} + +// TrafficSummaryOpenFlareAccessLogs returns window-level request/error/UV/bytes summary. +func TrafficSummaryOpenFlareAccessLogs(ctx context.Context, query model.OpenFlareAccessLogQuery) (model.OpenFlareAccessLogTrafficSummary, error) { + return currentAccessLogStore().TrafficSummary(ctx, query) +} + +// ValueCountsOpenFlareAccessLogs groups logs by status_code, host, path, remote_addr, or user_agent. +func ValueCountsOpenFlareAccessLogs(ctx context.Context, query model.OpenFlareAccessLogQuery, column string, limit int) ([]model.OpenFlareAccessLogValueCount, error) { + return currentAccessLogStore().ValueCounts(ctx, query, column, limit) +} + +// NodeAggregatesOpenFlareAccessLogs returns per-node request/error/UV for the window. +func NodeAggregatesOpenFlareAccessLogs(ctx context.Context, query model.OpenFlareAccessLogQuery) ([]model.OpenFlareAccessLogNodeAggregate, error) { + return currentAccessLogStore().NodeAggregates(ctx, query) +} + +// ListOpenFlareAccessLogRegionCounts returns region counts for access logs. +func ListOpenFlareAccessLogRegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareAccessLogRegionCount, error) { + return currentAccessLogStore().RegionCounts(ctx, nodeID, since, limit) +} + +// ListOpenFlareAccessLogBuckets lists folded access log buckets. +func ListOpenFlareAccessLogBuckets(ctx context.Context, query model.OpenFlareAccessLogBucketQuery) ([]*model.OpenFlareAccessLogBucketRow, error) { + return buildOpenFlareAccessLogBucketRows(ctx, query) +} + +// CountOpenFlareAccessLogBuckets counts folded access log buckets. +func CountOpenFlareAccessLogBuckets(ctx context.Context, query model.OpenFlareAccessLogBucketQuery) (int64, error) { + filter := openFlareAccessLogQueryFromBucket(query) + bucketSeconds := int64(query.FoldMinutes * secondsPerMinute) + if bucketSeconds <= 0 { + bucketSeconds = 180 + } + return currentAccessLogStore().CountBuckets(ctx, filter, bucketSeconds) +} + +// ListOpenFlareAccessLogBucketIPs lists folded IP rows for a bucket window. +func ListOpenFlareAccessLogBucketIPs(ctx context.Context, query model.OpenFlareAccessLogBucketIPQuery) ([]*model.OpenFlareAccessLogBucketIPRow, error) { + rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query) + if err != nil { + return nil, err + } + start, end := openFlareAccessLogPaginateBounds(len(rows), query.Page, query.PageSize) + if start >= len(rows) { + return []*model.OpenFlareAccessLogBucketIPRow{}, nil + } + return rows[start:end], nil +} + +// CountOpenFlareAccessLogBucketIPs counts folded IP rows for a bucket window. +func CountOpenFlareAccessLogBucketIPs(ctx context.Context, query model.OpenFlareAccessLogBucketIPQuery) (int64, error) { + rows, err := buildOpenFlareAccessLogBucketIPRows(ctx, query) + if err != nil { + return 0, err + } + return int64(len(rows)), nil +} + +// ListOpenFlareAccessLogIPSummaries lists IP summaries. +func ListOpenFlareAccessLogIPSummaries(ctx context.Context, query model.OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*analyticsmodel.NodeAccessLogIPSummary, error) { + return buildOpenFlareAccessLogIPSummaryRows(ctx, query, recentSince) +} + +// CountOpenFlareAccessLogIPSummaries counts IP summaries. +func CountOpenFlareAccessLogIPSummaries(ctx context.Context, query model.OpenFlareAccessLogIPSummaryQuery) (int64, error) { + filter := openFlareAccessLogQueryFromIPSummary(query) + return currentAccessLogStore().CountIPSummaries(ctx, filter) +} + +// ListOpenFlareAccessLogIPTrend lists IP trend points. +func ListOpenFlareAccessLogIPTrend(ctx context.Context, query model.OpenFlareAccessLogIPTrendQuery) ([]*analyticsmodel.NodeAccessLogIPTrend, error) { + remoteAddr := strings.TrimSpace(query.RemoteAddr) + if remoteAddr == "" { + return []*analyticsmodel.NodeAccessLogIPTrend{}, nil + } + filter := model.OpenFlareAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: remoteAddr, + Host: query.Host, + Since: query.Since, + } + bucketSeconds := int64(query.BucketMinutes * secondsPerMinute) + if bucketSeconds <= 0 { + bucketSeconds = 1800 + } + rows, err := currentAccessLogStore().IPTrend(ctx, filter, bucketSeconds) + if err != nil { + return nil, err + } + result := make([]*analyticsmodel.NodeAccessLogIPTrend, len(rows)) + for index, row := range rows { + result[index] = &analyticsmodel.NodeAccessLogIPTrend{ + BucketEpoch: row.BucketEpoch, + RequestCount: row.RequestCount, + } + } + return result, nil +} + +// DeleteAllOpenFlareAccessLogs deletes all access logs. +func DeleteAllOpenFlareAccessLogs(ctx context.Context) (int64, error) { + return currentAccessLogStore().DeleteAll(ctx) +} + +// DeleteOpenFlareAccessLogsBefore deletes access logs older than cutoff. +func DeleteOpenFlareAccessLogsBefore(ctx context.Context, cutoff time.Time) (int64, error) { + return currentAccessLogStore().DeleteBefore(ctx, cutoff) +} + +// DeleteOpenFlareAccessLogsByNodeBefore deletes access logs for a node older than cutoff. +func DeleteOpenFlareAccessLogsByNodeBefore(ctx context.Context, nodeID string, cutoff time.Time) (int64, error) { + return currentAccessLogStore().DeleteByNodeBefore(ctx, nodeID, cutoff) +} + +func buildOpenFlareAccessLogBucketRows(ctx context.Context, query model.OpenFlareAccessLogBucketQuery) ([]*model.OpenFlareAccessLogBucketRow, error) { + filter := openFlareAccessLogQueryFromBucket(query) + bucketSeconds := int64(query.FoldMinutes * secondsPerMinute) + if bucketSeconds <= 0 { + bucketSeconds = 180 + } + + partials, err := currentAccessLogStore().BucketAggregates(ctx, filter, bucketSeconds) + if err != nil { + return nil, err + } + rows := make([]*model.OpenFlareAccessLogBucketRow, 0, len(partials)) + for _, partial := range partials { + rows = append(rows, &model.OpenFlareAccessLogBucketRow{ + BucketEpoch: partial.BucketEpoch, + RequestCount: partial.RequestCount, + UniqueIPCount: partial.UniqueIPCount, + UniqueHostCount: partial.UniqueHostCount, + SuccessCount: partial.SuccessCount, + ClientErrorCount: partial.ClientErrorCount, + ServerErrorCount: partial.ServerErrorCount, + BytesSent: partial.BytesSent, + RequestLength: partial.RequestLength, + }) + } + return rows, nil +} + +func buildOpenFlareAccessLogBucketIPRows(ctx context.Context, query model.OpenFlareAccessLogBucketIPQuery) ([]*model.OpenFlareAccessLogBucketIPRow, error) { + if query.BucketStartedAt.IsZero() { + return []*model.OpenFlareAccessLogBucketIPRow{}, nil + } + foldMinutes := query.FoldMinutes + if foldMinutes <= 0 { + foldMinutes = 3 + } + bucketStartedAt := query.BucketStartedAt.UTC() + filter := model.OpenFlareAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: query.RemoteAddr, + Host: query.Host, + Path: query.Path, + Since: bucketStartedAt, + Until: bucketStartedAt.Add(time.Duration(foldMinutes) * time.Minute), + } + rows, err := queryOpenFlareAccessLogIPAggregateRows(ctx, filter, false) + if err != nil { + return nil, err + } + sortOpenFlareAccessLogBucketIPRows(rows, query.SortBy, query.SortOrder) + return rows, nil +} + +func buildOpenFlareAccessLogIPSummaryRows(ctx context.Context, query model.OpenFlareAccessLogIPSummaryQuery, recentSince time.Time) ([]*analyticsmodel.NodeAccessLogIPSummary, error) { + filter := openFlareAccessLogQueryFromIPSummary(query) + partials, err := currentAccessLogStore().IPSummaries(ctx, filter, recentSince) + if err != nil { + return nil, err + } + rows := make([]*analyticsmodel.NodeAccessLogIPSummary, 0, len(partials)) + for _, partial := range partials { + remoteAddr := strings.TrimSpace(partial.RemoteAddr) + if remoteAddr == "" { + continue + } + rows = append(rows, &analyticsmodel.NodeAccessLogIPSummary{ + RemoteAddr: remoteAddr, + Region: strings.TrimSpace(partial.Region), + TotalRequests: partial.TotalRequests, + Success2xxCount: partial.Success2xxCount, + SuccessRatio: partial.SuccessRatio, + BytesReceived: partial.BytesReceived, + BytesSent: partial.BytesSent, + RecentRequests: 0, + LastSeenEpoch: partial.LastSeenEpoch, + }) + } + return rows, nil +} + +func queryOpenFlareAccessLogIPAggregateRows(ctx context.Context, filter model.OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]*model.OpenFlareAccessLogBucketIPRow, error) { + partials, err := currentAccessLogStore().IPAggregates(ctx, filter, exactRemoteAddr) + if err != nil { + return nil, err + } + rows := make([]*model.OpenFlareAccessLogBucketIPRow, 0, len(partials)) + for _, partial := range partials { + remoteAddr := strings.TrimSpace(partial.RemoteAddr) + if remoteAddr == "" { + continue + } + rows = append(rows, &model.OpenFlareAccessLogBucketIPRow{ + RemoteAddr: remoteAddr, + RequestCount: partial.RequestCount, + SuccessCount: partial.SuccessCount, + ClientErrorCount: partial.ClientErrorCount, + ServerErrorCount: partial.ServerErrorCount, + LastSeenEpoch: partial.LastSeenEpoch, + }) + } + return rows, nil +} + +func openFlareAccessLogQueryFromBucket(query model.OpenFlareAccessLogBucketQuery) model.OpenFlareAccessLogQuery { + return model.OpenFlareAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: query.RemoteAddr, + Host: query.Host, + Hosts: query.Hosts, + Path: query.Path, + Since: query.Since, + Until: query.Until, + Page: query.Page, + PageSize: query.PageSize, + SortBy: query.SortBy, + SortOrder: query.SortOrder, + } +} + +func openFlareAccessLogQueryFromIPSummary(query model.OpenFlareAccessLogIPSummaryQuery) model.OpenFlareAccessLogQuery { + return model.OpenFlareAccessLogQuery{ + NodeID: query.NodeID, + RemoteAddr: query.RemoteAddr, + Host: query.Host, + Since: query.Since, + Until: query.Until, + Page: query.Page, + PageSize: query.PageSize, + SortBy: query.SortBy, + SortOrder: query.SortOrder, + } +} + +func sortOpenFlareAccessLogBucketIPRows(items []*model.OpenFlareAccessLogBucketIPRow, sortBy string, sortOrder string) { + desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc + sort.Slice(items, func(i, j int) bool { + left := items[i] + right := items[j] + if left == nil || right == nil { + return left != nil + } + var compare int + switch strings.TrimSpace(sortBy) { + case "last_seen_at": + compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) + case "remote_addr": + compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) + default: + compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount) + } + if compare == 0 { + compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) + } + if compare == 0 { + compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) + } + if desc { + return compare > 0 + } + return compare < 0 + }) +} + +func openFlareAccessLogPaginateBounds(total int, page int, pageSize int) (int, int) { + if page < 0 { + page = 0 + } + if pageSize <= 0 { + return 0, total + } + start := page * pageSize + if start > total { + start = total + } + end := start + pageSize + if end > total { + end = total + } + return start, end +} + +func openFlareAccessLogNormalizeSortOrder(sortOrder string) string { + if strings.EqualFold(strings.TrimSpace(sortOrder), sortOrderAsc) { + return sortOrderAsc + } + return "desc" +} + +func openFlareAccessLogCompareInt64(left int64, right int64) int { + switch { + case left > right: + return 1 + case left < right: + return -1 + default: + return 0 + } +} + +func openFlareAccessLogStatusCodeToInt32(code int) int32 { + switch { + case code > math.MaxInt32: + return math.MaxInt32 + case code < math.MinInt32: + return math.MinInt32 + default: + return int32(code) + } +} + +func sortOpenFlareAccessLogBucketRows(items []*model.OpenFlareAccessLogBucketRow, sortBy string, sortOrder string) { + desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc + sort.Slice(items, func(i, j int) bool { + left := items[i] + right := items[j] + if left == nil || right == nil { + return left != nil + } + var compare int + switch strings.TrimSpace(sortBy) { + case "request_count": + compare = openFlareAccessLogCompareInt64(left.RequestCount, right.RequestCount) + default: + compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch) + } + if compare == 0 { + compare = openFlareAccessLogCompareInt64(left.BucketEpoch, right.BucketEpoch) + } + if desc { + return compare > 0 + } + return compare < 0 + }) +} + +func sortOpenFlareAccessLogIPSummaryRows(items []*model.OpenFlareAccessLogIPSummaryRow, sortBy string, sortOrder string) { + desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc + sort.Slice(items, func(i, j int) bool { + left := items[i] + right := items[j] + if left == nil || right == nil { + return left != nil + } + var compare int + switch strings.TrimSpace(sortBy) { + case "request_length", "bytes_received": + compare = openFlareAccessLogCompareInt64(left.BytesReceived, right.BytesReceived) + case "bytes_sent": + compare = openFlareAccessLogCompareInt64(left.BytesSent, right.BytesSent) + case "success_ratio": + compare = openFlareAccessLogCompareFloat64(left.SuccessRatio, right.SuccessRatio) + case "last_seen_at": + compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) + case "remote_addr": + compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) + default: + compare = openFlareAccessLogCompareInt64(left.TotalRequests, right.TotalRequests) + } + if compare == 0 { + compare = openFlareAccessLogCompareInt64(left.LastSeenEpoch, right.LastSeenEpoch) + } + if compare == 0 { + compare = strings.Compare(left.RemoteAddr, right.RemoteAddr) + } + if desc { + return compare > 0 + } + return compare < 0 + }) +} + +func openFlareAccessLogCompareFloat64(left, right float64) int { + if left < right { + return -1 + } + if left > right { + return 1 + } + return 0 +} + +func openFlareAccessLogUintToInt64(value uint64) int64 { + if value > math.MaxInt64 { + return math.MaxInt64 + } + return int64(value) +} diff --git a/internal/model/openflare_access_log_store.go b/internal/repository/openflare_access_log_store.go similarity index 65% rename from internal/model/openflare_access_log_store.go rename to internal/repository/openflare_access_log_store.go index 7591771e..bca8b399 100644 --- a/internal/model/openflare_access_log_store.go +++ b/internal/repository/openflare_access_log_store.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" @@ -9,6 +9,8 @@ import ( "sync" "time" + "github.com/Rain-kl/Wavelet/internal/model" + analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics" analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" ) @@ -38,50 +40,26 @@ func currentAccessLogInsertHooks() AccessLogInsertHooks { } type accessLogStore interface { - InsertBatch(ctx context.Context, records []*OpenFlareAccessLog) error - List(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) - Count(ctx context.Context, query OpenFlareAccessLogQuery) (int64, int64, int64, error) - RegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error) - BucketAggregates(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error) - CountBuckets(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error) - BucketDimensions(ctx context.Context, filter OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error) - IPAggregates(ctx context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error) - WAFIPAggregates(ctx context.Context, filter OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error) - IPSummaries(ctx context.Context, filter OpenFlareAccessLogQuery, recentSince time.Time) ([]openFlareAccessLogIPSummaryRow, error) - CountIPSummaries(ctx context.Context, filter OpenFlareAccessLogQuery) (int64, error) - IPTrend(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error) - TrafficSummary(ctx context.Context, filter OpenFlareAccessLogQuery) (OpenFlareAccessLogTrafficSummary, error) - ValueCounts(ctx context.Context, filter OpenFlareAccessLogQuery, column string, limit int) ([]OpenFlareAccessLogValueCount, error) - NodeAggregates(ctx context.Context, filter OpenFlareAccessLogQuery) ([]OpenFlareAccessLogNodeAggregate, error) + InsertBatch(ctx context.Context, records []*model.OpenFlareAccessLog) error + List(ctx context.Context, query model.OpenFlareAccessLogQuery) ([]*model.OpenFlareAccessLog, error) + Count(ctx context.Context, query model.OpenFlareAccessLogQuery) (int64, int64, int64, error) + RegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareAccessLogRegionCount, error) + BucketAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error) + CountBuckets(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error) + BucketDimensions(ctx context.Context, filter model.OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error) + IPAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error) + WAFIPAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error) + IPSummaries(ctx context.Context, filter model.OpenFlareAccessLogQuery, recentSince time.Time) ([]openFlareAccessLogIPSummaryRow, error) + CountIPSummaries(ctx context.Context, filter model.OpenFlareAccessLogQuery) (int64, error) + IPTrend(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error) + TrafficSummary(ctx context.Context, filter model.OpenFlareAccessLogQuery) (model.OpenFlareAccessLogTrafficSummary, error) + ValueCounts(ctx context.Context, filter model.OpenFlareAccessLogQuery, column string, limit int) ([]model.OpenFlareAccessLogValueCount, error) + NodeAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery) ([]model.OpenFlareAccessLogNodeAggregate, error) DeleteAll(ctx context.Context) (int64, error) DeleteBefore(ctx context.Context, cutoff time.Time) (int64, error) DeleteByNodeBefore(ctx context.Context, nodeID string, before time.Time) (int64, error) } -// OpenFlareAccessLogTrafficSummary is a window-level traffic summary from access logs. -type OpenFlareAccessLogTrafficSummary struct { - RequestCount int64 - ErrorCount int64 - UniqueIPCount int64 - BytesSent int64 - RequestLength int64 - NodeCount int64 -} - -// OpenFlareAccessLogValueCount is a dimension value count. -type OpenFlareAccessLogValueCount struct { - Value string - Count int64 -} - -// OpenFlareAccessLogNodeAggregate is per-node traffic over a window. -type OpenFlareAccessLogNodeAggregate struct { - NodeID string - RequestCount int64 - ErrorCount int64 - UniqueIPCount int64 -} - var ( accessLogStoreMu sync.RWMutex accessLogStoreHolder accessLogStore @@ -112,13 +90,13 @@ func SetAccessLogStoreForTest(store accessLogStore) func() { // NewMemoryAccessLogStore returns an in-memory access log store for unit tests. func NewMemoryAccessLogStore() accessLogStore { return &memoryAccessLogStore{ - records: make([]*OpenFlareAccessLog, 0), + records: make([]*model.OpenFlareAccessLog, 0), } } type clickhouseAccessLogStore struct{} -func (clickhouseAccessLogStore) InsertBatch(_ context.Context, records []*OpenFlareAccessLog) error { +func (clickhouseAccessLogStore) InsertBatch(_ context.Context, records []*model.OpenFlareAccessLog) error { logs := make([]analyticsmodel.NodeAccessLog, 0, len(records)) for _, record := range records { if record == nil { @@ -132,7 +110,7 @@ func (clickhouseAccessLogStore) InsertBatch(_ context.Context, records []*OpenFl return nil } -func (clickhouseAccessLogStore) List(ctx context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) { +func (clickhouseAccessLogStore) List(ctx context.Context, query model.OpenFlareAccessLogQuery) ([]*model.OpenFlareAccessLog, error) { rows, err := analyticsrepo.ListNodeAccessLogs(ctx, toNodeAccessLogFilter(query)) if err != nil { return nil, err @@ -140,18 +118,18 @@ func (clickhouseAccessLogStore) List(ctx context.Context, query OpenFlareAccessL return fromAnalyticsNodeAccessLogs(rows), nil } -func (clickhouseAccessLogStore) Count(ctx context.Context, query OpenFlareAccessLogQuery) (int64, int64, int64, error) { +func (clickhouseAccessLogStore) Count(ctx context.Context, query model.OpenFlareAccessLogQuery) (int64, int64, int64, error) { return analyticsrepo.CountNodeAccessLogs(ctx, toNodeAccessLogFilter(query)) } -func (clickhouseAccessLogStore) RegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error) { +func (clickhouseAccessLogStore) RegionCounts(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareAccessLogRegionCount, error) { rows, err := analyticsrepo.RegionCountsNodeAccessLogs(ctx, nodeID, since, limit) if err != nil { return nil, err } - result := make([]*OpenFlareAccessLogRegionCount, len(rows)) + result := make([]*model.OpenFlareAccessLogRegionCount, len(rows)) for index, row := range rows { - result[index] = &OpenFlareAccessLogRegionCount{ + result[index] = &model.OpenFlareAccessLogRegionCount{ Region: row.Region, Count: row.Count, } @@ -159,35 +137,35 @@ func (clickhouseAccessLogStore) RegionCounts(ctx context.Context, nodeID string, return result, nil } -func (clickhouseAccessLogStore) BucketAggregates(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error) { +func (clickhouseAccessLogStore) BucketAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error) { return analyticsrepo.BucketAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), bucketSeconds) } -func (clickhouseAccessLogStore) CountBuckets(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error) { +func (clickhouseAccessLogStore) CountBuckets(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error) { return analyticsrepo.CountBucketAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), bucketSeconds) } -func (clickhouseAccessLogStore) BucketDimensions(ctx context.Context, filter OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error) { +func (clickhouseAccessLogStore) BucketDimensions(ctx context.Context, filter model.OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error) { return analyticsrepo.BucketDimensionsNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), column, bucketSeconds) } -func (clickhouseAccessLogStore) IPAggregates(ctx context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error) { +func (clickhouseAccessLogStore) IPAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error) { return analyticsrepo.IPAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), exactRemoteAddr) } -func (clickhouseAccessLogStore) IPSummaries(ctx context.Context, filter OpenFlareAccessLogQuery, recentSince time.Time) ([]openFlareAccessLogIPSummaryRow, error) { +func (clickhouseAccessLogStore) IPSummaries(ctx context.Context, filter model.OpenFlareAccessLogQuery, recentSince time.Time) ([]openFlareAccessLogIPSummaryRow, error) { return analyticsrepo.IPSummariesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), recentSince) } -func (clickhouseAccessLogStore) CountIPSummaries(ctx context.Context, filter OpenFlareAccessLogQuery) (int64, error) { +func (clickhouseAccessLogStore) CountIPSummaries(ctx context.Context, filter model.OpenFlareAccessLogQuery) (int64, error) { return analyticsrepo.CountIPSummaryNodeAccessLogs(ctx, toNodeAccessLogFilter(filter)) } -func (clickhouseAccessLogStore) WAFIPAggregates(ctx context.Context, filter OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error) { +func (clickhouseAccessLogStore) WAFIPAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error) { return analyticsrepo.IPAggregatesForWAFNodeAccessLogs(ctx, toNodeAccessLogFilter(filter)) } -func (clickhouseAccessLogStore) IPTrend(ctx context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error) { +func (clickhouseAccessLogStore) IPTrend(ctx context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error) { return analyticsrepo.IPTrendNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), bucketSeconds) } @@ -203,12 +181,12 @@ func (clickhouseAccessLogStore) DeleteByNodeBefore(ctx context.Context, nodeID s return analyticsrepo.DeleteNodeAccessLogsByNodeBefore(ctx, nodeID, before) } -func (clickhouseAccessLogStore) TrafficSummary(ctx context.Context, filter OpenFlareAccessLogQuery) (OpenFlareAccessLogTrafficSummary, error) { +func (clickhouseAccessLogStore) TrafficSummary(ctx context.Context, filter model.OpenFlareAccessLogQuery) (model.OpenFlareAccessLogTrafficSummary, error) { row, err := analyticsrepo.TrafficSummaryNodeAccessLogs(ctx, toNodeAccessLogFilter(filter)) if err != nil { - return OpenFlareAccessLogTrafficSummary{}, err + return model.OpenFlareAccessLogTrafficSummary{}, err } - return OpenFlareAccessLogTrafficSummary{ + return model.OpenFlareAccessLogTrafficSummary{ RequestCount: row.RequestCount, ErrorCount: row.ErrorCount, UniqueIPCount: row.UniqueIPCount, @@ -218,26 +196,26 @@ func (clickhouseAccessLogStore) TrafficSummary(ctx context.Context, filter OpenF }, nil } -func (clickhouseAccessLogStore) ValueCounts(ctx context.Context, filter OpenFlareAccessLogQuery, column string, limit int) ([]OpenFlareAccessLogValueCount, error) { +func (clickhouseAccessLogStore) ValueCounts(ctx context.Context, filter model.OpenFlareAccessLogQuery, column string, limit int) ([]model.OpenFlareAccessLogValueCount, error) { rows, err := analyticsrepo.ValueCountsNodeAccessLogs(ctx, toNodeAccessLogFilter(filter), column, limit) if err != nil { return nil, err } - result := make([]OpenFlareAccessLogValueCount, len(rows)) + result := make([]model.OpenFlareAccessLogValueCount, len(rows)) for i, row := range rows { - result[i] = OpenFlareAccessLogValueCount{Value: row.Value, Count: row.Count} + result[i] = model.OpenFlareAccessLogValueCount{Value: row.Value, Count: row.Count} } return result, nil } -func (clickhouseAccessLogStore) NodeAggregates(ctx context.Context, filter OpenFlareAccessLogQuery) ([]OpenFlareAccessLogNodeAggregate, error) { +func (clickhouseAccessLogStore) NodeAggregates(ctx context.Context, filter model.OpenFlareAccessLogQuery) ([]model.OpenFlareAccessLogNodeAggregate, error) { rows, err := analyticsrepo.NodeAggregatesNodeAccessLogs(ctx, toNodeAccessLogFilter(filter)) if err != nil { return nil, err } - result := make([]OpenFlareAccessLogNodeAggregate, len(rows)) + result := make([]model.OpenFlareAccessLogNodeAggregate, len(rows)) for i, row := range rows { - result[i] = OpenFlareAccessLogNodeAggregate{ + result[i] = model.OpenFlareAccessLogNodeAggregate{ NodeID: row.NodeID, RequestCount: row.RequestCount, ErrorCount: row.ErrorCount, @@ -247,7 +225,7 @@ func (clickhouseAccessLogStore) NodeAggregates(ctx context.Context, filter OpenF return result, nil } -func toNodeAccessLogFilter(query OpenFlareAccessLogQuery) analyticsrepo.NodeAccessLogFilter { +func toNodeAccessLogFilter(query model.OpenFlareAccessLogQuery) analyticsrepo.NodeAccessLogFilter { return analyticsrepo.NodeAccessLogFilter{ NodeID: query.NodeID, RemoteAddr: query.RemoteAddr, @@ -263,7 +241,7 @@ func toNodeAccessLogFilter(query OpenFlareAccessLogQuery) analyticsrepo.NodeAcce } } -func toAnalyticsNodeAccessLog(record *OpenFlareAccessLog) analyticsmodel.NodeAccessLog { +func toAnalyticsNodeAccessLog(record *model.OpenFlareAccessLog) analyticsmodel.NodeAccessLog { var bytesSent uint64 if record.BytesSent > 0 { bytesSent = uint64(record.BytesSent) @@ -294,8 +272,8 @@ func toAnalyticsNodeAccessLog(record *OpenFlareAccessLog) analyticsmodel.NodeAcc } } -func fromAnalyticsNodeAccessLogs(rows []analyticsmodel.NodeAccessLog) []*OpenFlareAccessLog { - result := make([]*OpenFlareAccessLog, len(rows)) +func fromAnalyticsNodeAccessLogs(rows []analyticsmodel.NodeAccessLog) []*model.OpenFlareAccessLog { + result := make([]*model.OpenFlareAccessLog, len(rows)) for index, row := range rows { var bytesSent int64 if row.BytesSent <= math.MaxInt64 { @@ -309,7 +287,7 @@ func fromAnalyticsNodeAccessLogs(rows []analyticsmodel.NodeAccessLog) []*OpenFla } else { requestLength = math.MaxInt64 } - result[index] = &OpenFlareAccessLog{ + result[index] = &model.OpenFlareAccessLog{ ID: row.ID, NodeID: row.NodeID, LoggedAt: row.LoggedAt, diff --git a/internal/model/openflare_access_log_store_memory.go b/internal/repository/openflare_access_log_store_memory.go similarity index 85% rename from internal/model/openflare_access_log_store_memory.go rename to internal/repository/openflare_access_log_store_memory.go index d754f898..99cfb964 100644 --- a/internal/model/openflare_access_log_store_memory.go +++ b/internal/repository/openflare_access_log_store_memory.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" @@ -14,6 +14,8 @@ import ( "sync" "time" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" ) @@ -27,10 +29,10 @@ const ( type memoryAccessLogStore struct { mu sync.RWMutex - records []*OpenFlareAccessLog + records []*model.OpenFlareAccessLog } -func (s *memoryAccessLogStore) InsertBatch(_ context.Context, records []*OpenFlareAccessLog) error { +func (s *memoryAccessLogStore) InsertBatch(_ context.Context, records []*model.OpenFlareAccessLog) error { s.mu.Lock() defer s.mu.Unlock() now := time.Now().UTC() @@ -52,7 +54,7 @@ func (s *memoryAccessLogStore) InsertBatch(_ context.Context, records []*OpenFla return nil } -func (s *memoryAccessLogStore) List(_ context.Context, query OpenFlareAccessLogQuery) ([]*OpenFlareAccessLog, error) { +func (s *memoryAccessLogStore) List(_ context.Context, query model.OpenFlareAccessLogQuery) ([]*model.OpenFlareAccessLog, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(query) @@ -64,7 +66,7 @@ func (s *memoryAccessLogStore) List(_ context.Context, query OpenFlareAccessLogQ return cloneAccessLogSlice(rows), nil } -func (s *memoryAccessLogStore) Count(_ context.Context, query OpenFlareAccessLogQuery) (int64, int64, int64, error) { +func (s *memoryAccessLogStore) Count(_ context.Context, query model.OpenFlareAccessLogQuery) (int64, int64, int64, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(query) @@ -81,10 +83,10 @@ func (s *memoryAccessLogStore) Count(_ context.Context, query OpenFlareAccessLog return int64(len(rows)), int64(len(ips)), totalBytes, nil } -func (s *memoryAccessLogStore) RegionCounts(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareAccessLogRegionCount, error) { +func (s *memoryAccessLogStore) RegionCounts(_ context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareAccessLogRegionCount, error) { s.mu.RLock() defer s.mu.RUnlock() - rows := s.filterRecords(OpenFlareAccessLogQuery{NodeID: nodeID, Since: since}) + rows := s.filterRecords(model.OpenFlareAccessLogQuery{NodeID: nodeID, Since: since}) counts := make(map[string]int64) for _, row := range rows { region := strings.TrimSpace(row.Region) @@ -93,9 +95,9 @@ func (s *memoryAccessLogStore) RegionCounts(_ context.Context, nodeID string, si } counts[region]++ } - result := make([]*OpenFlareAccessLogRegionCount, 0, len(counts)) + result := make([]*model.OpenFlareAccessLogRegionCount, 0, len(counts)) for region, count := range counts { - result = append(result, &OpenFlareAccessLogRegionCount{Region: region, Count: count}) + result = append(result, &model.OpenFlareAccessLogRegionCount{Region: region, Count: count}) } sort.Slice(result, func(i, j int) bool { if result[i].Count == result[j].Count { @@ -109,7 +111,7 @@ func (s *memoryAccessLogStore) RegionCounts(_ context.Context, nodeID string, si return result, nil } -func (s *memoryAccessLogStore) BucketAggregates(_ context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error) { +func (s *memoryAccessLogStore) BucketAggregates(_ context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogBucketAggregateRow, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(filter) @@ -154,9 +156,9 @@ func (s *memoryAccessLogStore) BucketAggregates(_ context.Context, filter OpenFl item.UniqueHostCount = int64(len(item.uniqueHosts)) result = append(result, item.openFlareAccessLogBucketAggregateRow) } - bucketRows := make([]*OpenFlareAccessLogBucketRow, len(result)) + bucketRows := make([]*model.OpenFlareAccessLogBucketRow, len(result)) for index := range result { - bucketRows[index] = &OpenFlareAccessLogBucketRow{ + bucketRows[index] = &model.OpenFlareAccessLogBucketRow{ BucketEpoch: result[index].BucketEpoch, RequestCount: result[index].RequestCount, UniqueIPCount: result[index].UniqueIPCount, @@ -189,7 +191,7 @@ func (s *memoryAccessLogStore) BucketAggregates(_ context.Context, filter OpenFl return result, nil } -func (s *memoryAccessLogStore) CountBuckets(_ context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error) { +func (s *memoryAccessLogStore) CountBuckets(_ context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) (int64, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(filter) @@ -200,7 +202,7 @@ func (s *memoryAccessLogStore) CountBuckets(_ context.Context, filter OpenFlareA return int64(len(seen)), nil } -func (s *memoryAccessLogStore) BucketDimensions(_ context.Context, filter OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error) { +func (s *memoryAccessLogStore) BucketDimensions(_ context.Context, filter model.OpenFlareAccessLogQuery, column string, bucketSeconds int64) ([]openFlareAccessLogBucketDimensionRow, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(filter) @@ -232,7 +234,7 @@ func (s *memoryAccessLogStore) BucketDimensions(_ context.Context, filter OpenFl return result, nil } -func (s *memoryAccessLogStore) IPAggregates(_ context.Context, filter OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error) { +func (s *memoryAccessLogStore) IPAggregates(_ context.Context, filter model.OpenFlareAccessLogQuery, exactRemoteAddr bool) ([]openFlareAccessLogIPAggregateRow, error) { s.mu.RLock() defer s.mu.RUnlock() if exactRemoteAddr && strings.TrimSpace(filter.RemoteAddr) == "" { @@ -274,7 +276,7 @@ func (s *memoryAccessLogStore) IPAggregates(_ context.Context, filter OpenFlareA return result, nil } -func (s *memoryAccessLogStore) IPSummaries(_ context.Context, filter OpenFlareAccessLogQuery, _ time.Time) ([]openFlareAccessLogIPSummaryRow, error) { +func (s *memoryAccessLogStore) IPSummaries(_ context.Context, filter model.OpenFlareAccessLogQuery, _ time.Time) ([]openFlareAccessLogIPSummaryRow, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(filter) @@ -314,13 +316,13 @@ func (s *memoryAccessLogStore) IPSummaries(_ context.Context, filter OpenFlareAc item.Region = strings.TrimSpace(row.Region) } } - summaryRows := make([]*OpenFlareAccessLogIPSummaryRow, 0, len(aggregates)) + summaryRows := make([]*model.OpenFlareAccessLogIPSummaryRow, 0, len(aggregates)) for _, item := range aggregates { ratio := 0.0 if item.TotalRequests > 0 { ratio = float64(item.Success2xxCount) / float64(item.TotalRequests) } - summaryRows = append(summaryRows, &OpenFlareAccessLogIPSummaryRow{ + summaryRows = append(summaryRows, &model.OpenFlareAccessLogIPSummaryRow{ RemoteAddr: item.RemoteAddr, Region: item.Region, TotalRequests: item.TotalRequests, @@ -354,7 +356,7 @@ func (s *memoryAccessLogStore) IPSummaries(_ context.Context, filter OpenFlareAc return result, nil } -func (s *memoryAccessLogStore) CountIPSummaries(_ context.Context, filter OpenFlareAccessLogQuery) (int64, error) { +func (s *memoryAccessLogStore) CountIPSummaries(_ context.Context, filter model.OpenFlareAccessLogQuery) (int64, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(filter) @@ -369,7 +371,7 @@ func (s *memoryAccessLogStore) CountIPSummaries(_ context.Context, filter OpenFl return int64(len(seen)), nil } -func (s *memoryAccessLogStore) WAFIPAggregates(_ context.Context, filter OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error) { +func (s *memoryAccessLogStore) WAFIPAggregates(_ context.Context, filter model.OpenFlareAccessLogQuery) ([]openFlareAccessLogWAFIPAggregateRow, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(filter) @@ -417,7 +419,7 @@ func (s *memoryAccessLogStore) WAFIPAggregates(_ context.Context, filter OpenFla return result, nil } -func (s *memoryAccessLogStore) IPTrend(_ context.Context, filter OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error) { +func (s *memoryAccessLogStore) IPTrend(_ context.Context, filter model.OpenFlareAccessLogQuery, bucketSeconds int64) ([]openFlareAccessLogIPTrendRow, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(filter) @@ -446,7 +448,7 @@ func (s *memoryAccessLogStore) DeleteBefore(_ context.Context, cutoff time.Time) s.mu.Lock() defer s.mu.Unlock() cutoff = cutoff.UTC() - remaining := make([]*OpenFlareAccessLog, 0, len(s.records)) + remaining := make([]*model.OpenFlareAccessLog, 0, len(s.records)) var deleted int64 for _, row := range s.records { if row.LoggedAt.Before(cutoff) { @@ -463,7 +465,7 @@ func (s *memoryAccessLogStore) DeleteByNodeBefore(_ context.Context, nodeID stri s.mu.Lock() defer s.mu.Unlock() before = before.UTC() - remaining := make([]*OpenFlareAccessLog, 0, len(s.records)) + remaining := make([]*model.OpenFlareAccessLog, 0, len(s.records)) var deleted int64 for _, row := range s.records { if row.NodeID == nodeID && row.LoggedAt.Before(before) { @@ -476,13 +478,13 @@ func (s *memoryAccessLogStore) DeleteByNodeBefore(_ context.Context, nodeID stri return deleted, nil } -func (s *memoryAccessLogStore) TrafficSummary(_ context.Context, filter OpenFlareAccessLogQuery) (OpenFlareAccessLogTrafficSummary, error) { +func (s *memoryAccessLogStore) TrafficSummary(_ context.Context, filter model.OpenFlareAccessLogQuery) (model.OpenFlareAccessLogTrafficSummary, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(filter) ips := make(map[string]struct{}) nodes := make(map[string]struct{}) - var summary OpenFlareAccessLogTrafficSummary + var summary model.OpenFlareAccessLogTrafficSummary for _, row := range rows { summary.RequestCount++ summary.BytesSent += row.BytesSent @@ -502,7 +504,7 @@ func (s *memoryAccessLogStore) TrafficSummary(_ context.Context, filter OpenFlar return summary, nil } -func (s *memoryAccessLogStore) ValueCounts(_ context.Context, filter OpenFlareAccessLogQuery, column string, limit int) ([]OpenFlareAccessLogValueCount, error) { +func (s *memoryAccessLogStore) ValueCounts(_ context.Context, filter model.OpenFlareAccessLogQuery, column string, limit int) ([]model.OpenFlareAccessLogValueCount, error) { s.mu.RLock() defer s.mu.RUnlock() col := strings.TrimSpace(strings.ToLower(column)) @@ -532,9 +534,9 @@ func (s *memoryAccessLogStore) ValueCounts(_ context.Context, filter OpenFlareAc } counts[value]++ } - result := make([]OpenFlareAccessLogValueCount, 0, len(counts)) + result := make([]model.OpenFlareAccessLogValueCount, 0, len(counts)) for value, count := range counts { - result = append(result, OpenFlareAccessLogValueCount{Value: value, Count: count}) + result = append(result, model.OpenFlareAccessLogValueCount{Value: value, Count: count}) } sort.Slice(result, func(i, j int) bool { if result[i].Count == result[j].Count { @@ -548,12 +550,12 @@ func (s *memoryAccessLogStore) ValueCounts(_ context.Context, filter OpenFlareAc return result, nil } -func (s *memoryAccessLogStore) NodeAggregates(_ context.Context, filter OpenFlareAccessLogQuery) ([]OpenFlareAccessLogNodeAggregate, error) { +func (s *memoryAccessLogStore) NodeAggregates(_ context.Context, filter model.OpenFlareAccessLogQuery) ([]model.OpenFlareAccessLogNodeAggregate, error) { s.mu.RLock() defer s.mu.RUnlock() rows := s.filterRecords(filter) type acc struct { - OpenFlareAccessLogNodeAggregate + model.OpenFlareAccessLogNodeAggregate ips map[string]struct{} } byNode := make(map[string]*acc) @@ -565,7 +567,7 @@ func (s *memoryAccessLogStore) NodeAggregates(_ context.Context, filter OpenFlar item := byNode[id] if item == nil { item = &acc{ - OpenFlareAccessLogNodeAggregate: OpenFlareAccessLogNodeAggregate{NodeID: id}, + OpenFlareAccessLogNodeAggregate: model.OpenFlareAccessLogNodeAggregate{NodeID: id}, ips: make(map[string]struct{}), } byNode[id] = item @@ -578,7 +580,7 @@ func (s *memoryAccessLogStore) NodeAggregates(_ context.Context, filter OpenFlar item.ips[ip] = struct{}{} } } - result := make([]OpenFlareAccessLogNodeAggregate, 0, len(byNode)) + result := make([]model.OpenFlareAccessLogNodeAggregate, 0, len(byNode)) for _, item := range byNode { item.UniqueIPCount = int64(len(item.ips)) result = append(result, item.OpenFlareAccessLogNodeAggregate) @@ -592,8 +594,8 @@ func (s *memoryAccessLogStore) NodeAggregates(_ context.Context, filter OpenFlar return result, nil } -func (s *memoryAccessLogStore) filterRecords(query OpenFlareAccessLogQuery) []*OpenFlareAccessLog { - result := make([]*OpenFlareAccessLog, 0, len(s.records)) +func (s *memoryAccessLogStore) filterRecords(query model.OpenFlareAccessLogQuery) []*model.OpenFlareAccessLog { + result := make([]*model.OpenFlareAccessLog, 0, len(s.records)) for _, row := range s.records { if !memoryAccessLogMatches(row, query) { continue @@ -603,7 +605,7 @@ func (s *memoryAccessLogStore) filterRecords(query OpenFlareAccessLogQuery) []*O return result } -func memoryAccessLogMatches(row *OpenFlareAccessLog, query OpenFlareAccessLogQuery) bool { +func memoryAccessLogMatches(row *model.OpenFlareAccessLog, query model.OpenFlareAccessLogQuery) bool { if row == nil { return false } @@ -661,8 +663,8 @@ func memoryAccessLogBucketEpoch(loggedAt time.Time, bucketSeconds int64) int64 { return (epoch / bucketSeconds) * bucketSeconds } -func cloneAccessLogSlice(rows []*OpenFlareAccessLog) []*OpenFlareAccessLog { - result := make([]*OpenFlareAccessLog, len(rows)) +func cloneAccessLogSlice(rows []*model.OpenFlareAccessLog) []*model.OpenFlareAccessLog { + result := make([]*model.OpenFlareAccessLog, len(rows)) for index, row := range rows { if row == nil { continue @@ -673,7 +675,7 @@ func cloneAccessLogSlice(rows []*OpenFlareAccessLog) []*OpenFlareAccessLog { return result } -func sortOpenFlareAccessLogRows(items []*OpenFlareAccessLog, sortBy string, sortOrder string) { +func sortOpenFlareAccessLogRows(items []*model.OpenFlareAccessLog, sortBy string, sortOrder string) { desc := openFlareAccessLogNormalizeSortOrder(sortOrder) != sortOrderAsc sort.Slice(items, func(i, j int) bool { left := items[i] diff --git a/internal/model/openflare_access_log_test.go b/internal/repository/openflare_access_log_test.go similarity index 87% rename from internal/model/openflare_access_log_test.go rename to internal/repository/openflare_access_log_test.go index f1f6a0d1..3b3d5410 100644 --- a/internal/model/openflare_access_log_test.go +++ b/internal/repository/openflare_access_log_test.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" @@ -9,6 +9,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -24,7 +26,7 @@ func setupOpenFlareAccessLogTestEnvironment(t *testing.T) (context.Context, func func seedOpenFlareAccessLogs(t *testing.T, ctx context.Context, now time.Time) { t.Helper() - records := []*OpenFlareAccessLog{ + records := []*model.OpenFlareAccessLog{ {NodeID: "node-a", LoggedAt: now.Add(-5 * time.Minute), RemoteAddr: "1.1.1.1", Region: "US", Host: "a.example.com", Path: "/alpha", StatusCode: 200}, {NodeID: "node-a", LoggedAt: now.Add(-4 * time.Minute), RemoteAddr: "2.2.2.2", Region: "US", Host: "a.example.com", Path: "/beta", StatusCode: 404}, {NodeID: "node-b", LoggedAt: now.Add(-3 * time.Minute), RemoteAddr: "1.1.1.1", Region: "EU", Host: "b.example.com", Path: "/gamma", StatusCode: 502}, @@ -40,7 +42,7 @@ func TestListOpenFlareAccessLogsPaginated(t *testing.T) { now := time.Now().UTC() for index := range 15 { - record := &OpenFlareAccessLog{ + record := &model.OpenFlareAccessLog{ NodeID: "node-page", LoggedAt: now.Add(-time.Duration(index) * time.Minute), RemoteAddr: fmt.Sprintf("203.0.113.%d", (index%5)+1), @@ -48,10 +50,10 @@ func TestListOpenFlareAccessLogsPaginated(t *testing.T) { Path: fmt.Sprintf("/path-%02d", index), StatusCode: 200, } - require.NoError(t, InsertOpenFlareAccessLogsBatch(ctx, []*OpenFlareAccessLog{record})) + require.NoError(t, InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{record})) } - query := OpenFlareAccessLogQuery{ + query := model.OpenFlareAccessLogQuery{ NodeID: "node-page", Since: now.Add(-24 * time.Hour), Page: 1, @@ -73,7 +75,7 @@ func TestCountOpenFlareAccessLogs(t *testing.T) { now := time.Now().UTC() seedOpenFlareAccessLogs(t, ctx, now) - query := OpenFlareAccessLogQuery{ + query := model.OpenFlareAccessLogQuery{ Since: now.Add(-10 * time.Minute), } totalRecords, totalIPs, _, err := CountOpenFlareAccessLogs(ctx, query) @@ -89,7 +91,7 @@ func TestListOpenFlareAccessLogsFiltersAndSort(t *testing.T) { now := time.Now().UTC() seedOpenFlareAccessLogs(t, ctx, now) - query := OpenFlareAccessLogQuery{ + query := model.OpenFlareAccessLogQuery{ NodeID: "node-a", Since: now.Add(-10 * time.Minute), SortBy: "status_code", @@ -113,7 +115,7 @@ func TestDeleteOpenFlareAccessLogsBefore(t *testing.T) { require.NoError(t, err) assert.Equal(t, int64(3), deleted) - totalRecords, _, _, err := CountOpenFlareAccessLogs(ctx, OpenFlareAccessLogQuery{Since: now.Add(-10 * time.Minute)}) + totalRecords, _, _, err := CountOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{Since: now.Add(-10 * time.Minute)}) require.NoError(t, err) assert.Equal(t, int64(2), totalRecords) } diff --git a/internal/repository/openflare_acme_account.go b/internal/repository/openflare_acme_account.go new file mode 100644 index 00000000..36533b2c --- /dev/null +++ b/internal/repository/openflare_acme_account.go @@ -0,0 +1,68 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + + "gorm.io/gorm" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// GetAcmeAccountByID 按 ID 查询 ACME 账号。 +func GetAcmeAccountByID(ctx context.Context, id uint) (*model.AcmeAccount, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var account model.AcmeAccount + if err := conn.First(&account, id).Error; err != nil { + return nil, err + } + return &account, nil +} + +// CreateAcmeAccountRecord 创建 ACME 账号。 +func CreateAcmeAccountRecord(ctx context.Context, account *model.AcmeAccount) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(account).Error +} + +// SaveAcmeAccount 保存 ACME 账号。 +func SaveAcmeAccount(ctx context.Context, account *model.AcmeAccount) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Save(account).Error +} + +// GetDefaultAcmeAccount 获取默认 ACME 账号,不存在时创建占位记录。 +func GetDefaultAcmeAccount(ctx context.Context) (*model.AcmeAccount, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var account model.AcmeAccount + err := conn.Order("id asc").First(&account).Error + if err == nil { + return &account, nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, err + } + account = model.AcmeAccount{ + Email: "admin@openflare.dev", + } + if err = conn.Create(&account).Error; err != nil { + return nil, err + } + return &account, nil +} diff --git a/internal/repository/openflare_apply_log.go b/internal/repository/openflare_apply_log.go new file mode 100644 index 00000000..a2691695 --- /dev/null +++ b/internal/repository/openflare_apply_log.go @@ -0,0 +1,162 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + "strings" + "time" + + "gorm.io/gorm" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// ListOpenFlareApplyLogs returns apply logs ordered by id desc with optional pagination. +func ListOpenFlareApplyLogs(ctx context.Context, query model.OpenFlareApplyLogQuery) ([]*model.OpenFlareApplyLog, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + + dbQuery := conn.Model(&model.OpenFlareApplyLog{}).Order("id desc") + if query.NodeID != "" { + dbQuery = dbQuery.Where("node_id = ?", query.NodeID) + } + if query.PageSize > 0 { + offset := 0 + if query.PageNo > 1 { + offset = (query.PageNo - 1) * query.PageSize + } + dbQuery = dbQuery.Limit(query.PageSize).Offset(offset) + } + + var logs []*model.OpenFlareApplyLog + if err := dbQuery.Find(&logs).Error; err != nil { + return nil, err + } + return logs, nil +} + +// CountOpenFlareApplyLogs returns total apply logs, optionally filtered by node_id. +func CountOpenFlareApplyLogs(ctx context.Context, nodeID string) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + + query := conn.Model(&model.OpenFlareApplyLog{}) + if nodeID != "" { + query = query.Where("node_id = ?", nodeID) + } + + var total int64 + if err := query.Count(&total).Error; err != nil { + return 0, err + } + return total, nil +} + +// GetLatestOpenFlareApplyLogByNodeID returns the most recent apply log for a node. +func GetLatestOpenFlareApplyLogByNodeID(ctx context.Context, nodeID string) (*model.OpenFlareApplyLog, error) { + nodeID = strings.TrimSpace(nodeID) + if nodeID == "" { + return nil, errors.New("node_id is required") + } + + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + + var log model.OpenFlareApplyLog + err := conn.Where("node_id = ?", nodeID).Order("id desc").First(&log).Error + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, nil + } + if err != nil { + return nil, err + } + return &log, nil +} + +// GetLatestOpenFlareApplyLogsByNodeIDs returns the latest apply log per node id. +func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string) (map[string]*model.OpenFlareApplyLog, error) { + result := make(map[string]*model.OpenFlareApplyLog) + if len(nodeIDs) == 0 { + return result, nil + } + + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + + var logs []*model.OpenFlareApplyLog + subQuery := conn.Model(&model.OpenFlareApplyLog{}). + Select("MAX(id) AS id"). + Where("node_id IN ?", nodeIDs). + Group("node_id") + if err := conn.Where("id IN (?)", subQuery).Find(&logs).Error; err != nil { + return nil, err + } + for _, log := range logs { + result[log.NodeID] = log + } + return result, nil +} + +// CreateOpenFlareApplyLog inserts an apply log row. +func CreateOpenFlareApplyLog(ctx context.Context, log *model.OpenFlareApplyLog) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(log).Error +} + +// CreateOpenFlareApplyLogAndUpdateNode creates an apply log and updates the node from the apply result in one transaction. +func CreateOpenFlareApplyLogAndUpdateNode(ctx context.Context, log *model.OpenFlareApplyLog, applyResult, version, message string) error { + if log == nil { + return errors.New("apply log is required") + } + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + now := log.CreatedAt + if now.IsZero() { + now = time.Now() + } + return conn.Transaction(func(tx *gorm.DB) error { + if err := tx.Create(log).Error; err != nil { + return err + } + return updateOpenFlareNodeFromApplyResultTx(tx, log.NodeID, applyResult, version, message, now) + }) +} + +// DeleteAllOpenFlareApplyLogs removes every apply log record. +func DeleteAllOpenFlareApplyLogs(ctx context.Context) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + + result := conn.Session(&gorm.Session{AllowGlobalUpdate: true}).Delete(&model.OpenFlareApplyLog{}) + return result.RowsAffected, result.Error +} + +// DeleteOpenFlareApplyLogsBefore removes apply logs created before the cutoff time. +func DeleteOpenFlareApplyLogsBefore(ctx context.Context, before time.Time) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + + result := conn.Where("created_at < ?", before).Delete(&model.OpenFlareApplyLog{}) + return result.RowsAffected, result.Error +} diff --git a/internal/model/openflare_apply_log_test.go b/internal/repository/openflare_apply_log_test.go similarity index 63% rename from internal/model/openflare_apply_log_test.go rename to internal/repository/openflare_apply_log_test.go index cddfdc44..9408f9cd 100644 --- a/internal/model/openflare_apply_log_test.go +++ b/internal/repository/openflare_apply_log_test.go @@ -1,10 +1,12 @@ -package model +package repository import ( "context" "testing" "time" + "github.com/Rain-kl/Wavelet/internal/model" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -19,7 +21,7 @@ func setupApplyLogModelTestDB(t *testing.T) func() { DisableForeignKeyConstraintWhenMigrating: true, }) require.NoError(t, err) - require.NoError(t, sqliteDB.AutoMigrate(&OpenFlareApplyLog{})) + require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareApplyLog{})) db.SetDB(sqliteDB) return func() { @@ -28,17 +30,17 @@ func setupApplyLogModelTestDB(t *testing.T) func() { } func TestIsRepeatSuccessApplyLog(t *testing.T) { - latest := &OpenFlareApplyLog{ + latest := &model.OpenFlareApplyLog{ Version: "20260615-001", Checksum: "checksum-a", Result: "success", } - assert.True(t, IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-a", "success")) - assert.False(t, IsRepeatSuccessApplyLog(latest, "20260615-002", "checksum-a", "success")) - assert.False(t, IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-b", "success")) - assert.False(t, IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-a", "failed")) - assert.False(t, IsRepeatSuccessApplyLog(nil, "20260615-001", "checksum-a", "success")) + assert.True(t, model.IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-a", "success")) + assert.False(t, model.IsRepeatSuccessApplyLog(latest, "20260615-002", "checksum-a", "success")) + assert.False(t, model.IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-b", "success")) + assert.False(t, model.IsRepeatSuccessApplyLog(latest, "20260615-001", "checksum-a", "failed")) + assert.False(t, model.IsRepeatSuccessApplyLog(nil, "20260615-001", "checksum-a", "success")) } func TestGetLatestOpenFlareApplyLogByNodeID(t *testing.T) { @@ -47,14 +49,14 @@ func TestGetLatestOpenFlareApplyLogByNodeID(t *testing.T) { ctx := context.Background() now := time.Now().UTC() - require.NoError(t, db.DB(ctx).Create(&OpenFlareApplyLog{ + require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{ NodeID: "node-1", Version: "v1", Result: "success", Checksum: "checksum-1", CreatedAt: now.Add(-time.Hour), }).Error) - require.NoError(t, db.DB(ctx).Create(&OpenFlareApplyLog{ + require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{ NodeID: "node-1", Version: "v2", Result: "success", diff --git a/internal/repository/openflare_config_version.go b/internal/repository/openflare_config_version.go new file mode 100644 index 00000000..00d12a97 --- /dev/null +++ b/internal/repository/openflare_config_version.go @@ -0,0 +1,135 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + + "gorm.io/gorm" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// ListConfigVersionSummaries returns config version summaries ordered by created_at desc. +func ListConfigVersionSummaries(ctx context.Context) ([]*model.ConfigVersionSummary, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var versions []*model.ConfigVersionSummary + err := conn.Model(&model.ConfigVersion{}). + Select("version", "checksum", "is_active", "created_by", "created_at"). + Order("created_at desc, version desc"). + Find(&versions).Error + return versions, err +} + +// GetConfigVersionByVersion returns a config version by version string. +func GetConfigVersionByVersion(ctx context.Context, version string) (*model.ConfigVersion, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var cv model.ConfigVersion + if err := conn.First(&cv, "version = ?", version).Error; err != nil { + return nil, err + } + return &cv, nil +} + +// GetActiveConfigVersion returns the currently active config version. +func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var version model.ConfigVersion + if err := conn.Where("is_active = ?", true).Order("version desc").First(&version).Error; err != nil { + return nil, err + } + return &version, nil +} + +// GetLatestConfigVersionByPrefix returns the latest version string matching a date prefix. +func GetLatestConfigVersionByPrefix(ctx context.Context, prefix string) (string, error) { + conn := db.DB(ctx) + if conn == nil { + return "", errors.New(errDatabaseNotInitialized) + } + var version model.ConfigVersion + err := conn.Model(&model.ConfigVersion{}). + Select("version"). + Where("version LIKE ?", prefix+"-%"). + Order("version desc"). + First(&version).Error + if err != nil { + return "", err + } + return version.Version, nil +} + +// CreateConfigVersion inserts a new config version record. +func CreateConfigVersion(ctx context.Context, version *model.ConfigVersion) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(version).Error +} + +// PublishConfigVersionTx deactivates all versions and creates a new active version. +func PublishConfigVersionTx(ctx context.Context, version *model.ConfigVersion) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Transaction(func(tx *gorm.DB) error { + if err := tx.Model(&model.ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil { + return err + } + return tx.Create(version).Error + }) +} + +// ActivateConfigVersionTx marks the given version active and deactivates others. +func ActivateConfigVersionTx(ctx context.Context, version string) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Transaction(func(tx *gorm.DB) error { + if err := tx.Model(&model.ConfigVersion{}).Where("is_active = ?", true).Update("is_active", false).Error; err != nil { + return err + } + return tx.Model(&model.ConfigVersion{}).Where("version = ?", version).Update("is_active", true).Error + }) +} + +// DeleteConfigVersionsByVersions removes config versions by versions. +func DeleteConfigVersionsByVersions(ctx context.Context, versions []string) (int64, error) { + if len(versions) == 0 { + return 0, nil + } + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + result := conn.Where("version IN ?", versions).Delete(&model.ConfigVersion{}) + return result.RowsAffected, result.Error +} + +// ListEnabledProxyRoutes returns enabled proxy routes ordered by id asc. +func ListEnabledProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var routes []*model.ProxyRoute + if err := conn.Where("enabled = ?", true).Order("id asc").Find(&routes).Error; err != nil { + return nil, err + } + return routes, nil +} diff --git a/internal/repository/openflare_dns_account.go b/internal/repository/openflare_dns_account.go new file mode 100644 index 00000000..69ac94fc --- /dev/null +++ b/internal/repository/openflare_dns_account.go @@ -0,0 +1,65 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// ListDNSAccounts 列出全部 DNS 账号(授权信息不通过 JSON 暴露)。 +func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var accounts []model.DNSAccount + if err := conn.Order("id desc").Find(&accounts).Error; err != nil { + return nil, err + } + return accounts, nil +} + +// GetDNSAccountByID 按 ID 查询 DNS 账号。 +func GetDNSAccountByID(ctx context.Context, id uint) (*model.DNSAccount, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var account model.DNSAccount + if err := conn.First(&account, id).Error; err != nil { + return nil, err + } + return &account, nil +} + +// CreateDNSAccountRecord 创建 DNS 账号。 +func CreateDNSAccountRecord(ctx context.Context, account *model.DNSAccount) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(account).Error +} + +// SaveDNSAccount 保存 DNS 账号。 +func SaveDNSAccount(ctx context.Context, account *model.DNSAccount) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Save(account).Error +} + +// DeleteDNSAccountRecord 删除 DNS 账号。 +func DeleteDNSAccountRecord(ctx context.Context, id uint) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Delete(&model.DNSAccount{}, id).Error +} diff --git a/internal/model/openflare_insert_hooks_test.go b/internal/repository/openflare_insert_hooks_test.go similarity index 87% rename from internal/model/openflare_insert_hooks_test.go rename to internal/repository/openflare_insert_hooks_test.go index fe6f45be..22359038 100644 --- a/internal/model/openflare_insert_hooks_test.go +++ b/internal/repository/openflare_insert_hooks_test.go @@ -1,13 +1,15 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" "testing" "time" + "github.com/Rain-kl/Wavelet/internal/model" + analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics" ) @@ -24,7 +26,7 @@ func TestObservabilityInsertHooksAreInvoked(t *testing.T) { SetObservabilityInsertHooks(ObservabilityInsertHooks{}) }) - record := &OpenFlareMetricSnapshot{ + record := &model.OpenFlareMetricSnapshot{ NodeID: "node-1", CapturedAt: time.Unix(100, 0).UTC(), } @@ -47,7 +49,7 @@ func TestAccessLogInsertHooksAreInvoked(t *testing.T) { SetAccessLogInsertHooks(AccessLogInsertHooks{}) }) - records := []*OpenFlareAccessLog{ + records := []*model.OpenFlareAccessLog{ {NodeID: "n1", Path: "/a"}, {NodeID: "n1", Path: "/b"}, } @@ -66,10 +68,10 @@ func TestInsertHooksNoopWhenUnset(t *testing.T) { SetObservabilityInsertHooks(ObservabilityInsertHooks{}) SetAccessLogInsertHooks(AccessLogInsertHooks{}) - if err := (clickhouseObservabilityStore{}).InsertMetricSnapshot(context.Background(), &OpenFlareMetricSnapshot{NodeID: "x"}); err != nil { + if err := (clickhouseObservabilityStore{}).InsertMetricSnapshot(context.Background(), &model.OpenFlareMetricSnapshot{NodeID: "x"}); err != nil { t.Fatalf("InsertMetricSnapshot with nil hook error = %v", err) } - if err := (clickhouseAccessLogStore{}).InsertBatch(context.Background(), []*OpenFlareAccessLog{{NodeID: "x"}}); err != nil { + if err := (clickhouseAccessLogStore{}).InsertBatch(context.Background(), []*model.OpenFlareAccessLog{{NodeID: "x"}}); err != nil { t.Fatalf("InsertBatch with nil hook error = %v", err) } } diff --git a/internal/repository/openflare_node.go b/internal/repository/openflare_node.go new file mode 100644 index 00000000..55bc954f --- /dev/null +++ b/internal/repository/openflare_node.go @@ -0,0 +1,167 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + "time" + + "gorm.io/gorm" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + openFlareNodeStatusOnline = "online" + openFlareApplyResultSuccess = "success" +) + +// ListOpenFlareNodes returns all nodes ordered by id desc. +func ListOpenFlareNodes(ctx context.Context) ([]model.OpenFlareNode, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var nodes []model.OpenFlareNode + if err := conn.Order("id desc").Find(&nodes).Error; err != nil { + return nil, err + } + return nodes, nil +} + +// ListOpenFlareNodesByNodeIDs returns nodes matching the given node ids. +func ListOpenFlareNodesByNodeIDs(ctx context.Context, nodeIDs []string) ([]model.OpenFlareNode, error) { + if len(nodeIDs) == 0 { + return []model.OpenFlareNode{}, nil + } + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var nodes []model.OpenFlareNode + if err := conn.Where("node_id IN ?", nodeIDs).Find(&nodes).Error; err != nil { + return nil, err + } + return nodes, nil +} + +// GetOpenFlareNodeByID returns a node by primary key. +func GetOpenFlareNodeByID(ctx context.Context, id uint) (*model.OpenFlareNode, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var node model.OpenFlareNode + if err := conn.First(&node, id).Error; err != nil { + return nil, err + } + return &node, nil +} + +// GetOpenFlareNodeByNodeID returns a node by node_id. +func GetOpenFlareNodeByNodeID(ctx context.Context, nodeID string) (*model.OpenFlareNode, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var node model.OpenFlareNode + if err := conn.Where("node_id = ?", nodeID).First(&node).Error; err != nil { + return nil, err + } + return &node, nil +} + +// GetOpenFlareNodeByAccessToken returns a node by access token. +func GetOpenFlareNodeByAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var node model.OpenFlareNode + if err := conn.Where("access_token = ?", token).First(&node).Error; err != nil { + return nil, err + } + return &node, nil +} + +// CreateOpenFlareNode inserts a new node. +func CreateOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(node).Error +} + +// SaveOpenFlareNode persists node changes. +func SaveOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Save(node).Error +} + +// UpdateOpenFlareNodeFields updates selected columns for a node. +func UpdateOpenFlareNodeFields(ctx context.Context, node *model.OpenFlareNode, fields ...string) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + if len(fields) == 0 { + return conn.Save(node).Error + } + return conn.Model(node).Select(fields).Updates(node).Error +} + +// UpdateOpenFlareNodeColumns updates node columns from a map of column values. +// Empty maps are no-ops. +func UpdateOpenFlareNodeColumns(ctx context.Context, node *model.OpenFlareNode, changes map[string]any) error { + if node == nil || len(changes) == 0 { + return nil + } + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Model(node).Updates(changes).Error +} + +// UpdateOpenFlareNodeFromApplyResult updates node status, last_seen, version and last_error after an apply report. +// When applyResult is "success", current_version is set and last_error is cleared; otherwise last_error is set to message. +func UpdateOpenFlareNodeFromApplyResult(ctx context.Context, nodeID, applyResult, version, message string, now time.Time) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return updateOpenFlareNodeFromApplyResultTx(conn, nodeID, applyResult, version, message, now) +} + +func updateOpenFlareNodeFromApplyResultTx(tx *gorm.DB, nodeID, applyResult, version, message string, now time.Time) error { + record := &model.OpenFlareNode{} + if err := tx.Where("node_id = ?", nodeID).First(record).Error; err != nil { + return err + } + record.Status = openFlareNodeStatusOnline + lastSeen := now + record.LastSeenAt = &lastSeen + if applyResult == openFlareApplyResultSuccess { + record.CurrentVersion = version + record.LastError = "" + } else { + record.LastError = message + } + return tx.Model(record).Select("status", "last_seen_at", "current_version", "last_error").Updates(record).Error +} + +// DeleteOpenFlareNode removes a node by primary key. +func DeleteOpenFlareNode(ctx context.Context, id uint) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Delete(&model.OpenFlareNode{}, id).Error +} diff --git a/internal/repository/openflare_observability.go b/internal/repository/openflare_observability.go new file mode 100644 index 00000000..b1fe60bf --- /dev/null +++ b/internal/repository/openflare_observability.go @@ -0,0 +1,537 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "encoding/json" + "errors" + "strings" + "time" + + "gorm.io/gorm" + "gorm.io/gorm/clause" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" + analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" +) + +const ( + openFlareHealthEventStatusActive = "active" + openFlareHealthEventStatusResolved = "resolved" + openFlareHealthSeverityInfo = "info" + openFlareHealthSeverityWarning = "warning" + openFlareHealthSeverityCritical = "critical" + openFlareHealthEventMessageMaxLen = 4096 +) + +// OpenFlareHealthEventInput describes a desired active health event for reconciliation. +type OpenFlareHealthEventInput struct { + EventType string + Severity string + Message string + TriggeredAtUnix int64 + Metadata map[string]string +} + +func isMissingTableError(err error) bool { + if err == nil { + return false + } + if errors.Is(err, gorm.ErrRecordNotFound) { + return false + } + msg := strings.ToLower(err.Error()) + return strings.Contains(msg, "no such table") || + strings.Contains(msg, "doesn't exist") || + strings.Contains(msg, "does not exist") +} + +// InsertOpenFlareMetricSnapshot inserts a metric snapshot into ClickHouse. +func InsertOpenFlareMetricSnapshot(ctx context.Context, record *model.OpenFlareMetricSnapshot) error { + return currentObservabilityStore().InsertMetricSnapshot(ctx, record) +} + +// InsertOpenFlareEdgeHealth inserts an L2 edge health snapshot into ClickHouse. +func InsertOpenFlareEdgeHealth(ctx context.Context, record *model.OpenFlareEdgeHealth) error { + return currentObservabilityStore().InsertEdgeHealth(ctx, record) +} + +// InsertOpenFlareNodeObservationFrps inserts an FRPS observation into ClickHouse. +func InsertOpenFlareNodeObservationFrps(ctx context.Context, record *model.OpenFlareNodeObservationFrps) error { + return currentObservabilityStore().InsertNodeObservationFrps(ctx, record) +} + +// InsertOpenFlareNodeObservationFrpc inserts an FRPC observation into ClickHouse. +func InsertOpenFlareNodeObservationFrpc(ctx context.Context, record *model.OpenFlareNodeObservationFrpc) error { + return currentObservabilityStore().InsertNodeObservationFrpc(ctx, record) +} + +// ListOpenFlareMetricSnapshotsSince returns metric snapshots since the given time. +func ListOpenFlareMetricSnapshotsSince(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareMetricSnapshot, error) { + return currentObservabilityStore().ListMetricSnapshots(ctx, nodeID, since, limit) +} + +// ListOpenFlareLatestMetricSnapshotsSince returns the latest metric snapshot per node. +// Prefer ClickHouse LIMIT 1 BY; on CH unavailability fall back to store list + reduce. +func ListOpenFlareLatestMetricSnapshotsSince(ctx context.Context, nodeID string, since time.Time) ([]*model.OpenFlareMetricSnapshot, error) { + rows, err := analyticsrepo.ListLatestNodeMetricSnapshots(ctx, analyticsrepo.NodeObservabilityFilter{ + NodeID: nodeID, + Since: since, + }) + if err == nil { + return fromAnalyticsNodeMetricSnapshots(rows), nil + } + // Fallback for unit tests (memory store) and environments without ClickHouse. + all, listErr := ListOpenFlareMetricSnapshotsSince(ctx, nodeID, since, 0) + if listErr != nil { + return nil, err + } + return openFlareLatestMetricSnapshots(all), nil +} + +func openFlareLatestMetricSnapshots(snapshots []*model.OpenFlareMetricSnapshot) []*model.OpenFlareMetricSnapshot { + latestByNode := make(map[string]*model.OpenFlareMetricSnapshot, len(snapshots)) + for _, snapshot := range snapshots { + if snapshot == nil || snapshot.NodeID == "" { + continue + } + if existing, ok := latestByNode[snapshot.NodeID]; ok && !snapshot.CapturedAt.After(existing.CapturedAt) { + continue + } + latestByNode[snapshot.NodeID] = snapshot + } + result := make([]*model.OpenFlareMetricSnapshot, 0, len(latestByNode)) + for _, snapshot := range latestByNode { + result = append(result, snapshot) + } + return result +} + +// ListOpenFlareTrafficHourlySince returns hourly traffic rollup rows since the given time. +// Source: of_access_log_hourly (M5). +func ListOpenFlareTrafficHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*model.OpenFlareTrafficHourly, error) { + rows, err := analyticsrepo.ListNodeTrafficHourly(ctx, analyticsrepo.NodeObservabilityFilter{ + NodeID: nodeID, + Since: since, + }) + if err != nil { + return nil, err + } + result := make([]*model.OpenFlareTrafficHourly, len(rows)) + for index, row := range rows { + result[index] = &model.OpenFlareTrafficHourly{ + NodeID: row.NodeID, + Hour: row.Hour, + RequestCount: row.RequestCount, + ErrorCount: row.ErrorCount, + UniqueVisitorCount: row.UniqueVisitorCount, + } + } + return result, nil +} + +// ListOpenFlareAccessLogHourlySince returns of_access_log_hourly rows since the given time. +func ListOpenFlareAccessLogHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*model.OpenFlareAccessLogHourly, error) { + rows, err := analyticsrepo.ListAccessLogHourly(ctx, analyticsrepo.NodeObservabilityFilter{ + NodeID: nodeID, + Since: since, + }) + if err != nil { + return nil, err + } + result := make([]*model.OpenFlareAccessLogHourly, len(rows)) + for index, row := range rows { + result[index] = &model.OpenFlareAccessLogHourly{ + NodeID: row.NodeID, + Hour: row.Hour, + Host: row.Host, + RequestCount: row.RequestCount, + ErrorCount: row.ErrorCount, + BytesSent: row.BytesSent, + RequestLength: row.RequestLength, + } + } + return result, nil +} + +// ListOpenFlareMetricHourlySince returns hourly metric aggregates since the given time. +func ListOpenFlareMetricHourlySince(ctx context.Context, nodeID string, since time.Time) ([]*model.OpenFlareMetricHourly, error) { + rows, err := analyticsrepo.ListNodeMetricHourly(ctx, analyticsrepo.NodeObservabilityFilter{ + NodeID: nodeID, + Since: since, + }) + if err != nil { + return nil, err + } + result := make([]*model.OpenFlareMetricHourly, len(rows)) + for index, row := range rows { + result[index] = &model.OpenFlareMetricHourly{ + Hour: row.Hour, + AverageCPUUsagePercent: row.AverageCPUUsagePercent, + AverageMemoryUsagePercent: row.AverageMemoryUsagePercent, + NetworkRxBytes: row.NetworkRxBytes, + NetworkTxBytes: row.NetworkTxBytes, + DiskReadBytes: row.DiskReadBytes, + DiskWriteBytes: row.DiskWriteBytes, + ReportedNodes: row.ReportedNodes, + } + } + return result, nil +} + +// ListOpenFlareActiveHealthEvents returns active health events across all nodes. +func ListOpenFlareActiveHealthEvents(ctx context.Context) ([]*model.OpenFlareHealthEvent, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var rows []*model.OpenFlareHealthEvent + if err := conn.Where("status = ?", "active").Order("last_triggered_at desc").Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*model.OpenFlareHealthEvent{}, nil + } + return nil, err + } + return rows, nil +} + +// ListOpenFlareHealthEvents returns health events for a node. +func ListOpenFlareHealthEvents(ctx context.Context, nodeID string, activeOnly bool, limit int) ([]*model.OpenFlareHealthEvent, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + query := conn.Model(&model.OpenFlareHealthEvent{}).Where("node_id = ?", nodeID).Order("last_triggered_at desc") + if activeOnly { + query = query.Where("status = ?", "active") + } + if limit > 0 { + query = query.Limit(limit) + } + var rows []*model.OpenFlareHealthEvent + if err := query.Find(&rows).Error; err != nil { + if isMissingTableError(err) { + return []*model.OpenFlareHealthEvent{}, nil + } + return nil, err + } + return rows, nil +} + +// DeleteOpenFlareMetricSnapshotsBefore deletes metric snapshots captured before cutoff. +func DeleteOpenFlareMetricSnapshotsBefore(ctx context.Context, cutoff time.Time) (int64, error) { + return currentObservabilityStore().DeleteMetricSnapshotsBefore(ctx, cutoff) +} + +// DeleteAllOpenFlareMetricSnapshots deletes all metric snapshots. +func DeleteAllOpenFlareMetricSnapshots(ctx context.Context) (int64, error) { + return currentObservabilityStore().DeleteAllMetricSnapshots(ctx) +} + +// DeleteOpenFlareEdgeHealthBefore deletes edge health rows captured before cutoff. +func DeleteOpenFlareEdgeHealthBefore(ctx context.Context, cutoff time.Time) (int64, error) { + return currentObservabilityStore().DeleteEdgeHealthBefore(ctx, cutoff) +} + +// DeleteAllOpenFlareEdgeHealth deletes all edge health snapshots. +func DeleteAllOpenFlareEdgeHealth(ctx context.Context) (int64, error) { + return currentObservabilityStore().DeleteAllEdgeHealth(ctx) +} + +// DeleteOpenFlareNodeObservationFrpsBefore deletes FRPS observations captured before cutoff. +func DeleteOpenFlareNodeObservationFrpsBefore(ctx context.Context, cutoff time.Time) (int64, error) { + return currentObservabilityStore().DeleteNodeObservationFrpsBefore(ctx, cutoff) +} + +// DeleteAllOpenFlareNodeObservationFrps deletes all FRPS observations. +func DeleteAllOpenFlareNodeObservationFrps(ctx context.Context) (int64, error) { + return currentObservabilityStore().DeleteAllNodeObservationFrps(ctx) +} + +// DeleteOpenFlareNodeObservationFrpcBefore deletes FRPC observations captured before cutoff. +func DeleteOpenFlareNodeObservationFrpcBefore(ctx context.Context, cutoff time.Time) (int64, error) { + return currentObservabilityStore().DeleteNodeObservationFrpcBefore(ctx, cutoff) +} + +// DeleteAllOpenFlareNodeObservationFrpc deletes all FRPC observations. +func DeleteAllOpenFlareNodeObservationFrpc(ctx context.Context) (int64, error) { + return currentObservabilityStore().DeleteAllNodeObservationFrpc(ctx) +} + +// DeleteOpenFlareHealthEventsByNodeID deletes all health events for a node. +func DeleteOpenFlareHealthEventsByNodeID(ctx context.Context, nodeID string) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + result := conn.Where("node_id = ?", nodeID).Delete(&model.OpenFlareHealthEvent{}) + if result.Error != nil { + if isMissingTableError(result.Error) { + return 0, nil + } + return 0, result.Error + } + return result.RowsAffected, nil +} + +// GetOpenFlareNodeSystemProfile returns the system profile for a node. +func GetOpenFlareNodeSystemProfile(ctx context.Context, nodeID string) (*model.OpenFlareNodeSystemProfile, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var profile model.OpenFlareNodeSystemProfile + if err := conn.Where("node_id = ?", nodeID).First(&profile).Error; err != nil { + if errors.Is(err, gorm.ErrRecordNotFound) || isMissingTableError(err) { + return nil, gorm.ErrRecordNotFound + } + return nil, err + } + return &profile, nil +} + +// UpsertOpenFlareNodeSystemProfile inserts or updates the latest system profile for a node. +func UpsertOpenFlareNodeSystemProfile(ctx context.Context, record *model.OpenFlareNodeSystemProfile) error { + if record == nil { + return nil + } + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return upsertOpenFlareNodeSystemProfileTx(conn, record) +} + +func upsertOpenFlareNodeSystemProfileTx(tx *gorm.DB, record *model.OpenFlareNodeSystemProfile) error { + return tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{{Name: "node_id"}}, + DoUpdates: clause.AssignmentColumns([]string{ + "hostname", + "os_name", + "os_version", + "kernel_version", + "architecture", + "cpu_model", + "cpu_cores", + "total_memory_bytes", + "total_disk_bytes", + "uptime_seconds", + "reported_at", + "updated_at", + }), + }).Create(record).Error +} + +// ReconcileOpenFlareHealthEvents reconciles active health events for a node. +// Desired active events are created or updated; previously active types not present are resolved. +// When managedEventTypes is non-empty, only those event types are considered. +// Runs inside a transaction so multi-row create/update/resolve stays atomic. +func ReconcileOpenFlareHealthEvents( + ctx context.Context, + nodeID string, + events []OpenFlareHealthEventInput, + reportedAt time.Time, + managedEventTypes map[string]struct{}, +) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Transaction(func(tx *gorm.DB) error { + return reconcileOpenFlareHealthEventsTx(tx, nodeID, events, reportedAt, managedEventTypes) + }) +} + +// PersistOpenFlareNodePGObservability upserts an optional system profile and optionally reconciles +// health events in a single transaction (Postgres-side heartbeat observability). +// When reconcileHealth is false, health events are left untouched. +func PersistOpenFlareNodePGObservability( + ctx context.Context, + profile *model.OpenFlareNodeSystemProfile, + nodeID string, + events []OpenFlareHealthEventInput, + reconcileHealth bool, + reportedAt time.Time, + managedEventTypes map[string]struct{}, +) error { + if profile == nil && !reconcileHealth { + return nil + } + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Transaction(func(tx *gorm.DB) error { + if profile != nil { + if err := upsertOpenFlareNodeSystemProfileTx(tx, profile); err != nil { + return err + } + } + if reconcileHealth { + if err := reconcileOpenFlareHealthEventsTx(tx, nodeID, events, reportedAt, managedEventTypes); err != nil { + return err + } + } + return nil + }) +} + +func reconcileOpenFlareHealthEventsTx( + tx *gorm.DB, + nodeID string, + events []OpenFlareHealthEventInput, + reportedAt time.Time, + managedEventTypes map[string]struct{}, +) error { + activeTypes := make(map[string]OpenFlareHealthEventInput, len(events)) + for _, event := range events { + eventType := normalizeOpenFlareHealthEventType(event.EventType) + if eventType == "" { + continue + } + if len(managedEventTypes) > 0 { + if _, ok := managedEventTypes[eventType]; !ok { + continue + } + } + event.EventType = eventType + event.Severity = normalizeOpenFlareHealthSeverity(event.Severity) + if event.TriggeredAtUnix <= 0 { + event.TriggeredAtUnix = reportedAt.Unix() + } + activeTypes[eventType] = event + } + + var activeEvents []*model.OpenFlareHealthEvent + query := tx.Where("node_id = ? AND status = ?", nodeID, openFlareHealthEventStatusActive) + if len(managedEventTypes) > 0 { + scopedTypes := make([]string, 0, len(managedEventTypes)) + for eventType := range managedEventTypes { + eventType = normalizeOpenFlareHealthEventType(eventType) + if eventType != "" { + scopedTypes = append(scopedTypes, eventType) + } + } + if len(scopedTypes) == 0 { + return nil + } + query = query.Where("event_type IN ?", scopedTypes) + } + if err := query.Find(&activeEvents).Error; err != nil { + return err + } + + activeByType := make(map[string]*model.OpenFlareHealthEvent, len(activeEvents)) + for _, event := range activeEvents { + activeByType[event.EventType] = event + } + + for eventType, event := range activeTypes { + triggeredAt := timeFromUnixSeconds(event.TriggeredAtUnix, reportedAt) + if existing, ok := activeByType[eventType]; ok { + existing.Severity = event.Severity + existing.Message = normalizeOpenFlareHealthEventMessage(event.Message) + existing.LastTriggeredAt = triggeredAt + existing.ReportedAt = reportedAt + existing.MetadataJSON = marshalOpenFlareHealthMetadata(event.Metadata) + existing.ResolvedAt = nil + if err := tx.Save(existing).Error; err != nil { + return err + } + continue + } + record := &model.OpenFlareHealthEvent{ + NodeID: nodeID, + EventType: eventType, + Severity: event.Severity, + Status: openFlareHealthEventStatusActive, + Message: normalizeOpenFlareHealthEventMessage(event.Message), + FirstTriggeredAt: triggeredAt, + LastTriggeredAt: triggeredAt, + ReportedAt: reportedAt, + MetadataJSON: marshalOpenFlareHealthMetadata(event.Metadata), + } + if err := tx.Create(record).Error; err != nil { + return err + } + } + + for _, existing := range activeEvents { + if _, ok := activeTypes[existing.EventType]; ok { + continue + } + resolvedAt := reportedAt + existing.Status = openFlareHealthEventStatusResolved + existing.ReportedAt = reportedAt + existing.ResolvedAt = &resolvedAt + if err := tx.Save(existing).Error; err != nil { + return err + } + } + + return nil +} + +func normalizeOpenFlareHealthEventType(eventType string) string { + eventType = strings.TrimSpace(strings.ToLower(eventType)) + eventType = strings.ReplaceAll(eventType, " ", "_") + return eventType +} + +func normalizeOpenFlareHealthSeverity(severity string) string { + switch strings.ToLower(strings.TrimSpace(severity)) { + case openFlareHealthSeverityCritical: + return openFlareHealthSeverityCritical + case openFlareHealthSeverityInfo: + return openFlareHealthSeverityInfo + default: + return openFlareHealthSeverityWarning + } +} + +func normalizeOpenFlareHealthEventMessage(message string) string { + if openFlareHealthEventMessageMaxLen <= 0 { + return "" + } + runes := []rune(strings.TrimSpace(message)) + if len(runes) <= openFlareHealthEventMessageMaxLen { + return string(runes) + } + return string(runes[:openFlareHealthEventMessageMaxLen]) +} + +func timeFromUnixSeconds(unixSeconds int64, fallback time.Time) time.Time { + if unixSeconds <= 0 { + return fallback + } + return time.Unix(unixSeconds, 0).UTC() +} + +func marshalOpenFlareHealthMetadata(value map[string]string) string { + if value == nil { + return "" + } + raw, err := json.Marshal(value) + if err != nil { + return "" + } + return string(raw) +} + +// ListOpenFlareEdgeHealth returns L2 edge health snapshots. +func ListOpenFlareEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareEdgeHealth, error) { + return currentObservabilityStore().ListEdgeHealth(ctx, nodeID, since, limit) +} + +// ListOpenFlareNodeObservationFrpc returns frpc observations. +func ListOpenFlareNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareNodeObservationFrpc, error) { + return currentObservabilityStore().ListNodeObservationFrpc(ctx, nodeID, since, limit) +} + +// ListOpenFlareNodeObservationFrps returns frps observations. +func ListOpenFlareNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareNodeObservationFrps, error) { + return currentObservabilityStore().ListNodeObservationFrps(ctx, nodeID, since, limit) +} diff --git a/internal/model/openflare_observability_store.go b/internal/repository/openflare_observability_store.go similarity index 82% rename from internal/model/openflare_observability_store.go rename to internal/repository/openflare_observability_store.go index 9e6d94f0..fd74a75d 100644 --- a/internal/model/openflare_observability_store.go +++ b/internal/repository/openflare_observability_store.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" @@ -10,6 +10,8 @@ import ( "sync" "time" + "github.com/Rain-kl/Wavelet/internal/model" + analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics" analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" ) @@ -42,23 +44,23 @@ func currentObservabilityInsertHooks() ObservabilityInsertHooks { } type observabilityStore interface { - InsertMetricSnapshot(ctx context.Context, record *OpenFlareMetricSnapshot) error - ListMetricSnapshots(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error) + InsertMetricSnapshot(ctx context.Context, record *model.OpenFlareMetricSnapshot) error + ListMetricSnapshots(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareMetricSnapshot, error) DeleteAllMetricSnapshots(ctx context.Context) (int64, error) DeleteMetricSnapshotsBefore(ctx context.Context, cutoff time.Time) (int64, error) - InsertEdgeHealth(ctx context.Context, record *OpenFlareEdgeHealth) error - ListEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareEdgeHealth, error) + InsertEdgeHealth(ctx context.Context, record *model.OpenFlareEdgeHealth) error + ListEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareEdgeHealth, error) DeleteAllEdgeHealth(ctx context.Context) (int64, error) DeleteEdgeHealthBefore(ctx context.Context, cutoff time.Time) (int64, error) - InsertNodeObservationFrps(ctx context.Context, record *OpenFlareNodeObservationFrps) error - ListNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error) + InsertNodeObservationFrps(ctx context.Context, record *model.OpenFlareNodeObservationFrps) error + ListNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareNodeObservationFrps, error) DeleteAllNodeObservationFrps(ctx context.Context) (int64, error) DeleteNodeObservationFrpsBefore(ctx context.Context, cutoff time.Time) (int64, error) - InsertNodeObservationFrpc(ctx context.Context, record *OpenFlareNodeObservationFrpc) error - ListNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error) + InsertNodeObservationFrpc(ctx context.Context, record *model.OpenFlareNodeObservationFrpc) error + ListNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareNodeObservationFrpc, error) DeleteAllNodeObservationFrpc(ctx context.Context) (int64, error) DeleteNodeObservationFrpcBefore(ctx context.Context, cutoff time.Time) (int64, error) } @@ -97,7 +99,7 @@ func NewMemoryObservabilityStore() observabilityStore { type clickhouseObservabilityStore struct{} -func (clickhouseObservabilityStore) InsertMetricSnapshot(_ context.Context, record *OpenFlareMetricSnapshot) error { +func (clickhouseObservabilityStore) InsertMetricSnapshot(_ context.Context, record *model.OpenFlareMetricSnapshot) error { if record == nil { return nil } @@ -107,7 +109,7 @@ func (clickhouseObservabilityStore) InsertMetricSnapshot(_ context.Context, reco return nil } -func (clickhouseObservabilityStore) ListMetricSnapshots(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error) { +func (clickhouseObservabilityStore) ListMetricSnapshots(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareMetricSnapshot, error) { rows, err := analyticsrepo.ListNodeMetricSnapshots(ctx, toNodeObservabilityFilter(nodeID, since, limit)) if err != nil { return nil, err @@ -133,7 +135,7 @@ func normalizeEdgeHealthStatus(status string) string { return status } -func (clickhouseObservabilityStore) InsertEdgeHealth(_ context.Context, record *OpenFlareEdgeHealth) error { +func (clickhouseObservabilityStore) InsertEdgeHealth(_ context.Context, record *model.OpenFlareEdgeHealth) error { if record == nil { return nil } @@ -143,7 +145,7 @@ func (clickhouseObservabilityStore) InsertEdgeHealth(_ context.Context, record * return nil } -func (clickhouseObservabilityStore) ListEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareEdgeHealth, error) { +func (clickhouseObservabilityStore) ListEdgeHealth(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareEdgeHealth, error) { rows, err := analyticsrepo.ListNodeEdgeHealth(ctx, toNodeObservabilityFilter(nodeID, since, limit)) if err != nil { return nil, err @@ -159,7 +161,7 @@ func (clickhouseObservabilityStore) DeleteEdgeHealthBefore(ctx context.Context, return analyticsrepo.DeleteNodeEdgeHealthBefore(ctx, cutoff) } -func (clickhouseObservabilityStore) InsertNodeObservationFrps(_ context.Context, record *OpenFlareNodeObservationFrps) error { +func (clickhouseObservabilityStore) InsertNodeObservationFrps(_ context.Context, record *model.OpenFlareNodeObservationFrps) error { if record == nil { return nil } @@ -169,7 +171,7 @@ func (clickhouseObservabilityStore) InsertNodeObservationFrps(_ context.Context, return nil } -func (clickhouseObservabilityStore) ListNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error) { +func (clickhouseObservabilityStore) ListNodeObservationFrps(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareNodeObservationFrps, error) { rows, err := analyticsrepo.ListNodeObsFrps(ctx, toNodeObservabilityFilter(nodeID, since, limit)) if err != nil { return nil, err @@ -185,7 +187,7 @@ func (clickhouseObservabilityStore) DeleteNodeObservationFrpsBefore(ctx context. return analyticsrepo.DeleteNodeObsFrpsBefore(ctx, cutoff) } -func (clickhouseObservabilityStore) InsertNodeObservationFrpc(_ context.Context, record *OpenFlareNodeObservationFrpc) error { +func (clickhouseObservabilityStore) InsertNodeObservationFrpc(_ context.Context, record *model.OpenFlareNodeObservationFrpc) error { if record == nil { return nil } @@ -195,7 +197,7 @@ func (clickhouseObservabilityStore) InsertNodeObservationFrpc(_ context.Context, return nil } -func (clickhouseObservabilityStore) ListNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error) { +func (clickhouseObservabilityStore) ListNodeObservationFrpc(ctx context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareNodeObservationFrpc, error) { rows, err := analyticsrepo.ListNodeObsFrpc(ctx, toNodeObservabilityFilter(nodeID, since, limit)) if err != nil { return nil, err @@ -219,7 +221,7 @@ func toNodeObservabilityFilter(nodeID string, since time.Time, limit int) analyt } } -func toAnalyticsNodeMetricSnapshot(record *OpenFlareMetricSnapshot) analyticsmodel.NodeMetricSnapshot { +func toAnalyticsNodeMetricSnapshot(record *model.OpenFlareMetricSnapshot) analyticsmodel.NodeMetricSnapshot { return analyticsmodel.NodeMetricSnapshot{ ID: uint64(record.ID), NodeID: record.NodeID, @@ -237,10 +239,10 @@ func toAnalyticsNodeMetricSnapshot(record *OpenFlareMetricSnapshot) analyticsmod } } -func fromAnalyticsNodeMetricSnapshots(rows []analyticsmodel.NodeMetricSnapshot) []*OpenFlareMetricSnapshot { - result := make([]*OpenFlareMetricSnapshot, len(rows)) +func fromAnalyticsNodeMetricSnapshots(rows []analyticsmodel.NodeMetricSnapshot) []*model.OpenFlareMetricSnapshot { + result := make([]*model.OpenFlareMetricSnapshot, len(rows)) for index, row := range rows { - result[index] = &OpenFlareMetricSnapshot{ + result[index] = &model.OpenFlareMetricSnapshot{ ID: uint(row.ID), NodeID: row.NodeID, CapturedAt: row.CapturedAt, @@ -259,7 +261,7 @@ func fromAnalyticsNodeMetricSnapshots(rows []analyticsmodel.NodeMetricSnapshot) return result } -func toAnalyticsNodeEdgeHealth(record *OpenFlareEdgeHealth) analyticsmodel.NodeEdgeHealth { +func toAnalyticsNodeEdgeHealth(record *model.OpenFlareEdgeHealth) analyticsmodel.NodeEdgeHealth { return analyticsmodel.NodeEdgeHealth{ ID: uint64(record.ID), NodeID: record.NodeID, @@ -270,10 +272,10 @@ func toAnalyticsNodeEdgeHealth(record *OpenFlareEdgeHealth) analyticsmodel.NodeE } } -func fromAnalyticsNodeEdgeHealth(rows []analyticsmodel.NodeEdgeHealth) []*OpenFlareEdgeHealth { - result := make([]*OpenFlareEdgeHealth, len(rows)) +func fromAnalyticsNodeEdgeHealth(rows []analyticsmodel.NodeEdgeHealth) []*model.OpenFlareEdgeHealth { + result := make([]*model.OpenFlareEdgeHealth, len(rows)) for index, row := range rows { - result[index] = &OpenFlareEdgeHealth{ + result[index] = &model.OpenFlareEdgeHealth{ ID: uint(row.ID), NodeID: row.NodeID, CapturedAt: row.CapturedAt, @@ -285,7 +287,7 @@ func fromAnalyticsNodeEdgeHealth(rows []analyticsmodel.NodeEdgeHealth) []*OpenFl return result } -func toAnalyticsNodeObsFrps(record *OpenFlareNodeObservationFrps) analyticsmodel.NodeObsFrps { +func toAnalyticsNodeObsFrps(record *model.OpenFlareNodeObservationFrps) analyticsmodel.NodeObsFrps { return analyticsmodel.NodeObsFrps{ ID: uint64(record.ID), NodeID: record.NodeID, @@ -298,10 +300,10 @@ func toAnalyticsNodeObsFrps(record *OpenFlareNodeObservationFrps) analyticsmodel } } -func fromAnalyticsNodeObsFrps(rows []analyticsmodel.NodeObsFrps) []*OpenFlareNodeObservationFrps { - result := make([]*OpenFlareNodeObservationFrps, len(rows)) +func fromAnalyticsNodeObsFrps(rows []analyticsmodel.NodeObsFrps) []*model.OpenFlareNodeObservationFrps { + result := make([]*model.OpenFlareNodeObservationFrps, len(rows)) for index, row := range rows { - result[index] = &OpenFlareNodeObservationFrps{ + result[index] = &model.OpenFlareNodeObservationFrps{ ID: uint(row.ID), NodeID: row.NodeID, CapturedAt: row.CapturedAt, @@ -315,7 +317,7 @@ func fromAnalyticsNodeObsFrps(rows []analyticsmodel.NodeObsFrps) []*OpenFlareNod return result } -func toAnalyticsNodeObsFrpc(record *OpenFlareNodeObservationFrpc) analyticsmodel.NodeObsFrpc { +func toAnalyticsNodeObsFrpc(record *model.OpenFlareNodeObservationFrpc) analyticsmodel.NodeObsFrpc { return analyticsmodel.NodeObsFrpc{ ID: uint64(record.ID), NodeID: record.NodeID, @@ -337,10 +339,10 @@ func openFlareObservabilityIntToInt32(value int) int32 { } } -func fromAnalyticsNodeObsFrpc(rows []analyticsmodel.NodeObsFrpc) []*OpenFlareNodeObservationFrpc { - result := make([]*OpenFlareNodeObservationFrpc, len(rows)) +func fromAnalyticsNodeObsFrpc(rows []analyticsmodel.NodeObsFrpc) []*model.OpenFlareNodeObservationFrpc { + result := make([]*model.OpenFlareNodeObservationFrpc, len(rows)) for index, row := range rows { - result[index] = &OpenFlareNodeObservationFrpc{ + result[index] = &model.OpenFlareNodeObservationFrpc{ ID: uint(row.ID), NodeID: row.NodeID, CapturedAt: row.CapturedAt, diff --git a/internal/model/openflare_observability_store_memory.go b/internal/repository/openflare_observability_store_memory.go similarity index 76% rename from internal/model/openflare_observability_store_memory.go rename to internal/repository/openflare_observability_store_memory.go index c18d1456..7e8b28e3 100644 --- a/internal/model/openflare_observability_store_memory.go +++ b/internal/repository/openflare_observability_store_memory.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" @@ -10,18 +10,20 @@ import ( "sync" "time" + "github.com/Rain-kl/Wavelet/internal/model" + "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" ) type memoryObservabilityStore struct { mu sync.RWMutex - metricSnapshots []*OpenFlareMetricSnapshot - edgeHealth []*OpenFlareEdgeHealth - frpsObs []*OpenFlareNodeObservationFrps - frpcObs []*OpenFlareNodeObservationFrpc + metricSnapshots []*model.OpenFlareMetricSnapshot + edgeHealth []*model.OpenFlareEdgeHealth + frpsObs []*model.OpenFlareNodeObservationFrps + frpcObs []*model.OpenFlareNodeObservationFrpc } -func (s *memoryObservabilityStore) InsertMetricSnapshot(_ context.Context, record *OpenFlareMetricSnapshot) error { +func (s *memoryObservabilityStore) InsertMetricSnapshot(_ context.Context, record *model.OpenFlareMetricSnapshot) error { if record == nil { return nil } @@ -35,7 +37,7 @@ func (s *memoryObservabilityStore) InsertMetricSnapshot(_ context.Context, recor return nil } -func (s *memoryObservabilityStore) ListMetricSnapshots(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareMetricSnapshot, error) { +func (s *memoryObservabilityStore) ListMetricSnapshots(_ context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareMetricSnapshot, error) { s.mu.RLock() defer s.mu.RUnlock() rows := memoryFilterMetricSnapshots(s.metricSnapshots, nodeID, since) @@ -55,7 +57,7 @@ func (s *memoryObservabilityStore) DeleteMetricSnapshotsBefore(_ context.Context s.mu.Lock() defer s.mu.Unlock() cutoff = cutoff.UTC() - remaining := make([]*OpenFlareMetricSnapshot, 0, len(s.metricSnapshots)) + remaining := make([]*model.OpenFlareMetricSnapshot, 0, len(s.metricSnapshots)) var deleted int64 for _, row := range s.metricSnapshots { if row.CapturedAt.Before(cutoff) { @@ -68,7 +70,7 @@ func (s *memoryObservabilityStore) DeleteMetricSnapshotsBefore(_ context.Context return deleted, nil } -func (s *memoryObservabilityStore) InsertEdgeHealth(_ context.Context, record *OpenFlareEdgeHealth) error { +func (s *memoryObservabilityStore) InsertEdgeHealth(_ context.Context, record *model.OpenFlareEdgeHealth) error { if record == nil { return nil } @@ -78,7 +80,7 @@ func (s *memoryObservabilityStore) InsertEdgeHealth(_ context.Context, record *O return nil } -func (s *memoryObservabilityStore) ListEdgeHealth(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareEdgeHealth, error) { +func (s *memoryObservabilityStore) ListEdgeHealth(_ context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareEdgeHealth, error) { s.mu.RLock() defer s.mu.RUnlock() rows := memoryFilterEdgeHealth(s.edgeHealth, nodeID, since) @@ -98,7 +100,7 @@ func (s *memoryObservabilityStore) DeleteEdgeHealthBefore(_ context.Context, cut s.mu.Lock() defer s.mu.Unlock() cutoff = cutoff.UTC() - remaining := make([]*OpenFlareEdgeHealth, 0, len(s.edgeHealth)) + remaining := make([]*model.OpenFlareEdgeHealth, 0, len(s.edgeHealth)) var deleted int64 for _, row := range s.edgeHealth { if row.CapturedAt.Before(cutoff) { @@ -111,7 +113,7 @@ func (s *memoryObservabilityStore) DeleteEdgeHealthBefore(_ context.Context, cut return deleted, nil } -func (s *memoryObservabilityStore) InsertNodeObservationFrps(_ context.Context, record *OpenFlareNodeObservationFrps) error { +func (s *memoryObservabilityStore) InsertNodeObservationFrps(_ context.Context, record *model.OpenFlareNodeObservationFrps) error { if record == nil { return nil } @@ -121,7 +123,7 @@ func (s *memoryObservabilityStore) InsertNodeObservationFrps(_ context.Context, return nil } -func (s *memoryObservabilityStore) ListNodeObservationFrps(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrps, error) { +func (s *memoryObservabilityStore) ListNodeObservationFrps(_ context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareNodeObservationFrps, error) { s.mu.RLock() defer s.mu.RUnlock() rows := memoryFilterFrpsObservations(s.frpsObs, nodeID, since) @@ -141,7 +143,7 @@ func (s *memoryObservabilityStore) DeleteNodeObservationFrpsBefore(_ context.Con s.mu.Lock() defer s.mu.Unlock() cutoff = cutoff.UTC() - remaining := make([]*OpenFlareNodeObservationFrps, 0, len(s.frpsObs)) + remaining := make([]*model.OpenFlareNodeObservationFrps, 0, len(s.frpsObs)) var deleted int64 for _, row := range s.frpsObs { if row.CapturedAt.Before(cutoff) { @@ -154,7 +156,7 @@ func (s *memoryObservabilityStore) DeleteNodeObservationFrpsBefore(_ context.Con return deleted, nil } -func (s *memoryObservabilityStore) InsertNodeObservationFrpc(_ context.Context, record *OpenFlareNodeObservationFrpc) error { +func (s *memoryObservabilityStore) InsertNodeObservationFrpc(_ context.Context, record *model.OpenFlareNodeObservationFrpc) error { if record == nil { return nil } @@ -164,7 +166,7 @@ func (s *memoryObservabilityStore) InsertNodeObservationFrpc(_ context.Context, return nil } -func (s *memoryObservabilityStore) ListNodeObservationFrpc(_ context.Context, nodeID string, since time.Time, limit int) ([]*OpenFlareNodeObservationFrpc, error) { +func (s *memoryObservabilityStore) ListNodeObservationFrpc(_ context.Context, nodeID string, since time.Time, limit int) ([]*model.OpenFlareNodeObservationFrpc, error) { s.mu.RLock() defer s.mu.RUnlock() rows := memoryFilterFrpcObservations(s.frpcObs, nodeID, since) @@ -184,7 +186,7 @@ func (s *memoryObservabilityStore) DeleteNodeObservationFrpcBefore(_ context.Con s.mu.Lock() defer s.mu.Unlock() cutoff = cutoff.UTC() - remaining := make([]*OpenFlareNodeObservationFrpc, 0, len(s.frpcObs)) + remaining := make([]*model.OpenFlareNodeObservationFrpc, 0, len(s.frpcObs)) var deleted int64 for _, row := range s.frpcObs { if row.CapturedAt.Before(cutoff) { @@ -197,8 +199,8 @@ func (s *memoryObservabilityStore) DeleteNodeObservationFrpcBefore(_ context.Con return deleted, nil } -func memoryFilterMetricSnapshots(rows []*OpenFlareMetricSnapshot, nodeID string, since time.Time) []*OpenFlareMetricSnapshot { - result := make([]*OpenFlareMetricSnapshot, 0, len(rows)) +func memoryFilterMetricSnapshots(rows []*model.OpenFlareMetricSnapshot, nodeID string, since time.Time) []*model.OpenFlareMetricSnapshot { + result := make([]*model.OpenFlareMetricSnapshot, 0, len(rows)) for _, row := range rows { if !memoryObservabilityMatchesNodeID(row.NodeID, nodeID) { continue @@ -211,8 +213,8 @@ func memoryFilterMetricSnapshots(rows []*OpenFlareMetricSnapshot, nodeID string, return result } -func memoryFilterEdgeHealth(rows []*OpenFlareEdgeHealth, nodeID string, since time.Time) []*OpenFlareEdgeHealth { - result := make([]*OpenFlareEdgeHealth, 0, len(rows)) +func memoryFilterEdgeHealth(rows []*model.OpenFlareEdgeHealth, nodeID string, since time.Time) []*model.OpenFlareEdgeHealth { + result := make([]*model.OpenFlareEdgeHealth, 0, len(rows)) for _, row := range rows { if !memoryObservabilityMatchesNodeID(row.NodeID, nodeID) { continue @@ -225,8 +227,8 @@ func memoryFilterEdgeHealth(rows []*OpenFlareEdgeHealth, nodeID string, since ti return result } -func memoryFilterFrpsObservations(rows []*OpenFlareNodeObservationFrps, nodeID string, since time.Time) []*OpenFlareNodeObservationFrps { - result := make([]*OpenFlareNodeObservationFrps, 0, len(rows)) +func memoryFilterFrpsObservations(rows []*model.OpenFlareNodeObservationFrps, nodeID string, since time.Time) []*model.OpenFlareNodeObservationFrps { + result := make([]*model.OpenFlareNodeObservationFrps, 0, len(rows)) for _, row := range rows { if !memoryObservabilityMatchesNodeID(row.NodeID, nodeID) { continue @@ -239,8 +241,8 @@ func memoryFilterFrpsObservations(rows []*OpenFlareNodeObservationFrps, nodeID s return result } -func memoryFilterFrpcObservations(rows []*OpenFlareNodeObservationFrpc, nodeID string, since time.Time) []*OpenFlareNodeObservationFrpc { - result := make([]*OpenFlareNodeObservationFrpc, 0, len(rows)) +func memoryFilterFrpcObservations(rows []*model.OpenFlareNodeObservationFrpc, nodeID string, since time.Time) []*model.OpenFlareNodeObservationFrpc { + result := make([]*model.OpenFlareNodeObservationFrpc, 0, len(rows)) for _, row := range rows { if !memoryObservabilityMatchesNodeID(row.NodeID, nodeID) { continue @@ -261,7 +263,7 @@ func memoryObservabilityMatchesNodeID(rowNodeID string, nodeID string) bool { return rowNodeID == trimmed } -func memoryMetricSnapshotExists(rows []*OpenFlareMetricSnapshot, nodeID string, capturedAt time.Time) bool { +func memoryMetricSnapshotExists(rows []*model.OpenFlareMetricSnapshot, nodeID string, capturedAt time.Time) bool { capturedAt = capturedAt.UTC() for _, row := range rows { if row.NodeID == nodeID && row.CapturedAt.UTC().Equal(capturedAt) { @@ -271,7 +273,7 @@ func memoryMetricSnapshotExists(rows []*OpenFlareMetricSnapshot, nodeID string, return false } -func sortOpenFlareMetricSnapshots(items []*OpenFlareMetricSnapshot) { +func sortOpenFlareMetricSnapshots(items []*model.OpenFlareMetricSnapshot) { sort.Slice(items, func(i, j int) bool { left := items[i] right := items[j] @@ -285,7 +287,7 @@ func sortOpenFlareMetricSnapshots(items []*OpenFlareMetricSnapshot) { }) } -func sortOpenFlareEdgeHealth(items []*OpenFlareEdgeHealth) { +func sortOpenFlareEdgeHealth(items []*model.OpenFlareEdgeHealth) { sort.Slice(items, func(i, j int) bool { left := items[i] right := items[j] @@ -299,7 +301,7 @@ func sortOpenFlareEdgeHealth(items []*OpenFlareEdgeHealth) { }) } -func sortOpenFlareNodeObservationFrps(items []*OpenFlareNodeObservationFrps) { +func sortOpenFlareNodeObservationFrps(items []*model.OpenFlareNodeObservationFrps) { sort.Slice(items, func(i, j int) bool { left := items[i] right := items[j] @@ -313,7 +315,7 @@ func sortOpenFlareNodeObservationFrps(items []*OpenFlareNodeObservationFrps) { }) } -func sortOpenFlareNodeObservationFrpc(items []*OpenFlareNodeObservationFrpc) { +func sortOpenFlareNodeObservationFrpc(items []*model.OpenFlareNodeObservationFrpc) { sort.Slice(items, func(i, j int) bool { left := items[i] right := items[j] @@ -338,7 +340,7 @@ func memoryLimitObservabilityRows[T any](rows []T, limit int) []T { return result } -func cloneOpenFlareMetricSnapshot(record *OpenFlareMetricSnapshot) *OpenFlareMetricSnapshot { +func cloneOpenFlareMetricSnapshot(record *model.OpenFlareMetricSnapshot) *model.OpenFlareMetricSnapshot { copyRecord := *record if copyRecord.ID == 0 { copyRecord.ID = uint(idgen.NextUint64ID()) @@ -352,7 +354,7 @@ func cloneOpenFlareMetricSnapshot(record *OpenFlareMetricSnapshot) *OpenFlareMet return ©Record } -func cloneOpenFlareEdgeHealth(record *OpenFlareEdgeHealth) *OpenFlareEdgeHealth { +func cloneOpenFlareEdgeHealth(record *model.OpenFlareEdgeHealth) *model.OpenFlareEdgeHealth { copyRecord := *record if copyRecord.ID == 0 { copyRecord.ID = uint(idgen.NextUint64ID()) @@ -372,7 +374,7 @@ func cloneOpenFlareEdgeHealth(record *OpenFlareEdgeHealth) *OpenFlareEdgeHealth return ©Record } -func cloneOpenFlareNodeObservationFrps(record *OpenFlareNodeObservationFrps) *OpenFlareNodeObservationFrps { +func cloneOpenFlareNodeObservationFrps(record *model.OpenFlareNodeObservationFrps) *model.OpenFlareNodeObservationFrps { copyRecord := *record if copyRecord.ID == 0 { copyRecord.ID = uint(idgen.NextUint64ID()) @@ -389,7 +391,7 @@ func cloneOpenFlareNodeObservationFrps(record *OpenFlareNodeObservationFrps) *Op return ©Record } -func cloneOpenFlareNodeObservationFrpc(record *OpenFlareNodeObservationFrpc) *OpenFlareNodeObservationFrpc { +func cloneOpenFlareNodeObservationFrpc(record *model.OpenFlareNodeObservationFrpc) *model.OpenFlareNodeObservationFrpc { copyRecord := *record if copyRecord.ID == 0 { copyRecord.ID = uint(idgen.NextUint64ID()) diff --git a/internal/repository/openflare_origin.go b/internal/repository/openflare_origin.go new file mode 100644 index 00000000..634f7ee5 --- /dev/null +++ b/internal/repository/openflare_origin.go @@ -0,0 +1,127 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + + "gorm.io/gorm" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// WithOriginTx runs fn inside a database transaction for origin multi-step work. +func WithOriginTx(ctx context.Context, fn func(tx *gorm.DB) error) error { + return db.DB(ctx).Transaction(fn) +} + +// HasProxyRoutesTable 判断代理规则表是否已迁移。 +func HasProxyRoutesTable(ctx context.Context) bool { + return db.DB(ctx).Migrator().HasTable(&model.OriginProxyRoute{}) +} + +// ListOrigins 列出全部源站。 +func ListOrigins(ctx context.Context) ([]model.Origin, error) { + var origins []model.Origin + if err := db.DB(ctx).Order("id desc").Find(&origins).Error; err != nil { + return nil, err + } + return origins, nil +} + +// GetOriginByID 按 ID 查询源站。 +func GetOriginByID(ctx context.Context, id uint) (*model.Origin, error) { + var origin model.Origin + if err := db.DB(ctx).First(&origin, id).Error; err != nil { + return nil, err + } + return &origin, nil +} + +// GetOriginByAddress 按地址查询源站。 +func GetOriginByAddress(ctx context.Context, address string) (*model.Origin, error) { + var origin model.Origin + if err := db.DB(ctx).Where("address = ?", address).First(&origin).Error; err != nil { + return nil, err + } + return &origin, nil +} + +// CreateOriginRecord 创建源站。 +func CreateOriginRecord(ctx context.Context, origin *model.Origin) error { + return db.DB(ctx).Create(origin).Error +} + +// SaveOrigin 保存源站。 +func SaveOrigin(ctx context.Context, origin *model.Origin) error { + return SaveOriginTx(db.DB(ctx), origin) +} + +// SaveOriginTx saves an origin within an existing transaction. +func SaveOriginTx(tx *gorm.DB, origin *model.Origin) error { + return tx.Save(origin).Error +} + +// DeleteOriginRecord 删除源站。 +func DeleteOriginRecord(ctx context.Context, id uint) error { + return db.DB(ctx).Delete(&model.Origin{}, id).Error +} + +// ListOriginRouteCounts 统计各源站关联的代理规则数量。 +func ListOriginRouteCounts(ctx context.Context) ([]model.OriginRouteCount, error) { + if !HasProxyRoutesTable(ctx) { + return nil, nil + } + result := make([]model.OriginRouteCount, 0) + err := db.DB(ctx).Model(&model.OriginProxyRoute{}). + Select("origin_id, COUNT(*) AS route_count"). + Where("origin_id IS NOT NULL"). + Group("origin_id"). + Scan(&result).Error + return result, err +} + +// ListProxyRoutesByOriginID 列出源站关联的代理规则。 +func ListProxyRoutesByOriginID(ctx context.Context, originID uint) ([]model.OriginProxyRoute, error) { + if !HasProxyRoutesTable(ctx) { + return nil, nil + } + var routes []model.OriginProxyRoute + if err := db.DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil { + return nil, err + } + return routes, nil +} + +// ListProxyRoutesByOriginIDAscTx lists origin-linked proxy routes ordered by id asc within a transaction. +func ListProxyRoutesByOriginIDAscTx(tx *gorm.DB, originID uint) ([]model.OriginProxyRoute, error) { + var routes []model.OriginProxyRoute + if err := tx.Where("origin_id = ?", originID).Order("id asc").Find(&routes).Error; err != nil { + return nil, err + } + return routes, nil +} + +// UpdateProxyRouteOriginAddressTx updates a proxy route's origin_url and upstreams within a transaction. +func UpdateProxyRouteOriginAddressTx(tx *gorm.DB, routeID uint, originURL, upstreamsJSON string) error { + return tx.Model(&model.OriginProxyRoute{}). + Where("id = ?", routeID). + Updates(map[string]any{ + "origin_url": originURL, + "upstreams": upstreamsJSON, + }).Error +} + +// CountProxyRoutesByOriginID 统计源站关联的代理规则数量。 +func CountProxyRoutesByOriginID(ctx context.Context, originID uint) (int64, error) { + if !HasProxyRoutesTable(ctx) { + return 0, nil + } + var count int64 + if err := db.DB(ctx).Model(&model.OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} diff --git a/internal/repository/openflare_pages.go b/internal/repository/openflare_pages.go new file mode 100644 index 00000000..250bbef8 --- /dev/null +++ b/internal/repository/openflare_pages.go @@ -0,0 +1,96 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// HasPagesProjectsTable 判断 Pages 项目表是否已迁移。 +func HasPagesProjectsTable(ctx context.Context) bool { + return db.DB(ctx).Migrator().HasTable(&model.PagesProject{}) +} + +// ListPagesProjects 列出全部 Pages 项目。 +func ListPagesProjects(ctx context.Context) ([]model.PagesProject, error) { + var projects []model.PagesProject + if err := db.DB(ctx).Order("id desc").Find(&projects).Error; err != nil { + return nil, err + } + return projects, nil +} + +// GetPagesProjectByID 按 ID 查询 Pages 项目。 +func GetPagesProjectByID(ctx context.Context, id uint) (*model.PagesProject, error) { + var project model.PagesProject + if err := db.DB(ctx).First(&project, id).Error; err != nil { + return nil, err + } + return &project, nil +} + +// GetPagesProjectBySlug 按 slug 查询 Pages 项目。 +func GetPagesProjectBySlug(ctx context.Context, slug string) (*model.PagesProject, error) { + var project model.PagesProject + if err := db.DB(ctx).Where("slug = ?", slug).First(&project).Error; err != nil { + return nil, err + } + return &project, nil +} + +// CreatePagesProjectRecord 创建 Pages 项目。 +func CreatePagesProjectRecord(ctx context.Context, project *model.PagesProject) error { + return db.DB(ctx).Create(project).Error +} + +// ListPagesDeployments 列出项目的全部部署。 +func ListPagesDeployments(ctx context.Context, projectID uint) ([]model.PagesDeployment, error) { + var deployments []model.PagesDeployment + if err := db.DB(ctx).Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil { + return nil, err + } + return deployments, nil +} + +// GetPagesDeploymentByID 按 ID 查询 Pages 部署。 +func GetPagesDeploymentByID(ctx context.Context, id uint) (*model.PagesDeployment, error) { + var deployment model.PagesDeployment + if err := db.DB(ctx).First(&deployment, id).Error; err != nil { + return nil, err + } + return &deployment, nil +} + +// ListPagesDeploymentFiles 列出部署文件清单。 +func ListPagesDeploymentFiles(ctx context.Context, deploymentID uint) ([]model.PagesDeploymentFile, error) { + var files []model.PagesDeploymentFile + if err := db.DB(ctx).Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error; err != nil { + return nil, err + } + return files, nil +} + +// CountPagesDeploymentsByProjectID 统计项目部署数量。 +func CountPagesDeploymentsByProjectID(ctx context.Context, projectID uint) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", projectID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// CountProxyRoutesByPagesProjectID 统计引用 Pages 项目的代理规则数量。 +func CountProxyRoutesByPagesProjectID(ctx context.Context, projectID uint) (int64, error) { + if !HasProxyRoutesTable(ctx) { + return 0, nil + } + var count int64 + if err := db.DB(ctx).Model(&model.ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} diff --git a/internal/repository/openflare_pages_cleanup.go b/internal/repository/openflare_pages_cleanup.go new file mode 100644 index 00000000..21e28851 --- /dev/null +++ b/internal/repository/openflare_pages_cleanup.go @@ -0,0 +1,58 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// ListPagesOrphanUploadCandidates returns at most 100 unreferenced, isolated +// Pages V2 upload records. Callers must still lock and recheck every condition +// before deleting a candidate. +func ListPagesOrphanUploadCandidates( + ctx context.Context, + input model.PagesOrphanUploadCandidateQuery, +) ([]model.Upload, error) { + if input.SystemUserID == 0 || input.UploadType == "" || input.Marker == "" || input.CreatedBefore.IsZero() { + return nil, errors.New("invalid pages orphan upload candidate query") + } + markerPredicate, err := pagesOrphanMarkerPredicate(db.DB(ctx).Name()) + if err != nil { + return nil, err + } + + deploymentTable := (model.PagesDeployment{}).TableName() + uploadTable := (model.Upload{}).TableName() + var candidates []model.Upload + err = db.DB(ctx). + Model(&model.Upload{}). + Where(uploadTable+".status = ?", model.UploadStatusUsed). + Where(uploadTable+".user_id = ?", input.SystemUserID). + Where(uploadTable+".type = ?", input.UploadType). + Where(uploadTable+".created_at < ?", input.CreatedBefore). + Where(markerPredicate, input.Marker). + Where("NOT EXISTS (SELECT 1 FROM " + deploymentTable + " WHERE " + deploymentTable + ".upload_id = " + uploadTable + ".id)"). + Order(uploadTable + ".id ASC"). + Limit(model.PagesOrphanUploadCandidateLimit). + Find(&candidates).Error + if err != nil { + return nil, err + } + return candidates, nil +} + +func pagesOrphanMarkerPredicate(dialect string) (string, error) { + switch dialect { + case "postgres": + return model.PagesOrphanMarkerPredicatePostgres, nil + case "sqlite": + return model.PagesOrphanMarkerPredicateSQLite, nil + default: + return "", errors.New("unsupported database dialect for Pages orphan cleanup") + } +} diff --git a/internal/model/openflare_pages_cleanup_test.go b/internal/repository/openflare_pages_cleanup_test.go similarity index 74% rename from internal/model/openflare_pages_cleanup_test.go rename to internal/repository/openflare_pages_cleanup_test.go index ccf38e91..10309be4 100644 --- a/internal/model/openflare_pages_cleanup_test.go +++ b/internal/repository/openflare_pages_cleanup_test.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" @@ -9,6 +9,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/model" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/glebarez/sqlite" "gorm.io/gorm" @@ -56,53 +58,53 @@ func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) { gormDB := setupPagesCleanupModelTestDB(t) cutoff := time.Now().UTC().Add(-2 * time.Hour) old := cutoff.Add(-time.Minute) - marker := UploadMetadata{Extra: map[string]any{ + marker := model.UploadMetadata{Extra: map[string]any{ "pages_ingest_marker": "pages_deployment_v2", "pages_project_id": "1", }} - valid := make([]Upload, 0, PagesOrphanUploadCandidateLimit+1) - for index := 0; index < PagesOrphanUploadCandidateLimit+1; index++ { - valid = append(valid, pagesCleanupModelUpload(uint64(index+100), 999, "openflare_pages_deployment", UploadStatusUsed, old, marker)) + valid := make([]model.Upload, 0, model.PagesOrphanUploadCandidateLimit+1) + for index := 0; index < model.PagesOrphanUploadCandidateLimit+1; index++ { + valid = append(valid, pagesCleanupModelUpload(uint64(index+100), 999, "openflare_pages_deployment", model.UploadStatusUsed, old, marker)) } if err := gormDB.Create(&valid).Error; err != nil { t.Fatalf("create valid candidates error = %v, want nil", err) } - referenced := pagesCleanupModelUpload(1, 999, "openflare_pages_deployment", UploadStatusUsed, old, marker) - wrongOwner := pagesCleanupModelUpload(2, 1000, "openflare_pages_deployment", UploadStatusUsed, old, marker) - wrongType := pagesCleanupModelUpload(3, 999, "generic", UploadStatusUsed, old, marker) - wrongStatus := pagesCleanupModelUpload(4, 999, "openflare_pages_deployment", UploadStatusPending, old, marker) - fresh := pagesCleanupModelUpload(5, 999, "openflare_pages_deployment", UploadStatusUsed, cutoff, marker) - wrongMarker := pagesCleanupModelUpload(6, 999, "openflare_pages_deployment", UploadStatusUsed, old, UploadMetadata{Extra: map[string]any{ + referenced := pagesCleanupModelUpload(1, 999, "openflare_pages_deployment", model.UploadStatusUsed, old, marker) + wrongOwner := pagesCleanupModelUpload(2, 1000, "openflare_pages_deployment", model.UploadStatusUsed, old, marker) + wrongType := pagesCleanupModelUpload(3, 999, "generic", model.UploadStatusUsed, old, marker) + wrongStatus := pagesCleanupModelUpload(4, 999, "openflare_pages_deployment", model.UploadStatusPending, old, marker) + fresh := pagesCleanupModelUpload(5, 999, "openflare_pages_deployment", model.UploadStatusUsed, cutoff, marker) + wrongMarker := pagesCleanupModelUpload(6, 999, "openflare_pages_deployment", model.UploadStatusUsed, old, model.UploadMetadata{Extra: map[string]any{ "pages_ingest_marker": "pages_deployment_v1", "pages_project_id": "1", }}) - for _, upload := range []Upload{referenced, wrongOwner, wrongType, wrongStatus, fresh, wrongMarker} { + for _, upload := range []model.Upload{referenced, wrongOwner, wrongType, wrongStatus, fresh, wrongMarker} { if err := gormDB.Create(&upload).Error; err != nil { t.Fatalf("create filtered upload %d error = %v, want nil", upload.ID, err) } } - if err := gormDB.Create(&PagesDeployment{ + if err := gormDB.Create(&model.PagesDeployment{ ProjectID: 1, DeploymentNumber: 1, Checksum: "referenced", - Status: PagesDeploymentStatusUploaded, + Status: model.PagesDeploymentStatusUploaded, UploadID: referenced.ID, }).Error; err != nil { t.Fatalf("create referenced deployment error = %v, want nil", err) } - invalidJSON := pagesCleanupModelUpload(7, 999, "openflare_pages_deployment", UploadStatusUsed, old, marker) + invalidJSON := pagesCleanupModelUpload(7, 999, "openflare_pages_deployment", model.UploadStatusUsed, old, marker) if err := gormDB.Create(&invalidJSON).Error; err != nil { t.Fatalf("create invalid JSON upload error = %v, want nil", err) } - if err := gormDB.Table((Upload{}).TableName()).Where("id = ?", invalidJSON.ID). + if err := gormDB.Table((model.Upload{}).TableName()).Where("id = ?", invalidJSON.ID). UpdateColumn("metadata", "{invalid").Error; err != nil { t.Fatalf("corrupt upload metadata error = %v, want nil", err) } - got, err := ListPagesOrphanUploadCandidates(ctx, PagesOrphanUploadCandidateQuery{ + got, err := ListPagesOrphanUploadCandidates(ctx, model.PagesOrphanUploadCandidateQuery{ SystemUserID: 999, UploadType: "openflare_pages_deployment", Marker: "pages_deployment_v2", @@ -111,8 +113,8 @@ func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) { if err != nil { t.Fatalf("ListPagesOrphanUploadCandidates() error = %v, want nil", err) } - if len(got) != PagesOrphanUploadCandidateLimit { - t.Fatalf("ListPagesOrphanUploadCandidates() count = %d, want %d", len(got), PagesOrphanUploadCandidateLimit) + if len(got) != model.PagesOrphanUploadCandidateLimit { + t.Fatalf("ListPagesOrphanUploadCandidates() count = %d, want %d", len(got), model.PagesOrphanUploadCandidateLimit) } for index, candidate := range got { wantID := uint64(index + 100) @@ -126,16 +128,16 @@ func TestListPagesOrphanUploadCandidatesSkipsInvalidSQLiteJSON(t *testing.T) { ctx := context.Background() gormDB := setupPagesCleanupModelTestDB(t) cutoff := time.Now().UTC().Add(-2 * time.Hour) - upload := pagesCleanupModelUpload(1, 999, "openflare_pages_deployment", UploadStatusUsed, cutoff.Add(-time.Minute), UploadMetadata{}) + upload := pagesCleanupModelUpload(1, 999, "openflare_pages_deployment", model.UploadStatusUsed, cutoff.Add(-time.Minute), model.UploadMetadata{}) if err := gormDB.Create(&upload).Error; err != nil { t.Fatalf("create invalid JSON candidate error = %v, want nil", err) } - if err := gormDB.Table((Upload{}).TableName()).Where("id = ?", upload.ID). + if err := gormDB.Table((model.Upload{}).TableName()).Where("id = ?", upload.ID). UpdateColumn("metadata", "{invalid").Error; err != nil { t.Fatalf("corrupt upload metadata error = %v, want nil", err) } - got, err := ListPagesOrphanUploadCandidates(ctx, PagesOrphanUploadCandidateQuery{ + got, err := ListPagesOrphanUploadCandidates(ctx, model.PagesOrphanUploadCandidateQuery{ SystemUserID: 999, UploadType: "openflare_pages_deployment", Marker: "pages_deployment_v2", @@ -158,7 +160,7 @@ func setupPagesCleanupModelTestDB(t *testing.T) *gorm.DB { if err != nil { t.Fatalf("open Pages cleanup model test database error = %v, want nil", err) } - if err := gormDB.AutoMigrate(&Upload{}, &PagesDeployment{}); err != nil { + if err := gormDB.AutoMigrate(&model.Upload{}, &model.PagesDeployment{}); err != nil { t.Fatalf("migrate Pages cleanup model test database error = %v, want nil", err) } db.SetDB(gormDB) @@ -170,11 +172,11 @@ func pagesCleanupModelUpload( id uint64, userID uint64, uploadType string, - status UploadStatus, + status model.UploadStatus, createdAt time.Time, - metadata UploadMetadata, -) Upload { - return Upload{ + metadata model.UploadMetadata, +) model.Upload { + return model.Upload{ ID: id, UserID: userID, FileName: "site.zip", diff --git a/internal/repository/openflare_pages_source.go b/internal/repository/openflare_pages_source.go new file mode 100644 index 00000000..74ca2370 --- /dev/null +++ b/internal/repository/openflare_pages_source.go @@ -0,0 +1,417 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "time" + + "gorm.io/gorm" + "gorm.io/gorm/clause" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +const pagesRowLockStrength = "UPDATE" + +// WithPagesTx runs fn inside a database transaction for Pages multi-step work. +func WithPagesTx(ctx context.Context, fn func(tx *gorm.DB) error) error { + return db.DB(ctx).Transaction(fn) +} + +// GetPagesProjectSourceByID loads a project source by primary key. +func GetPagesProjectSourceByID(ctx context.Context, id uint) (*model.PagesProjectSource, error) { + var source model.PagesProjectSource + if err := db.DB(ctx).Where("id = ?", id).First(&source).Error; err != nil { + return nil, err + } + return &source, nil +} + +// GetPagesProjectSourceByProjectID loads the unique source for a project. +func GetPagesProjectSourceByProjectID(ctx context.Context, projectID uint) (*model.PagesProjectSource, error) { + var source model.PagesProjectSource + if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { + return nil, err + } + return &source, nil +} + +// GetPagesProjectSourceByIDAndConfigVersion loads a source matching both id and config version. +func GetPagesProjectSourceByIDAndConfigVersion( + ctx context.Context, + id uint, + configVersion int, +) (*model.PagesProjectSource, error) { + var source model.PagesProjectSource + if err := db.DB(ctx).Where("id = ? AND config_version = ?", id, configVersion).First(&source).Error; err != nil { + return nil, err + } + return &source, nil +} + +// GetPagesProjectSourceRuntimeBySourceID loads runtime for a source. +func GetPagesProjectSourceRuntimeBySourceID( + ctx context.Context, + sourceID uint, +) (*model.PagesProjectSourceRuntime, error) { + var runtime model.PagesProjectSourceRuntime + if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { + return nil, err + } + return &runtime, nil +} + +// GetPagesProjectSourceAndRuntimeByProjectID loads source and its runtime for a project. +func GetPagesProjectSourceAndRuntimeByProjectID( + ctx context.Context, + projectID uint, +) (*model.PagesProjectSource, *model.PagesProjectSourceRuntime, error) { + source, err := GetPagesProjectSourceByProjectID(ctx, projectID) + if err != nil { + return nil, nil, err + } + runtime, err := GetPagesProjectSourceRuntimeBySourceID(ctx, source.ID) + if err != nil { + return nil, nil, err + } + return source, runtime, nil +} + +// CreatePagesProjectSourceTx creates a source row inside an existing transaction. +func CreatePagesProjectSourceTx(tx *gorm.DB, source *model.PagesProjectSource) error { + return tx.Create(source).Error +} + +// CreatePagesProjectSourceRuntimeTx creates a runtime row inside an existing transaction. +func CreatePagesProjectSourceRuntimeTx(tx *gorm.DB, runtime *model.PagesProjectSourceRuntime) error { + return tx.Create(runtime).Error +} + +// UpdatePagesProjectSourceTx applies partial updates to a source inside a transaction. +func UpdatePagesProjectSourceTx(tx *gorm.DB, source *model.PagesProjectSource, updates map[string]any) error { + if len(updates) == 0 { + return nil + } + return tx.Model(source).Updates(updates).Error +} + +// UpdatePagesProjectSourceRuntimeTx applies partial updates to a runtime inside a transaction. +func UpdatePagesProjectSourceRuntimeTx( + tx *gorm.DB, + runtime *model.PagesProjectSourceRuntime, + updates map[string]any, +) error { + if len(updates) == 0 { + return nil + } + return tx.Model(runtime).Updates(updates).Error +} + +// UpdatePagesProjectSourceRuntimeFieldTx updates a single column on a runtime row. +func UpdatePagesProjectSourceRuntimeFieldTx( + tx *gorm.DB, + runtime *model.PagesProjectSourceRuntime, + column string, + value any, +) error { + return tx.Model(runtime).Update(column, value).Error +} + +// DeletePagesProjectSourceRuntimeBySourceIDTx deletes runtime rows for a source. +func DeletePagesProjectSourceRuntimeBySourceIDTx(tx *gorm.DB, sourceID uint) error { + return tx.Where("source_id = ?", sourceID).Delete(&model.PagesProjectSourceRuntime{}).Error +} + +// DeletePagesProjectSourceTx deletes a source row inside a transaction. +func DeletePagesProjectSourceTx(tx *gorm.DB, source *model.PagesProjectSource) error { + return tx.Delete(source).Error +} + +// LockPagesProjectByIDTx locks a project row for update. +func LockPagesProjectByIDTx(tx *gorm.DB, id uint) (*model.PagesProject, error) { + var project model.PagesProject + if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, id).Error; err != nil { + return nil, err + } + return &project, nil +} + +// LockPagesProjectSourceByProjectIDTx locks the source for a project. +func LockPagesProjectSourceByProjectIDTx(tx *gorm.DB, projectID uint) (*model.PagesProjectSource, error) { + var source model.PagesProjectSource + if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). + Where("project_id = ?", projectID). + First(&source).Error; err != nil { + return nil, err + } + return &source, nil +} + +// LockPagesProjectSourceByIDTx locks a source by id. +func LockPagesProjectSourceByIDTx(tx *gorm.DB, sourceID uint) (*model.PagesProjectSource, error) { + var source model.PagesProjectSource + if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). + Where("id = ?", sourceID). + First(&source).Error; err != nil { + return nil, err + } + return &source, nil +} + +// LockPagesProjectSourceByIDAndProjectIDTx locks a source matching both identifiers. +func LockPagesProjectSourceByIDAndProjectIDTx( + tx *gorm.DB, + sourceID uint, + projectID uint, +) (*model.PagesProjectSource, error) { + var source model.PagesProjectSource + if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). + Where("id = ? AND project_id = ?", sourceID, projectID). + First(&source).Error; err != nil { + return nil, err + } + return &source, nil +} + +// LockPagesProjectSourceRuntimeBySourceIDTx locks runtime for a source. +func LockPagesProjectSourceRuntimeBySourceIDTx( + tx *gorm.DB, + sourceID uint, +) (*model.PagesProjectSourceRuntime, error) { + var runtime model.PagesProjectSourceRuntime + if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). + Where("source_id = ?", sourceID). + First(&runtime).Error; err != nil { + return nil, err + } + return &runtime, nil +} + +// GetPagesProjectSourceByIDTx loads a source by id without locking. +func GetPagesProjectSourceByIDTx(tx *gorm.DB, sourceID uint) (*model.PagesProjectSource, error) { + var source model.PagesProjectSource + if err := tx.Where("id = ?", sourceID).First(&source).Error; err != nil { + return nil, err + } + return &source, nil +} + +// TryAcquirePagesSourceRuntimeLease conditionally claims an idle/expired lease when config matches. +func TryAcquirePagesSourceRuntimeLease( + ctx context.Context, + sourceID uint, + expectedConfigVersion int, + now time.Time, + updates map[string]any, +) (int64, error) { + result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + Where("source_id = ?", sourceID). + Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). + Where( + "EXISTS (SELECT 1 FROM of_pages_project_sources source WHERE source.id = ? AND source.config_version = ?)", + sourceID, + expectedConfigVersion, + ). + Updates(updates) + return result.RowsAffected, result.Error +} + +// RenewPagesSourceRuntimeLease extends an active lease held by the given token. +func RenewPagesSourceRuntimeLease( + ctx context.Context, + sourceID uint, + token string, + now time.Time, + expiresAt time.Time, +) (int64, error) { + result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", sourceID, token, now). + Updates(map[string]any{"lease_expires_at": expiresAt}) + return result.RowsAffected, result.Error +} + +// UpdatePagesSourceRuntimeByActiveLease updates runtime while the caller still owns the lease. +func UpdatePagesSourceRuntimeByActiveLease( + ctx context.Context, + sourceID uint, + token string, + now time.Time, + updates map[string]any, +) (int64, error) { + return UpdatePagesSourceRuntimeByActiveLeaseTx(db.DB(ctx), sourceID, token, now, updates) +} + +// UpdatePagesSourceRuntimeByActiveLeaseTx updates runtime under an active lease inside a transaction. +func UpdatePagesSourceRuntimeByActiveLeaseTx( + tx *gorm.DB, + sourceID uint, + token string, + now time.Time, + updates map[string]any, +) (int64, error) { + result := tx.Model(&model.PagesProjectSourceRuntime{}). + Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", sourceID, token, now). + Updates(updates) + return result.RowsAffected, result.Error +} + +// RecoverExpiredPagesSourceRuntimeLease clears one exact expired lease owner. +func RecoverExpiredPagesSourceRuntimeLease( + ctx context.Context, + sourceID uint, + token string, + expiresAt time.Time, + status string, + now time.Time, + updates map[string]any, +) (int64, error) { + result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + Where("source_id = ?", sourceID). + Where("lease_token = ?", token). + Where("lease_expires_at = ? AND lease_expires_at <= ?", expiresAt, now). + Where("sync_status = ?", status). + Updates(updates) + return result.RowsAffected, result.Error +} + +// MarkPagesSourceInitialCheckDispatchFailed marks runtime failed when config still matches and lease is free. +func MarkPagesSourceInitialCheckDispatchFailed( + ctx context.Context, + sourceID uint, + configVersion int, + now time.Time, + updates map[string]any, +) (int64, error) { + result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + Where("source_id = ?", sourceID). + Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). + Where( + "EXISTS (SELECT 1 FROM of_pages_project_sources source WHERE source.id = ? AND source.config_version = ?)", + sourceID, + configVersion, + ). + Updates(updates) + return result.RowsAffected, result.Error +} + +// RecordPagesSourceAutoDispatchFailure records a failed auto-sync dispatch while status still matches. +func RecordPagesSourceAutoDispatchFailure( + ctx context.Context, + sourceID uint, + configVersion int, + sourceType string, + releaseSelector string, + revision string, + updateAvailableStatus string, + now time.Time, + updates map[string]any, +) (int64, error) { + result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + Where("source_id = ?", sourceID). + Where("sync_status = ? AND last_seen_revision = ?", updateAvailableStatus, revision). + Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). + Where(`EXISTS ( + SELECT 1 FROM of_pages_project_sources AS source + WHERE source.id = ? AND source.config_version = ? + AND source.source_type = ? AND source.release_selector = ? + AND source.auto_update_enabled = ? + )`, + sourceID, + configVersion, + sourceType, + releaseSelector, + true, + ). + Updates(updates) + return result.RowsAffected, result.Error +} + +// ListExpiredPagesSourceLeaseCandidates returns expired checking/syncing leases for recovery. +func ListExpiredPagesSourceLeaseCandidates( + ctx context.Context, + now time.Time, + syncStatuses []string, +) ([]model.PagesExpiredSourceLeaseCandidate, error) { + var candidates []model.PagesExpiredSourceLeaseCandidate + err := db.DB(ctx). + Table("of_pages_project_source_runtime AS runtime"). + Select(`runtime.source_id, runtime.lease_token, runtime.lease_expires_at, + runtime.sync_status, source.source_type, source.release_selector`). + Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id"). + Where("runtime.lease_token <> ''"). + Where("runtime.lease_expires_at IS NOT NULL AND runtime.lease_expires_at <= ?", now). + Where("runtime.sync_status IN ?", syncStatuses). + Order("runtime.source_id ASC"). + Scan(&candidates).Error + if err != nil { + return nil, err + } + return candidates, nil +} + +// CountDueGitHubPagesSourceChecks counts due latest GitHub sources. +func CountDueGitHubPagesSourceChecks( + ctx context.Context, + now time.Time, + sourceType string, + releaseSelector string, +) (int64, error) { + var count int64 + err := dueGitHubPagesSourceQuery(ctx, now, sourceType, releaseSelector).Count(&count).Error + return count, err +} + +// ListDueGitHubPagesSourceChecks lists a batch of due latest GitHub sources in stable order. +func ListDueGitHubPagesSourceChecks( + ctx context.Context, + now time.Time, + sourceType string, + releaseSelector string, + limit int, +) ([]model.PagesDueGitHubSourceCandidate, error) { + var candidates []model.PagesDueGitHubSourceCandidate + err := dueGitHubPagesSourceQuery(ctx, now, sourceType, releaseSelector). + Select("source.id AS source_id, source.config_version"). + Order("runtime.next_check_at ASC"). + Order("source.id ASC"). + Limit(limit). + Scan(&candidates).Error + if err != nil { + return nil, err + } + return candidates, nil +} + +func dueGitHubPagesSourceQuery( + ctx context.Context, + now time.Time, + sourceType string, + releaseSelector string, +) *gorm.DB { + return db.DB(ctx). + Table("of_pages_project_source_runtime AS runtime"). + Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id"). + Where("source.source_type = ?", sourceType). + Where("source.release_selector = ?", releaseSelector). + Where("runtime.next_check_at IS NOT NULL AND runtime.next_check_at <= ?", now) +} + +// GetPagesDeploymentBySourceRevision loads a deployment by project source identity and revision. +func GetPagesDeploymentBySourceRevision( + ctx context.Context, + projectID uint, + sourceIdentity string, + revision string, +) (*model.PagesDeployment, error) { + var deployment model.PagesDeployment + err := db.DB(ctx). + Where("project_id = ? AND source_identity = ? AND source_revision = ?", projectID, sourceIdentity, revision). + First(&deployment).Error + if err != nil { + return nil, err + } + return &deployment, nil +} diff --git a/internal/repository/openflare_proxy_route.go b/internal/repository/openflare_proxy_route.go new file mode 100644 index 00000000..f9a6e833 --- /dev/null +++ b/internal/repository/openflare_proxy_route.go @@ -0,0 +1,154 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + + "gorm.io/gorm" + "gorm.io/gorm/clause" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// Zone domain binding sentinel errors (shared by proxy_route and zone binding helpers). +var ( + // ErrZoneDomainBoundToAnotherRoute is returned when a domain is already bound to a different route. + ErrZoneDomainBoundToAnotherRoute = errors.New("zone domain is already bound to another proxy route") + // ErrZoneDomainNotFound is returned when one or more requested domain IDs do not exist. + ErrZoneDomainNotFound = errors.New("one or more zone domains do not exist") +) + +// WithProxyRouteTx runs fn inside a database transaction for proxy-route multi-step work. +func WithProxyRouteTx(ctx context.Context, fn func(tx *gorm.DB) error) error { + return db.DB(ctx).Transaction(fn) +} + +// ListProxyRoutes 列出全部代理规则。 +func ListProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) { + var routes []*model.ProxyRoute + if err := db.DB(ctx).Order("id desc").Find(&routes).Error; err != nil { + return nil, err + } + return routes, nil +} + +// GetProxyRouteByID 按 ID 查询代理规则。 +func GetProxyRouteByID(ctx context.Context, id uint) (*model.ProxyRoute, error) { + var route model.ProxyRoute + if err := db.DB(ctx).First(&route, id).Error; err != nil { + return nil, err + } + return &route, nil +} + +// CreateProxyRouteRecord 创建代理规则。 +func CreateProxyRouteRecord(ctx context.Context, route *model.ProxyRoute) error { + return CreateProxyRouteRecordTx(db.DB(ctx), route) +} + +// CreateProxyRouteRecordTx creates a proxy route within an existing transaction. +func CreateProxyRouteRecordTx(tx *gorm.DB, route *model.ProxyRoute) error { + return tx.Create(route).Error +} + +// UpdateProxyRouteRecord 更新代理规则。 +func UpdateProxyRouteRecord(ctx context.Context, route *model.ProxyRoute) error { + return UpdateProxyRouteRecordTx(db.DB(ctx), route) +} + +// UpdateProxyRouteRecordTx updates a proxy route within an existing transaction. +func UpdateProxyRouteRecordTx(tx *gorm.DB, route *model.ProxyRoute) error { + return tx.Model(&model.ProxyRoute{}).Where("id = ?", route.ID).Updates(proxyRouteUpdateMap(route)).Error +} + +func proxyRouteUpdateMap(route *model.ProxyRoute) map[string]any { + return map[string]any{ + "site_name": route.SiteName, + "origin_id": route.OriginID, + "origin_url": route.OriginURL, + "origin_host": route.OriginHost, + "upstreams": route.Upstreams, + colEnabled: route.Enabled, + "enable_https": route.EnableHTTPS, + "redirect_http": route.RedirectHTTP, + "limit_conn_per_server": route.LimitConnPerServer, + "limit_conn_per_ip": route.LimitConnPerIP, + "limit_rate": route.LimitRate, + "limit_req_per_ip": route.LimitReqPerIP, + "cache_enabled": route.CacheEnabled, + "cache_policy": route.CachePolicy, + "cache_rules": route.CacheRules, + "custom_headers": route.CustomHeaders, + "basic_auth_enabled": route.BasicAuthEnabled, + "basic_auth_username": route.BasicAuthUsername, + "basic_auth_password": route.BasicAuthPassword, + "upstream_type": route.UpstreamType, + "tunnel_node_id": route.TunnelNodeID, + "tunnel_target_addr": route.TunnelTargetAddr, + "tunnel_target_protocol": route.TunnelTargetProtocol, + "pages_project_id": route.PagesProjectID, + } +} + +// DeleteProxyRouteRecord 删除代理规则。 +func DeleteProxyRouteRecord(ctx context.Context, id uint) error { + return DeleteProxyRouteRecordTx(db.DB(ctx), id) +} + +// DeleteProxyRouteRecordTx deletes a proxy route within an existing transaction. +func DeleteProxyRouteRecordTx(tx *gorm.DB, id uint) error { + return tx.Delete(&model.ProxyRoute{}, id).Error +} + +// ClearZoneDomainProxyRouteBindingsTx unbinds every zone domain from a proxy route. +func ClearZoneDomainProxyRouteBindingsTx(tx *gorm.DB, routeID uint) error { + return tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ?", routeID).Update("proxy_route_id", nil).Error +} + +// ReplaceZoneDomainRouteBindingsTx replaces every ZoneDomain binding for a proxy route +// inside the caller's transaction (with row locks on requested domains). +func ReplaceZoneDomainRouteBindingsTx(tx *gorm.DB, routeID uint, domainIDs []uint) error { + var requested []model.ZoneDomain + if len(domainIDs) > 0 { + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). + Where("id IN ?", domainIDs). + Find(&requested).Error; err != nil { + return err + } + if len(requested) != len(uniqueZoneDomainIDs(domainIDs)) { + return ErrZoneDomainNotFound + } + for _, domain := range requested { + if domain.ProxyRouteID != nil && *domain.ProxyRouteID != routeID { + return ErrZoneDomainBoundToAnotherRoute + } + } + } + + current := tx.Model(&model.ZoneDomain{}).Where("proxy_route_id = ?", routeID) + if len(domainIDs) > 0 { + current = current.Where("id NOT IN ?", domainIDs) + } + if err := current.Update("proxy_route_id", nil).Error; err != nil { + return err + } + + if len(domainIDs) == 0 { + return nil + } + return tx.Model(&model.ZoneDomain{}).Where("id IN ?", domainIDs).Update("proxy_route_id", routeID).Error +} + +// DeleteProxyRouteAndUnbind clears domain bindings then deletes the proxy route in one transaction. +func DeleteProxyRouteAndUnbind(ctx context.Context, id uint) error { + return WithProxyRouteTx(ctx, func(tx *gorm.DB) error { + if err := ClearZoneDomainProxyRouteBindingsTx(tx, id); err != nil { + return err + } + return DeleteProxyRouteRecordTx(tx, id) + }) +} diff --git a/internal/repository/openflare_tls.go b/internal/repository/openflare_tls.go new file mode 100644 index 00000000..318883f7 --- /dev/null +++ b/internal/repository/openflare_tls.go @@ -0,0 +1,95 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// HasTLSProxyRoutesTable 判断代理规则表是否已迁移。 +func HasTLSProxyRoutesTable(ctx context.Context) bool { + return db.DB(ctx).Migrator().HasTable(&model.TLSProxyRouteRef{}) +} + +// ListTLSCertificates 列出全部证书(不含 PEM 敏感字段的 JSON 暴露由 struct tag 控制)。 +func ListTLSCertificates(ctx context.Context) ([]model.TLSCertificate, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var certificates []model.TLSCertificate + if err := conn.Order("id desc").Find(&certificates).Error; err != nil { + return nil, err + } + return certificates, nil +} + +// GetTLSCertificateByID 按 ID 查询证书。 +func GetTLSCertificateByID(ctx context.Context, id uint) (*model.TLSCertificate, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + var certificate model.TLSCertificate + if err := conn.First(&certificate, id).Error; err != nil { + return nil, err + } + return &certificate, nil +} + +// CreateTLSCertificateRecord 创建证书记录。 +func CreateTLSCertificateRecord(ctx context.Context, certificate *model.TLSCertificate) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Create(certificate).Error +} + +// SaveTLSCertificate 保存证书记录。 +func SaveTLSCertificate(ctx context.Context, certificate *model.TLSCertificate) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Save(certificate).Error +} + +// DeleteTLSCertificateRecord 删除证书记录。 +func DeleteTLSCertificateRecord(ctx context.Context, id uint) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + return conn.Delete(&model.TLSCertificate{}, id).Error +} + +// CountTLSCertificatesByDNSAccountID 统计引用指定 DNS 账号的证书数量。 +func CountTLSCertificatesByDNSAccountID(ctx context.Context, dnsAccountID uint) (int64, error) { + conn := db.DB(ctx) + if conn == nil { + return 0, errors.New(errDatabaseNotInitialized) + } + var count int64 + if err := conn.Model(&model.TLSCertificate{}).Where("dns_account_id = ?", dnsAccountID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// ListTLSProxyRouteRefs 列出代理规则证书引用字段。 +func ListTLSProxyRouteRefs(ctx context.Context) ([]model.TLSProxyRouteRef, error) { + if !HasTLSProxyRoutesTable(ctx) { + return nil, nil + } + var routes []model.TLSProxyRouteRef + if err := db.DB(ctx).Order("id asc").Find(&routes).Error; err != nil { + return nil, err + } + return routes, nil +} diff --git a/internal/repository/openflare_waf.go b/internal/repository/openflare_waf.go new file mode 100644 index 00000000..54f2e1dc --- /dev/null +++ b/internal/repository/openflare_waf.go @@ -0,0 +1,357 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "errors" + "time" + + "gorm.io/gorm" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +func wafDB(ctx context.Context) (*gorm.DB, error) { + conn := db.DB(ctx) + if conn == nil { + return nil, errors.New(errDatabaseNotInitialized) + } + return conn, nil +} + +// ListOpenFlareWAFRuleGroups returns all rule groups. +func ListOpenFlareWAFRuleGroups(ctx context.Context) ([]*model.OpenFlareWAFRuleGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var groups []*model.OpenFlareWAFRuleGroup + if err = conn.Order("is_global desc").Order("id asc").Find(&groups).Error; err != nil { + return nil, err + } + return groups, nil +} + +// GetOpenFlareWAFRuleGroupByID returns a rule group by id. +func GetOpenFlareWAFRuleGroupByID(ctx context.Context, id uint) (*model.OpenFlareWAFRuleGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var group model.OpenFlareWAFRuleGroup + if err = conn.First(&group, id).Error; err != nil { + return nil, err + } + return &group, nil +} + +// GetGlobalOpenFlareWAFRuleGroup returns the global rule group if present. +func GetGlobalOpenFlareWAFRuleGroup(ctx context.Context) (*model.OpenFlareWAFRuleGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var group model.OpenFlareWAFRuleGroup + if err = conn.Where("is_global = ?", true).Order("id asc").First(&group).Error; err != nil { + return nil, err + } + return &group, nil +} + +// CreateOpenFlareWAFRuleGroup inserts a rule group. +func CreateOpenFlareWAFRuleGroup(ctx context.Context, group *model.OpenFlareWAFRuleGroup) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Create(group).Error +} + +// UpdateOpenFlareWAFRuleGroup updates mutable rule group fields. +func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *model.OpenFlareWAFRuleGroup) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Model(&model.OpenFlareWAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ + "name": group.Name, + colEnabled: group.Enabled, + "is_global": group.IsGlobal, + }).Error +} + +// UpdateOpenFlareWAFRuleGraph atomically replaces a graph when revision is current. +func UpdateOpenFlareWAFRuleGraph(ctx context.Context, id uint, revision uint64, graph string) (uint64, error) { + conn, err := wafDB(ctx) + if err != nil { + return 0, err + } + result := conn.Model(&model.OpenFlareWAFRuleGroup{}). + Where("id = ? AND revision = ?", id, revision). + Updates(map[string]any{"graph": graph, "revision": gorm.Expr("revision + 1")}) + if result.Error != nil { + return 0, result.Error + } + if result.RowsAffected != 1 { + return 0, model.ErrWAFRuleRevisionConflict + } + return revision + 1, nil +} + +// DeleteOpenFlareWAFRuleGroup removes a rule group. +func DeleteOpenFlareWAFRuleGroup(ctx context.Context, id uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Delete(&model.OpenFlareWAFRuleGroup{}, id).Error +} + +// ListOpenFlareWAFIPGroups returns all IP groups. +func ListOpenFlareWAFIPGroups(ctx context.Context) ([]*model.OpenFlareWAFIPGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var groups []*model.OpenFlareWAFIPGroup + if err = conn.Order("type asc").Order("id asc").Find(&groups).Error; err != nil { + return nil, err + } + return groups, nil +} + +// ListOpenFlareWAFIPGroupsByIDs returns IP groups for the given ids. +func ListOpenFlareWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*model.OpenFlareWAFIPGroup, error) { + if len(ids) == 0 { + return []*model.OpenFlareWAFIPGroup{}, nil + } + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var groups []*model.OpenFlareWAFIPGroup + if err = conn.Where("id IN ?", ids).Find(&groups).Error; err != nil { + return nil, err + } + return groups, nil +} + +// GetOpenFlareWAFIPGroupByID returns an IP group by id. +func GetOpenFlareWAFIPGroupByID(ctx context.Context, id uint) (*model.OpenFlareWAFIPGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var group model.OpenFlareWAFIPGroup + if err = conn.First(&group, id).Error; err != nil { + return nil, err + } + return &group, nil +} + +// CreateOpenFlareWAFIPGroup inserts an IP group. +func CreateOpenFlareWAFIPGroup(ctx context.Context, group *model.OpenFlareWAFIPGroup) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Create(group).Error +} + +// UpdateOpenFlareWAFIPGroup updates mutable IP group fields. +func UpdateOpenFlareWAFIPGroup(ctx context.Context, group *model.OpenFlareWAFIPGroup) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Model(&model.OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ + "name": group.Name, + "type": group.Type, + colEnabled: group.Enabled, + "ip_list": group.IPList, + "auto_config": group.AutoConfig, + "ext_ips": group.ExtIPs, + "subscription_url": group.SubscriptionURL, + "subscription_format": group.SubscriptionFormat, + "subscription_mapping_rule": group.SubscriptionMappingRule, + "sync_interval_minutes": group.SyncIntervalMinutes, + "next_sync_at": group.NextSyncAt, + "last_sync_status": group.LastSyncStatus, + "last_sync_message": group.LastSyncMessage, + }).Error +} + +// ListDueOpenFlareWAFIPGroups returns enabled automatic/subscription groups due for sync. +func ListDueOpenFlareWAFIPGroups(ctx context.Context, now time.Time) ([]*model.OpenFlareWAFIPGroup, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var groups []*model.OpenFlareWAFIPGroup + err = conn.Where( + "enabled = ? AND (type = ? OR (type = ? AND subscription_url <> '')) AND (next_sync_at IS NULL OR next_sync_at <= ?)", + true, "automatic", "subscription", now, + ).Order("id asc").Find(&groups).Error + return groups, err +} + +// UpdateOpenFlareWAFIPGroupSyncResult persists IP group sync outcome fields. +func UpdateOpenFlareWAFIPGroupSyncResult(ctx context.Context, group *model.OpenFlareWAFIPGroup) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Model(&model.OpenFlareWAFIPGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ + "ip_list": group.IPList, + "ext_ips": group.ExtIPs, + "last_synced_at": group.LastSyncedAt, + "next_sync_at": group.NextSyncAt, + "last_sync_status": group.LastSyncStatus, + "last_sync_message": group.LastSyncMessage, + "subscription_format": group.SubscriptionFormat, + }).Error +} + +// DeleteOpenFlareWAFIPGroup removes an IP group. +func DeleteOpenFlareWAFIPGroup(ctx context.Context, id uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Delete(&model.OpenFlareWAFIPGroup{}, id).Error +} + +// ListOpenFlareWAFRuleGroupBindings returns all bindings. +func ListOpenFlareWAFRuleGroupBindings(ctx context.Context) ([]model.OpenFlareWAFRuleGroupBinding, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var bindings []model.OpenFlareWAFRuleGroupBinding + if err = conn.Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil { + return nil, err + } + return bindings, nil +} + +// ListOpenFlareWAFRuleGroupBindingsByRouteID returns bindings for a proxy route. +func ListOpenFlareWAFRuleGroupBindingsByRouteID(ctx context.Context, routeID uint) ([]model.OpenFlareWAFRuleGroupBinding, error) { + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var bindings []model.OpenFlareWAFRuleGroupBinding + if err = conn.Where("proxy_route_id = ?", routeID).Order("sequence asc").Order("id asc").Find(&bindings).Error; err != nil { + return nil, err + } + return bindings, nil +} + +func syncWAFBindingIDSequence(tx *gorm.DB) error { + if tx == nil || tx.Dialector.Name() != "postgres" { //nolint:staticcheck // QF1008: keep explicit Dialector field access + return nil + } + return tx.Exec(` + SELECT setval( + pg_get_serial_sequence('of_waf_rule_group_bindings', 'id'), + GREATEST(COALESCE((SELECT MAX(id) FROM of_waf_rule_group_bindings), 0), 1), + COALESCE((SELECT MAX(id) FROM of_waf_rule_group_bindings), 0) > 0 + ) + `).Error +} + +func insertOpenFlareWAFRuleGroupBindings(tx *gorm.DB, bindings []model.OpenFlareWAFRuleGroupBinding) error { + if len(bindings) == 0 { + return nil + } + if err := syncWAFBindingIDSequence(tx); err != nil { + return err + } + return tx.Create(&bindings).Error +} + +// ReplaceOpenFlareWAFRuleGroupBindings replaces bindings for a rule group. +func ReplaceOpenFlareWAFRuleGroupBindings(ctx context.Context, groupID uint, routeIDs []uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Transaction(func(tx *gorm.DB) error { + if err = tx.Where("rule_group_id = ?", groupID).Delete(&model.OpenFlareWAFRuleGroupBinding{}).Error; err != nil { + return err + } + bindings := make([]model.OpenFlareWAFRuleGroupBinding, 0, len(routeIDs)) + for index, routeID := range routeIDs { + bindings = append(bindings, model.OpenFlareWAFRuleGroupBinding{ + RuleGroupID: groupID, + ProxyRouteID: routeID, + Sequence: index, + }) + } + return insertOpenFlareWAFRuleGroupBindings(tx, bindings) + }) +} + +// ReplaceOpenFlareWAFSiteRuleGroupBindings replaces bindings for a proxy route. +func ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx context.Context, routeID uint, groupIDs []uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Transaction(func(tx *gorm.DB) error { + if err = tx.Where("proxy_route_id = ?", routeID).Delete(&model.OpenFlareWAFRuleGroupBinding{}).Error; err != nil { + return err + } + bindings := make([]model.OpenFlareWAFRuleGroupBinding, 0, len(groupIDs)) + for index, groupID := range groupIDs { + bindings = append(bindings, model.OpenFlareWAFRuleGroupBinding{ + RuleGroupID: groupID, + ProxyRouteID: routeID, + Sequence: index, + }) + } + return insertOpenFlareWAFRuleGroupBindings(tx, bindings) + }) +} + +// DeleteOpenFlareWAFRuleGroupBindingsByGroupID removes bindings for a rule group. +func DeleteOpenFlareWAFRuleGroupBindingsByGroupID(ctx context.Context, groupID uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Where("rule_group_id = ?", groupID).Delete(&model.OpenFlareWAFRuleGroupBinding{}).Error +} + +// DeleteOpenFlareWAFRuleGroupWithBindings removes a rule group and its bindings. +func DeleteOpenFlareWAFRuleGroupWithBindings(ctx context.Context, groupID uint) error { + conn, err := wafDB(ctx) + if err != nil { + return err + } + return conn.Transaction(func(tx *gorm.DB) error { + if err = tx.Where("rule_group_id = ?", groupID).Delete(&model.OpenFlareWAFRuleGroupBinding{}).Error; err != nil { + return err + } + return tx.Delete(&model.OpenFlareWAFRuleGroup{}, groupID).Error + }) +} + +// GetOpenFlareProxyRouteByID returns a proxy route by id when the table exists. +func GetOpenFlareProxyRouteByID(ctx context.Context, id uint) (*model.OriginProxyRoute, error) { + if !HasProxyRoutesTable(ctx) { + return nil, gorm.ErrRecordNotFound + } + conn, err := wafDB(ctx) + if err != nil { + return nil, err + } + var route model.OriginProxyRoute + if err = conn.First(&route, id).Error; err != nil { + return nil, err + } + return &route, nil +} diff --git a/internal/model/openflare_waf_bindings_test.go b/internal/repository/openflare_waf_bindings_test.go similarity index 82% rename from internal/model/openflare_waf_bindings_test.go rename to internal/repository/openflare_waf_bindings_test.go index a16fa271..8245d79a 100644 --- a/internal/model/openflare_waf_bindings_test.go +++ b/internal/repository/openflare_waf_bindings_test.go @@ -1,12 +1,14 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" "testing" + "github.com/Rain-kl/Wavelet/internal/model" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -21,7 +23,7 @@ func setupWAFBindingsTestDB(t *testing.T) func() { DisableForeignKeyConstraintWhenMigrating: true, }) require.NoError(t, err) - require.NoError(t, sqliteDB.AutoMigrate(&OpenFlareWAFRuleGroupBinding{})) + require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareWAFRuleGroupBinding{})) db.SetDB(sqliteDB) return func() { @@ -36,7 +38,7 @@ func TestReplaceOpenFlareWAFRuleGroupBindingsAfterExplicitHighID(t *testing.T) { conn := db.DB(ctx) require.NotNil(t, conn) - require.NoError(t, conn.Create(&OpenFlareWAFRuleGroupBinding{ + require.NoError(t, conn.Create(&model.OpenFlareWAFRuleGroupBinding{ ID: 50, RuleGroupID: 1, ProxyRouteID: 1, @@ -44,7 +46,7 @@ func TestReplaceOpenFlareWAFRuleGroupBindingsAfterExplicitHighID(t *testing.T) { require.NoError(t, ReplaceOpenFlareWAFRuleGroupBindings(ctx, 2, []uint{2, 3})) - var bindings []OpenFlareWAFRuleGroupBinding + var bindings []model.OpenFlareWAFRuleGroupBinding require.NoError(t, conn.Where("rule_group_id = ?", 2).Order("proxy_route_id asc").Find(&bindings).Error) require.Len(t, bindings, 2) assert.Equal(t, uint(2), bindings[0].ProxyRouteID) diff --git a/internal/model/openflare_waf_graph_test.go b/internal/repository/openflare_waf_graph_test.go similarity index 88% rename from internal/model/openflare_waf_graph_test.go rename to internal/repository/openflare_waf_graph_test.go index d1fd8087..a4cedbdd 100644 --- a/internal/model/openflare_waf_graph_test.go +++ b/internal/repository/openflare_waf_graph_test.go @@ -1,7 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" @@ -12,6 +12,8 @@ import ( "testing" "testing/fstest" + "github.com/Rain-kl/Wavelet/internal/model" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/glebarez/sqlite" "github.com/pressly/goose/v3" @@ -26,7 +28,7 @@ func wafMigrationFS(t *testing.T) fs.FS { t.Helper() _, filename, _, ok := runtime.Caller(0) require.True(t, ok) - dir := filepath.Join(filepath.Dir(filename), "..", "db", "migrator", "goose", "sqlite") + dir := filepath.Join(filepath.Dir(filename), "..", "infra", "persistence", "migrator", "goose", "sqlite") migrations := fstest.MapFS{} for _, name := range []string{ "202607150001_orchestrate_waf_rules.sql", @@ -55,7 +57,7 @@ func TestOpenFlareWAFGraphMigrationResetsGraphsAndOrdersBindings(t *testing.T) { require.NoError(t, goose.SetDialect("sqlite3")) require.NoError(t, goose.Up(sqlDB, ".")) - var groups []OpenFlareWAFRuleGroup + var groups []model.OpenFlareWAFRuleGroup require.NoError(t, conn.Order("id asc").Find(&groups).Error) require.Len(t, groups, 2) for _, group := range groups { @@ -67,12 +69,12 @@ func TestOpenFlareWAFGraphMigrationResetsGraphsAndOrdersBindings(t *testing.T) { } require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (name) VALUES ('new')`).Error) - var newGroup OpenFlareWAFRuleGroup + var newGroup model.OpenFlareWAFRuleGroup require.NoError(t, conn.First(&newGroup, 3).Error) assert.Empty(t, newGroup.Graph) assert.Equal(t, uint64(1), newGroup.Revision) - var bindings []OpenFlareWAFRuleGroupBinding + var bindings []model.OpenFlareWAFRuleGroupBinding require.NoError(t, conn.Where("proxy_route_id = ?", 7).Order("sequence asc").Order("id asc").Find(&bindings).Error) require.Len(t, bindings, 2) assert.Equal(t, []int{0, 1}, []int{bindings[0].Sequence, bindings[1].Sequence}) @@ -82,11 +84,11 @@ func TestOpenFlareWAFGraphMigrationResetsGraphsAndOrdersBindings(t *testing.T) { func TestOpenFlareWAFGraphOptimisticUpdate(t *testing.T) { conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) require.NoError(t, err) - require.NoError(t, conn.AutoMigrate(&OpenFlareWAFRuleGroup{})) + require.NoError(t, conn.AutoMigrate(&model.OpenFlareWAFRuleGroup{})) db.SetDB(conn) t.Cleanup(func() { db.SetDB(nil) }) - group := OpenFlareWAFRuleGroup{Name: "rule", Graph: defaultWAFRuleGraph, Revision: 1} + group := model.OpenFlareWAFRuleGroup{Name: "rule", Graph: defaultWAFRuleGraph, Revision: 1} require.NoError(t, conn.Create(&group).Error) nextRevision, err := UpdateOpenFlareWAFRuleGraph(context.Background(), group.ID, 1, `{"schema_version":1}`) @@ -94,7 +96,7 @@ func TestOpenFlareWAFGraphOptimisticUpdate(t *testing.T) { assert.Equal(t, uint64(2), nextRevision) _, err = UpdateOpenFlareWAFRuleGraph(context.Background(), group.ID, 1, defaultWAFRuleGraph) - assert.ErrorIs(t, err, ErrWAFRuleRevisionConflict) + assert.ErrorIs(t, err, model.ErrWAFRuleRevisionConflict) } func TestReplaceOpenFlareWAFRuleGroupBindingsPreservesInputOrder(t *testing.T) { @@ -113,10 +115,10 @@ func TestReplaceOpenFlareWAFRuleGroupBindingsPreservesInputOrder(t *testing.T) { func TestLegacyWAFColumnsRemoved(t *testing.T) { conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) require.NoError(t, err) - require.NoError(t, conn.AutoMigrate(&OpenFlareWAFRuleGroup{})) + require.NoError(t, conn.AutoMigrate(&model.OpenFlareWAFRuleGroup{})) legacy := []string{"block_status_code", "block_response_body", "ip_whitelist", "ip_blacklist", "ip_whitelist_groups", "ip_blacklist_groups", "country_whitelist", "country_blacklist", "region_whitelist", "region_blacklist", "pow_enabled", "pow_config"} for _, column := range legacy { - if conn.Migrator().HasColumn(&OpenFlareWAFRuleGroup{}, column) { + if conn.Migrator().HasColumn(&model.OpenFlareWAFRuleGroup{}, column) { t.Fatalf("legacy WAF column %s still exists", column) } } diff --git a/internal/repository/openflare_zone.go b/internal/repository/openflare_zone.go new file mode 100644 index 00000000..9f99c877 --- /dev/null +++ b/internal/repository/openflare_zone.go @@ -0,0 +1,164 @@ +package repository + +import ( + "context" + "errors" + "fmt" + + "gorm.io/gorm" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// ListZones returns all zones ordered by domain ascending. +func ListZones(ctx context.Context) ([]model.Zone, error) { + var zones []model.Zone + if err := db.DB(ctx).Order("domain asc").Find(&zones).Error; err != nil { + return nil, err + } + return zones, nil +} + +// GetZoneByID returns a zone by primary key. +func GetZoneByID(ctx context.Context, id uint) (*model.Zone, error) { + var zone model.Zone + if err := db.DB(ctx).First(&zone, id).Error; err != nil { + return nil, err + } + return &zone, nil +} + +// CreateZone creates a zone record. +func CreateZone(ctx context.Context, zone *model.Zone) error { + return db.DB(ctx).Create(zone).Error +} + +// SaveZone persists zone updates. +func SaveZone(ctx context.Context, zone *model.Zone) error { + return db.DB(ctx).Save(zone).Error +} + +// DeleteZone deletes a zone by primary key. +func DeleteZone(ctx context.Context, id uint) error { + return db.DB(ctx).Delete(&model.Zone{}, id).Error +} + +// ListZoneDomainCounts returns per-zone domain counts for list cards. +func ListZoneDomainCounts(ctx context.Context) ([]model.ZoneDomainCount, error) { + var rows []model.ZoneDomainCount + if err := db.DB(ctx).Model(&model.ZoneDomain{}). + Select("zone_id, count(*) as count"). + Group("zone_id"). + Scan(&rows).Error; err != nil { + return nil, err + } + return rows, nil +} + +// ListZoneDomainsByZoneID returns domains under a zone ordered by domain ascending. +func ListZoneDomainsByZoneID(ctx context.Context, zoneID uint) ([]model.ZoneDomain, error) { + var domains []model.ZoneDomain + if err := db.DB(ctx).Where("zone_id = ?", zoneID).Order("domain asc").Find(&domains).Error; err != nil { + return nil, err + } + return domains, nil +} + +// CountZoneDomainsByZoneID counts domains under a zone. +func CountZoneDomainsByZoneID(ctx context.Context, zoneID uint) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.ZoneDomain{}).Where("zone_id = ?", zoneID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// GetZoneDomainByZoneAndID returns a domain scoped to a zone. +func GetZoneDomainByZoneAndID(ctx context.Context, zoneID, id uint) (*model.ZoneDomain, error) { + var item model.ZoneDomain + if err := db.DB(ctx).Where("id = ? AND zone_id = ?", id, zoneID).First(&item).Error; err != nil { + return nil, err + } + return &item, nil +} + +// CreateZoneDomain creates a zone domain record. +func CreateZoneDomain(ctx context.Context, domain *model.ZoneDomain) error { + return db.DB(ctx).Create(domain).Error +} + +// SaveZoneDomain persists zone domain updates. +func SaveZoneDomain(ctx context.Context, domain *model.ZoneDomain) error { + return db.DB(ctx).Save(domain).Error +} + +// DeleteZoneDomain deletes a zone domain record. +func DeleteZoneDomain(ctx context.Context, domain *model.ZoneDomain) error { + return db.DB(ctx).Delete(domain).Error +} + +// ListZoneDomainsByRouteID returns the domains bound to a proxy route. +func ListZoneDomainsByRouteID(ctx context.Context, routeID uint) ([]model.ZoneDomain, error) { + var domains []model.ZoneDomain + if err := db.DB(ctx).Where("proxy_route_id = ?", routeID).Order("id asc").Find(&domains).Error; err != nil { + return nil, err + } + return domains, nil +} + +// ListZoneDomainsByIDs returns explicit domains in the requested ID order. +func ListZoneDomainsByIDs(ctx context.Context, domainIDs []uint) ([]model.ZoneDomain, error) { + if len(domainIDs) == 0 { + return []model.ZoneDomain{}, nil + } + var domains []model.ZoneDomain + if err := db.DB(ctx).Where("id IN ?", domainIDs).Find(&domains).Error; err != nil { + return nil, err + } + byID := make(map[uint]model.ZoneDomain, len(domains)) + for _, domain := range domains { + byID[domain.ID] = domain + } + ordered := make([]model.ZoneDomain, 0, len(domainIDs)) + for _, id := range domainIDs { + domain, ok := byID[id] + if !ok { + return nil, fmt.Errorf("one or more zone domains do not exist") + } + ordered = append(ordered, domain) + } + return ordered, nil +} + +// CountZoneDomainsByCertificateID reports whether a certificate is assigned to a model.Zone domain. +func CountZoneDomainsByCertificateID(ctx context.Context, certificateID uint) (int64, error) { + var count int64 + err := db.DB(ctx).Model(&model.ZoneDomain{}).Where("cert_id = ?", certificateID).Count(&count).Error + return count, err +} + +// ReplaceZoneDomainRouteBindings replaces every model.ZoneDomain binding for a proxy route. +func ReplaceZoneDomainRouteBindings(ctx context.Context, routeID uint, domainIDs []uint) error { + conn := db.DB(ctx) + if conn == nil { + return errors.New(errDatabaseNotInitialized) + } + + return conn.Transaction(func(tx *gorm.DB) error { + return ReplaceZoneDomainRouteBindingsTx(tx, routeID, domainIDs) + }) +} + +func uniqueZoneDomainIDs(domainIDs []uint) []uint { + seen := make(map[uint]struct{}, len(domainIDs)) + ids := make([]uint, 0, len(domainIDs)) + for _, id := range domainIDs { + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + ids = append(ids, id) + } + return ids +} diff --git a/internal/model/openflare_zone_test.go b/internal/repository/openflare_zone_test.go similarity index 74% rename from internal/model/openflare_zone_test.go rename to internal/repository/openflare_zone_test.go index ce50449d..98800d11 100644 --- a/internal/model/openflare_zone_test.go +++ b/internal/repository/openflare_zone_test.go @@ -1,12 +1,14 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" "testing" + "github.com/Rain-kl/Wavelet/internal/model" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/glebarez/sqlite" "github.com/stretchr/testify/require" @@ -20,7 +22,7 @@ func setupZoneTestDB(t *testing.T) *gorm.DB { DisableForeignKeyConstraintWhenMigrating: true, }) require.NoError(t, err) - require.NoError(t, sqliteDB.AutoMigrate(&Zone{}, &ZoneDomain{})) + require.NoError(t, sqliteDB.AutoMigrate(&model.Zone{}, &model.ZoneDomain{})) db.SetDB(sqliteDB) t.Cleanup(func() { db.SetDB(nil) }) return sqliteDB @@ -30,10 +32,10 @@ func TestReplaceZoneDomainRouteBindingsRejectsForeignDomain(t *testing.T) { conn := setupZoneTestDB(t) ctx := context.Background() - zone := Zone{Domain: "example.com"} + zone := model.Zone{Domain: "example.com"} require.NoError(t, conn.Create(&zone).Error) foreignRouteID := uint(11) - domain := ZoneDomain{ + domain := model.ZoneDomain{ ZoneID: zone.ID, ProxyRouteID: &foreignRouteID, Domain: "api.example.com", @@ -43,7 +45,7 @@ func TestReplaceZoneDomainRouteBindingsRejectsForeignDomain(t *testing.T) { err := ReplaceZoneDomainRouteBindings(ctx, 12, []uint{domain.ID}) require.Error(t, err) - var got ZoneDomain + var got model.ZoneDomain require.NoError(t, conn.First(&got, domain.ID).Error) require.Equal(t, &foreignRouteID, got.ProxyRouteID) } @@ -52,17 +54,17 @@ func TestReplaceZoneDomainRouteBindingsReplacesCurrentRouteBindings(t *testing.T conn := setupZoneTestDB(t) ctx := context.Background() - zone := Zone{Domain: "example.com"} + zone := model.Zone{Domain: "example.com"} require.NoError(t, conn.Create(&zone).Error) routeID := uint(21) - boundDomain := ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &routeID, Domain: "old.example.com"} - requestedDomain := ZoneDomain{ZoneID: zone.ID, Domain: "new.example.com"} + boundDomain := model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &routeID, Domain: "old.example.com"} + requestedDomain := model.ZoneDomain{ZoneID: zone.ID, Domain: "new.example.com"} require.NoError(t, conn.Create(&boundDomain).Error) require.NoError(t, conn.Create(&requestedDomain).Error) require.NoError(t, ReplaceZoneDomainRouteBindings(ctx, routeID, []uint{requestedDomain.ID})) - var domains []ZoneDomain + var domains []model.ZoneDomain require.NoError(t, conn.Order("id asc").Find(&domains).Error) require.Len(t, domains, 2) require.Nil(t, domains[0].ProxyRouteID) @@ -73,11 +75,11 @@ func TestListZoneDomainsByRouteID(t *testing.T) { conn := setupZoneTestDB(t) ctx := context.Background() - zone := Zone{Domain: "example.com"} + zone := model.Zone{Domain: "example.com"} require.NoError(t, conn.Create(&zone).Error) routeID := uint(31) - boundDomain := ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &routeID, Domain: "api.example.com"} - unboundDomain := ZoneDomain{ZoneID: zone.ID, Domain: "www.example.com"} + boundDomain := model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &routeID, Domain: "api.example.com"} + unboundDomain := model.ZoneDomain{ZoneID: zone.ID, Domain: "www.example.com"} require.NoError(t, conn.Create(&boundDomain).Error) require.NoError(t, conn.Create(&unboundDomain).Error) diff --git a/internal/repository/push_history.go b/internal/repository/push_history.go index f9ef9fbd..1ad29560 100644 --- a/internal/repository/push_history.go +++ b/internal/repository/push_history.go @@ -5,10 +5,12 @@ package repository import ( "context" + "time" + + "gorm.io/gorm" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" - "gorm.io/gorm" ) // PushHistoryListFilter filters push history pagination queries. @@ -48,6 +50,19 @@ func CreatePushHistory(ctx context.Context, history *model.PushHistory) error { return db.DB(ctx).Create(history).Error } +// CountPushHistoriesCreatedBefore returns how many push history rows were created before cutoff. +func CountPushHistoriesCreatedBefore(ctx context.Context, cutoff time.Time) (int64, error) { + var count int64 + err := db.DB(ctx).Model(&model.PushHistory{}).Where("created_at < ?", cutoff).Count(&count).Error + return count, err +} + +// DeletePushHistoriesCreatedBefore deletes push history rows created before cutoff. +func DeletePushHistoriesCreatedBefore(ctx context.Context, cutoff time.Time) (int64, error) { + result := db.DB(ctx).Where("created_at < ?", cutoff).Delete(&model.PushHistory{}) + return result.RowsAffected, result.Error +} + // PushHistoryQuery returns a scoped query builder for push histories. func PushHistoryQuery(ctx context.Context) *gorm.DB { return db.DB(ctx).Model(&model.PushHistory{}) diff --git a/internal/repository/schedule.go b/internal/repository/schedule.go new file mode 100644 index 00000000..7195928b --- /dev/null +++ b/internal/repository/schedule.go @@ -0,0 +1,53 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/model" +) + +// CreateSchedule 创建定时任务 +func CreateSchedule(ctx context.Context, schedule *model.Schedule) error { + return db.DB(ctx).Create(schedule).Error +} + +// UpdateSchedule 更新定时任务 +func UpdateSchedule(ctx context.Context, schedule *model.Schedule) error { + return db.DB(ctx).Save(schedule).Error +} + +// DeleteSchedule 删除定时任务 +func DeleteSchedule(ctx context.Context, id uint64) error { + return db.DB(ctx).Delete(&model.Schedule{}, id).Error +} + +// GetScheduleByID 根据 ID 获取定时任务 +func GetScheduleByID(ctx context.Context, id uint64) (*model.Schedule, error) { + var schedule model.Schedule + if err := db.DB(ctx).Where("id = ?", id).First(&schedule).Error; err != nil { + return nil, err + } + return &schedule, nil +} + +// ListSchedules 获取所有定时任务 +func ListSchedules(ctx context.Context) ([]model.Schedule, error) { + var schedules []model.Schedule + if err := db.DB(ctx).Order("id DESC").Find(&schedules).Error; err != nil { + return nil, err + } + return schedules, nil +} + +// ListActiveSchedules 获取所有启用的定时任务 +func ListActiveSchedules(ctx context.Context) ([]model.Schedule, error) { + var schedules []model.Schedule + if err := db.DB(ctx).Where("is_active = ?", true).Find(&schedules).Error; err != nil { + return nil, err + } + return schedules, nil +} diff --git a/internal/repository/system_config.go b/internal/repository/system_config.go index 804a7220..4e816678 100644 --- a/internal/repository/system_config.go +++ b/internal/repository/system_config.go @@ -18,14 +18,7 @@ import ( "github.com/Rain-kl/Wavelet/pkg/cache/ram" ) -const ( - configTypeSystem = "system" - errDatabaseNotInitialized = "database not initialized" - errConfigIntParseFailed = "配置 %s 的值 '%s' 无法转换为整数: %w" - errConfigDecimalParseFailed = "配置 %s 的值 '%s' 无法转换为decimal: %w" - errConfigBoolParseFailed = "配置 %s 的值 '%s' 无法转换为布尔值: %w" - errParseMenuDisplayConfigFailed = "解析目录显示配置失败: %w" -) +const configTypeSystem = "system" // PreheatSystemConfigs loads all system configs from database. // This function strictly performs database read and does not perform any cache read or write operations. diff --git a/internal/repository/system_config_admin.go b/internal/repository/system_config_admin.go index acb3ff3d..cdc3f7a6 100644 --- a/internal/repository/system_config_admin.go +++ b/internal/repository/system_config_admin.go @@ -54,7 +54,12 @@ func CreateSystemConfig(ctx context.Context, config *model.SystemConfig) error { // UpdateSystemConfigFields applies partial updates to a system config row. func UpdateSystemConfigFields(ctx context.Context, config *model.SystemConfig, updates map[string]any) error { - return db.DB(ctx).Model(config).Updates(updates).Error + return UpdateSystemConfigFieldsTx(db.DB(ctx), config, updates) +} + +// UpdateSystemConfigFieldsTx applies partial updates within an existing transaction. +func UpdateSystemConfigFieldsTx(tx *gorm.DB, config *model.SystemConfig, updates map[string]any) error { + return tx.Model(config).Updates(updates).Error } // SaveOrUpdateSystemConfig creates or updates a config row and invalidates cache. diff --git a/internal/repository/task_execution.go b/internal/repository/task_execution.go new file mode 100644 index 00000000..aa07d9db --- /dev/null +++ b/internal/repository/task_execution.go @@ -0,0 +1,281 @@ +package repository + +import ( + "context" + "errors" + "fmt" + "strings" + "time" + + "github.com/redis/go-redis/v9" + "gorm.io/gorm" + + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" + "github.com/Rain-kl/Wavelet/internal/model" +) + +const ( + taskExecutionLogRedisKeyPrefix = "task:execution:log:" + taskExecutionLogExpiration = 24 * time.Hour + taskExecutionLogMaxLines = 1000 +) + +// CreateTaskExecution 创建任务执行记录 +func CreateTaskExecution(ctx context.Context, execution *model.TaskExecution) error { + execution.ID = idgen.NextUint64ID() + return db.DB(ctx).Create(execution).Error +} + +// UpdateTaskExecution 更新任务执行记录,忽略由 Redis 缓冲和归档流程管理的 log 字段。 +func UpdateTaskExecution(ctx context.Context, execution *model.TaskExecution) error { + return db.DB(ctx).Omit("log").Save(execution).Error +} + +// GetTaskExecutionByTaskID 根据 TaskID 获取执行记录 +func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*model.TaskExecution, error) { + var execution model.TaskExecution + if err := db.DB(ctx).Where("task_id = ?", taskID).First(&execution).Error; err != nil { + return nil, err + } + if err := loadTaskExecutionLog(ctx, &execution); err != nil { + return nil, err + } + return &execution, nil +} + +// GetTaskExecutionByID 根据 ID 获取执行记录 +func GetTaskExecutionByID(ctx context.Context, id uint64) (*model.TaskExecution, error) { + var execution model.TaskExecution + if err := db.DB(ctx).Where("id = ?", id).First(&execution).Error; err != nil { + return nil, err + } + if err := loadTaskExecutionLog(ctx, &execution); err != nil { + return nil, err + } + return &execution, nil +} + +// GetLatestTaskExecutionByTaskType returns the most recent execution for a task type. +// ok is false when no row exists. +func GetLatestTaskExecutionByTaskType(ctx context.Context, taskType string) (*model.TaskExecution, bool, error) { + var execution model.TaskExecution + err := db.DB(ctx). + Where("task_type = ?", taskType). + Order("id DESC"). + First(&execution).Error + if err == nil { + if loadErr := loadTaskExecutionLog(ctx, &execution); loadErr != nil { + return nil, false, loadErr + } + return &execution, true, nil + } + if errors.Is(err, gorm.ErrRecordNotFound) { + return nil, false, nil + } + return nil, false, err +} + +// AppendTaskExecutionLog 将日志追加到 Redis 缓冲,任务完成后再持久化到数据库。 +func AppendTaskExecutionLog(ctx context.Context, taskID string, logLine string) error { + if db.Redis == nil { + return errors.New("redis client is not initialized") + } + + now := time.Now().Format("15:04:05") + line := fmt.Sprintf("[%s] %s\n", now, logLine) + key := taskExecutionLogRedisKey(taskID) + + _, err := db.Redis.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + pipe.RPush(ctx, key, line) + pipe.LTrim(ctx, key, -taskExecutionLogMaxLines, -1) + pipe.Expire(ctx, key, taskExecutionLogExpiration) + return nil + }) + if err != nil { + return fmt.Errorf("append task execution log to redis: %w", err) + } + return nil +} + +// FlushTaskExecutionLog 将 Redis 中的完整任务日志写入数据库,并在成功后清理缓存。 +func FlushTaskExecutionLog(ctx context.Context, taskID string) error { + if db.Redis == nil { + return errors.New("redis client is not initialized") + } + + key := taskExecutionLogRedisKey(taskID) + logLines, err := db.Redis.LRange(ctx, key, 0, -1).Result() + if err != nil { + return fmt.Errorf("get task execution log from redis: %w", err) + } + if len(logLines) == 0 { + return nil + } + logText := strings.Join(logLines, "") + + result := db.DB(ctx).Model(&model.TaskExecution{}). + Where("task_id = ?", taskID). + Update("log", logText) + if result.Error != nil { + return fmt.Errorf("persist task execution log: %w", result.Error) + } + if result.RowsAffected == 0 { + return fmt.Errorf("persist task execution log: task %q not found", taskID) + } + + if err := db.Redis.Del(ctx, key).Err(); err != nil { + return fmt.Errorf("delete persisted task execution log from redis: %w", err) + } + return nil +} + +// ListTaskExecutions 分页查询任务执行记录 +func ListTaskExecutions(ctx context.Context, req model.ListTaskExecutionsRequest) ([]model.TaskExecution, int64, error) { + if req.Page <= 0 { + req.Page = 1 + } + if req.PageSize <= 0 { + req.PageSize = 20 + } + + query := db.DB(ctx).Model(&model.TaskExecution{}) + + if req.Status != "" { + query = query.Where("status = ?", req.Status) + } + if req.TaskType != "" { + query = query.Where("task_type = ?", req.TaskType) + } + + var total int64 + if err := query.Count(&total).Error; err != nil { + return nil, 0, err + } + + var executions []model.TaskExecution + offset := (req.Page - 1) * req.PageSize + if err := query.Order("id DESC").Offset(offset).Limit(req.PageSize).Find(&executions).Error; err != nil { + return nil, 0, err + } + if err := loadTaskExecutionLogs(ctx, executions); err != nil { + return nil, 0, err + } + + return executions, total, nil +} + +// MarkFailedTaskExecutionsSucceededTx marks failed executions of a task type as succeeded within a transaction. +func MarkFailedTaskExecutionsSucceededTx( + tx *gorm.DB, + taskType string, + result string, + finishedAt time.Time, +) error { + return tx.Model(&model.TaskExecution{}). + Where("task_type = ? AND status = ?", taskType, model.TaskExecutionStatusFailed). + Updates(map[string]any{ + "status": model.TaskExecutionStatusSucceeded, + "result": result, + "finished_at": finishedAt, + }).Error +} + +// CleanupTaskExecutionLogs removes finished task execution logs according to frequency-based retention. +func CleanupTaskExecutionLogs(ctx context.Context, now time.Time) (model.TaskExecutionCleanupStats, error) { + const ( + frequencyWindowDays = 30 + highFrequencyThreshold = frequencyWindowDays + ) + + frequencyWindowStart := now.AddDate(0, 0, -frequencyWindowDays) + highFrequencyCutoff := now.AddDate(0, 0, -3) + lowFrequencyCutoff := now.AddDate(0, 0, -30) + terminalStatuses := []model.TaskExecutionStatus{model.TaskExecutionStatusSucceeded, model.TaskExecutionStatusFailed} + + var highFrequencyTaskTypes []string + if err := db.DB(ctx). + Model(&model.TaskExecution{}). + Select("task_type"). + Where("created_at >= ?", frequencyWindowStart). + Group("task_type"). + Having("COUNT(*) > ?", highFrequencyThreshold). + Pluck("task_type", &highFrequencyTaskTypes).Error; err != nil { + return model.TaskExecutionCleanupStats{}, fmt.Errorf("query high-frequency task types: %w", err) + } + + var highFrequencyDeleted int64 + if len(highFrequencyTaskTypes) > 0 { + highFrequencyResult := db.DB(ctx). + Where("status IN ?", terminalStatuses). + Where("created_at < ?", highFrequencyCutoff). + Where("task_type IN ?", highFrequencyTaskTypes). + Delete(&model.TaskExecution{}) + if highFrequencyResult.Error != nil { + return model.TaskExecutionCleanupStats{}, fmt.Errorf("delete high-frequency task execution logs: %w", highFrequencyResult.Error) + } + highFrequencyDeleted = highFrequencyResult.RowsAffected + } + + lowFrequencyQuery := db.DB(ctx). + Where("status IN ?", terminalStatuses). + Where("created_at < ?", lowFrequencyCutoff) + if len(highFrequencyTaskTypes) > 0 { + lowFrequencyQuery = lowFrequencyQuery.Where("task_type NOT IN ?", highFrequencyTaskTypes) + } + lowFrequencyResult := lowFrequencyQuery.Delete(&model.TaskExecution{}) + if lowFrequencyResult.Error != nil { + return model.TaskExecutionCleanupStats{}, fmt.Errorf("delete low-frequency task execution logs: %w", lowFrequencyResult.Error) + } + + return model.TaskExecutionCleanupStats{ + HighFrequencyDeleted: highFrequencyDeleted, + LowFrequencyDeleted: lowFrequencyResult.RowsAffected, + }, nil +} + +func taskExecutionLogRedisKey(taskID string) string { + return db.PrefixedKey(taskExecutionLogRedisKeyPrefix + taskID) +} + +func loadTaskExecutionLog(ctx context.Context, execution *model.TaskExecution) error { + if db.Redis == nil { + return nil + } + + logLines, err := db.Redis.LRange(ctx, taskExecutionLogRedisKey(execution.TaskID), 0, -1).Result() + if err != nil { + return fmt.Errorf("get task execution log from redis: %w", err) + } + if len(logLines) == 0 { + return nil + } + + execution.Log = strings.Join(logLines, "") + return nil +} + +func loadTaskExecutionLogs(ctx context.Context, executions []model.TaskExecution) error { + if db.Redis == nil || len(executions) == 0 { + return nil + } + + commands := make([]*redis.StringSliceCmd, len(executions)) + _, err := db.Redis.Pipelined(ctx, func(pipe redis.Pipeliner) error { + for i := range executions { + commands[i] = pipe.LRange(ctx, taskExecutionLogRedisKey(executions[i].TaskID), 0, -1) + } + return nil + }) + if err != nil { + return fmt.Errorf("get task execution logs from redis: %w", err) + } + + for i := range executions { + logLines := commands[i].Val() + if len(logLines) > 0 { + executions[i].Log = strings.Join(logLines, "") + } + } + return nil +} diff --git a/internal/model/task_execution_test.go b/internal/repository/task_execution_test.go similarity index 79% rename from internal/model/task_execution_test.go rename to internal/repository/task_execution_test.go index 4fcc681f..b80b94cc 100644 --- a/internal/model/task_execution_test.go +++ b/internal/repository/task_execution_test.go @@ -2,7 +2,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package repository import ( "context" @@ -10,6 +10,8 @@ import ( "testing" "time" + "github.com/Rain-kl/Wavelet/internal/model" + db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/alicebob/miniredis/v2" "github.com/glebarez/sqlite" @@ -26,7 +28,7 @@ func setupTaskExecutionTestEnvironment(t *testing.T) func() { }) require.NoError(t, err) - err = sqliteDB.AutoMigrate(&TaskExecution{}) + err = sqliteDB.AutoMigrate(&model.TaskExecution{}) require.NoError(t, err) miniRedis, err := miniredis.Run() @@ -54,11 +56,11 @@ func TestCreateTaskExecution(t *testing.T) { defer cleanup() ctx := context.Background() - execution := &TaskExecution{ + execution := &model.TaskExecution{ TaskID: "manual_cleanup_123", TaskType: "system:cleanup", TaskName: "清理未使用上传", - Status: TaskExecutionStatusPending, + Status: model.TaskExecutionStatusPending, Retryable: true, MaxRetry: 3, RetryCount: 0, @@ -79,11 +81,11 @@ func TestGetTaskExecutionByTaskID(t *testing.T) { ctx := context.Background() // 创建记录 - execution := &TaskExecution{ + execution := &model.TaskExecution{ TaskID: "test_task_id_001", TaskType: "system:cleanup", TaskName: "清理未使用上传", - Status: TaskExecutionStatusPending, + Status: model.TaskExecutionStatusPending, Retryable: true, MaxRetry: 3, TriggeredBy: "manual", @@ -96,7 +98,7 @@ func TestGetTaskExecutionByTaskID(t *testing.T) { require.NoError(t, err) assert.Equal(t, execution.ID, found.ID) assert.Equal(t, "test_task_id_001", found.TaskID) - assert.Equal(t, TaskExecutionStatusPending, found.Status) + assert.Equal(t, model.TaskExecutionStatusPending, found.Status) assert.True(t, found.Retryable) assert.Equal(t, 3, found.MaxRetry) @@ -110,11 +112,11 @@ func TestGetTaskExecutionByID(t *testing.T) { defer cleanup() ctx := context.Background() - execution := &TaskExecution{ + execution := &model.TaskExecution{ TaskID: "test_by_id_001", TaskType: "system:cleanup", TaskName: "清理未使用上传", - Status: TaskExecutionStatusPending, + Status: model.TaskExecutionStatusPending, TriggeredBy: "system", } err := CreateTaskExecution(ctx, execution) @@ -132,11 +134,11 @@ func TestUpdateTaskExecution(t *testing.T) { ctx := context.Background() // 创建记录 - execution := &TaskExecution{ + execution := &model.TaskExecution{ TaskID: "test_update_001", TaskType: "system:cleanup", TaskName: "清理未使用上传", - Status: TaskExecutionStatusPending, + Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", } err := CreateTaskExecution(ctx, execution) @@ -144,7 +146,7 @@ func TestUpdateTaskExecution(t *testing.T) { // 更新状态为 running now := time.Now() - execution.Status = TaskExecutionStatusRunning + execution.Status = model.TaskExecutionStatusRunning execution.StartedAt = &now err = UpdateTaskExecution(ctx, execution) require.NoError(t, err) @@ -152,12 +154,12 @@ func TestUpdateTaskExecution(t *testing.T) { // 验证更新 found, err := GetTaskExecutionByTaskID(ctx, "test_update_001") require.NoError(t, err) - assert.Equal(t, TaskExecutionStatusRunning, found.Status) + assert.Equal(t, model.TaskExecutionStatusRunning, found.Status) assert.NotNil(t, found.StartedAt) // 更新为 succeeded finishTime := time.Now() - execution.Status = TaskExecutionStatusSucceeded + execution.Status = model.TaskExecutionStatusSucceeded execution.FinishedAt = &finishTime execution.Duration = 1500 execution.Result = "共清理 50 个文件" @@ -166,7 +168,7 @@ func TestUpdateTaskExecution(t *testing.T) { found, err = GetTaskExecutionByTaskID(ctx, "test_update_001") require.NoError(t, err) - assert.Equal(t, TaskExecutionStatusSucceeded, found.Status) + assert.Equal(t, model.TaskExecutionStatusSucceeded, found.Status) assert.Equal(t, int64(1500), found.Duration) assert.Equal(t, "共清理 50 个文件", found.Result) } @@ -176,11 +178,11 @@ func TestUpdateTaskExecutionFailed(t *testing.T) { defer cleanup() ctx := context.Background() - execution := &TaskExecution{ + execution := &model.TaskExecution{ TaskID: "test_fail_001", TaskType: "system:cleanup", TaskName: "清理未使用上传", - Status: TaskExecutionStatusPending, + Status: model.TaskExecutionStatusPending, Retryable: true, MaxRetry: 3, TriggeredBy: "manual", @@ -190,7 +192,7 @@ func TestUpdateTaskExecutionFailed(t *testing.T) { // 标记为失败 now := time.Now() - execution.Status = TaskExecutionStatusFailed + execution.Status = model.TaskExecutionStatusFailed execution.StartedAt = &now execution.FinishedAt = &now execution.Duration = 200 @@ -200,7 +202,7 @@ func TestUpdateTaskExecutionFailed(t *testing.T) { found, err := GetTaskExecutionByTaskID(ctx, "test_fail_001") require.NoError(t, err) - assert.Equal(t, TaskExecutionStatusFailed, found.Status) + assert.Equal(t, model.TaskExecutionStatusFailed, found.Status) assert.Equal(t, "S3 连接超时", found.ErrorMessage) assert.Equal(t, int64(200), found.Duration) } @@ -210,11 +212,11 @@ func TestUpdateTaskExecutionDoesNotPersistBufferedLog(t *testing.T) { defer cleanup() ctx := context.Background() - execution := &TaskExecution{ + execution := &model.TaskExecution{ TaskID: "test_omit_log_001", TaskType: "system:cleanup", TaskName: "清理未使用上传", - Status: TaskExecutionStatusPending, + Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", } err := CreateTaskExecution(ctx, execution) @@ -226,15 +228,15 @@ func TestUpdateTaskExecutionDoesNotPersistBufferedLog(t *testing.T) { assert.Empty(t, execution.Log) - execution.Status = TaskExecutionStatusSucceeded + execution.Status = model.TaskExecutionStatusSucceeded execution.Duration = 100 err = UpdateTaskExecution(ctx, execution) require.NoError(t, err) - var persisted TaskExecution + var persisted model.TaskExecution err = db.DB(ctx).Where("task_id = ?", "test_omit_log_001").First(&persisted).Error require.NoError(t, err) - assert.Equal(t, TaskExecutionStatusSucceeded, persisted.Status) + assert.Equal(t, model.TaskExecutionStatusSucceeded, persisted.Status) assert.Empty(t, persisted.Log) found, err := GetTaskExecutionByTaskID(ctx, "test_omit_log_001") @@ -247,11 +249,11 @@ func TestAppendTaskExecutionLog(t *testing.T) { defer cleanup() ctx := context.Background() - execution := &TaskExecution{ + execution := &model.TaskExecution{ TaskID: "test_log_001", TaskType: "system:cleanup", TaskName: "清理未使用上传", - Status: TaskExecutionStatusPending, + Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", } err := CreateTaskExecution(ctx, execution) @@ -274,7 +276,7 @@ func TestAppendTaskExecutionLog(t *testing.T) { assert.Contains(t, found.Log, "本批次找到 42 个待清理文件") assert.Contains(t, found.Log, "清理完成,共删除 42 个文件") - var persisted TaskExecution + var persisted model.TaskExecution err = db.DB(ctx).Where("task_id = ?", "test_log_001").First(&persisted).Error require.NoError(t, err) assert.Empty(t, persisted.Log) @@ -332,11 +334,11 @@ func TestGetTaskExecutionLogPrefersRedis(t *testing.T) { defer cleanup() ctx := context.Background() - execution := &TaskExecution{ + execution := &model.TaskExecution{ TaskID: "redis_priority_001", TaskType: "system:cleanup", TaskName: "清理未使用上传", - Status: TaskExecutionStatusRunning, + Status: model.TaskExecutionStatusRunning, Log: "数据库旧日志", TriggeredBy: "manual", } @@ -357,12 +359,12 @@ func TestListTaskExecutions(t *testing.T) { ctx := context.Background() // 创建多条记录,包含不同状态和类型 - records := []*TaskExecution{ - {TaskID: "list_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusSucceeded, TriggeredBy: "manual"}, - {TaskID: "list_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusFailed, TriggeredBy: "system"}, - {TaskID: "list_003", TaskType: "other:task", TaskName: "其他任务", Status: TaskExecutionStatusPending, TriggeredBy: "manual"}, - {TaskID: "list_004", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: TaskExecutionStatusRunning, TriggeredBy: "manual"}, - {TaskID: "list_005", TaskType: "other:task", TaskName: "其他任务", Status: TaskExecutionStatusSucceeded, TriggeredBy: "system"}, + records := []*model.TaskExecution{ + {TaskID: "list_001", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "manual"}, + {TaskID: "list_002", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusFailed, TriggeredBy: "system"}, + {TaskID: "list_003", TaskType: "other:task", TaskName: "其他任务", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual"}, + {TaskID: "list_004", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusRunning, TriggeredBy: "manual"}, + {TaskID: "list_005", TaskType: "other:task", TaskName: "其他任务", Status: model.TaskExecutionStatusSucceeded, TriggeredBy: "system"}, } for _, r := range records { err := CreateTaskExecution(ctx, r) @@ -372,7 +374,7 @@ func TestListTaskExecutions(t *testing.T) { require.NoError(t, err) // 查询全部(分页) - items, total, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 1, PageSize: 10}) + items, total, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 1, PageSize: 10}) require.NoError(t, err) assert.Equal(t, int64(5), total) assert.Len(t, items, 5) @@ -383,24 +385,24 @@ func TestListTaskExecutions(t *testing.T) { } // 按状态筛选:failed - items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Status: "failed", Page: 1, PageSize: 10}) + items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Status: "failed", Page: 1, PageSize: 10}) require.NoError(t, err) assert.Equal(t, int64(1), total) assert.Len(t, items, 1) assert.Equal(t, "list_002", items[0].TaskID) // 按类型筛选 - _, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{TaskType: "other:task", Page: 1, PageSize: 10}) + _, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{TaskType: "other:task", Page: 1, PageSize: 10}) require.NoError(t, err) assert.Equal(t, int64(2), total) // 分页测试 - items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 1, PageSize: 2}) + items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 1, PageSize: 2}) require.NoError(t, err) assert.Equal(t, int64(5), total) assert.Len(t, items, 2) - items2, total2, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{Page: 2, PageSize: 2}) + items2, total2, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Page: 2, PageSize: 2}) require.NoError(t, err) assert.Equal(t, int64(5), total2) assert.Len(t, items2, 2) @@ -409,7 +411,7 @@ func TestListTaskExecutions(t *testing.T) { assert.NotEqual(t, items[0].ID, items2[0].ID) // 状态 + 类型组合筛选 - items, total, err = ListTaskExecutions(ctx, ListTaskExecutionsRequest{Status: "succeeded", TaskType: "system:cleanup", Page: 1, PageSize: 10}) + items, total, err = ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{Status: "succeeded", TaskType: "system:cleanup", Page: 1, PageSize: 10}) require.NoError(t, err) assert.Equal(t, int64(1), total) assert.Equal(t, "list_001", items[0].TaskID) @@ -421,7 +423,7 @@ func TestListTaskExecutionsDefaultPaging(t *testing.T) { ctx := context.Background() // 不传分页参数,应使用默认值 page=1, pageSize=20 - items, total, err := ListTaskExecutions(ctx, ListTaskExecutionsRequest{}) + items, total, err := ListTaskExecutions(ctx, model.ListTaskExecutionsRequest{}) require.NoError(t, err) assert.Equal(t, int64(0), total) assert.Len(t, items, 0) @@ -434,14 +436,14 @@ func TestCleanupTaskExecutionLogs(t *testing.T) { now := time.Date(2026, 6, 17, 12, 0, 0, 0, time.UTC) for i := 0; i < 31; i++ { - createTaskExecutionForCleanup(t, ctx, fmt.Sprintf("high_recent_%02d", i), "high:task", TaskExecutionStatusSucceeded, now.Add(-2*time.Hour)) + createTaskExecutionForCleanup(t, ctx, fmt.Sprintf("high_recent_%02d", i), "high:task", model.TaskExecutionStatusSucceeded, now.Add(-2*time.Hour)) } - createTaskExecutionForCleanup(t, ctx, "high_old_4d", "high:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -4)) - createTaskExecutionForCleanup(t, ctx, "high_old_40d", "high:task", TaskExecutionStatusFailed, now.AddDate(0, 0, -40)) - createTaskExecutionForCleanup(t, ctx, "high_running_old", "high:task", TaskExecutionStatusRunning, now.AddDate(0, 0, -10)) - createTaskExecutionForCleanup(t, ctx, "low_old_31d", "low:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -31)) - createTaskExecutionForCleanup(t, ctx, "low_recent_29d", "low:task", TaskExecutionStatusSucceeded, now.AddDate(0, 0, -29)) - createTaskExecutionForCleanup(t, ctx, "low_pending_old", "low:task", TaskExecutionStatusPending, now.AddDate(0, 0, -45)) + createTaskExecutionForCleanup(t, ctx, "high_old_4d", "high:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -4)) + createTaskExecutionForCleanup(t, ctx, "high_old_40d", "high:task", model.TaskExecutionStatusFailed, now.AddDate(0, 0, -40)) + createTaskExecutionForCleanup(t, ctx, "high_running_old", "high:task", model.TaskExecutionStatusRunning, now.AddDate(0, 0, -10)) + createTaskExecutionForCleanup(t, ctx, "low_old_31d", "low:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -31)) + createTaskExecutionForCleanup(t, ctx, "low_recent_29d", "low:task", model.TaskExecutionStatusSucceeded, now.AddDate(0, 0, -29)) + createTaskExecutionForCleanup(t, ctx, "low_pending_old", "low:task", model.TaskExecutionStatusPending, now.AddDate(0, 0, -45)) stats, err := CleanupTaskExecutionLogs(ctx, now) require.NoError(t, err) @@ -450,27 +452,27 @@ func TestCleanupTaskExecutionLogs(t *testing.T) { for _, taskID := range []string{"high_old_4d", "high_old_40d", "low_old_31d"} { var count int64 - err := db.DB(ctx).Model(&TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error + err := db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error require.NoError(t, err) assert.Equal(t, int64(0), count, "CleanupTaskExecutionLogs(%s) should delete expired log", taskID) } for _, taskID := range []string{"high_recent_00", "high_running_old", "low_recent_29d", "low_pending_old"} { var count int64 - err := db.DB(ctx).Model(&TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error + err := db.DB(ctx).Model(&model.TaskExecution{}).Where("task_id = ?", taskID).Count(&count).Error require.NoError(t, err) assert.Equal(t, int64(1), count, "CleanupTaskExecutionLogs(%s) should keep retained log", taskID) } } func TestTaskExecutionTableName(t *testing.T) { - execution := TaskExecution{} + execution := model.TaskExecution{} assert.Equal(t, "w_task_executions", execution.TableName()) } -func createTaskExecutionForCleanup(t *testing.T, ctx context.Context, taskID string, taskType string, status TaskExecutionStatus, createdAt time.Time) { +func createTaskExecutionForCleanup(t *testing.T, ctx context.Context, taskID string, taskType string, status model.TaskExecutionStatus, createdAt time.Time) { t.Helper() - execution := &TaskExecution{ + execution := &model.TaskExecution{ TaskID: taskID, TaskType: taskType, TaskName: taskType, diff --git a/internal/repository/upload.go b/internal/repository/upload.go index 4d74a745..3bad47e7 100644 --- a/internal/repository/upload.go +++ b/internal/repository/upload.go @@ -6,10 +6,13 @@ package repository import ( "context" "strings" + "time" + + "gorm.io/gorm" + "gorm.io/gorm/clause" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" - "gorm.io/gorm" ) // UploadListFilter filters paginated upload queries. @@ -22,6 +25,22 @@ type UploadListFilter struct { PageSize int } +// UploadStorageObject is a distinct active object path with aggregated metadata for migration. +type UploadStorageObject struct { + FilePath string `gorm:"column:file_path"` + FileSize int64 `gorm:"column:file_size"` + MimeType string `gorm:"column:mime_type"` + Hash string `gorm:"column:hash"` +} + +// RunInTransaction executes fn inside a database transaction. +// Prefer domain-specific repository methods when the full operation can live in repository. +// Upload package multi-step flows (lock + soft-delete + stats) use this boundary so apps +// do not call db.DB directly. +func RunInTransaction(ctx context.Context, fn func(tx *gorm.DB) error) error { + return db.DB(ctx).Transaction(fn) +} + // ListUploads returns paginated upload records matching the filter. func ListUploads(ctx context.Context, filter UploadListFilter) (int64, []model.Upload, error) { query := db.DB(ctx).Model(&model.Upload{}). @@ -62,6 +81,28 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (model.Upload, error) { return upload, nil } +// GetCacheableUploadByID loads a pending or used upload by ID (for metadata cache DB fallback). +func GetCacheableUploadByID(ctx context.Context, id uint64) (model.Upload, error) { + 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 + } + return upload, nil +} + +// GetUploadByIDForUpdateTx loads and row-locks an upload by ID within an existing transaction. +func GetUploadByIDForUpdateTx(tx *gorm.DB, id uint64) (model.Upload, error) { + var upload model.Upload + if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}). + Where("id = ?", id). + First(&upload).Error; err != nil { + return model.Upload{}, err + } + return upload, nil +} + // SoftDeleteUpload marks an active upload as deleted and reports whether the row transitioned. // External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this. func SoftDeleteUpload(ctx context.Context, upload *model.Upload) (int64, error) { @@ -131,6 +172,108 @@ func ListUploadsByIDs(ctx context.Context, ids []uint64) ([]model.Upload, error) return uploads, nil } +// CountActiveUploads returns the number of non-deleted upload records. +func CountActiveUploads(ctx context.Context) (int64, error) { + var count int64 + err := db.DB(ctx).Model(&model.Upload{}). + Where("status != ?", model.UploadStatusDeleted). + Count(&count).Error + return count, err +} + +// ListPendingUploadsOlderThan returns pending uploads created before olderThan, after lastID, ordered by id. +func ListPendingUploadsOlderThan(ctx context.Context, lastID uint64, olderThan time.Time, limit int) ([]model.Upload, error) { + var uploads []model.Upload + err := db.DB(ctx). + Where("id > ? AND status = ? AND created_at < ?", lastID, model.UploadStatusPending, olderThan). + Order("id ASC"). + Limit(limit). + Find(&uploads).Error + return uploads, err +} + +// ListActiveImageUploadsAfterID returns non-deleted image uploads with id greater than lastID. +func ListActiveImageUploadsAfterID(ctx context.Context, lastID uint64, limit int) ([]model.Upload, error) { + var uploads []model.Upload + err := db.DB(ctx). + Where("id > ? AND status != ? AND (LOWER(mime_type) LIKE ? OR LOWER(extension) IN ?)", + lastID, + model.UploadStatusDeleted, + "image/%", + []string{"jpg", "jpeg", "png", "webp", "gif"}, + ). + Order("id ASC"). + Limit(limit). + Find(&uploads).Error + return uploads, err +} + +// CountDistinctActiveFilePaths returns the number of distinct non-deleted upload file paths. +func CountDistinctActiveFilePaths(ctx context.Context) (int64, error) { + var count int64 + err := db.DB(ctx).Model(&model.Upload{}). + Where("status != ?", model.UploadStatusDeleted). + Distinct("file_path"). + Count(&count).Error + return count, err +} + +// ListDistinctActiveStorageObjects returns a page of distinct active file paths ordered by path. +// When afterFilePath is non-empty, only paths strictly greater than it are returned. +func ListDistinctActiveStorageObjects(ctx context.Context, afterFilePath string, limit int) ([]UploadStorageObject, error) { + var objects []UploadStorageObject + query := db.DB(ctx).Model(&model.Upload{}). + Select("file_path, MAX(file_size) AS file_size, MAX(mime_type) AS mime_type, MAX(hash) AS hash"). + Where("status != ?", model.UploadStatusDeleted) + if afterFilePath != "" { + query = query.Where("file_path > ?", afterFilePath) + } + err := query.Group("file_path"). + Order("file_path ASC"). + Limit(limit). + Scan(&objects).Error + return objects, err +} + +// UpdateActiveUploadsFilePath rewrites file_path for all non-deleted uploads matching oldPath. +func UpdateActiveUploadsFilePath(ctx context.Context, oldPath, newPath string) error { + return db.DB(ctx).Model(&model.Upload{}). + Where("file_path = ? AND status != ?", oldPath, model.UploadStatusDeleted). + Update("file_path", newPath).Error +} + +// MarkActiveUploadsDeletedByFilePath marks all non-deleted uploads with the given path as deleted +// and returns the rows that transitioned (for stats adjustment). +func MarkActiveUploadsDeletedByFilePath(ctx context.Context, filePath string) ([]model.Upload, error) { + var affected []model.Upload + err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := tx. + Where("file_path = ? AND status != ?", filePath, model.UploadStatusDeleted). + Find(&affected).Error; err != nil { + return err + } + if len(affected) == 0 { + return nil + } + return tx.Model(&model.Upload{}). + Where("file_path = ?", filePath). + Update("status", model.UploadStatusDeleted).Error + }) + if err != nil { + return nil, err + } + return affected, nil +} + +// ListActiveUploadsTx returns all non-deleted uploads within an existing transaction. +func ListActiveUploadsTx(tx *gorm.DB) ([]model.Upload, error) { + var uploads []model.Upload + if err := tx.Where("status != ?", model.UploadStatusDeleted).Find(&uploads).Error; err != nil { + return nil, err + } + return uploads, nil +} + // UploadQuery returns a scoped GORM query for uploads. func UploadQuery(ctx context.Context) *gorm.DB { return db.DB(ctx).Model(&model.Upload{}) diff --git a/internal/repository/upload_stat.go b/internal/repository/upload_stat.go index eec59276..7e224e19 100644 --- a/internal/repository/upload_stat.go +++ b/internal/repository/upload_stat.go @@ -5,6 +5,10 @@ package repository import ( "context" + "time" + + "gorm.io/gorm" + "gorm.io/gorm/clause" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" @@ -18,3 +22,76 @@ func ListUploadStats(ctx context.Context) ([]model.UploadStat, error) { } return stats, nil } + +// GetTotalUploadStat returns the aggregate total-dimension stats row. +func GetTotalUploadStat(ctx context.Context) (model.UploadStat, error) { + var total model.UploadStat + if err := db.DB(ctx). + Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, ""). + First(&total).Error; err != nil { + return model.UploadStat{}, err + } + return total, nil +} + +// ListUploadStatsByDimension returns stats rows for a single dimension. +func ListUploadStatsByDimension(ctx context.Context, dimension string) ([]model.UploadStat, error) { + var rows []model.UploadStat + if err := db.DB(ctx).Where("dimension = ?", dimension).Find(&rows).Error; err != nil { + return nil, err + } + return rows, nil +} + +// DeleteAllUploadStatsTx removes every row from w_upload_stats within a transaction. +func DeleteAllUploadStatsTx(tx *gorm.DB) error { + return tx.Where("1 = 1").Delete(&model.UploadStat{}).Error +} + +// UpsertUploadStatDeltaTx applies an incremental count/size delta for one dimension key. +func UpsertUploadStatDeltaTx(tx *gorm.DB, dimension, key string, countDelta, sizeDelta int64) error { + return tx.Clauses(clause.OnConflict{ + Columns: []clause.Column{ + {Name: "dimension"}, + {Name: "stat_key"}, + }, + DoUpdates: clause.Assignments(map[string]any{ + "file_count": gorm.Expr( + "CASE WHEN w_upload_stats.file_count + ? < 0 THEN 0 ELSE w_upload_stats.file_count + ? END", + countDelta, + countDelta, + ), + "file_size": gorm.Expr( + "CASE WHEN w_upload_stats.file_size + ? < 0 THEN 0 ELSE w_upload_stats.file_size + ? END", + sizeDelta, + sizeDelta, + ), + "updated_at": time.Now(), + }), + }).Create(&model.UploadStat{ + Dimension: dimension, + StatKey: key, + FileCount: countDelta, + FileSize: sizeDelta, + }).Error +} + +// RebuildUploadStats clears w_upload_stats and re-applies deltas for every active upload +// inside a single transaction. applyDelta should apply +1 stats for one upload row. +func RebuildUploadStats(ctx context.Context, applyDelta func(tx *gorm.DB, upload *model.Upload) error) error { + return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := DeleteAllUploadStatsTx(tx); err != nil { + return err + } + uploads, err := ListActiveUploadsTx(tx) + if err != nil { + return err + } + for i := range uploads { + if err := applyDelta(tx, &uploads[i]); err != nil { + return err + } + } + return nil + }) +} diff --git a/internal/repository/user.go b/internal/repository/user.go index f1dc1050..1e54051b 100644 --- a/internal/repository/user.go +++ b/internal/repository/user.go @@ -5,8 +5,11 @@ package repository import ( "context" + "errors" + "time" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" + "github.com/Rain-kl/Wavelet/internal/infra/persistence/idgen" "github.com/Rain-kl/Wavelet/internal/model" "gorm.io/gorm" ) @@ -175,7 +178,117 @@ func ListUsersByIDs(ctx context.Context, ids []uint64) ([]model.User, error) { return users, nil } +// ListUserIDsByUsernameContains returns user IDs whose username contains the given fragment. +func ListUserIDsByUsernameContains(ctx context.Context, username string) ([]uint64, error) { + if username == "" { + return []uint64{}, nil + } + var userIDs []uint64 + if err := db.DB(ctx).Model(&model.User{}). + Where("username LIKE ?", "%"+username+"%"). + Pluck("id", &userIDs).Error; err != nil { + return nil, err + } + return userIDs, nil +} + // UpdateUser updates all fields of an existing user. func UpdateUser(ctx context.Context, user *model.User) error { return db.DB(ctx).Save(user).Error } + +// CreateUserFromOAuth creates a user from OAuth profile data and fills userOut. +func CreateUserFromOAuth(ctx context.Context, userOut *model.User, oauthInfo *model.OAuthUserInfo) error { + now := time.Now() + userID := oauthInfo.GetID() + newUser := model.User{ + ID: userID, + Username: oauthInfo.Username, + Nickname: oauthInfo.Name, + Email: oauthInfo.Email, + AvatarURL: oauthInfo.AvatarURL, + IsActive: oauthInfo.Active, + LastLoginAt: now, + IsAdmin: false, + } + if newUser.ID == 0 { + newUser.ID = idgen.NextUint64ID() + } + if err := db.DB(ctx).Create(&newUser).Error; err != nil { + return err + } + *userOut = newUser + return nil +} + +// ListUsernamesMatchingBase returns usernames equal to base or prefixed with base+"-". +func ListUsernamesMatchingBase(ctx context.Context, base string) ([]string, error) { + var names []string + if err := db.DB(ctx).Model(&model.User{}). + Where("username = ? OR username LIKE ?", base, base+"-%"). + Pluck("username", &names).Error; err != nil { + return nil, err + } + return names, nil +} + +// GetActiveUserByID loads a user by ID who is active. +func GetActiveUserByID(ctx context.Context, id uint64) (model.User, error) { + var user model.User + if err := db.DB(ctx).Where("id = ? AND is_active = ?", id, true).First(&user).Error; err != nil { + return model.User{}, err + } + return user, nil +} + +// GetUserByUsernameOrEmail loads a user by username or email. +func GetUserByUsernameOrEmail(ctx context.Context, input string) (model.User, error) { + var user model.User + if err := db.DB(ctx).Where("username = ? OR email = ?", input, input).First(&user).Error; err != nil { + return model.User{}, err + } + return user, nil +} + +// CountUsersByEmailExceptID counts users with the email excluding a given user id. +func CountUsersByEmailExceptID(ctx context.Context, email string, exceptID uint64) (int64, error) { + var count int64 + if err := db.DB(ctx).Model(&model.User{}).Where("email = ? AND id != ?", email, exceptID).Count(&count).Error; err != nil { + return 0, err + } + return count, nil +} + +// UpdateUserLastLoginAt updates only last_login_at for a user. +func UpdateUserLastLoginAt(ctx context.Context, userID uint64, at time.Time) error { + return db.DB(ctx).Model(&model.User{}).Where("id = ?", userID).Update("last_login_at", at).Error +} + +// UpdateUserPassword updates only the password hash for a user. +func UpdateUserPassword(ctx context.Context, userID uint64, passwordHash string) error { + return db.DB(ctx).Model(&model.User{}).Where("id = ?", userID).Update("password", passwordHash).Error +} + +// RegisterUserWithChecks validates username/email uniqueness then creates the user. +func RegisterUserWithChecks(ctx context.Context, user *model.User) error { + var count int64 + if err := db.DB(ctx).Model(&model.User{}).Where("username = ?", user.Username).Count(&count).Error; err != nil { + return err + } + if count > 0 { + return errors.New("用户名已存在") + } + if user.Email != "" { + var emailCount int64 + if err := db.DB(ctx).Model(&model.User{}).Where("email = ?", user.Email).Count(&emailCount).Error; err != nil { + return err + } + if emailCount > 0 { + return errors.New("该邮箱已被其他账号绑定") + } + } + if user.ID == 0 { + user.ID = idgen.NextUint64ID() + } + return db.DB(ctx).Create(user).Error +} diff --git a/scripts/live_ch_smoke/main.go b/scripts/live_ch_smoke/main.go index 3cdfa9a7..e0ffa5a2 100644 --- a/scripts/live_ch_smoke/main.go +++ b/scripts/live_ch_smoke/main.go @@ -10,6 +10,8 @@ import ( "os" "time" + "github.com/Rain-kl/Wavelet/internal/repository" + "github.com/Rain-kl/Wavelet/internal/apps/openflare/chwriter" db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/model" @@ -38,7 +40,7 @@ func run() error { now := time.Now().UTC() nodeID := "e2e-app-" + now.Format("150405") - if err := model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ + if err := repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ NodeID: nodeID, CapturedAt: now, CPUUsagePercent: 41.2, MemoryUsedBytes: 123, MemoryTotalBytes: 1000, StorageUsedBytes: 456, StorageTotalBytes: 2000, @@ -64,7 +66,7 @@ func run() error { func waitForSnapshot(ctx context.Context, nodeID string, now time.Time) error { deadline := time.Now().Add(flushWaitTimeout) for time.Now().Before(deadline) { - rows, err := model.ListOpenFlareMetricSnapshotsSince(ctx, nodeID, now.Add(-time.Minute), listLimit) + rows, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, nodeID, now.Add(-time.Minute), listLimit) if err != nil { return fmt.Errorf("list: %w", err) } @@ -78,7 +80,7 @@ func waitForSnapshot(ctx context.Context, nodeID string, now time.Time) error { } func assertLatestIncludes(ctx context.Context, nodeID string, now time.Time) error { - latest, err := model.ListOpenFlareLatestMetricSnapshotsSince(ctx, "", now.Add(-time.Hour)) + latest, err := repository.ListOpenFlareLatestMetricSnapshotsSince(ctx, "", now.Add(-time.Hour)) if err != nil { return fmt.Errorf("latest: %w", err) }