refactor(repository): 收敛 model/repository 分层为唯一持久化入口

将 OpenFlare 与平台业务的数据访问从 model 与 apps 直连迁入 repository,
model 仅保留实体与无 IO 规则;补充 code-check 架构守卫与开发规范。
This commit is contained in:
ryan
2026-07-24 17:00:17 +08:00
parent 23a5488203
commit 943818f7d4
184 changed files with 5592 additions and 4364 deletions
+1 -1
View File
@@ -115,7 +115,7 @@ internal/
## 核心开发步骤 (Step-by-Step Flow) ## 核心开发步骤 (Step-by-Step Flow)
### 步骤 1:数据库定义与迁移 ### 步骤 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`) ### 步骤 2:在模块内实现业务逻辑 (`logics.go` / `service.go`)
业务逻辑逻辑应当实现于 `internal/apps/custom/` 目录下: 业务逻辑逻辑应当实现于 `internal/apps/custom/` 目录下:
+3 -2
View File
@@ -19,7 +19,8 @@ description: "Wavelet 项目专用:新增或修改 Asynq 异步任务、后台
- `internal/infra/task/worker/worker.go`:Worker 路由和队列 - `internal/infra/task/worker/worker.go`:Worker 路由和队列
- `internal/infra/task/scheduler/scheduler.go`:定时调度 - `internal/infra/task/scheduler/scheduler.go`:定时调度
- `internal/apps/admin/task/routers.go`:Admin 任务 API - `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)。 需要模板时阅读 [references/CODE-EXAMPLES.md](references/CODE-EXAMPLES.md)。
@@ -42,7 +43,7 @@ description: "Wavelet 项目专用:新增或修改 Asynq 异步任务、后台
- 成功返回 `&task.TaskResult{Message: ..., Detail: ...}`。 - 成功返回 `&task.TaskResult{Message: ..., Detail: ...}`。
- 失败返回 error,由任务框架处理状态和重试。 - 失败返回 error,由任务框架处理状态和重试。
- 不要吞掉关键错误。 - 不要吞掉关键错误。
- 复杂 SQL 放到 `internal/model/` 或模块内的业务服务层(如 `internal/apps/<module>/service.go` 或 `logics.go`)。 - 持久化只通过 `internal/repository/`(唯一入口);业务编排放模块内 `logics.go` / `service.go`。`internal/model` 仅实体/DTO,禁止 CRUD 与 DB 访问。
### 注册 ### 注册
+11 -9
View File
@@ -14,7 +14,7 @@ description: "Wavelet 项目专用:当新增或修改启动时设置、数据
Wavelet 当前有两套设置入口: Wavelet 当前有两套设置入口:
- 启动时设置:来自 `config.yaml` 或环境变量,适合进程启动前必须确定、通常不热更新的基础配置。 - 启动时设置:来自 `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/goose/postgres/*.sql` 和 `internal/infra/persistence/migrator/goose/sqlite/*.sql`: `system_configs` 表结构、初始化 seed、后续升级迁移。
- `internal/infra/persistence/migrator/migrator.go`: goose 迁移入口和 PostgreSQL/SQLite 方言选择。 - `internal/infra/persistence/migrator/migrator.go`: goose 迁移入口和 PostgreSQL/SQLite 方言选择。
- `internal/testhelper/test_helper.go`: Go 测试用默认系统配置 seed。 - `internal/testhelper/test_helper.go`: Go 测试用默认系统配置 seed。
@@ -61,12 +62,13 @@ Wavelet 当前有两套设置入口:
- 如果相关 Go 包测试依赖默认配置,同步 `internal/testhelper/test_helper.go` 的 `seedDefaultConfigs` 和公共 key 列表。 - 如果相关 Go 包测试依赖默认配置,同步 `internal/testhelper/test_helper.go` 的 `seedDefaultConfigs` 和公共 key 列表。
3. 读取配置。 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`。 - 运行时可热更新的规则不要放进 `config.Config`;启动时设置才走 `internal/infra/config/model.go` 和 `config.example.yaml`。
- 不要在 handler 或业务代码里直接读 `os.Getenv()`。 - 不要在 handler 或业务代码里直接读 `os.Getenv()`。
4. 如果前端需要未登录或全局消费,暴露为公共可见配置。 4. 如果前端需要未登录或全局消费,暴露为公共可见配置。
- 把该配置的 `visibility` 设为 `1`,`GetPublicConfig` 会通过 `model.ListVisibleSystemConfigs` 返回所有可见 key/value。 - 把该配置的 `visibility` 设为 `1`,`GetPublicConfig` 会通过 `repository.ListVisibleSystemConfigs` 返回所有可见 key/value。
- `/api/v1/config/public` 的 `data` 是动态对象:后端返回 `map[string]string`,前端类型是 `Record<string, string | undefined>`。 - `/api/v1/config/public` 的 `data` 是动态对象:后端返回 `map[string]string`,前端类型是 `Record<string, string | undefined>`。
- 前端读取时按配置 key 访问,必要时在消费侧把字符串转换为 boolean/number/JSON。 - 前端读取时按配置 key 访问,必要时在消费侧把字符串转换为 boolean/number/JSON。
- 检查使用方的 query key,更新后需要 invalidate `["public-config"]`。 - 检查使用方的 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`。 - 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"`。 - 公共响应:`/api/v1/config/public` 的 `data.feature_enabled` 为字符串 `"true"` 或 `"false"`。
- 前端图形控件:`Switch`,保存时写 `"true"` / `"false"`。 - 前端图形控件:`Switch`,保存时写 `"true"` / `"false"`。
@@ -109,13 +111,13 @@ Wavelet 当前有两套设置入口:
- model key:`ConfigKeyMaxSomething = "max_something"`。 - model key:`ConfigKeyMaxSomething = "max_something"`。
- goose SQL 默认值:例如 `"5"`,`type` 通常为 `"business"`,只有前端公共消费时才设 `visibility=1`。 - goose SQL 默认值:例如 `"5"`,`type` 通常为 `"business"`,只有前端公共消费时才设 `visibility=1`。
- 后端读取:`model.GetIntByKey` 或 `model.GetDecimalByKey`。 - 后端读取:`repository.GetIntByKey` 或 `repository.GetDecimalByKey`。
- 前端图形控件:`Input type="number"` 或合适的 shadcn 数值控件;保存前做最小必要校验,错误用 toast。 - 前端图形控件:`Input type="number"` 或合适的 shadcn 数值控件;保存前做最小必要校验,错误用 toast。
### JSON 设置 ### JSON 设置
- 默认值使用合法 JSON,例如 `"{}"` 或 `"[]"`。 - 默认值使用合法 JSON,例如 `"{}"` 或 `"[]"`。
- 在 model 或 service 层提供解析函数,像 `GetMenuDisplayConfig` 一样把 JSON 解析错误包装成清晰错误。 - 在 repository 或业务 logics 中提供解析函数,像 `repository.GetMenuDisplayConfig` 一样把 JSON 解析错误包装成清晰错误;不要在 model 中做 IO。
- 前端不要直接拼接 JSON 字符串;用 `JSON.stringify` 写入,用类型化对象在组件中操作。 - 前端不要直接拼接 JSON 字符串;用 `JSON.stringify` 写入,用类型化对象在组件中操作。
## 验证 ## 验证
@@ -125,7 +127,7 @@ Wavelet 当前有两套设置入口:
- 新增或修改系统配置默认值、visibility 或公共配置读取:至少运行相关 Go 包测试,例如: - 新增或修改系统配置默认值、visibility 或公共配置读取:至少运行相关 Go 包测试,例如:
```bash ```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。 - 新增 goose 迁移后,至少用当前数据库方言跑一次迁移;如果 SQL 同时改了 PostgreSQL 和 SQLite,尽量覆盖两种方言。涉及 schema/seed 的任务还应遵循 database-migration skill。
+13 -4
View File
@@ -59,6 +59,13 @@
- 在完成代码开发后或者 git 提交前必须运行 `make format` 格式化代码。 - 在完成代码开发后或者 git 提交前必须运行 `make format` 格式化代码。
- 需要缓存或文件管理能力时,必须复用现有平台实现,禁止在业务包中自行创建缓存目录、直接管理缓存文件或重复封装存储后端。 - 需要缓存或文件管理能力时,必须复用现有平台实现,禁止在业务包中自行创建缓存目录、直接管理缓存文件或重复封装存储后端。
- 文件摄取必须通过 `upload.Ingest`(`upload.PolicyCreate` / `PolicyDedupNewRecord` / `PolicyResolveExisting`);删除必须通过 `upload.Remove` 或 `upload.RemoveOwned`。禁止业务模块直接调用 `repository.CreateUpload` / `repository.SoftDeleteUpload`,禁止 `db.Create(&model.Upload{})` 旁路写 `w_uploads`。 - 文件摄取必须通过 `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` 入口显式装配。 - 禁止在 `init()` 中注册跨模块集成(任务 Handler、推送内置事件、域事件监听器、任务完成钩子)。统一通过 `internal/platform/bootstrap` 在 `internal/cmd` 入口显式装配。
- `internal/router/router.go` 的 `Serve()` 仅负责 HTTP 路由与中间件,禁止在其中执行 `SyncEvents`、`InitLogWriter` 等进程级运行时初始化。 - `internal/router/router.go` 的 `Serve()` 仅负责 HTTP 路由与中间件,禁止在其中执行 `SyncEvents`、`InitLogWriter` 等进程级运行时初始化。
- 核心业务模块(如 `oauth`、`user`)禁止直接 `import` `internal/apps/admin/push` 或 `custom_events` 触发通知;应通过 `internal/listener` 发射域事件,由 push 模块在 bootstrap 阶段订阅。 - 核心业务模块(如 `oauth`、`user`)禁止直接 `import` `internal/apps/admin/push` 或 `custom_events` 触发通知;应通过 `internal/listener` 发射域事件,由 push 模块在 bootstrap 阶段订阅。
@@ -115,7 +122,8 @@
- `internal/router/`:唯一的 HTTP 路由注册点。 - `internal/router/`:唯一的 HTTP 路由注册点。
- `internal/apps/`:按功能(Feature-based)组织的 HTTP Handler、中间件、内部服务与模块逻辑。移除全局 service 层,模块内部业务逻辑(如验证码业务逻辑管理器 `internal/apps/cap/manager.go`)均收敛于各自模块中;管理端模块位于 `internal/apps/admin/`。 - `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/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/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/diskcache/`:平台级磁盘字节缓存,通过 `diskcache.GetGlobalCache()` 提供 TTL、最大空间限制、LRU 淘汰、清空、状态统计和配置热更新。写入时使用 `DefaultExpiration`(全局默认 TTL)、正数 `time.Duration`(业务 TTL)或 `NoExpiration`(无 TTL,仍受空间限制和 LRU 淘汰)。
- `internal/infra/objectstore/`:S3 兼容对象存储适配,提供对象上传、读取、删除、CDN/代理读取及远端对象本地缓存。 - `internal/infra/objectstore/`:S3 兼容对象存储适配,提供对象上传、读取、删除、CDN/代理读取及远端对象本地缓存。
@@ -303,9 +311,10 @@ func doSomething(c *gin.Context) { response.AbortBadRequest(c, "...") }
数据库操作: 数据库操作:
- 简单查询可以直接从 model 层使用 GORM。 - **持久化只通过 `internal/repository`**(或 analytics 子包)。apps / logics 不要直接 `db.DB(ctx).Where...` 拼复杂查询;简单事务编排可在 logics 中调用多个 repository 方法。
- 管理员代码应首选 `db.DB(ctx)` 以获得链路追踪感知的 DB 访问。 - repository 内管理员/业务查询应使用 `db.DB(ctx)` 以获得链路追踪感知的 DB 访问。
- 不要在 Handler 中放置复杂的 SQL;将其移至 `internal/model/` 或模块内的业务服务层(如 `internal/apps/<module>/service.go` 或 `logics.go`)。 - 不要在 Handler 中放置 SQL;复杂查询放 `internal/repository/`,业务编排放 `internal/apps/<module>/logics.go`(或 `service.go`)。
- `internal/model` 只定义实体与无 IO 规则,不访问数据库。
- 在 `internal/infra/persistence/migrator/goose/` 下使用 goose SQL 迁移;不要添加基于 GORM AutoMigrate 的 Schema 升级。 - 在 `internal/infra/persistence/migrator/goose/` 下使用 goose SQL 迁移;不要添加基于 GORM AutoMigrate 的 Schema 升级。
- 不要创建物理数据库外键。改为关系字段添加显式索引。 - 不要创建物理数据库外键。改为关系字段添加显式索引。
- 数据库默认值必须与 Go 模型零值(`nil`、`0`、`false`、`""`)匹配,以避免意外的插入。 - 数据库默认值必须与 Go 模型零值(`nil`、`0`、`false`、`""`)匹配,以避免意外的插入。
+6
View File
@@ -38,6 +38,12 @@ build-embedded:
main.go main.go
code-check: 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 golangci-lint run
cd frontend && pnpm tsc --noEmit --jsx preserve && npx eslint . --max-warnings 0 cd frontend && pnpm tsc --noEmit --jsx preserve && npx eslint . --max-warnings 0
+6
View File
@@ -22,6 +22,12 @@ sidebar: false
## [unreleased] ## [unreleased]
### 改进
- 统一数据访问分层:业务持久化经 `internal/repository`,`internal/model` 仅保留实体与无 IO 领域规则,避免双轨 CRUD 与职责混淆。
- 构建检查增加 `internal/model` 禁止直接访问数据库/Redis 的架构守卫,并收敛 model 与 repository 的错误文案定义边界。
## [v3.4.3] - 2026-07-24 ## [v3.4.3] - 2026-07-24
### 新增 ### 新增
+2 -2
View File
@@ -94,9 +94,9 @@ OpenFlare 已收敛为**单 monorepo**(Go 模块 `github.com/Rain-kl/Wavelet`
| `internal/apps/openflare/` | OpenFlare 控制面业务域(`routers.go` + `logics.go`) | | `internal/apps/openflare/` | OpenFlare 控制面业务域(`routers.go` + `logics.go`) |
| `internal/apps/{admin,user,oauth,upload,cap,...}/` | Wavelet 平台能力(用户、认证、任务、推送等) | | `internal/apps/{admin,user,oauth,upload,cap,...}/` | Wavelet 平台能力(用户、认证、任务、推送等) |
| `internal/apps/openflare/{agent,relay,flared}/` | **Server 侧**边缘协议处理器(鉴权、心跳、WS) | | `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/infra/persistence/migrator/goose/` | goose SQL 迁移(PostgreSQL / SQLite / ClickHouse) |
| `internal/repository/` | 平台域数据访问层 | | `internal/repository/` | 数据访问层(平台 + OpenFlare 业务 CRUD、缓存、ClickHouse 分析读写);**唯一**持久化入口 |
| `internal/infra/task/` | Asynq 异步任务(Worker + Scheduler) | | `internal/infra/task/` | Asynq 异步任务(Worker + Scheduler) |
| `internal/infra/config/` | Viper 配置加载 | | `internal/infra/config/` | Viper 配置加载 |
| `internal/shared/` | 统一 API 响应封装(`response/`) | | `internal/shared/` | 统一 API 响应封装(`response/`) |
@@ -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`(在可接受时间内)
+2
View File
@@ -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 边界;生产环境验收边界见计划内验证记录。 * [Pages 项目部署源与 GitHub Releases 自动更新 V2](./20260719-pages-source-sync-v2.md):已完成 Remote URL / GitHub Release 来源、不可变部署、自动检查更新与安全回滚,并预留独立仓库构建 Provider 边界;生产环境验收边界见计划内验证记录。
## 使用建议 ## 使用建议
+8 -8
View File
@@ -49,7 +49,7 @@ type ToggleAuthSourceRequest struct {
// @Failure 500 {object} response.Any "内部错误" // @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/admin/auth-sources [get] // @Router /api/v1/admin/auth-sources [get]
func ListAuthSources(c *gin.Context) { func ListAuthSources(c *gin.Context) {
sources, err := model.GetAuthSources(c.Request.Context()) sources, err := repository.GetAuthSources(c.Request.Context())
if err != nil { if err != nil {
response.AbortInternal(c, err.Error()) response.AbortInternal(c, err.Error())
return return
@@ -88,7 +88,7 @@ func CreateAuthSource(c *gin.Context) {
Scopes: req.Scopes, Scopes: req.Scopes,
IconURL: req.IconURL, 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()) response.AbortBadRequest(c, err.Error())
return return
} }
@@ -126,7 +126,7 @@ func UpdateAuthSource(c *gin.Context) {
} }
// 记录更新前的 Discovery URL,以便更新成功后清除旧缓存条目。 // 记录更新前的 Discovery URL,以便更新成功后清除旧缓存条目。
existing, _ := model.GetAuthSourceByID(c.Request.Context(), id) existing, _ := repository.GetAuthSourceByID(c.Request.Context(), id)
source := model.AuthSource{ source := model.AuthSource{
ID: id, ID: id,
@@ -141,7 +141,7 @@ func UpdateAuthSource(c *gin.Context) {
IconURL: req.IconURL, IconURL: req.IconURL,
} }
keepSecret := source.ClientSecret == "" 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()) response.AbortBadRequest(c, err.Error())
return return
} }
@@ -154,7 +154,7 @@ func UpdateAuthSource(c *gin.Context) {
oauth.InvalidateOIDCProviderCache(normalizeIssuer(req.OpenIDDiscoveryURL)) oauth.InvalidateOIDCProviderCache(normalizeIssuer(req.OpenIDDiscoveryURL))
_ = repository.InvalidateAuthSourceCache(c.Request.Context()) _ = repository.InvalidateAuthSourceCache(c.Request.Context())
updated, err := model.GetAuthSourceByID(c.Request.Context(), id) updated, err := repository.GetAuthSourceByID(c.Request.Context(), id)
if err != nil { if err != nil {
response.AbortInternal(c, err.Error()) response.AbortInternal(c, err.Error())
return return
@@ -190,7 +190,7 @@ func ToggleAuthSource(c *gin.Context) {
return 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()) response.AbortBadRequest(c, err.Error())
return return
} }
@@ -216,7 +216,7 @@ func DeleteAuthSource(c *gin.Context) {
response.AbortBadRequest(c, err.Error()) response.AbortBadRequest(c, err.Error())
return 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()) response.AbortBadRequest(c, err.Error())
return return
} }
@@ -229,7 +229,7 @@ func parseSourceID(c *gin.Context) (uint64, error) {
if raw == "" { if raw == "" {
return 0, errors.New(admin.InvalidAuthSourceID) 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 { if err == nil {
return source.ID, nil return source.ID, nil
} }
+4 -9
View File
@@ -16,7 +16,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/admin" "github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/config"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" 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" analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -163,10 +163,7 @@ func buildAccessLogFilter(ctx context.Context, c *gin.Context) (analyticsrepo.Ac
username := c.Query("username") username := c.Query("username")
if username != "" { if username != "" {
var userIDs []uint64 userIDs, err := repository.ListUserIDsByUsernameContains(ctx, username)
err := db.DB(ctx).Model(&model.User{}).
Where("username LIKE ?", "%"+username+"%").
Pluck("id", &userIDs).Error
if err != nil { if err != nil {
return filter, fmt.Errorf("查询用户信息失败: %w", err) 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 }) userMap := make(map[uint64]struct{ Username, Nickname string })
var users []model.User if users, err := repository.ListUsersByIDs(ctx, userIDs); err == nil {
if err := db.DB(ctx).Where("id IN ?", userIDs).Find(&users).Error; err == nil {
for _, u := range users { for _, u := range users {
userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname} userMap[u.ID] = struct{ Username, Nickname string }{Username: u.Username, Nickname: u.Nickname}
} }
@@ -402,8 +398,7 @@ func GetLogsAnalytics(c *gin.Context) {
Username string Username string
Nickname string Nickname string
}) })
var users []model.User if users, errProfile := repository.ListUsersByIDs(ctx, userIDs); errProfile == nil {
if errProfile := db.DB(ctx).Where("id IN ?", userIDs).Find(&users).Error; errProfile == nil {
for _, u := range users { for _, u := range users {
userProfileMap[u.ID] = struct { userProfileMap[u.ID] = struct {
Username string Username string
+8 -10
View File
@@ -10,7 +10,6 @@ import (
"time" "time"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore" "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/model"
"github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
@@ -72,7 +71,7 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
req.Value = validatedVal 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{ updates := map[string]any{
"description": req.Description, "description": req.Description,
} }
@@ -84,7 +83,7 @@ func updateSystemConfig(ctx context.Context, key string, req UpdateSystemConfigR
updates["value"] = req.Value updates["value"] = req.Value
config.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 return err
} }
resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value) resolveStorageMigrationTasksOnDirectDriverUpdate(ctx, tx, key, originalDriver, req.Value)
@@ -116,13 +115,12 @@ func resolveStorageMigrationTasksOnDirectDriverUpdate(
return return
} }
if err := tx.Model(&model.TaskExecution{}). if err := repository.MarkFailedTaskExecutionsSucceededTx(
Where("task_type = ? AND status = ?", "storage:migrate", model.TaskExecutionStatusFailed). tx,
Updates(map[string]any{ "storage:migrate",
"status": model.TaskExecutionStatusSucceeded, "存储配置直接更新,故障迁移任务自动标记为已解决",
"result": "存储配置直接更新,故障迁移任务自动标记为已解决", time.Now(),
"finished_at": time.Now(), ); err != nil {
}).Error; err != nil {
logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err) logger.ErrorF(ctx, "自动更新迁移任务状态失败: %v", err)
} }
} }
+2 -5
View File
@@ -18,7 +18,6 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/cap" "github.com/Rain-kl/Wavelet/internal/apps/cap"
"github.com/Rain-kl/Wavelet/internal/apps/upload" "github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/infra/objectstore" "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/model"
"github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared/response" "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 { func validateMergedStorageConfig(ctx context.Context, currentCfg, newCfg, targetCfg objectstore.Config) error {
if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver { if newCfg.Driver != "" && newCfg.Driver != currentCfg.Driver {
var uploadCount int64 uploadCount, err := repository.CountActiveUploads(ctx)
if err := db.DB(ctx).Model(&model.Upload{}). if err != nil {
Where("status != ?", model.UploadStatusDeleted).
Count(&uploadCount).Error; err != nil {
return fmt.Errorf("检查存量文件失败: %w", err) return fmt.Errorf("检查存量文件失败: %w", err)
} }
if uploadCount > 0 { if uploadCount > 0 {
+10 -8
View File
@@ -11,6 +11,8 @@ import (
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/admin" "github.com/Rain-kl/Wavelet/internal/apps/admin"
"github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/infra/task/scheduler" "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 { if err != nil {
response.AbortInternal(c, err.Error()) response.AbortInternal(c, err.Error())
return return
@@ -153,7 +155,7 @@ func GetTaskExecution(c *gin.Context) {
return return
} }
execution, err := model.GetTaskExecutionByID(c.Request.Context(), id) execution, err := repository.GetTaskExecutionByID(c.Request.Context(), id)
if err != nil { if err != nil {
response.AbortNotFound(c, TaskNotFound) response.AbortNotFound(c, TaskNotFound)
return return
@@ -211,7 +213,7 @@ func RetryTask(c *gin.Context) {
// @Failure 403 {object} response.Any "无管理员权限" // @Failure 403 {object} response.Any "无管理员权限"
// @Router /api/v1/admin/tasks/schedules [get] // @Router /api/v1/admin/tasks/schedules [get]
func ListSchedules(c *gin.Context) { func ListSchedules(c *gin.Context) {
schedules, err := model.ListSchedules(c.Request.Context()) schedules, err := repository.ListSchedules(c.Request.Context())
if err != nil { if err != nil {
response.AbortInternal(c, err.Error()) response.AbortInternal(c, err.Error())
return return
@@ -289,7 +291,7 @@ func CreateSchedule(c *gin.Context) {
IsActive: *req.IsActive, 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)) response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return 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 { if err != nil {
response.AbortNotFound(c, ScheduleNotFound) response.AbortNotFound(c, ScheduleNotFound)
return return
@@ -381,7 +383,7 @@ func UpdateSchedule(c *gin.Context) {
schedule.Payload = string(validated) schedule.Payload = string(validated)
schedule.IsActive = *req.IsActive 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)) response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleSaveFailed, err))
return return
} }
@@ -422,7 +424,7 @@ func DeleteSchedule(c *gin.Context) {
response.AbortBadRequest(c, "无效的定时任务ID") response.AbortBadRequest(c, "无效的定时任务ID")
return return
} }
schedule, err := model.GetScheduleByID(c.Request.Context(), id) schedule, err := repository.GetScheduleByID(c.Request.Context(), id)
if err != nil { if err != nil {
response.AbortNotFound(c, ScheduleNotFound) response.AbortNotFound(c, ScheduleNotFound)
return return
@@ -432,7 +434,7 @@ func DeleteSchedule(c *gin.Context) {
return 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)) response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
return return
} }
+18 -16
View File
@@ -14,6 +14,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/apps/oauth"
uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task" uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task"
"github.com/Rain-kl/Wavelet/internal/apps/user" "github.com/Rain-kl/Wavelet/internal/apps/user"
@@ -154,8 +156,8 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
Payload: "{}", Payload: "{}",
IsActive: true, IsActive: true,
} }
require.NoError(t, model.CreateSchedule(ctx, internalSchedule)) require.NoError(t, repository.CreateSchedule(ctx, internalSchedule))
require.NoError(t, model.CreateSchedule(ctx, publicSchedule)) require.NoError(t, repository.CreateSchedule(ctx, publicSchedule))
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/tasks/schedules", nil) req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/tasks/schedules", nil)
w := httptest.NewRecorder() w := httptest.NewRecorder()
@@ -215,7 +217,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
Cron: "0 * * * *", Cron: "0 * * * *",
IsActive: true, IsActive: true,
} }
require.NoError(t, model.CreateSchedule(ctx, schedule)) require.NoError(t, repository.CreateSchedule(ctx, schedule))
isActive := false isActive := false
body, err := json.Marshal(UpdateScheduleRequest{ body, err := json.Marshal(UpdateScheduleRequest{
Name: "尝试修改内部排程", Name: "尝试修改内部排程",
@@ -235,7 +237,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
router.ServeHTTP(w, req) router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code) 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) require.NoError(t, err)
assert.Equal(t, "系统内部排程", unchanged.Name) assert.Equal(t, "系统内部排程", unchanged.Name)
assert.Equal(t, testInternalOnlyTaskType, unchanged.TaskType) assert.Equal(t, testInternalOnlyTaskType, unchanged.TaskType)
@@ -249,7 +251,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
Cron: "0 * * * *", Cron: "0 * * * *",
IsActive: true, IsActive: true,
} }
require.NoError(t, model.CreateSchedule(ctx, schedule)) require.NoError(t, repository.CreateSchedule(ctx, schedule))
isActive := true isActive := true
body, err := json.Marshal(UpdateScheduleRequest{ body, err := json.Marshal(UpdateScheduleRequest{
Name: "尝试切入内部任务", Name: "尝试切入内部任务",
@@ -269,7 +271,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
router.ServeHTTP(w, req) router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code) 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) require.NoError(t, err)
assert.Equal(t, "公开排程", unchanged.Name) assert.Equal(t, "公开排程", unchanged.Name)
assert.Equal(t, uploadtask.TaskTypeSystemCleanup, unchanged.TaskType) assert.Equal(t, uploadtask.TaskTypeSystemCleanup, unchanged.TaskType)
@@ -282,7 +284,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
Cron: "*/5 * * * *", Cron: "*/5 * * * *",
IsActive: true, IsActive: true,
} }
require.NoError(t, model.CreateSchedule(ctx, schedule)) require.NoError(t, repository.CreateSchedule(ctx, schedule))
req := httptest.NewRequest( req := httptest.NewRequest(
http.MethodDelete, http.MethodDelete,
fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID), fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID),
@@ -296,7 +298,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
var resp response.Any var resp response.Any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp)) require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, InvalidTaskType, resp.ErrorMsg) assert.Equal(t, InvalidTaskType, resp.ErrorMsg)
preserved, err := model.GetScheduleByID(ctx, schedule.ID) preserved, err := repository.GetScheduleByID(ctx, schedule.ID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, testInternalOnlyTaskType, preserved.TaskType) assert.Equal(t, testInternalOnlyTaskType, preserved.TaskType)
}) })
@@ -320,7 +322,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
Cron: "0 * * * *", Cron: "0 * * * *",
IsActive: false, IsActive: false,
} }
require.NoError(t, model.CreateSchedule(ctx, schedule)) require.NoError(t, repository.CreateSchedule(ctx, schedule))
req := httptest.NewRequest( req := httptest.NewRequest(
http.MethodDelete, http.MethodDelete,
fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID), fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID),
@@ -331,7 +333,7 @@ func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
router.ServeHTTP(w, req) router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code) assert.Equal(t, http.StatusOK, w.Code)
_, err := model.GetScheduleByID(ctx, schedule.ID) _, err := repository.GetScheduleByID(ctx, schedule.ID)
assert.Error(t, err) 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}, {TaskID: "exec_003", TaskType: "system:cleanup", TaskName: "系统垃圾清理", Status: model.TaskExecutionStatusPending, TriggeredBy: "manual", Retryable: true, MaxRetry: 3},
} }
for _, r := range records { for _, r := range records {
err := model.CreateTaskExecution(ctx, r) err := repository.CreateTaskExecution(ctx, r)
require.NoError(t, err) require.NoError(t, err)
} }
@@ -581,7 +583,7 @@ func TestGetTaskExecution(t *testing.T) {
MaxRetry: 3, MaxRetry: 3,
TriggeredBy: "manual", TriggeredBy: "manual",
} }
err := model.CreateTaskExecution(ctx, execution) err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err) require.NoError(t, err)
t.Run("get existing execution", func(t *testing.T) { t.Run("get existing execution", func(t *testing.T) {
@@ -646,7 +648,7 @@ func TestRetryTask(t *testing.T) {
StartedAt: &now, StartedAt: &now,
FinishedAt: &now, FinishedAt: &now,
} }
err := model.CreateTaskExecution(ctx, execution) err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err) require.NoError(t, err)
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID) 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.True(t, ok)
assert.NotEmpty(t, newTaskID) assert.NotEmpty(t, newTaskID)
newExecution, err := model.GetTaskExecutionByTaskID(ctx, newTaskID) newExecution, err := repository.GetTaskExecutionByTaskID(ctx, newTaskID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, 1, newExecution.RetryCount) assert.Equal(t, 1, newExecution.RetryCount)
assert.Equal(t, "retry", newExecution.TriggeredBy) assert.Equal(t, "retry", newExecution.TriggeredBy)
@@ -682,7 +684,7 @@ func TestRetryTask(t *testing.T) {
MaxRetry: 3, MaxRetry: 3,
TriggeredBy: "manual", TriggeredBy: "manual",
} }
err := model.CreateTaskExecution(ctx, execution) err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err) require.NoError(t, err)
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID) url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
@@ -702,7 +704,7 @@ func TestRetryTask(t *testing.T) {
Retryable: false, Retryable: false,
TriggeredBy: "manual", TriggeredBy: "manual",
} }
err := model.CreateTaskExecution(ctx, execution) err := repository.CreateTaskExecution(ctx, execution)
require.NoError(t, err) require.NoError(t, err)
url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID) url := fmt.Sprintf("/api/v1/admin/tasks/executions/%d/retry", execution.ID)
+3 -5
View File
@@ -10,7 +10,6 @@ import (
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth" "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/infra/persistence/idgen"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/repository"
@@ -41,7 +40,7 @@ func updateUserStatus(ctx context.Context, id uint64, active bool) error {
var tokens []model.AccessToken var tokens []model.AccessToken
if !active { if !active {
_ = db.DB(ctx).Where("user_id = ?", id).Find(&tokens).Error tokens, _ = repository.ListAccessTokensByUserID(ctx, id)
} }
err = repository.UpdateUserActive(ctx, id, active) err = repository.UpdateUserActive(ctx, id, active)
@@ -68,8 +67,7 @@ func deleteUser(ctx context.Context, currentUserID, targetID uint64) error {
return errors.New(cannotDelete) return errors.New(cannotDelete)
} }
var tokens []model.AccessToken tokens, _ := repository.ListAccessTokensByUserID(ctx, targetID)
_ = db.DB(ctx).Where("user_id = ?", targetID).Find(&tokens).Error
err = repository.DeleteUserWithRelations(ctx, targetID) err = repository.DeleteUserWithRelations(ctx, targetID)
if err == nil { if err == nil {
@@ -181,7 +179,7 @@ func updateUser(ctx context.Context, currentUserID uint64, param updateUserParam
needRevokeTokens := (param.Password != "") || (targetUser.IsAdmin && !param.IsAdmin) needRevokeTokens := (param.Password != "") || (targetUser.IsAdmin && !param.IsAdmin)
var tokens []model.AccessToken var tokens []model.AccessToken
if needRevokeTokens { if needRevokeTokens {
_ = db.DB(ctx).Where("user_id = ?", param.ID).Find(&tokens).Error tokens, _ = repository.ListAccessTokensByUserID(ctx, param.ID)
} }
// 更新字段 // 更新字段
+12 -10
View File
@@ -130,12 +130,12 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS
response.AbortUnauthorized(c, shared.UnAuthorized) response.AbortUnauthorized(c, shared.UnAuthorized)
return return
} }
var user model.User user, err := repository.GetUserByID(ctx, userID)
if err := db.DB(ctx).First(&user, "id = ?", userID).Error; err != nil { if err != nil {
response.AbortInternal(c, err.Error()) response.AbortInternal(c, err.Error())
return return
} }
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{
AuthSourceID: source.ID, AuthSourceID: source.ID,
UserID: user.ID, UserID: user.ID,
ExternalID: userInfo.Sub, ExternalID: userInfo.Sub,
@@ -146,7 +146,7 @@ func handleCallbackBind(ctx context.Context, c *gin.Context, source *model.AuthS
return return
} }
user.LastLoginAt = time.Now() 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"))) 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) { func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.AuthSource, userInfo *model.OAuthUserInfo) {
var user model.User var user model.User
account, err := model.FindExternalAccount(ctx, source.ID, userInfo.Sub) account, err := repository.FindExternalAccount(ctx, source.ID, userInfo.Sub)
switch { switch {
case err == nil: case err == nil:
if err := db.DB(ctx).First(&user, "id = ?", account.UserID).Error; err != nil { loaded, loadErr := repository.GetUserByID(ctx, account.UserID)
response.AbortInternal(c, err.Error()) if loadErr != nil {
response.AbortInternal(c, loadErr.Error())
return return
} }
user = loaded
case errors.Is(err, gorm.ErrRecordNotFound): case errors.Is(err, gorm.ErrRecordNotFound):
newUser, ok := handleCallbackRegister(ctx, c, source, userInfo) newUser, ok := handleCallbackRegister(ctx, c, source, userInfo)
if !ok { if !ok {
@@ -173,7 +175,7 @@ func handleCallbackLogin(ctx context.Context, c *gin.Context, source *model.Auth
} }
user.LastLoginAt = time.Now() 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 { if err := setLoginSession(ctx, c, &user); err != nil {
response.AbortInternal(c, err.Error()) response.AbortInternal(c, err.Error())
return return
@@ -209,11 +211,11 @@ func handleCallbackRegister(ctx context.Context, c *gin.Context, source *model.A
userInfo.Username = username userInfo.Username = username
var user model.User 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()) response.AbortInternal(c, err.Error())
return model.User{}, false return model.User{}, false
} }
if err := model.BindExternalAccount(ctx, &model.ExternalAccount{ if err := repository.BindExternalAccount(ctx, &model.ExternalAccount{
AuthSourceID: source.ID, AuthSourceID: source.ID,
UserID: user.ID, UserID: user.ID,
ExternalID: userInfo.Sub, ExternalID: userInfo.Sub,
@@ -8,7 +8,8 @@ import (
"strconv" "strconv"
"strings" "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"
"github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -26,7 +27,7 @@ import (
// @Router /api/v1/oauth/external-accounts [get] // @Router /api/v1/oauth/external-accounts [get]
func ListExternalAccounts(c *gin.Context) { func ListExternalAccounts(c *gin.Context) {
userID := GetUserIDFromContext(c) userID := GetUserIDFromContext(c)
accounts, err := model.ListExternalAccountsByUserID(c.Request.Context(), userID) accounts, err := repository.ListExternalAccountsByUserID(c.Request.Context(), userID)
if err != nil { if err != nil {
response.AbortInternal(c, err.Error()) response.AbortInternal(c, err.Error())
return return
@@ -57,7 +58,7 @@ func DeleteExternalAccount(c *gin.Context) {
response.AbortBadRequest(c, errInvalidExternalAccountBindingID) response.AbortBadRequest(c, errInvalidExternalAccountBindingID)
return 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()) response.AbortBadRequest(c, err.Error())
return return
} }
+8 -9
View File
@@ -8,8 +8,8 @@ import (
"context" "context"
"errors" "errors"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/shared" "github.com/Rain-kl/Wavelet/internal/shared"
"github.com/Rain-kl/Wavelet/internal/shared/response" "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) tokenHash := model.HashToken(tokenStr)
tokenRecord, err := GetCachedToken(ctx, tokenHash) tokenRecord, err := GetCachedToken(ctx, tokenHash)
if err != nil { if err != nil {
var dbToken model.AccessToken dbToken, err := repository.GetAccessTokenByHash(ctx, tokenHash)
if err := db.DB(ctx).Where("token_hash = ?", tokenHash).First(&dbToken).Error; err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
tokenRecord = &dbToken tokenRecord = &dbToken
@@ -43,8 +43,8 @@ func getUserByToken(ctx context.Context, tokenStr string) (*model.User, *model.A
user, err := GetCachedUser(ctx, tokenRecord.UserID) user, err := GetCachedUser(ctx, tokenRecord.UserID)
if err != nil || !user.IsActive { if err != nil || !user.IsActive {
var dbUser model.User dbUser, err := repository.GetActiveUserByID(ctx, tokenRecord.UserID)
if err := db.DB(ctx).Where("id = ? AND is_active = ?", tokenRecord.UserID, true).First(&dbUser).Error; err != nil { if err != nil {
return nil, nil, err return nil, nil, err
} }
user = &dbUser user = &dbUser
@@ -87,11 +87,10 @@ func GetUserFromRequest(c *gin.Context) (*model.User, error) {
user, err := GetCachedUser(ctx, userID) user, err := GetCachedUser(ctx, userID)
if err != nil || !user.IsActive { if err != nil || !user.IsActive {
var dbUser model.User
// load user from db to make sure is active // load user from db to make sure is active
tx := db.DB(ctx).Where("id = ? AND is_active = ?", userID, true).First(&dbUser) dbUser, loadErr := repository.GetActiveUserByID(ctx, userID)
if tx.Error != nil { if loadErr != nil {
return nil, tx.Error return nil, loadErr
} }
user = &dbUser user = &dbUser
SetCachedUser(ctx, userID, user) SetCachedUser(ctx, userID, user)
+3 -5
View File
@@ -9,8 +9,8 @@ import (
"fmt" "fmt"
"strings" "strings"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/coreos/go-oidc/v3/oidc" "github.com/coreos/go-oidc/v3/oidc"
"golang.org/x/oauth2" "golang.org/x/oauth2"
) )
@@ -21,10 +21,8 @@ func uniqueUsername(ctx context.Context, base string) (string, error) {
base = "user" base = "user"
} }
var existingUsernames []string existingUsernames, err := repository.ListUsernamesMatchingBase(ctx, base)
if err := db.DB(ctx).Model(&model.User{}). if err != nil {
Where("username = ? OR username LIKE ?", base, base+"-%").
Pluck("username", &existingUsernames).Error; err != nil {
return "", err return "", err
} }
+3 -1
View File
@@ -10,6 +10,8 @@ import (
"sync" "sync"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm" "gorm.io/gorm"
) )
@@ -39,7 +41,7 @@ func newAccessTokenAuthCache() *accessTokenAuthCache {
positive: make(map[string]cachedAgentNode), positive: make(map[string]cachedAgentNode),
negative: make(map[string]time.Time), negative: make(map[string]time.Time),
now: time.Now, now: time.Now,
loadNodeByToken: model.GetOpenFlareNodeByAccessToken, loadNodeByToken: repository.GetOpenFlareNodeByAccessToken,
} }
} }
+4 -3
View File
@@ -9,14 +9,15 @@ import (
"errors" "errors"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/pages" "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" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"gorm.io/gorm" "gorm.io/gorm"
) )
func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) { func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
version, err := model.GetActiveConfigVersion(ctx) version, err := repository.GetActiveConfigVersion(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -27,7 +28,7 @@ func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
} }
func getActiveConfigForAgent(ctx context.Context) (*ConfigResponse, error) { func getActiveConfigForAgent(ctx context.Context) (*ConfigResponse, error) {
version, err := model.GetActiveConfigVersion(ctx) version, err := repository.GetActiveConfigVersion(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
+8 -45
View File
@@ -9,11 +9,11 @@ import (
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip" ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/node" "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" "github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
) )
// RegisterWithAccessToken registers an agent on a reserved node token. // RegisterWithAccessToken registers an agent on a reserved node token.
@@ -27,7 +27,7 @@ func RegisterWithAccessToken(ctx context.Context, authNode *model.OpenFlareNode,
return nil, err return nil, err
} }
applyNodeRuntime(ctx, authNode, payload, true) applyNodeRuntime(ctx, authNode, payload, true)
if err := model.SaveOpenFlareNode(ctx, authNode); err != nil { if err := repository.SaveOpenFlareNode(ctx, authNode); err != nil {
return nil, err return nil, err
} }
RefreshAccessTokenCache(ctx, authNode) RefreshAccessTokenCache(ctx, authNode)
@@ -71,7 +71,7 @@ func RegisterWithDiscovery(ctx context.Context, payload NodePayload) (*Registrat
} }
applyNodeRuntime(ctx, record, payload, false) applyNodeRuntime(ctx, record, payload, false)
if err = model.CreateOpenFlareNode(ctx, record); err != nil { if err = repository.CreateOpenFlareNode(ctx, record); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errNodeIDConflict) return nil, errors.New(errNodeIDConflict)
} }
@@ -115,7 +115,7 @@ func HeartbeatNode(ctx context.Context, authNode *model.OpenFlareNode, payload N
for field := range changes { for field := range changes {
fields = append(fields, field) 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 return nil, err
} }
} }
@@ -181,12 +181,12 @@ func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFl
return nil, errors.New(errInvalidApplyResult) return nil, errors.New(errInvalidApplyResult)
} }
latest, err := model.GetLatestOpenFlareApplyLogByNodeID(ctx, payload.NodeID) latest, err := repository.GetLatestOpenFlareApplyLogByNodeID(ctx, payload.NodeID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
if model.IsRepeatSuccessApplyLog(latest, payload.Version, payload.Checksum, payload.Result) { 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 nil, err
} }
return latest, nil return latest, nil
@@ -204,49 +204,12 @@ func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFl
CreatedAt: now, CreatedAt: now,
} }
conn := db.DB(ctx) if err := repository.CreateOpenFlareApplyLogAndUpdateNode(ctx, log, payload.Result, payload.Version, payload.Message); err != nil {
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 {
return nil, err return nil, err
} }
return log, nil 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. // ValidateDiscoveryToken delegates to the node package discovery token helper.
func ValidateDiscoveryToken(ctx context.Context, token string) error { func ValidateDiscoveryToken(ctx context.Context, token string) error {
return node.ValidateDiscoveryToken(ctx, token) return node.ValidateDiscoveryToken(ctx, token)
+37 -160
View File
@@ -10,23 +10,16 @@ import (
"strings" "strings"
"time" "time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" "github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"go.uber.org/zap" "go.uber.org/zap"
"gorm.io/gorm"
"gorm.io/gorm/clause"
) )
const ( const (
healthEventStatusActive = "active"
healthEventStatusResolved = "resolved"
healthSeverityInfo = "info"
healthSeverityWarning = "warning"
healthSeverityCritical = "critical"
accessLogPathMaxLength = 100 accessLogPathMaxLength = 100
accessLogUserAgentMaxLength = 512 accessLogUserAgentMaxLength = 512
accessLogCacheStatusMaxLength = 32 accessLogCacheStatusMaxLength = 32
healthEventMessageMaxLength = 4096
) )
// PersistHeartbeatObservability stores profile, host metrics, edge health, and access logs. // PersistHeartbeatObservability stores profile, host metrics, edge health, and access logs.
@@ -43,28 +36,23 @@ func PersistHeartbeatObservability(ctx context.Context, nodeID string, payload N
return return
} }
conn := db.DB(ctx)
if conn == nil {
return
}
accessLogRecords, err := buildNodeAccessLogRecords(nodeID, payload.AccessLogs, payload.Buffered, reportedAt) accessLogRecords, err := buildNodeAccessLogRecords(nodeID, payload.AccessLogs, payload.Buffered, reportedAt)
if err != nil { if err != nil {
zap.L().Error("build heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err)) zap.L().Error("build heartbeat access logs failed", zap.String("node_id", nodeID), zap.Error(err))
return return
} }
if err := conn.Transaction(func(tx *gorm.DB) error { profile := buildNodeSystemProfileModel(nodeID, payload.Profile, reportedAt)
if err := persistNodeSystemProfile(tx, nodeID, payload.Profile, reportedAt); err != nil { healthEvents := healthEventInputs(payload.HealthEvents)
return err if err := repository.PersistOpenFlareNodePGObservability(
} ctx,
if payload.HealthEvents != nil { profile,
if err := reconcileNodeHealthEvents(tx, nodeID, payload.HealthEvents, reportedAt); err != nil { nodeID,
return err healthEvents,
} payload.HealthEvents != nil,
} reportedAt,
return nil nil,
}); err != nil { ); err != nil {
zap.L().Error("persist heartbeat observability failed", zap.String("node_id", nodeID), zap.Error(err)) zap.L().Error("persist heartbeat observability failed", zap.String("node_id", nodeID), zap.Error(err))
return return
} }
@@ -107,7 +95,7 @@ func persistNodeEdgeHealth(ctx context.Context, nodeID string, health *NodeEdgeH
if status == "" { if status == "" {
status = openrestyStatusUnknown status = openrestyStatusUnknown
} }
return model.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{ return repository.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{
NodeID: nodeID, NodeID: nodeID,
CapturedAt: timeFromUnix(health.CapturedAtUnix, reportedAt), CapturedAt: timeFromUnix(health.CapturedAtUnix, reportedAt),
Status: status, 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 { if profile == nil {
return nil return nil
} }
record := &model.OpenFlareNodeSystemProfile{ return &model.OpenFlareNodeSystemProfile{
NodeID: nodeID, NodeID: nodeID,
Hostname: strings.TrimSpace(profile.Hostname), Hostname: strings.TrimSpace(profile.Hostname),
OSName: strings.TrimSpace(profile.OSName), OSName: strings.TrimSpace(profile.OSName),
@@ -133,23 +121,23 @@ func persistNodeSystemProfile(tx *gorm.DB, nodeID string, profile *NodeSystemPro
UptimeSeconds: profile.UptimeSeconds, UptimeSeconds: profile.UptimeSeconds,
ReportedAt: timeFromUnix(profile.ReportedAtUnix, reportedAt), ReportedAt: timeFromUnix(profile.ReportedAtUnix, reportedAt),
} }
return tx.Clauses(clause.OnConflict{ }
Columns: []clause.Column{{Name: "node_id"}},
DoUpdates: clause.AssignmentColumns([]string{ func healthEventInputs(events []NodeHealthEvent) []repository.OpenFlareHealthEventInput {
"hostname", if events == nil {
"os_name", return nil
"os_version", }
"kernel_version", out := make([]repository.OpenFlareHealthEventInput, 0, len(events))
"architecture", for _, event := range events {
"cpu_model", out = append(out, repository.OpenFlareHealthEventInput{
"cpu_cores", EventType: event.EventType,
"total_memory_bytes", Severity: event.Severity,
"total_disk_bytes", Message: event.Message,
"uptime_seconds", TriggeredAtUnix: event.TriggeredAtUnix,
"reported_at", Metadata: event.Metadata,
"updated_at", })
}), }
}).Create(record).Error return out
} }
func persistNodeMetricSnapshot(ctx context.Context, nodeID string, snapshot *NodeMetricSnapshot, reportedAt time.Time) error { 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, DiskWriteBytes: snapshot.DiskWriteBytes,
// NetworkRx/Tx no longer collected from agents; CH columns remain 0. // 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) { 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 { if len(records) == 0 {
return nil return nil
} }
return model.InsertOpenFlareAccessLogsBatch(ctx, records) return repository.InsertOpenFlareAccessLogsBatch(ctx, records)
}
func reconcileNodeHealthEvents(tx *gorm.DB, nodeID string, events []NodeHealthEvent, reportedAt time.Time) error {
return ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, nil)
} }
// ReconcileScopedNodeHealthEvents reconciles health events, optionally scoped to managed event types. // 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 { func ReconcileScopedNodeHealthEvents(ctx context.Context, nodeID string, events []NodeHealthEvent, reportedAt time.Time, managedEventTypes map[string]struct{}) error {
activeTypes := make(map[string]NodeHealthEvent, len(events)) return repository.ReconcileOpenFlareHealthEvents(ctx, nodeID, healthEventInputs(events), reportedAt, managedEventTypes)
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 timeFromUnix(unixSeconds int64, fallback time.Time) time.Time { func timeFromUnix(unixSeconds int64, fallback time.Time) time.Time {
@@ -15,6 +15,8 @@ import (
"strconv" "strconv"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/protocol" "github.com/Rain-kl/Wavelet/pkg/protocol"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" 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 return []WAFIPGroup{}, nil
} }
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) 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 { if err != nil {
return nil, err return nil, err
} }
@@ -158,7 +160,7 @@ func checksumAgentWAFIPGroup(group WAFIPGroup) string {
} }
func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) { func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
version, err := model.GetActiveConfigVersion(ctx) version, err := repository.GetActiveConfigVersion(ctx)
if err != nil { if err != nil {
if isActiveConfigNotFound(err) { if isActiveConfigNotFound(err) {
return []uint{}, nil return []uint{}, nil
@@ -10,6 +10,8 @@ import (
"strings" "strings"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/protocol" "github.com/Rain-kl/Wavelet/pkg/protocol"
@@ -124,7 +126,7 @@ func TestChangedWAFIPGroupsForAgentDiscoversGraphReferences(t *testing.T) {
Enabled: true, Enabled: true,
IPList: `["192.0.2.88"]`, IPList: `["192.0.2.88"]`,
} }
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFGraphIPGroup(t, ctx, ipGroup.ID) seedActiveConfigWithWAFGraphIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil) groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
@@ -163,7 +165,7 @@ func TestChangedWAFIPGroupsForAgentRejectsOversizedSnapshotBeforeChecksumDelta(t
Enabled: true, Enabled: true,
IPList: `[]`, IPList: `[]`,
} }
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
agentGroup, err := buildAgentWAFIPGroup(ipGroup) agentGroup, err := buildAgentWAFIPGroup(ipGroup)
require.NoError(t, err) require.NoError(t, err)
@@ -185,7 +187,7 @@ func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
Enabled: true, Enabled: true,
IPList: `["203.0.113.44"]`, IPList: `["203.0.113.44"]`,
} }
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID) seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil) groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
@@ -201,7 +203,7 @@ func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
assert.Empty(t, same) assert.Empty(t, same)
ipGroup.IPList = `["203.0.113.45"]` 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}) delta, err := ChangedWAFIPGroupsForAgent(ctx, nil, map[string]string{groupKey: groups[0].Checksum})
require.NoError(t, err) require.NoError(t, err)
@@ -222,7 +224,7 @@ func TestSyncWAFIPGroupsReturnsChangedGroups(t *testing.T) {
Enabled: true, Enabled: true,
IPList: `["198.51.100.10"]`, IPList: `["198.51.100.10"]`,
} }
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID) seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
result, err := SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{ result, err := SyncWAFIPGroups(ctx, WAFIPGroupSyncInput{
@@ -255,9 +257,9 @@ func TestChangedWAFIPGroupsForAgentDisabledGroupClearsIPList(t *testing.T) {
Enabled: true, Enabled: true,
IPList: `["203.0.113.10"]`, IPList: `["203.0.113.10"]`,
} }
require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup)) require.NoError(t, repository.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
ipGroup.Enabled = false ipGroup.Enabled = false
require.NoError(t, model.UpdateOpenFlareWAFIPGroup(ctx, ipGroup)) require.NoError(t, repository.UpdateOpenFlareWAFIPGroup(ctx, ipGroup))
seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID) seedActiveConfigWithWAFIPGroup(t, ctx, ipGroup.ID)
groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil) groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
+3 -2
View File
@@ -8,8 +8,9 @@ import (
"encoding/json" "encoding/json"
"log/slog" "log/slog"
"github.com/Rain-kl/Wavelet/internal/repository"
ofws "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" 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). // 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 return
} }
authNode, err := model.GetOpenFlareNodeByNodeID(ctx, nodeID) authNode, err := repository.GetOpenFlareNodeByNodeID(ctx, nodeID)
if err != nil { if err != nil {
slog.Debug("agent ws status reload node failed", "node_id", nodeID, "error", err) slog.Debug("agent ws status reload node failed", "node_id", nodeID, "error", err)
return return
+6 -4
View File
@@ -9,6 +9,8 @@ import (
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
) )
@@ -53,7 +55,7 @@ func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) {
pageSize := normalizePageSize(input.PageSize) pageSize := normalizePageSize(input.PageSize)
nodeID := strings.TrimSpace(input.NodeID) nodeID := strings.TrimSpace(input.NodeID)
rows, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ rows, err := repository.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
NodeID: nodeID, NodeID: nodeID,
PageNo: pageNo, PageNo: pageNo,
PageSize: pageSize, PageSize: pageSize,
@@ -62,7 +64,7 @@ func ListPage(ctx context.Context, input ListQuery) (*ListResult, error) {
return nil, err return nil, err
} }
total, err := model.CountOpenFlareApplyLogs(ctx, nodeID) total, err := repository.CountOpenFlareApplyLogs(ctx, nodeID)
if err != nil { if err != nil {
return nil, err 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. // Cleanup removes old apply logs or deletes all records.
func Cleanup(ctx context.Context, input CleanupInput) (*CleanupResult, error) { func Cleanup(ctx context.Context, input CleanupInput) (*CleanupResult, error) {
if input.DeleteAll { if input.DeleteAll {
deleted, err := model.DeleteAllOpenFlareApplyLogs(ctx) deleted, err := repository.DeleteAllOpenFlareApplyLogs(ctx)
if err != nil { if err != nil {
return nil, err 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) 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 { if err != nil {
return nil, err return nil, err
} }
@@ -8,6 +8,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite" "github.com/glebarez/sqlite"
@@ -67,7 +69,7 @@ func TestListPageAndCleanup(t *testing.T) {
assert.Equal(t, int64(1), cleanupResult.DeletedCount) assert.Equal(t, int64(1), cleanupResult.DeletedCount)
assert.NotNil(t, cleanupResult.Cutoff) assert.NotNil(t, cleanupResult.Cutoff)
remaining, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ remaining, err := repository.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
NodeID: "node-logs", NodeID: "node-logs",
PageNo: 1, PageNo: 1,
PageSize: 10, PageSize: 10,
@@ -80,7 +82,7 @@ func TestListPageAndCleanup(t *testing.T) {
assert.Equal(t, int64(2), cleanupAll.DeletedCount) assert.Equal(t, int64(2), cleanupAll.DeletedCount)
assert.True(t, cleanupAll.DeleteAll) assert.True(t, cleanupAll.DeleteAll)
finalLogs, err := model.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{ finalLogs, err := repository.ListOpenFlareApplyLogs(ctx, model.OpenFlareApplyLogQuery{
NodeID: "node-logs", NodeID: "node-logs",
PageNo: 1, PageNo: 1,
PageSize: 10, PageSize: 10,
+4 -4
View File
@@ -42,8 +42,8 @@ func TestDatabaseAutoCleanupHandlerDeletesRowsWhenEnabled(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{})) require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{}))
db.SetDB(sqliteDB) db.SetDB(sqliteDB)
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore())
resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore())
t.Cleanup(func() { t.Cleanup(func() {
resetObservabilityStore() resetObservabilityStore()
resetAccessLogStore() resetAccessLogStore()
@@ -52,7 +52,7 @@ func TestDatabaseAutoCleanupHandlerDeletesRowsWhenEnabled(t *testing.T) {
ctx := context.Background() ctx := context.Background()
now := time.Now().UTC() now := time.Now().UTC()
require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{{ require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{{
NodeID: "node-a", NodeID: "node-a",
LoggedAt: now.Add(-95 * 24 * time.Hour), LoggedAt: now.Add(-95 * 24 * time.Hour),
RemoteAddr: "203.0.113.10", RemoteAddr: "203.0.113.10",
@@ -69,7 +69,7 @@ func TestDatabaseAutoCleanupHandlerDeletesRowsWhenEnabled(t *testing.T) {
require.NotNil(t, result) require.NotNil(t, result)
assert.Contains(t, result.Message, "共删除") 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) require.NoError(t, err)
assert.Empty(t, rows) assert.Empty(t, rows)
} }
@@ -10,6 +10,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/chwriter" "github.com/Rain-kl/Wavelet/internal/apps/openflare/chwriter"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
@@ -27,7 +29,7 @@ func TestLiveAppWritePath(t *testing.T) {
now := time.Now().UTC() now := time.Now().UTC()
nodeID := "e2e-app-write-" + now.Format("150405") nodeID := "e2e-app-write-" + now.Format("150405")
if err := model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ if err := repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{
NodeID: nodeID, NodeID: nodeID,
CapturedAt: now, CapturedAt: now,
CPUUsagePercent: 33.3, CPUUsagePercent: 33.3,
@@ -46,7 +48,7 @@ func TestLiveAppWritePath(t *testing.T) {
deadline := time.Now().Add(45 * time.Second) deadline := time.Now().Add(45 * time.Second)
var found bool var found bool
for time.Now().Before(deadline) { 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 { if err != nil {
t.Fatalf("ListOpenFlareMetricSnapshotsSince: %v", err) 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") 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 { if err != nil {
t.Fatalf("ListOpenFlareLatestMetricSnapshotsSince: %v", err) t.Fatalf("ListOpenFlareLatestMetricSnapshotsSince: %v", err)
} }
+4 -3
View File
@@ -11,9 +11,10 @@ import (
"sync" "sync"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/config"
"github.com/Rain-kl/Wavelet/internal/infra/persistence/batchwriter" "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" analyticsmodel "github.com/Rain-kl/Wavelet/internal/model/analytics"
"github.com/Rain-kl/Wavelet/internal/platform/lifecycle" "github.com/Rain-kl/Wavelet/internal/platform/lifecycle"
analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" 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() { func wireModelInsertHooks() {
model.SetObservabilityInsertHooks(model.ObservabilityInsertHooks{ repository.SetObservabilityInsertHooks(repository.ObservabilityInsertHooks{
QueueMetricSnapshot: QueueMetricSnapshot, QueueMetricSnapshot: QueueMetricSnapshot,
QueueEdgeHealth: QueueEdgeHealth, QueueEdgeHealth: QueueEdgeHealth,
QueueFrpsObservation: QueueFrpsObservation, QueueFrpsObservation: QueueFrpsObservation,
QueueFrpcObservation: QueueFrpcObservation, QueueFrpcObservation: QueueFrpcObservation,
}) })
model.SetAccessLogInsertHooks(model.AccessLogInsertHooks{ repository.SetAccessLogInsertHooks(repository.AccessLogInsertHooks{
QueueNodeAccessLogs: QueueNodeAccessLogs, QueueNodeAccessLogs: QueueNodeAccessLogs,
}) })
} }
@@ -15,6 +15,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
oftls "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls" oftls "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
"github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/config"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
@@ -76,7 +78,7 @@ func TestBuildSnapshotReadsZoneDomainCertificates(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
route := &model.ProxyRoute{SiteName: "tls-site", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true, EnableHTTPS: true} 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"} zone := &model.Zone{Domain: "example.com"}
require.NoError(t, db.DB(ctx).Create(zone).Error) 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) require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &route.ID, Domain: "one.example.com", CertID: &first.ID}).Error)
@@ -11,6 +11,8 @@ import (
"strconv" "strconv"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
) )
@@ -208,7 +210,7 @@ func relayAgentAddress(node *model.OpenFlareNode) string {
} }
func resolveTunnelOpenRestyUpstreamURL(ctx context.Context) string { func resolveTunnelOpenRestyUpstreamURL(ctx context.Context) string {
nodes, err := model.ListOpenFlareNodes(ctx) nodes, err := repository.ListOpenFlareNodes(ctx)
if err == nil { if err == nil {
for index := range nodes { for index := range nodes {
node := &nodes[index] node := &nodes[index]
@@ -230,7 +232,7 @@ func listWAFIPGroupsByIDs(ctx context.Context, ids []uint) ([]*model.OpenFlareWA
} }
groups := make([]*model.OpenFlareWAFIPGroup, 0, len(ids)) groups := make([]*model.OpenFlareWAFIPGroup, 0, len(ids))
for _, id := range ids { for _, id := range ids {
group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id) group, err := repository.GetOpenFlareWAFIPGroupByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -14,6 +14,8 @@ import (
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol" pkgprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
@@ -75,17 +77,17 @@ type CleanupResult struct {
// ListConfigVersions returns all config version summaries. // ListConfigVersions returns all config version summaries.
func ListConfigVersions(ctx context.Context) ([]*model.ConfigVersionSummary, error) { func ListConfigVersions(ctx context.Context) ([]*model.ConfigVersionSummary, error) {
return model.ListConfigVersionSummaries(ctx) return repository.ListConfigVersionSummaries(ctx)
} }
// GetConfigVersionDetail returns a config version by version. // GetConfigVersionDetail returns a config version by version.
func GetConfigVersionDetail(ctx context.Context, version string) (*model.ConfigVersion, error) { 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. // GetActiveConfigVersion returns the active config version.
func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) { func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) {
return model.GetActiveConfigVersion(ctx) return repository.GetActiveConfigVersion(ctx)
} }
// PreviewConfigVersion renders the current draft configuration. // PreviewConfigVersion renders the current draft configuration.
@@ -123,7 +125,7 @@ func DiffConfigVersion(ctx context.Context) (*ConfigDiffResult, error) {
ChangedOptionDetails: []ConfigOptionDiffItem{}, ChangedOptionDetails: []ConfigOptionDiffItem{},
CurrentWebsiteCount: len(bundle.SnapshotRoutes), CurrentWebsiteCount: len(bundle.SnapshotRoutes),
} }
activeVersion, err := model.GetActiveConfigVersion(ctx) activeVersion, err := repository.GetActiveConfigVersion(ctx)
if err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
for _, route := range bundle.SnapshotRoutes { for _, route := range bundle.SnapshotRoutes {
@@ -203,7 +205,7 @@ func PublishConfigVersion(ctx context.Context, createdBy string, force bool) (*m
if len(bundle.Routes) == 0 { if len(bundle.Routes) == 0 {
return nil, errors.New(errNoEnabledRoutes) return nil, errors.New(errNoEnabledRoutes)
} }
activeVersion, err := model.GetActiveConfigVersion(ctx) activeVersion, err := repository.GetActiveConfigVersion(ctx)
if !force && err == nil && activeVersion.Checksum == bundle.Checksum { if !force && err == nil && activeVersion.Checksum == bundle.Checksum {
return nil, errors.New(errNoChangesToPublish) return nil, errors.New(errNoChangesToPublish)
} }
@@ -228,7 +230,7 @@ func PublishConfigVersion(ctx context.Context, createdBy string, force bool) (*m
IsActive: true, IsActive: true,
CreatedBy: createdBy, CreatedBy: createdBy,
} }
if err = model.PublishConfigVersionTx(ctx, record); err != nil { if err = repository.PublishConfigVersionTx(ctx, record); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errVersionConflict) 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. // ActivateConfigVersion activates an existing config version.
func ActivateConfigVersion(ctx context.Context, versionStr string) (*model.ConfigVersion, error) { 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 { if err != nil {
return nil, err return nil, err
} }
if err = model.ActivateConfigVersionTx(ctx, versionStr); err != nil { if err = repository.ActivateConfigVersionTx(ctx, versionStr); err != nil {
return nil, err return nil, err
} }
version.IsActive = true version.IsActive = true
@@ -263,7 +265,7 @@ func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult,
if keepCount < minConfigVersionKeepCount { if keepCount < minConfigVersionKeepCount {
keepCount = minConfigVersionKeepCount keepCount = minConfigVersionKeepCount
} }
versions, err := model.ListConfigVersionSummaries(ctx) versions, err := repository.ListConfigVersionSummaries(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -283,7 +285,7 @@ func CleanupConfigVersions(ctx context.Context, keepCount int) (*CleanupResult,
if len(deleteVersions) == 0 { if len(deleteVersions) == 0 {
return &CleanupResult{DeletedCount: 0, Message: cleanupSuccessMessage}, nil return &CleanupResult{DeletedCount: 0, Message: cleanupSuccessMessage}, nil
} }
deletedCount, err := model.DeleteConfigVersionsByVersions(ctx, deleteVersions) deletedCount, err := repository.DeleteConfigVersionsByVersions(ctx, deleteVersions)
if err != nil { if err != nil {
return nil, err 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) { func nextVersionNumber(ctx context.Context, now time.Time) (string, error) {
prefix := now.Format("20060102") prefix := now.Format("20060102")
latest, err := model.GetLatestConfigVersionByPrefix(ctx, prefix) latest, err := repository.GetLatestConfigVersionByPrefix(ctx, prefix)
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return fmt.Sprintf("%s-%03d", prefix, 1), nil return fmt.Sprintf("%s-%03d", prefix, 1), nil
} }
@@ -10,6 +10,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf" "github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
@@ -101,7 +103,7 @@ func TestPublishConfigVersionCreatesVersion(t *testing.T) {
Upstreams: `["http://origin.publish.example.com:8080"]`, Upstreams: `["http://origin.publish.example.com:8080"]`,
Enabled: true, Enabled: true,
} }
require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, "publish.example.com") createSnapshotZoneDomains(t, ctx, route, "publish.example.com")
version, err := PublishConfigVersion(ctx, "tester", false) version, err := PublishConfigVersion(ctx, "tester", false)
@@ -145,15 +147,15 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
Upstreams: `["http://origin.example.com:8080"]`, Upstreams: `["http://origin.example.com:8080"]`,
Enabled: true, 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") createSnapshotZoneDomains(t, ctx, route, "example.com", "www.example.com")
require.NoError(t, waf.EnsureDefaultRuleGroup(ctx)) require.NoError(t, waf.EnsureDefaultRuleGroup(ctx))
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx) globalGroup, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err) require.NoError(t, err)
customGroup := createSnapshotRule(t, ctx, "pow-group", waf.DefaultRuleGraph()) 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) bundle, err := buildCurrentConfigBundle(ctx, true)
require.NoError(t, err) require.NoError(t, err)
@@ -203,11 +205,11 @@ func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testi
Upstreams: `["http://origin.example.com:8080"]`, Upstreams: `["http://origin.example.com:8080"]`,
Enabled: true, Enabled: true,
} }
require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, "pow-global.example.com") createSnapshotZoneDomains(t, ctx, route, "pow-global.example.com")
require.NoError(t, waf.EnsureDefaultRuleGroup(ctx)) require.NoError(t, waf.EnsureDefaultRuleGroup(ctx))
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx) globalGroup, err := repository.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err) require.NoError(t, err)
graphJSON, err := json.Marshal(snapshotPoWGraph()) graphJSON, err := json.Marshal(snapshotPoWGraph())
require.NoError(t, err) require.NoError(t, err)
@@ -10,6 +10,8 @@ import (
"path" "path"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive" "github.com/Rain-kl/Wavelet/pkg/pagesarchive"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
@@ -26,13 +28,13 @@ func buildPagesRouteSnapshot(
if route == nil { if route == nil {
return "", nil, nil, nil, errors.New("pages 路由配置无效") 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) return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 模块不可用", route.SiteName)
} }
if route.PagesProjectID == nil || *route.PagesProjectID == 0 { if route.PagesProjectID == nil || *route.PagesProjectID == 0 {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: 未绑定 Pages 项目", route.SiteName) 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 err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目不存在", route.SiteName) return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目不存在", route.SiteName)
@@ -45,7 +47,7 @@ func buildPagesRouteSnapshot(
if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 { if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目没有激活部署", route.SiteName) 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 err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 激活部署不存在", route.SiteName) return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 激活部署不存在", route.SiteName)
@@ -8,6 +8,8 @@ import (
"encoding/json" "encoding/json"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
@@ -52,7 +54,7 @@ func TestBuildSnapshotRoutesPages(t *testing.T) {
UpstreamType: "pages", UpstreamType: "pages",
PagesProjectID: &project.ID, PagesProjectID: &project.ID,
} }
require.NoError(t, model.CreateProxyRouteRecord(ctx, route)) require.NoError(t, repository.CreateProxyRouteRecord(ctx, route))
createSnapshotZoneDomains(t, ctx, route, "speedtest.arctel.net") createSnapshotZoneDomains(t, ctx, route, "speedtest.arctel.net")
bundle, err := buildCurrentConfigBundle(ctx, true) bundle, err := buildCurrentConfigBundle(ctx, true)
@@ -162,7 +162,7 @@ type configBundle struct {
} }
func buildCurrentConfigBundle(ctx context.Context, requireRoutes bool) (*configBundle, error) { func buildCurrentConfigBundle(ctx context.Context, requireRoutes bool) (*configBundle, error) {
routes, err := model.ListEnabledProxyRoutes(ctx) routes, err := repository.ListEnabledProxyRoutes(ctx)
if err != nil { if err != nil {
return nil, err 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) { func buildSnapshotRoutes(ctx context.Context, routes []*model.ProxyRoute) ([]snapshotRoute, error) {
items := make([]snapshotRoute, 0, len(routes)) items := make([]snapshotRoute, 0, len(routes))
for _, route := range routes { for _, route := range routes {
zoneDomains, err := model.ListZoneDomainsByRouteID(ctx, route.ID) zoneDomains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -310,7 +310,7 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
if err := waf.EnsureDefaultRuleGroup(ctx); err != nil { if err := waf.EnsureDefaultRuleGroup(ctx); err != nil {
return snapshotWAFDocument{}, err return snapshotWAFDocument{}, err
} }
groups, err := model.ListOpenFlareWAFRuleGroups(ctx) groups, err := repository.ListOpenFlareWAFRuleGroups(ctx)
if err != nil { if err != nil {
return snapshotWAFDocument{}, err return snapshotWAFDocument{}, err
} }
@@ -349,7 +349,7 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
if route == nil { if route == nil {
continue continue
} }
domains, domainErr := model.ListZoneDomainsByRouteID(ctx, route.ID) domains, domainErr := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if domainErr != nil { if domainErr != nil {
return snapshotWAFDocument{}, domainErr return snapshotWAFDocument{}, domainErr
} }
@@ -358,7 +358,7 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
} }
enabledRouteSiteNames[route.ID] = route.SiteName enabledRouteSiteNames[route.ID] = route.SiteName
} }
rawBindings, err := model.ListOpenFlareWAFRuleGroupBindings(ctx) rawBindings, err := repository.ListOpenFlareWAFRuleGroupBindings(ctx)
if err != nil { if err != nil {
return snapshotWAFDocument{}, err 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) { 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) { if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil 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] }) sort.Slice(certIDs, func(i, j int) bool { return certIDs[i] < certIDs[j] })
files := make([]SupportFile, 0, len(certIDs)*supportFilesPerCertificate) files := make([]SupportFile, 0, len(certIDs)*supportFilesPerCertificate)
for _, certID := range certIDs { for _, certID := range certIDs {
certificate, err := model.GetTLSCertificateByID(ctx, certID) certificate, err := repository.GetTLSCertificateByID(ctx, certID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -9,6 +9,8 @@ import (
"strings" "strings"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf" "github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
@@ -51,7 +53,7 @@ func TestWAFGraphSnapshotPreservesOrderAndGraphReferences(t *testing.T) {
ctx := context.Background() ctx := context.Background()
route := &model.ProxyRoute{SiteName: "ordered.example.com", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true} 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) createSnapshotZoneDomains(t, ctx, route, route.SiteName)
referenced := &model.OpenFlareWAFIPGroup{Name: "referenced", Type: "manual", Enabled: true, IPList: `["192.0.2.1"]`} 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()) customA := createSnapshotRule(t, ctx, "custom-a", waf.DefaultRuleGraph())
customB := createSnapshotRule(t, ctx, "custom-b", snapshotIPMatchGraph(referenced.ID)) 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}) snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route})
require.NoError(t, err) require.NoError(t, err)
@@ -91,7 +93,7 @@ func TestWAFGraphSnapshotEncodesEmptyBindingsAsArrays(t *testing.T) {
ctx := context.Background() ctx := context.Background()
route := &model.ProxyRoute{SiteName: "empty-binding.example.com", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true} 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) createSnapshotZoneDomains(t, ctx, route, route.SiteName)
snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route}) snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route})
+9 -7
View File
@@ -8,6 +8,8 @@ import (
"sort" "sort"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/observability" "github.com/Rain-kl/Wavelet/internal/apps/openflare/observability"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
) )
@@ -113,24 +115,24 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
now := time.Now() now := time.Now()
since := now.Add(-24 * time.Hour) since := now.Add(-24 * time.Hour)
nodes, err := model.ListOpenFlareNodes(ctx) nodes, err := repository.ListOpenFlareNodes(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
// Latest-per-node health: dedicated LIMIT 1 BY queries (not a global raw LIMIT). // 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 { if err != nil {
return nil, err return nil, err
} }
snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, "", since, dashboardOverviewSnapshotLimit) snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, "", since, dashboardOverviewSnapshotLimit)
if err != nil { if err != nil {
return nil, err return nil, err
} }
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, "", since, dashboardDistributionLimit) accessLogRegions, err := repository.ListOpenFlareAccessLogRegionCounts(ctx, "", since, dashboardDistributionLimit)
if err != nil { if err != nil {
return nil, err return nil, err
} }
activeEvents, err := model.ListOpenFlareActiveHealthEvents(ctx) activeEvents, err := repository.ListOpenFlareActiveHealthEvents(ctx)
if err != nil { if err != nil {
return nil, err 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). // 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, Since: since,
Until: now, Until: now,
}); sumErr == nil { }); sumErr == nil {
@@ -164,7 +166,7 @@ func buildOverviewView(ctx context.Context) (*OverviewView, error) {
} }
nodeTraffic := map[string]model.OpenFlareAccessLogNodeAggregate{} nodeTraffic := map[string]model.OpenFlareAccessLogNodeAggregate{}
if aggregates, aggErr := model.NodeAggregatesOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ if aggregates, aggErr := repository.NodeAggregatesOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{
Since: since, Since: since,
Until: now, Until: now,
}); aggErr == nil { }); aggErr == nil {
@@ -8,6 +8,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite" "github.com/glebarez/sqlite"
@@ -24,8 +26,8 @@ func setupDashboardTestDB(t *testing.T) func() {
require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareNode{})) require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareNode{}))
db.SetDB(sqliteDB) db.SetDB(sqliteDB)
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore())
resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore())
return func() { return func() {
resetObservabilityStore() resetObservabilityStore()
resetAccessLogStore() resetAccessLogStore()
@@ -59,14 +61,14 @@ func TestGetOverviewStructure(t *testing.T) {
}).Error) }).Error)
// Seed older + newer snapshots per node; health must use latest-per-node, not a global raw limit. // 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", NodeID: "node-dashboard-1",
CapturedAt: now.Add(-2 * time.Hour), CapturedAt: now.Add(-2 * time.Hour),
CPUUsagePercent: 10, CPUUsagePercent: 10,
MemoryUsedBytes: 1, MemoryUsedBytes: 1,
MemoryTotalBytes: 10, MemoryTotalBytes: 10,
})) }))
require.NoError(t, model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{
NodeID: "node-dashboard-1", NodeID: "node-dashboard-1",
CapturedAt: now.Add(-time.Minute), CapturedAt: now.Add(-time.Minute),
CPUUsagePercent: 55, CPUUsagePercent: 55,
@@ -101,7 +103,7 @@ func TestGetOverviewStructure(t *testing.T) {
StatusCode: 502, StatusCode: 502,
BytesSent: 10, BytesSent: 10,
}) })
require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, logs)) require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, logs))
overview, err := GetOverview(ctx) overview, err := GetOverview(ctx)
require.NoError(t, err) require.NoError(t, err)
+4 -2
View File
@@ -10,6 +10,8 @@ import (
"strconv" "strconv"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/relay" "github.com/Rain-kl/Wavelet/internal/apps/openflare/relay"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
) )
@@ -51,7 +53,7 @@ func normalizeFlaredHeartbeatPayload(payload HeartbeatPayload) HeartbeatPayload
} }
func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) { func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
version, err := model.GetActiveConfigVersion(ctx) version, err := repository.GetActiveConfigVersion(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -62,7 +64,7 @@ func getActiveConfigMeta(ctx context.Context) (*ActiveConfigMeta, error) {
} }
func listTunnelRelayNodes(ctx context.Context) ([]model.OpenFlareNode, error) { func listTunnelRelayNodes(ctx context.Context) ([]model.OpenFlareNode, error) {
nodes, err := model.ListOpenFlareNodes(ctx) nodes, err := repository.ListOpenFlareNodes(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
+8 -40
View File
@@ -10,8 +10,9 @@ import (
"strings" "strings"
"time" "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/agent"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm" "gorm.io/gorm"
) )
@@ -74,7 +75,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
node.LastSeenAt = &lastSeen node.LastSeenAt = &lastSeen
node.Status = nodeStatusOnline 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) return nil, fmt.Errorf("update flared heartbeat: %w", err)
} }
agent.RefreshAccessTokenCache(ctx, node) 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) return nil, fmt.Errorf("no active config version: %w", err)
} }
routes, err := model.ListProxyRoutes(ctx) routes, err := repository.ListProxyRoutes(ctx)
if err != nil { if err != nil {
return nil, fmt.Errorf("get proxy routes: %w", err) 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 { if !route.Enabled {
continue continue
} }
zoneDomains, domainErr := model.ListZoneDomainsByRouteID(ctx, route.ID) zoneDomains, domainErr := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if domainErr != nil || len(zoneDomains) == 0 { if domainErr != nil || len(zoneDomains) == 0 {
continue continue
} }
@@ -177,12 +178,12 @@ func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFl
return nil, errors.New("result 仅支持 success、warning 或 failed") 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 { if err != nil {
return nil, err return nil, err
} }
if model.IsRepeatSuccessApplyLog(latest, payload.Version, payload.Checksum, payload.Result) { 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 nil, err
} }
return latest, nil return latest, nil
@@ -200,45 +201,12 @@ func ReportApplyLog(ctx context.Context, payload ApplyLogPayload) (*model.OpenFl
CreatedAt: now, CreatedAt: now,
} }
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { if err := repository.CreateOpenFlareApplyLogAndUpdateNode(ctx, log, payload.Result, payload.Version, payload.Message); err != nil {
if err := tx.Create(log).Error; err != nil {
return err
}
return updateFlaredNodeFromApplyLogTx(tx, payload, now)
})
if err != nil {
return nil, err return nil, err
} }
return log, nil 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 { func normalizeApplyLogPayload(payload ApplyLogPayload) ApplyLogPayload {
payload.NodeID = strings.TrimSpace(payload.NodeID) payload.NodeID = strings.TrimSpace(payload.NodeID)
payload.Version = strings.TrimSpace(payload.Version) payload.Version = strings.TrimSpace(payload.Version)
+3 -1
View File
@@ -8,6 +8,8 @@ import (
"errors" "errors"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -38,7 +40,7 @@ func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlar
if token == "" { if token == "" {
return nil, errors.New("missing tunnel token") return nil, errors.New("missing tunnel token")
} }
node, err := model.GetOpenFlareNodeByAccessToken(ctx, token) node, err := repository.GetOpenFlareNodeByAccessToken(ctx, token)
if err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("invalid tunnel token") return nil, errors.New("invalid tunnel token")
@@ -9,6 +9,8 @@ import (
"net/http/httptest" "net/http/httptest"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/internal/shared/response"
@@ -44,7 +46,7 @@ func seedFlaredNode(t *testing.T, nodeType, accessToken string) *model.OpenFlare
NodeType: nodeType, NodeType: nodeType,
AccessToken: accessToken, AccessToken: accessToken,
} }
require.NoError(t, model.CreateOpenFlareNode(ctx, node)) require.NoError(t, repository.CreateOpenFlareNode(ctx, node))
return node return node
} }
@@ -10,7 +10,6 @@ import (
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"go.uber.org/zap" "go.uber.org/zap"
) )
@@ -40,11 +39,7 @@ func persistFlaredObservability(ctx context.Context, nodeID string, payload Hear
}, },
}) })
} }
conn := db.DB(ctx) if err := agent.ReconcileScopedNodeHealthEvents(ctx, nodeID, events, reportedAt, managedTypes); err != nil {
if conn == nil {
return
}
if err := agent.ReconcileScopedNodeHealthEvents(conn, nodeID, events, reportedAt, managedTypes); err != nil {
zap.L().Error("persist flared health events failed", zap.String("node_id", nodeID), zap.Error(err)) zap.L().Error("persist flared health events failed", zap.String("node_id", nodeID), zap.Error(err))
} }
} }
@@ -7,6 +7,8 @@ import (
"context" "context"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
@@ -62,7 +64,7 @@ func TestHeartbeatFlaredEmitsHealthEventOnUnhealthy(t *testing.T) {
}) })
require.NoError(t, err) 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.NoError(t, err)
require.Len(t, events, 1) require.Len(t, events, 1)
assert.Equal(t, flaredRuntimeUnhealthyEventType, events[0].EventType) assert.Equal(t, flaredRuntimeUnhealthyEventType, events[0].EventType)
@@ -8,6 +8,8 @@ import (
"net/http" "net/http"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
ofnode "github.com/Rain-kl/Wavelet/internal/apps/openflare/node" ofnode "github.com/Rain-kl/Wavelet/internal/apps/openflare/node"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
@@ -38,8 +40,8 @@ func setupProtocolTestEnv(t *testing.T) (*gin.Engine, func()) {
db.SetDB(sqliteDB) db.SetDB(sqliteDB)
agent.ResetAuthCacheForTest() agent.ResetAuthCacheForTest()
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore())
resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore())
engine := testhelper.NewTestGinEngine() engine := testhelper.NewTestGinEngine()
mountOpenFlareTestRoutes(engine) mountOpenFlareTestRoutes(engine)
@@ -110,7 +112,7 @@ func TestAgentRelayFlaredProtocol(t *testing.T) {
assert.NotNil(t, heartbeatData.RelayConfig) assert.NotNil(t, heartbeatData.RelayConfig)
assert.NotNil(t, heartbeatData.RelaySettings) assert.NotNil(t, heartbeatData.RelaySettings)
stored, err := model.GetOpenFlareNodeByNodeID(ctx, relayNode.NodeID) stored, err := repository.GetOpenFlareNodeByNodeID(ctx, relayNode.NodeID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "online", stored.Status) assert.Equal(t, "online", stored.Status)
assert.Equal(t, "healthy", stored.RelayStatus) assert.Equal(t, "healthy", stored.RelayStatus)
@@ -134,7 +136,7 @@ func TestAgentRelayFlaredProtocol(t *testing.T) {
assert.Equal(t, http.StatusOK, rec.Code) assert.Equal(t, http.StatusOK, rec.Code)
requireAPIOK(t, rec) requireAPIOK(t, rec)
stored, err := model.GetOpenFlareNodeByNodeID(ctx, clientNode.NodeID) stored, err := repository.GetOpenFlareNodeByNodeID(ctx, clientNode.NodeID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "online", stored.Status) assert.Equal(t, "online", stored.Status)
assert.Equal(t, "v0.2.0", stored.Version) assert.Equal(t, "v0.2.0", stored.Version)
@@ -161,7 +163,7 @@ func TestAgentRelayFlaredProtocol(t *testing.T) {
assert.NotEmpty(t, registration.AccessToken) assert.NotEmpty(t, registration.AccessToken)
assert.Equal(t, "discovered-edge", registration.Name) 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) require.NoError(t, err)
assert.Equal(t, "online", stored.Status) assert.Equal(t, "online", stored.Status)
assert.Equal(t, registration.AccessToken, stored.AccessToken) 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, "success", applyLog.Result)
assert.Equal(t, "20260618-001", applyLog.Version) 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) require.NoError(t, err)
assert.Equal(t, "online", stored.Status) assert.Equal(t, "online", stored.Status)
assert.Equal(t, "20260618-001", stored.CurrentVersion) assert.Equal(t, "20260618-001", stored.CurrentVersion)
+14 -14
View File
@@ -134,7 +134,7 @@ type HealthEventCleanupResult = observability.HealthEventCleanupResult
// ListNodes returns all node views with latest apply log metadata. // ListNodes returns all node views with latest apply log metadata.
func ListNodes(ctx context.Context) ([]*View, error) { func ListNodes(ctx context.Context) ([]*View, error) {
nodes, err := model.ListOpenFlareNodes(ctx) nodes, err := repository.ListOpenFlareNodes(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -142,7 +142,7 @@ func ListNodes(ctx context.Context) ([]*View, error) {
for _, node := range nodes { for _, node := range nodes {
nodeIDs = append(nodeIDs, node.NodeID) nodeIDs = append(nodeIDs, node.NodeID)
} }
latestLogs, err := model.GetLatestOpenFlareApplyLogsByNodeIDs(ctx, nodeIDs) latestLogs, err := repository.GetLatestOpenFlareApplyLogsByNodeIDs(ctx, nodeIDs)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -209,7 +209,7 @@ func CreateNode(ctx context.Context, input Input) (*View, error) {
node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL) node.RelayClientProxyURL = strings.TrimSpace(input.RelayClientProxyURL)
node.RelayWebServerEnabled = input.RelayWebServerEnabled node.RelayWebServerEnabled = input.RelayWebServerEnabled
} }
if err = model.CreateOpenFlareNode(ctx, node); err != nil { if err = repository.CreateOpenFlareNode(ctx, node); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errNodeIDConflict) return nil, errors.New(errNodeIDConflict)
} }
@@ -227,7 +227,7 @@ func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) {
if name == "" { if name == "" {
return nil, errors.New(errNodeNameRequired) return nil, errors.New(errNodeNameRequired)
} }
node, err := model.GetOpenFlareNodeByID(ctx, id) node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -252,7 +252,7 @@ func UpdateNode(ctx context.Context, id uint, input Input) (*View, error) {
node.RelayVhostHTTPPort = input.RelayVhostHTTPPort node.RelayVhostHTTPPort = input.RelayVhostHTTPPort
} }
} }
if err = model.SaveOpenFlareNode(ctx, node); err != nil { if err = repository.SaveOpenFlareNode(ctx, node); err != nil {
return nil, err return nil, err
} }
return buildNodeView(node), nil 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. // DeleteNode removes a node by id.
func DeleteNode(ctx context.Context, id uint) error { 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 err
} }
return model.DeleteOpenFlareNode(ctx, id) return repository.DeleteOpenFlareNode(ctx, id)
} }
// GetBootstrapToken returns the global discovery token, creating one if missing. // 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. // GetAgentRelease checks the latest agent release for a node.
func GetAgentRelease(ctx context.Context, id uint, channel string) (*AgentReleaseInfo, error) { 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 { if err != nil {
return nil, err 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. // RequestAgentUpdate marks a node for manual agent update.
func RequestAgentUpdate(ctx context.Context, id uint, input AgentUpdateInput) (*View, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -323,7 +323,7 @@ func RequestAgentUpdate(ctx context.Context, id uint, input AgentUpdateInput) (*
node.UpdateRequested = true node.UpdateRequested = true
node.UpdateChannel = channel.String() node.UpdateChannel = channel.String()
node.UpdateTag = tagName 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 nil, err
} }
return buildNodeView(node), nil return buildNodeView(node), nil
@@ -331,12 +331,12 @@ func RequestAgentUpdate(ctx context.Context, id uint, input AgentUpdateInput) (*
// RequestOpenrestyRestart marks a node for openresty restart. // RequestOpenrestyRestart marks a node for openresty restart.
func RequestOpenrestyRestart(ctx context.Context, id uint) (*View, error) { func RequestOpenrestyRestart(ctx context.Context, id uint) (*View, error) {
node, err := model.GetOpenFlareNodeByID(ctx, id) node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
node.RestartOpenrestyRequested = true 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 nil, err
} }
return buildNodeView(node), nil 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. // RequestForceSync pushes force_sync_config to a connected agent websocket.
func RequestForceSync(ctx context.Context, id uint) (*View, error) { func RequestForceSync(ctx context.Context, id uint) (*View, error) {
node, err := model.GetOpenFlareNodeByID(ctx, id) node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
activeConfig, err := model.GetActiveConfigVersion(ctx) activeConfig, err := repository.GetActiveConfigVersion(ctx)
if err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, fmt.Errorf("无法获取当前激活的配置版本:%s", errNoActiveConfigVersion) return nil, fmt.Errorf("无法获取当前激活的配置版本:%s", errNoActiveConfigVersion)
+4 -4
View File
@@ -40,8 +40,8 @@ func setupNodeTestDB(t *testing.T) func() {
)) ))
db.SetDB(sqliteDB) db.SetDB(sqliteDB)
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore())
resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore())
return func() { return func() {
resetObservabilityStore() resetObservabilityStore()
@@ -83,7 +83,7 @@ func TestCreateTunnelRelayNode(t *testing.T) {
assert.Equal(t, 7000, view.RelayBindPort) assert.Equal(t, 7000, view.RelayBindPort)
assert.Equal(t, 8080, view.RelayVhostHTTPPort) assert.Equal(t, 8080, view.RelayVhostHTTPPort)
stored, err := model.GetOpenFlareNodeByID(ctx, view.ID) stored, err := repository.GetOpenFlareNodeByID(ctx, view.ID)
require.NoError(t, err) require.NoError(t, err)
assert.NotEmpty(t, stored.RelayAuthToken) assert.NotEmpty(t, stored.RelayAuthToken)
} }
@@ -139,7 +139,7 @@ func TestDeleteNode(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, DeleteNode(ctx, created.ID)) require.NoError(t, DeleteNode(ctx, created.ID))
_, err = model.GetOpenFlareNodeByID(ctx, created.ID) _, err = repository.GetOpenFlareNodeByID(ctx, created.ID)
require.Error(t, err) require.Error(t, err)
assert.ErrorIs(t, err, gorm.ErrRecordNotFound) assert.ErrorIs(t, err, gorm.ErrRecordNotFound)
} }
@@ -8,6 +8,8 @@ import (
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics" analyticsrepo "github.com/Rain-kl/Wavelet/internal/repository/analytics"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
@@ -309,7 +311,7 @@ func GetAccessLogOverview(ctx context.Context, input AccessLogOverviewQuery) (*A
Until: now, Until: now,
} }
summaryRow, err := model.TrafficSummaryOpenFlareAccessLogs(ctx, query) summaryRow, err := repository.TrafficSummaryOpenFlareAccessLogs(ctx, query)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -408,7 +410,7 @@ func valueCountDistribution(
column string, column string,
limit int, limit int,
) []DistributionItem { ) []DistributionItem {
rows, err := model.ValueCountsOpenFlareAccessLogs(ctx, query, column, limit) rows, err := repository.ValueCountsOpenFlareAccessLogs(ctx, query, column, limit)
if err != nil { if err != nil {
logger.ErrorF(ctx, "[AccessLog] ValueCountsOpenFlareAccessLogs failed for column %s: %v", column, err) logger.ErrorF(ctx, "[AccessLog] ValueCountsOpenFlareAccessLogs failed for column %s: %v", column, err)
return []DistributionItem{} return []DistributionItem{}
@@ -496,7 +498,7 @@ func buildAccessLogOverviewTrends(
bandwidth[index].BucketStartedAt = bucketAt bandwidth[index].BucketStartedAt = bucketAt
} }
buckets, err := model.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{
NodeID: query.NodeID, NodeID: query.NodeID,
Host: query.Host, Host: query.Host,
Hosts: query.Hosts, Hosts: query.Hosts,
@@ -532,11 +534,11 @@ func buildAccessLogOverviewTrends(
func ListAccessLogs(ctx context.Context, input AccessLogQuery) (*AccessLogList, error) { func ListAccessLogs(ctx context.Context, input AccessLogQuery) (*AccessLogList, error) {
normalized := normalizeAccessLogQuery(input) normalized := normalizeAccessLogQuery(input)
modelQuery := buildModelAccessLogQuery(normalized) modelQuery := buildModelAccessLogQuery(normalized)
logs, err := model.ListOpenFlareAccessLogs(ctx, modelQuery) logs, err := repository.ListOpenFlareAccessLogs(ctx, modelQuery)
if err != nil { if err != nil {
return nil, err return nil, err
} }
totalRecords, totalIPs, _, err := model.CountOpenFlareAccessLogs(ctx, modelQuery) totalRecords, totalIPs, _, err := repository.CountOpenFlareAccessLogs(ctx, modelQuery)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -597,15 +599,15 @@ func ListFoldedAccessLogs(ctx context.Context, input AccessLogQuery) (*FoldedAcc
SortOrder: normalized.SortOrder, SortOrder: normalized.SortOrder,
FoldMinutes: foldMinutes, FoldMinutes: foldMinutes,
} }
items, err := model.ListOpenFlareAccessLogBuckets(ctx, bucketQuery) items, err := repository.ListOpenFlareAccessLogBuckets(ctx, bucketQuery)
if err != nil { if err != nil {
return nil, err return nil, err
} }
totalBuckets, err := model.CountOpenFlareAccessLogBuckets(ctx, bucketQuery) totalBuckets, err := repository.CountOpenFlareAccessLogBuckets(ctx, bucketQuery)
if err != nil { if err != nil {
return nil, err return nil, err
} }
totalRecords, totalIPs, _, err := model.CountOpenFlareAccessLogs(ctx, modelQuery) totalRecords, totalIPs, _, err := repository.CountOpenFlareAccessLogs(ctx, modelQuery)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -654,11 +656,11 @@ func ListFoldedAccessLogIPs(ctx context.Context, input FoldedAccessLogIPQuery) (
SortBy: normalized.SortBy, SortBy: normalized.SortBy,
SortOrder: normalized.SortOrder, SortOrder: normalized.SortOrder,
} }
items, err := model.ListOpenFlareAccessLogBucketIPs(ctx, modelQuery) items, err := repository.ListOpenFlareAccessLogBucketIPs(ctx, modelQuery)
if err != nil { if err != nil {
return nil, err return nil, err
} }
totalIP, err := model.CountOpenFlareAccessLogBucketIPs(ctx, modelQuery) totalIP, err := repository.CountOpenFlareAccessLogBucketIPs(ctx, modelQuery)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -706,11 +708,11 @@ func ListAccessLogIPSummaries(ctx context.Context, input AccessLogIPSummaryQuery
SortBy: normalized.SortBy, SortBy: normalized.SortBy,
SortOrder: normalized.SortOrder, SortOrder: normalized.SortOrder,
} }
items, err := model.ListOpenFlareAccessLogIPSummaries(ctx, query, time.Time{}) items, err := repository.ListOpenFlareAccessLogIPSummaries(ctx, query, time.Time{})
if err != nil { if err != nil {
return nil, err return nil, err
} }
totalIP, err := model.CountOpenFlareAccessLogIPSummaries(ctx, query) totalIP, err := repository.CountOpenFlareAccessLogIPSummaries(ctx, query)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -751,7 +753,7 @@ func GetAccessLogIPTrend(ctx context.Context, input AccessLogIPTrendQuery) (*Acc
if err != nil { if err != nil {
return nil, err return nil, err
} }
points, err := model.ListOpenFlareAccessLogIPTrend(ctx, model.OpenFlareAccessLogIPTrendQuery{ points, err := repository.ListOpenFlareAccessLogIPTrend(ctx, model.OpenFlareAccessLogIPTrendQuery{
NodeID: strings.TrimSpace(normalized.NodeID), NodeID: strings.TrimSpace(normalized.NodeID),
RemoteAddr: strings.TrimSpace(normalized.RemoteAddr), RemoteAddr: strings.TrimSpace(normalized.RemoteAddr),
Host: strings.TrimSpace(normalized.Host), Host: strings.TrimSpace(normalized.Host),
@@ -802,7 +804,7 @@ func GetAccessLogIPAnalysis(ctx context.Context, input AccessLogIPAnalysisQuery)
Until: now, Until: now,
} }
summaryRow, err := model.TrafficSummaryOpenFlareAccessLogs(ctx, query) summaryRow, err := repository.TrafficSummaryOpenFlareAccessLogs(ctx, query)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -868,7 +870,7 @@ func CleanupAccessLogs(ctx context.Context, input AccessLogCleanupInput) (*Acces
return nil, errors.New("retention_days 必须在 1 到 90 之间") return nil, errors.New("retention_days 必须在 1 到 90 之间")
} }
cutoff := time.Now().UTC().Add(-time.Duration(input.RetentionDays) * 24 * time.Hour) 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 { if err != nil {
return nil, err return nil, err
} }
@@ -913,7 +915,7 @@ func listNodeNameMap(ctx context.Context, logs []*model.OpenFlareAccessLog) (map
if len(nodeIDs) == 0 { if len(nodeIDs) == 0 {
return map[string]string{}, nil return map[string]string{}, nil
} }
nodes, err := model.ListOpenFlareNodesByNodeIDs(ctx, nodeIDs) nodes, err := repository.ListOpenFlareNodesByNodeIDs(ctx, nodeIDs)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -9,6 +9,8 @@ import (
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
) )
@@ -135,7 +137,7 @@ func buildTrafficWindowSummaryFromAccessLogs(
nodeID string, nodeID string,
since, until time.Time, since, until time.Time,
) *TrafficWindowSummary { ) *TrafficWindowSummary {
row, err := model.TrafficSummaryOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{ row, err := repository.TrafficSummaryOpenFlareAccessLogs(ctx, model.OpenFlareAccessLogQuery{
NodeID: nodeID, NodeID: nodeID,
Since: since, Since: since,
Until: until, Until: until,
@@ -203,7 +205,7 @@ func BuildTrafficDistributionsFromAccessLogs(
topDomains := make(distributionAccumulator) topDomains := make(distributionAccumulator)
query := model.OpenFlareAccessLogQuery{Since: since, Until: until} 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 { for _, row := range statusRows {
if strings.TrimSpace(row.Value) == "" || row.Count <= 0 { if strings.TrimSpace(row.Value) == "" || row.Count <= 0 {
continue continue
@@ -211,7 +213,7 @@ func BuildTrafficDistributionsFromAccessLogs(
statusCodes[row.Value] = row.Count 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 { for _, row := range hostRows {
if strings.TrimSpace(row.Value) == "" || row.Count <= 0 { if strings.TrimSpace(row.Value) == "" || row.Count <= 0 {
continue continue
@@ -287,7 +289,7 @@ func BuildNodeTrends(
applyAccessLogBytesToNetworkTrend(ctx, now, nodeID, trendSince, networkTrend) applyAccessLogBytesToNetworkTrend(ctx, now, nodeID, trendSince, networkTrend)
diskIOTrend := BuildDiskIOTrendPoints(now, snapshots) diskIOTrend := BuildDiskIOTrendPoints(now, snapshots)
metricHourly, metricErr := model.ListOpenFlareMetricHourlySince(ctx, nodeID, trendSince) metricHourly, metricErr := repository.ListOpenFlareMetricHourlySince(ctx, nodeID, trendSince)
if metricErr == nil && len(metricHourly) > 0 { if metricErr == nil && len(metricHourly) > 0 {
capacityTrend = BuildCapacityTrendPointsFromHourly(now, metricHourly) capacityTrend = BuildCapacityTrendPointsFromHourly(now, metricHourly)
diskIOTrend = BuildDiskIOTrendPointsFromHourly(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) 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 { for _, row := range hourly {
if row == nil { if row == nil {
continue continue
@@ -327,7 +329,7 @@ func BuildTrafficTrendPointsFromAccessLogs(ctx context.Context, now time.Time, n
return points return points
} }
buckets, err := model.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{
NodeID: nodeID, NodeID: nodeID,
Since: since, Since: since,
Until: now, Until: now,
@@ -373,7 +375,7 @@ func applyAccessLogBytesToNetworkTrend(ctx context.Context, now time.Time, nodeI
return return
} }
buckets, err := model.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{ buckets, err := repository.ListOpenFlareAccessLogBuckets(ctx, model.OpenFlareAccessLogBucketQuery{
NodeID: nodeID, NodeID: nodeID,
Since: since, Since: since,
Until: now, Until: now,
@@ -406,7 +408,7 @@ type accessLogHourBytes struct {
} }
func analyticsListAccessLogHourlyBytes(ctx context.Context, nodeID string, since time.Time) (map[int64]accessLogHourBytes, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -10,6 +10,8 @@ import (
"sync" "sync"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm" "gorm.io/gorm"
) )
@@ -94,7 +96,7 @@ type HealthEventCleanupResult struct {
// GetNodeObservability returns observability details for a node. // GetNodeObservability returns observability details for a node.
func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeView, error) { func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeView, error) {
now := time.Now() now := time.Now()
node, err := model.GetOpenFlareNodeByID(ctx, id) node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -105,7 +107,7 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV
limit := normalizeObservabilityLimit(query.Limit) limit := normalizeObservabilityLimit(query.Limit)
since := now.Add(-normalizeObservabilityWindow(query.Hours)) 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) { if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err return nil, err
} }
@@ -113,19 +115,19 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV
profile = nil profile = nil
} }
snapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, since, limit) snapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, node.NodeID, since, limit)
if err != nil { if err != nil {
return nil, err return nil, err
} }
edgeHealth, err := model.ListOpenFlareEdgeHealth(ctx, node.NodeID, since, limit) edgeHealth, err := repository.ListOpenFlareEdgeHealth(ctx, node.NodeID, since, limit)
if err != nil { if err != nil {
return nil, err return nil, err
} }
accessLogRegions, err := model.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, defaultTrafficDistributionLimit) accessLogRegions, err := repository.ListOpenFlareAccessLogRegionCounts(ctx, node.NodeID, since, defaultTrafficDistributionLimit)
if err != nil { if err != nil {
return nil, err return nil, err
} }
events, err := model.ListOpenFlareHealthEvents(ctx, node.NodeID, false, limit) events, err := repository.ListOpenFlareHealthEvents(ctx, node.NodeID, false, limit)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -146,7 +148,7 @@ func GetNodeObservability(ctx context.Context, id uint, query NodeQuery) (*NodeV
Trends: BuildNodeTrends(ctx, now, node.NodeID, snapshots), Trends: BuildNodeTrends(ctx, now, node.NodeID, snapshots),
} }
if node.NodeType == "tunnel_relay" { 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 { if frpsErr != nil {
return nil, frpsErr return nil, frpsErr
} }
@@ -187,11 +189,11 @@ func setCachedNodeObservability(nodeID string, view *NodeView) {
// CleanupHealthEvents removes all health events for a node. // CleanupHealthEvents removes all health events for a node.
func CleanupHealthEvents(ctx context.Context, id uint) (*HealthEventCleanupResult, error) { func CleanupHealthEvents(ctx context.Context, id uint) (*HealthEventCleanupResult, error) {
node, err := model.GetOpenFlareNodeByID(ctx, id) node, err := repository.GetOpenFlareNodeByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
deletedCount, err := model.DeleteOpenFlareHealthEventsByNodeID(ctx, node.NodeID) deletedCount, err := repository.DeleteOpenFlareHealthEventsByNodeID(ctx, node.NodeID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
+1 -1
View File
@@ -149,7 +149,7 @@ func getStatus(ctx context.Context, baseAPIPath string) (*statusView, error) {
} }
func publicAuthSources(ctx context.Context, baseAPIPath string) ([]publicAuthSourceView, error) { func publicAuthSources(ctx context.Context, baseAPIPath string) ([]publicAuthSourceView, error) {
sources, err := model.GetActiveAuthSources(ctx) sources, err := repository.GetActiveAuthSources(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -123,11 +123,11 @@ func TestCleanupDatabaseObservabilityDeletesRows(t *testing.T) {
defer cleanup() defer cleanup()
ctx := context.Background() ctx := context.Background()
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore())
defer resetAccessLogStore() defer resetAccessLogStore()
now := time.Now().UTC() now := time.Now().UTC()
require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{ require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{
{ {
NodeID: "node-a", NodeID: "node-a",
LoggedAt: now.Add(-10 * 24 * time.Hour), LoggedAt: now.Add(-10 * 24 * time.Hour),
@@ -165,7 +165,7 @@ func TestCleanupDatabaseObservabilityDeletesRows(t *testing.T) {
assert.True(t, result.DeleteAll) assert.True(t, result.DeleteAll)
assert.Equal(t, "truncate", result.CleanupMode) 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) require.NoError(t, err)
assert.Empty(t, rows) assert.Empty(t, rows)
} }
+17 -22
View File
@@ -11,8 +11,8 @@ import (
"sort" "sort"
"strings" "strings"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm" "gorm.io/gorm"
) )
@@ -51,7 +51,7 @@ type DetailView struct {
// ListOrigins 列出全部源站。 // ListOrigins 列出全部源站。
func ListOrigins(ctx context.Context) ([]View, error) { func ListOrigins(ctx context.Context) ([]View, error) {
origins, err := model.ListOrigins(ctx) origins, err := repository.ListOrigins(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -60,7 +60,7 @@ func ListOrigins(ctx context.Context) ([]View, error) {
// GetOriginDetail 获取源站详情。 // GetOriginDetail 获取源站详情。
func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) { func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) {
origin, err := model.GetOriginByID(ctx, id) origin, err := repository.GetOriginByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -68,13 +68,13 @@ func GetOriginDetail(ctx context.Context, id uint) (*DetailView, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
routes, err := model.ListProxyRoutesByOriginID(ctx, id) routes, err := repository.ListProxyRoutesByOriginID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
items := make([]RouteSummary, 0, len(routes)) items := make([]RouteSummary, 0, len(routes))
for _, route := range routes { for _, route := range routes {
domains, err := model.ListZoneDomainsByRouteID(ctx, route.ID) domains, err := repository.ListZoneDomainsByRouteID(ctx, route.ID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -105,7 +105,7 @@ func CreateOrigin(ctx context.Context, input Input) (*model.Origin, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
if err = model.CreateOriginRecord(ctx, origin); err != nil { if err = repository.CreateOriginRecord(ctx, origin); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errOriginAddressExists) return nil, errors.New(errOriginAddressExists)
} }
@@ -116,7 +116,7 @@ func CreateOrigin(ctx context.Context, input Input) (*model.Origin, error) {
// UpdateOrigin 更新源站。 // UpdateOrigin 更新源站。
func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -125,8 +125,8 @@ func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, err
if err != nil { if err != nil {
return nil, err return nil, err
} }
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { err = repository.WithOriginTx(ctx, func(tx *gorm.DB) error {
if err := tx.Save(nextOrigin).Error; err != nil { if err := repository.SaveOriginTx(tx, nextOrigin); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return errors.New(errOriginAddressExists) return errors.New(errOriginAddressExists)
} }
@@ -145,17 +145,17 @@ func UpdateOrigin(ctx context.Context, id uint, input Input) (*model.Origin, err
// DeleteOrigin 删除源站。 // DeleteOrigin 删除源站。
func DeleteOrigin(ctx context.Context, id uint) error { func DeleteOrigin(ctx context.Context, id uint) error {
count, err := model.CountProxyRoutesByOriginID(ctx, id) count, err := repository.CountProxyRoutesByOriginID(ctx, id)
if err != nil { if err != nil {
return err return err
} }
if count > 0 { if count > 0 {
return errors.New(errOriginDeleteReferenced) return errors.New(errOriginDeleteReferenced)
} }
if _, err = model.GetOriginByID(ctx, id); err != nil { if _, err = repository.GetOriginByID(ctx, id); err != nil {
return err return err
} }
return model.DeleteOriginRecord(ctx, id) return repository.DeleteOriginRecord(ctx, id)
} }
func buildOrigin(existing *model.Origin, input Input) (*model.Origin, error) { 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) { func buildOriginViews(ctx context.Context, origins []model.Origin) ([]View, error) {
countRows, err := model.ListOriginRouteCounts(ctx) countRows, err := repository.ListOriginRouteCounts(ctx)
if err != nil { if err != nil {
return nil, err 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 { func updateRoutesForOriginAddress(ctx context.Context, tx *gorm.DB, originID uint, address string) error {
if !model.HasProxyRoutesTable(ctx) { if !repository.HasProxyRoutesTable(ctx) {
return nil return nil
} }
var routes []model.OriginProxyRoute routes, err := repository.ListProxyRoutesByOriginIDAscTx(tx, originID)
if err := tx.Where("origin_id = ?", originID).Order("id asc").Find(&routes).Error; err != nil { if err != nil {
return fmt.Errorf("query routes for origin update failed: %w", err) return fmt.Errorf("query routes for origin update failed: %w", err)
} }
for _, route := range routes { for _, route := range routes {
@@ -224,12 +224,7 @@ func updateRoutesForOriginAddress(ctx context.Context, tx *gorm.DB, originID uin
if err != nil { if err != nil {
return fmt.Errorf("encode route %d upstreams failed: %w", route.ID, err) return fmt.Errorf("encode route %d upstreams failed: %w", route.ID, err)
} }
if err := tx.Model(&model.OriginProxyRoute{}). if err := repository.UpdateProxyRouteOriginAddressTx(tx, route.ID, rewrittenOriginURL, string(upstreamsJSON)); err != nil {
Where("id = ?", route.ID).
Updates(map[string]any{
"origin_url": rewrittenOriginURL,
"upstreams": string(upstreamsJSON),
}).Error; err != nil {
return fmt.Errorf("update route %d origin address failed: %w", route.ID, err) return fmt.Errorf("update route %d origin address failed: %w", route.ID, err)
} }
} }
+13 -21
View File
@@ -17,11 +17,10 @@ import (
"unicode" "unicode"
"unicode/utf8" "unicode/utf8"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm" "gorm.io/gorm"
"gorm.io/gorm/clause"
) )
const ( const (
@@ -213,8 +212,7 @@ func unsafeGitHubInputRune(character rune) bool {
} }
func updateGitHubSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) (bool, error) { func updateGitHubSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) (bool, error) {
var project model.PagesProject if _, err := repository.LockPagesProjectByIDTx(tx, projectID); err != nil {
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return false, err return false, err
} }
existing, hasExisting, err := loadProjectSourceForUpdate(tx, projectID) existing, hasExisting, err := loadProjectSourceForUpdate(tx, projectID)
@@ -231,16 +229,15 @@ func updateGitHubSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput)
if !githubSourceConfigChanged(existing, config) { if !githubSourceConfigChanged(existing, config) {
return false, nil return false, nil
} }
var runtime model.PagesProjectSourceRuntime runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, existing.ID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
Where("source_id = ?", existing.ID).First(&runtime).Error; err != nil {
return false, err return false, err
} }
identityChanged := existing.SourceIdentity != config.SourceIdentity 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 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 false, err
} }
return true, nil return true, nil
@@ -259,7 +256,7 @@ func createGitHubSourceTx(tx *gorm.DB, projectID uint, config githubSourceConfig
ConfigVersion: 1, ConfigVersion: 1,
SourceIdentity: config.SourceIdentity, SourceIdentity: config.SourceIdentity,
} }
if err := tx.Create(source).Error; err != nil { if err := repository.CreatePagesProjectSourceTx(tx, source); err != nil {
return err return err
} }
runtime := &model.PagesProjectSourceRuntime{SourceID: source.ID, SyncStatus: pagesSourceStatusIdle} 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) next := nextGitHubCheckAt(time.Now(), source.ID, config.CheckInterval)
runtime.NextCheckAt = &next runtime.NextCheckAt = &next
} }
return tx.Create(runtime).Error return repository.CreatePagesProjectSourceRuntimeTx(tx, runtime)
} }
func githubSourceUpdates(config githubSourceConfig, version int) map[string]any { func githubSourceUpdates(config githubSourceConfig, version int) map[string]any {
@@ -308,7 +305,7 @@ func resetRuntimeAfterGitHubUpdate(
next := nextGitHubCheckAt(time.Now(), runtime.SourceID, config.CheckInterval) next := nextGitHubCheckAt(time.Now(), runtime.SourceID, config.CheckInterval)
nextCheckAt = &next 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 { 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, sourceRuntimeColumnSyncStatus: pagesSourceStatusFailed,
sourceRuntimeColumnLastError: errPagesSourceInitialCheckWarning, sourceRuntimeColumnLastError: errPagesSourceInitialCheckWarning,
} }
var source model.PagesProjectSource source, err := repository.GetPagesProjectSourceByIDAndConfigVersion(ctx, sourceID, configVersion)
if err := db.DB(ctx).Where("id = ? AND config_version = ?", sourceID, configVersion).First(&source).Error; err != nil { if err != nil {
if !errors.Is(err, gorm.ErrRecordNotFound) { if !errors.Is(err, gorm.ErrRecordNotFound) {
logger.ErrorF(ctx, "[PagesSource] load initial check source snapshot failed: source_id=%d error=%v", sourceID, err) 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 updates["next_check_at"] = &next
} }
now := time.Now() now := time.Now()
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). if _, err := repository.MarkPagesSourceInitialCheckDispatchFailed(ctx, sourceID, configVersion, now, updates); err != nil {
Where("source_id = ?", sourceID). logger.ErrorF(ctx, "[PagesSource] mark initial check dispatch failure: source_id=%d error=%v", sourceID, err)
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)
} }
} }
@@ -13,14 +13,13 @@ import (
"strings" "strings"
"time" "time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/integration/githubrelease" "github.com/Rain-kl/Wavelet/internal/integration/githubrelease"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive" "github.com/Rain-kl/Wavelet/pkg/pagesarchive"
"gorm.io/gorm" "gorm.io/gorm"
"gorm.io/gorm/clause"
) )
const githubSourceDetailProvider = "github" const githubSourceDetailProvider = "github"
@@ -201,7 +200,7 @@ func finishGitHubCheckNotModified(
) (string, string, error) { ) (string, string, error) {
var revision string var revision string
var status 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) runtime, now, err := lockOwnedSourceRuntime(tx, snapshot)
if err != nil { if err != nil {
return err return err
@@ -211,7 +210,7 @@ func finishGitHubCheckNotModified(
updates := githubCheckTerminalUpdates(snapshot, now, result.RetryAt) updates := githubCheckTerminalUpdates(snapshot, now, result.RetryAt)
updates["etag"] = result.ETag updates["etag"] = result.ETag
updates[sourceRuntimeColumnSyncStatus] = status updates[sourceRuntimeColumnSyncStatus] = status
return tx.Model(runtime).Updates(updates).Error return repository.UpdatePagesProjectSourceRuntimeTx(tx, runtime, updates)
}) })
return revision, status, err return revision, status, err
} }
@@ -223,7 +222,7 @@ func finishGitHubCheckTarget(
target *githubSourceTarget, target *githubSourceTarget,
) (string, error) { ) (string, error) {
status := pagesSourceStatusIdle 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) runtime, now, err := lockOwnedSourceRuntime(tx, snapshot)
if err != nil { if err != nil {
return err return err
@@ -234,7 +233,7 @@ func finishGitHubCheckTarget(
updates["last_seen_revision"] = target.Revision updates["last_seen_revision"] = target.Revision
updates["last_seen_detail"] = target.DetailJSON updates["last_seen_detail"] = target.DetailJSON
updates[sourceRuntimeColumnSyncStatus] = status updates[sourceRuntimeColumnSyncStatus] = status
return tx.Model(runtime).Updates(updates).Error return repository.UpdatePagesProjectSourceRuntimeTx(tx, runtime, updates)
}) })
return status, err return status, err
} }
@@ -273,9 +272,8 @@ func lockOwnedSourceRuntime(
tx *gorm.DB, tx *gorm.DB,
snapshot *sourceExecutionSnapshot, snapshot *sourceExecutionSnapshot,
) (*model.PagesProjectSourceRuntime, time.Time, error) { ) (*model.PagesProjectSourceRuntime, time.Time, error) {
var runtime model.PagesProjectSourceRuntime runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, snapshot.SourceID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
Where("source_id = ?", snapshot.SourceID).First(&runtime).Error; err != nil {
return nil, time.Time{}, err return nil, time.Time{}, err
} }
now := time.Now() now := time.Now()
@@ -283,7 +281,7 @@ func lockOwnedSourceRuntime(
!runtime.LeaseExpiresAt.After(now) { !runtime.LeaseExpiresAt.After(now) {
return nil, time.Time{}, errSourceFinalFence return nil, time.Time{}, errSourceFinalFence
} }
return &runtime, now, nil return runtime, now, nil
} }
func failGitHubCheckLease( func failGitHubCheckLease(
@@ -309,13 +307,13 @@ func failGitHubCheckLease(
} else { } else {
updates[sourceRuntimeColumnNextCheckAt] = nil updates[sourceRuntimeColumnNextCheckAt] = nil
} }
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). rows, err := repository.UpdatePagesSourceRuntimeByActiveLease(
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now). ctx, snapshot.SourceID, snapshot.LeaseToken, now, updates,
Updates(updates) )
if result.Error != nil { if err != nil {
return result.Error return err
} }
if result.RowsAffected != 1 { if rows != 1 {
return errSourceFinalFence return errSourceFinalFence
} }
return nil return nil
@@ -338,11 +336,11 @@ func targetRuntimeStatus(
} }
func preflightGitHubSyncConfirmation(ctx context.Context, sourceID uint, confirmedRevision string) error { func preflightGitHubSyncConfirmation(ctx context.Context, sourceID uint, confirmedRevision string) error {
var runtime model.PagesProjectSourceRuntime runtime, err := repository.GetPagesProjectSourceRuntimeBySourceID(ctx, sourceID)
if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { if err != nil {
return err return err
} }
replacement := sourceHasSameReleaseReplacement(&runtime) replacement := sourceHasSameReleaseReplacement(runtime)
if replacement && confirmedRevision == "" { if replacement && confirmedRevision == "" {
return errors.New(errPagesSourceConfirmationNeeded) return errors.New(errPagesSourceConfirmationNeeded)
} }
@@ -673,13 +671,13 @@ func releaseGitHubSyncWithoutActivation(
sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil, sourceRuntimeColumnLeaseExpiresAt: nil,
} }
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). rows, err := repository.UpdatePagesSourceRuntimeByActiveLease(
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now). ctx, snapshot.SourceID, snapshot.LeaseToken, now, updates,
Updates(updates) )
if result.Error != nil { if err != nil {
return result.Error return err
} }
if result.RowsAffected != 1 { if rows != 1 {
return errSourceFinalFence return errSourceFinalFence
} }
return nil return nil
@@ -16,7 +16,10 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" 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/integration/githubrelease"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/hibiken/asynq" "github.com/hibiken/asynq"
@@ -100,7 +103,7 @@ func TestGitHubSourceValidationNormalizationAndProviderSwitch(t *testing.T) {
if result.CheckTask == nil || result.CheckTask.Action != sourceActionCheck || result.Warning != "" { if result.CheckTask == nil || result.CheckTask.Action != sourceActionCheck || result.Warning != "" {
t.Errorf("UpdateSourceAs(GitHub) result = %+v, want initial check receipt without warning", result) 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 { if err != nil {
t.Fatalf("GetTaskExecutionByTaskID(%q) error = %v, want nil", result.CheckTask.TaskID, err) 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) { func TestGitHubSourceSaveSurvivesInitialCheckDispatchFailure(t *testing.T) {
ctx := setupPagesSourceTest(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") project := mustCreatePagesSourceProject(t, ctx, "github-dispatch-warning")
result, err := UpdateSourceAs(ctx, project.ID, SourceUpdateInput{ result, err := UpdateSourceAs(ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease, SourceType: PagesSourceTypeGitHubRelease,
@@ -679,7 +687,7 @@ func TestGitHubSyncActivatesWithMetadataRevisionAndPackageChecksum(t *testing.T)
if outcome == nil || outcome.Deployment == nil || outcome.Stale { if outcome == nil || outcome.Deployment == nil || outcome.Stale {
t.Fatalf("syncGitHubSource() = %+v, want active deployment", outcome) 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 { if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", outcome.Deployment.ID, err) 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 { if err == nil || err.Error() != errPagesSourceConfirmationStale {
t.Errorf("syncGitHubSource(stale confirmation) error = %v, want %q", err, 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 { if loadErr != nil {
t.Fatalf("GetPagesProjectByID() error = %v, want nil", loadErr) t.Fatalf("GetPagesProjectByID() error = %v, want nil", loadErr)
} }
+3 -5
View File
@@ -19,7 +19,6 @@ import (
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload" "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/model"
"github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
@@ -325,11 +324,10 @@ func removePagesUploadIfUnreferenced(ctx context.Context, projectID uint, upload
if uploadID == 0 { if uploadID == 0 {
return nil return nil
} }
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error {
if projectID != 0 { if projectID != 0 {
var project model.PagesProject if _, projectErr := repository.LockPagesProjectByIDTx(tx, projectID); projectErr != nil &&
projectErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error !errors.Is(projectErr, gorm.ErrRecordNotFound) {
if projectErr != nil && !errors.Is(projectErr, gorm.ErrRecordNotFound) {
return projectErr return projectErr
} }
} }
+45 -59
View File
@@ -17,8 +17,8 @@ import (
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/upload" "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/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive" "github.com/Rain-kl/Wavelet/pkg/pagesarchive"
"gorm.io/gorm" "gorm.io/gorm"
@@ -101,7 +101,7 @@ type View struct {
// ListProjects 列出全部 Pages 项目。 // ListProjects 列出全部 Pages 项目。
func ListProjects(ctx context.Context) ([]View, error) { func ListProjects(ctx context.Context) ([]View, error) {
projects, err := model.ListPagesProjects(ctx) projects, err := repository.ListPagesProjects(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -118,7 +118,7 @@ func ListProjects(ctx context.Context) ([]View, error) {
// GetProject 获取 Pages 项目详情。 // GetProject 获取 Pages 项目详情。
func GetProject(ctx context.Context, id uint) (*View, error) { func GetProject(ctx context.Context, id uint) (*View, error) {
project, err := model.GetPagesProjectByID(ctx, id) project, err := repository.GetPagesProjectByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -131,7 +131,7 @@ func CreateProject(ctx context.Context, input Input) (*View, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
if err = model.CreatePagesProjectRecord(ctx, project); err != nil { if err = repository.CreatePagesProjectRecord(ctx, project); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errPagesSlugExists) return nil, errors.New(errPagesSlugExists)
} }
@@ -143,7 +143,7 @@ func CreateProject(ctx context.Context, input Input) (*View, error) {
// UpdateProject 更新 Pages 项目。 // UpdateProject 更新 Pages 项目。
func UpdateProject(ctx context.Context, id uint, input Input) (*View, error) { func UpdateProject(ctx context.Context, id uint, input Input) (*View, error) {
var project *model.PagesProject 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 var existing model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&existing, id).Error; err != nil { if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&existing, id).Error; err != nil {
return err return err
@@ -177,10 +177,7 @@ func UpdateProject(ctx context.Context, id uint, input Input) (*View, error) {
} }
if contentConfigChanged { if contentConfigChanged {
updates["content_config_version"] = existing.ContentConfigVersion + 1 updates["content_config_version"] = existing.ContentConfigVersion + 1
var source model.PagesProjectSource source, sourceErr := repository.LockPagesProjectSourceByProjectIDTx(tx, existing.ID)
sourceErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", existing.ID).
First(&source).Error
if sourceErr != nil && !errors.Is(sourceErr, gorm.ErrRecordNotFound) { if sourceErr != nil && !errors.Is(sourceErr, gorm.ErrRecordNotFound) {
return sourceErr return sourceErr
} }
@@ -220,12 +217,12 @@ func ensureDeploymentEntry(conn *gorm.DB, deploymentID uint, rootDir, entryFile
// DeleteProject 删除 Pages 项目。 // DeleteProject 删除 Pages 项目。
func DeleteProject(ctx context.Context, id uint) error { func DeleteProject(ctx context.Context, id uint) error {
project, err := model.GetPagesProjectByID(ctx, id) project, err := repository.GetPagesProjectByID(ctx, id)
if err != nil { if err != nil {
return err return err
} }
var deployments []model.PagesDeployment 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 var lockedProject model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedProject, project.ID).Error; err != nil { if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedProject, project.ID).Error; err != nil {
return err return err
@@ -241,25 +238,19 @@ func DeleteProject(ctx context.Context, id uint) error {
return errors.New(errPagesDeleteReferenced) return errors.New(errPagesDeleteReferenced)
} }
} }
var source model.PagesProjectSource source, sourceErr := repository.LockPagesProjectSourceByProjectIDTx(tx, project.ID)
sourceErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", project.ID).
First(&source).Error
if sourceErr != nil && !errors.Is(sourceErr, gorm.ErrRecordNotFound) { if sourceErr != nil && !errors.Is(sourceErr, gorm.ErrRecordNotFound) {
return sourceErr return sourceErr
} }
if sourceErr == nil { if sourceErr == nil {
var runtime model.PagesProjectSourceRuntime if _, runtimeErr := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID); runtimeErr != nil &&
runtimeErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). !errors.Is(runtimeErr, gorm.ErrRecordNotFound) {
Where("source_id = ?", source.ID).
First(&runtime).Error
if runtimeErr != nil && !errors.Is(runtimeErr, gorm.ErrRecordNotFound) {
return runtimeErr 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 return err
} }
if err := tx.Delete(&source).Error; err != nil { if err := repository.DeletePagesProjectSourceTx(tx, source); err != nil {
return err return err
} }
} }
@@ -291,10 +282,10 @@ func DeleteProject(ctx context.Context, id uint) error {
// ListProjectDeployments 列出项目的全部部署。 // ListProjectDeployments 列出项目的全部部署。
func ListProjectDeployments(ctx context.Context, projectID uint) ([]DeploymentView, error) { 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 return nil, err
} }
deployments, err := model.ListPagesDeployments(ctx, projectID) deployments, err := repository.ListPagesDeployments(ctx, projectID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -307,10 +298,10 @@ func ListProjectDeployments(ctx context.Context, projectID uint) ([]DeploymentVi
// ListDeploymentFiles 列出部署文件清单。 // ListDeploymentFiles 列出部署文件清单。
func ListDeploymentFiles(ctx context.Context, deploymentID uint) ([]DeploymentFileView, error) { 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 return nil, err
} }
files, err := model.ListPagesDeploymentFiles(ctx, deploymentID) files, err := repository.ListPagesDeploymentFiles(ctx, deploymentID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -330,7 +321,7 @@ func ListDeploymentFiles(ctx context.Context, deploymentID uint) ([]DeploymentFi
// UploadDeployment 上传 Pages 部署包(本地 multipart 文件)。 // UploadDeployment 上传 Pages 部署包(本地 multipart 文件)。
func UploadDeployment(ctx context.Context, projectID uint, fileHeader *multipart.FileHeader, createdBy string) (*DeploymentView, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -364,7 +355,7 @@ type UploadFromURLInput struct {
// UploadDeploymentFromURL downloads a package from url and creates a deployment. // UploadDeploymentFromURL downloads a package from url and creates a deployment.
func UploadDeploymentFromURL(ctx context.Context, projectID uint, rawURL string, createdBy string) (*DeploymentView, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -438,7 +429,7 @@ func createDeploymentFromTempPackage(
} }
}() }()
deployment := &model.PagesDeployment{} 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 var lockedProject model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedProject, project.ID).Error; err != nil { if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedProject, project.ID).Error; err != nil {
return err return err
@@ -539,7 +530,7 @@ func pruneProjectDeploymentHistory(ctx context.Context, projectID uint, keepCoun
// Returns the number of deployments deleted from the database. // Returns the number of deployments deleted from the database.
func pruneProjectDeploymentHistoryOnce(ctx context.Context, projectID uint, keepCount int, preserveCandidateID uint) (int, error) { func pruneProjectDeploymentHistoryOnce(ctx context.Context, projectID uint, keepCount int, preserveCandidateID uint) (int, error) {
var deletedDeployments []model.PagesDeployment 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 var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil { if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return fmt.Errorf("load pages project: %w", err) 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 { 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 { if err != nil {
return err return err
} }
@@ -725,7 +716,7 @@ func activateDeploymentTransaction(
now time.Time, now time.Time,
) (deploymentActivationAudit, error) { ) (deploymentActivationAudit, error) {
audit := deploymentActivationAudit{} 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 var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil { if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return err return err
@@ -755,23 +746,18 @@ func activateDeploymentTransaction(
} }
func lockDeploymentActivationSource(tx *gorm.DB, projectID uint) (*deploymentActivationSource, error) { func lockDeploymentActivationSource(tx *gorm.DB, projectID uint) (*deploymentActivationSource, error) {
var source model.PagesProjectSource source, err := repository.LockPagesProjectSourceByProjectIDTx(tx, projectID)
err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", projectID).
First(&source).Error
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil return nil, nil
} }
if err != nil { if err != nil {
return nil, err return nil, err
} }
var runtime model.PagesProjectSourceRuntime runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
Where("source_id = ?", source.ID).
First(&runtime).Error; err != nil {
return nil, err return nil, err
} }
return &deploymentActivationSource{Source: &source, Runtime: &runtime}, nil return &deploymentActivationSource{Source: source, Runtime: runtime}, nil
} }
func loadDeploymentActivationTarget( func loadDeploymentActivationTarget(
@@ -821,10 +807,10 @@ func fenceDeploymentActivationSource(
audit.SourceType = state.Source.SourceType audit.SourceType = state.Source.SourceType
audit.SourceIdentity = state.Source.SourceIdentity audit.SourceIdentity = state.Source.SourceIdentity
audit.AutoDisabled = state.Source.AutoUpdateEnabled 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, sourceColumnConfigVersion: state.Source.ConfigVersion + 1,
sourceColumnAutoUpdateEnabled: false, sourceColumnAutoUpdateEnabled: false,
}).Error; err != nil { }); err != nil {
return err return err
} }
if deployment.SourceIdentity != nil && *deployment.SourceIdentity == state.Source.SourceIdentity && if deployment.SourceIdentity != nil && *deployment.SourceIdentity == state.Source.SourceIdentity &&
@@ -835,13 +821,13 @@ func fenceDeploymentActivationSource(
state.Runtime.LastAppliedRevision = "" state.Runtime.LastAppliedRevision = ""
state.Runtime.LastAppliedDetail = "" 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_revision": state.Runtime.LastAppliedRevision,
"last_applied_detail": state.Runtime.LastAppliedDetail, "last_applied_detail": state.Runtime.LastAppliedDetail,
sourceRuntimeColumnSyncStatus: normalizedSourceRuntimeStatus(state.Runtime), sourceRuntimeColumnSyncStatus: normalizedSourceRuntimeStatus(state.Runtime),
sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil, sourceRuntimeColumnLeaseExpiresAt: nil,
}).Error })
} }
func switchActiveDeploymentTx( func switchActiveDeploymentTx(
@@ -867,7 +853,7 @@ func switchActiveDeploymentTx(
// GetDeploymentPackageHash returns the upload SHA-256 hash of the deployment package. // GetDeploymentPackageHash returns the upload SHA-256 hash of the deployment package.
// Prefer GetProjectLatestPackageHash for Agent latest-pointer pulls. // Prefer GetProjectLatestPackageHash for Agent latest-pointer pulls.
func GetDeploymentPackageHash(ctx context.Context, deploymentID uint) (string, error) { func GetDeploymentPackageHash(ctx context.Context, deploymentID uint) (string, error) {
deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) deployment, err := repository.GetPagesDeploymentByID(ctx, deploymentID)
if err != nil { if err != nil {
return "", err 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. // OpenDeploymentPackage opens the deployment artifact from the upload storage framework.
func OpenDeploymentPackage(ctx context.Context, deploymentID uint) (DeploymentPackage, error) { func OpenDeploymentPackage(ctx context.Context, deploymentID uint) (DeploymentPackage, error) {
deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) deployment, err := repository.GetPagesDeploymentByID(ctx, deploymentID)
if err != nil { if err != nil {
return DeploymentPackage{}, err return DeploymentPackage{}, err
} }
@@ -926,7 +912,7 @@ func resolveProjectActiveDeploymentForAgent(ctx context.Context, projectID uint)
if projectID == 0 { if projectID == 0 {
return nil, errors.New(errPagesProjectNotFound) return nil, errors.New(errPagesProjectNotFound)
} }
project, err := model.GetPagesProjectByID(ctx, projectID) project, err := repository.GetPagesProjectByID(ctx, projectID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -939,7 +925,7 @@ func resolveProjectActiveDeploymentForAgent(ctx context.Context, projectID uint)
if err := ensureProjectInActiveConfig(ctx, project.ID); err != nil { if err := ensureProjectInActiveConfig(ctx, project.ID); err != nil {
return nil, err return nil, err
} }
deployment, err := model.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) deployment, err := repository.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -957,7 +943,7 @@ func deploymentPackageHash(ctx context.Context, deployment *model.PagesDeploymen
if err := ensureDeploymentUploadRecord(ctx, deployment); err != nil { if err := ensureDeploymentUploadRecord(ctx, deployment); err != nil {
return "", err return "", err
} }
reloaded, err := model.GetPagesDeploymentByID(ctx, deployment.ID) reloaded, err := repository.GetPagesDeploymentByID(ctx, deployment.ID)
if err != nil { if err != nil {
return "", err return "", err
} }
@@ -1012,7 +998,7 @@ func ensureDeploymentUploadRecord(ctx context.Context, deployment *model.PagesDe
if deployment.UploadID > 0 { if deployment.UploadID > 0 {
return nil return nil
} }
project, err := model.GetPagesProjectByID(ctx, deployment.ProjectID) project, err := repository.GetPagesProjectByID(ctx, deployment.ProjectID)
if err != nil { if err != nil {
return err return err
} }
@@ -1089,7 +1075,7 @@ func attachLegacyDeploymentUpload(
uploadID uint64, uploadID uint64,
) (uint64, error) { ) (uint64, error) {
winnerUploadID := uint64(0) 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 var err error
winnerUploadID, err = attachLegacyDeploymentUploadTx(tx, projectID, deploymentID, uploadID) winnerUploadID, err = attachLegacyDeploymentUploadTx(tx, projectID, deploymentID, uploadID)
return err return err
@@ -1146,11 +1132,11 @@ func attachLegacyDeploymentUploadTx(
// it is the project's current active deployment and the project is used by the // it is the project's current active deployment and the project is used by the
// active main config. // active main config.
func ensureDeploymentInActiveSnapshot(ctx context.Context, deploymentID uint) error { func ensureDeploymentInActiveSnapshot(ctx context.Context, deploymentID uint) error {
deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID) deployment, err := repository.GetPagesDeploymentByID(ctx, deploymentID)
if err != nil { if err != nil {
return err return err
} }
project, err := model.GetPagesProjectByID(ctx, deployment.ProjectID) project, err := repository.GetPagesProjectByID(ctx, deployment.ProjectID)
if err != nil { if err != nil {
return err return err
} }
@@ -1163,7 +1149,7 @@ func ensureDeploymentInActiveSnapshot(ctx context.Context, deploymentID uint) er
// ensureProjectInActiveConfig checks that the Pages project is referenced by at // ensureProjectInActiveConfig checks that the Pages project is referenced by at
// least one pages route in the active main config snapshot. // least one pages route in the active main config snapshot.
func ensureProjectInActiveConfig(ctx context.Context, projectID uint) error { func ensureProjectInActiveConfig(ctx context.Context, projectID uint) error {
version, err := model.GetActiveConfigVersion(ctx) version, err := repository.GetActiveConfigVersion(ctx)
if err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errPagesPackageNotInActiveConfig) return errors.New(errPagesPackageNotInActiveConfig)
@@ -1192,7 +1178,7 @@ func ensureProjectInActiveConfig(ctx context.Context, projectID uint) error {
} }
if route.PagesDeployment.DeploymentID != 0 { if route.PagesDeployment.DeploymentID != 0 {
// Historical snapshot may only pin deployment_id. // 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 { if snapErr != nil {
continue continue
} }
@@ -1239,7 +1225,7 @@ func parseSnapshotRoutes(snapshotJSON string) ([]snapshotRouteRef, error) {
// DeleteDeployment 删除 Pages 部署。 // DeleteDeployment 删除 Pages 部署。
func DeleteDeployment(ctx context.Context, projectID uint, deploymentID uint) error { func DeleteDeployment(ctx context.Context, projectID uint, deploymentID uint) error {
var removed model.PagesDeployment 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 var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil { if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return err return err
@@ -1354,13 +1340,13 @@ func buildProjectView(ctx context.Context, project *model.PagesProject) (*View,
CreatedAt: project.CreatedAt, CreatedAt: project.CreatedAt,
UpdatedAt: project.UpdatedAt, UpdatedAt: project.UpdatedAt,
} }
count, err := model.CountPagesDeploymentsByProjectID(ctx, project.ID) count, err := repository.CountPagesDeploymentsByProjectID(ctx, project.ID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
view.DeploymentCount = count view.DeploymentCount = count
if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID != 0 { if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID != 0 {
deployment, err := model.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID) deployment, err := repository.GetPagesDeploymentByID(ctx, *project.ActiveDeploymentID)
if err == nil { if err == nil {
active := buildDeploymentView(deployment) active := buildDeploymentView(deployment)
view.ActiveDeployment = &active view.ActiveDeployment = &active
+11 -11
View File
@@ -243,7 +243,7 @@ func TestUpdateProjectValidatesActiveDeploymentEntry(t *testing.T) {
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), errPagesEntryFileMissing) assert.Contains(t, err.Error(), errPagesEntryFileMissing)
stored, err := model.GetPagesProjectByID(ctx, project.ID) stored, err := repository.GetPagesProjectByID(ctx, project.ID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "dist", stored.RootDir) assert.Equal(t, "dist", stored.RootDir)
assert.Equal(t, "index.html", stored.EntryFile) assert.Equal(t, "index.html", stored.EntryFile)
@@ -291,7 +291,7 @@ func TestUploadDeploymentStoresPackageInUploadFramework(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
assert.NotZero(t, deployment.UploadID) assert.NotZero(t, deployment.UploadID)
storedDeployment, err := model.GetPagesDeploymentByID(ctx, deployment.ID) storedDeployment, err := repository.GetPagesDeploymentByID(ctx, deployment.ID)
require.NoError(t, err) require.NoError(t, err)
assert.NotZero(t, storedDeployment.UploadID) assert.NotZero(t, storedDeployment.UploadID)
assert.Empty(t, storedDeployment.ArtifactPath) assert.Empty(t, storedDeployment.ArtifactPath)
@@ -371,7 +371,7 @@ func TestOpenDeploymentPackageHydratesLegacyArtifactPath(t *testing.T) {
require.Len(t, reader.File, 1) require.Len(t, reader.File, 1)
assert.Equal(t, "index.html", reader.File[0].Name) 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) require.NoError(t, err)
assert.NotZero(t, storedDeployment.UploadID) assert.NotZero(t, storedDeployment.UploadID)
assert.Empty(t, storedDeployment.ArtifactPath) assert.Empty(t, storedDeployment.ArtifactPath)
@@ -607,7 +607,7 @@ func TestPruneProjectDeploymentHistory(t *testing.T) {
ids = append(ids, deployment.ID) ids = append(ids, deployment.ID)
} }
// After 3 uploads with keep=2 and no active: only 2 newest remain. // 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.NoError(t, err)
require.Len(t, deployments, 2) require.Len(t, deployments, 2)
assert.Equal(t, ids[2], deployments[0].ID) assert.Equal(t, ids[2], deployments[0].ID)
@@ -622,11 +622,11 @@ func TestPruneProjectDeploymentHistory(t *testing.T) {
})), "root") })), "root")
require.NoError(t, err) require.NoError(t, err)
deployments, err = model.ListPagesDeployments(ctx, project.ID) deployments, err = repository.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err) require.NoError(t, err)
require.Len(t, deployments, 2, "must be at most N=2, not active+N newest") 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.NoError(t, err)
require.NotNil(t, storedProject.ActiveDeploymentID) require.NotNil(t, storedProject.ActiveDeploymentID)
assert.Equal(t, ids[1], *storedProject.ActiveDeploymentID) assert.Equal(t, ids[1], *storedProject.ActiveDeploymentID)
@@ -666,7 +666,7 @@ func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) {
"index.html": "v2", "index.html": "v2",
})), "user:1") })), "user:1")
require.NoError(t, err) require.NoError(t, err)
deployments, err := model.ListPagesDeployments(ctx, project.ID) deployments, err := repository.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err) require.NoError(t, err)
require.Len(t, deployments, 2) require.Len(t, deployments, 2)
@@ -674,7 +674,7 @@ func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) {
"index.html": "v3", "index.html": "v3",
})), "user:1") })), "user:1")
require.NoError(t, err) require.NoError(t, err)
deployments, err = model.ListPagesDeployments(ctx, project.ID) deployments, err = repository.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err) require.NoError(t, err)
require.Len(t, deployments, 2) require.Len(t, deployments, 2)
kept := map[uint]bool{} kept := map[uint]bool{}
@@ -690,7 +690,7 @@ func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) {
_, err = ActivateDeployment(ctx, project.ID, newCandidate.ID) _, err = ActivateDeployment(ctx, project.ID, newCandidate.ID)
require.NoError(t, err) require.NoError(t, err)
deployments, err = model.ListPagesDeployments(ctx, project.ID) deployments, err = repository.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err) require.NoError(t, err)
require.Len(t, deployments, 1) require.Len(t, deployments, 1)
assert.Equal(t, newCandidate.ID, deployments[0].ID) 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) deleted, err := pruneProjectDeploymentHistoryOnce(ctx, project.ID, 1, staleCandidate.ID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, 1, deleted) assert.Equal(t, 1, deleted)
deployments, err := model.ListPagesDeployments(ctx, project.ID) deployments, err := repository.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err) require.NoError(t, err)
require.Len(t, deployments, 2) require.Len(t, deployments, 2)
kept := map[uint]bool{} kept := map[uint]bool{}
@@ -762,7 +762,7 @@ func TestDeleteDeploymentAndProjectSoftDeleteUnreferencedArtifacts(t *testing.T)
var firstUpload model.Upload var firstUpload model.Upload
require.NoError(t, db.DB(ctx).First(&firstUpload, first.UploadID).Error) require.NoError(t, db.DB(ctx).First(&firstUpload, first.UploadID).Error)
assert.Equal(t, model.UploadStatusDeleted, firstUpload.Status) assert.Equal(t, model.UploadStatusDeleted, firstUpload.Status)
_, err = model.GetPagesProjectByID(ctx, project.ID) _, err = repository.GetPagesProjectByID(ctx, project.ID)
assert.Error(t, err) assert.Error(t, err)
} }
@@ -9,8 +9,9 @@ import (
"fmt" "fmt"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/upload" "github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/model"
) )
// ProjectLatestPackageMetadata describes the active package limits published to Agents. // 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 { if err := ensureDeploymentUploadRecord(ctx, deployment); err != nil {
return nil, err return nil, err
} }
deployment, err = model.GetPagesDeploymentByID(ctx, deployment.ID) deployment, err = repository.GetPagesDeploymentByID(ctx, deployment.ID)
if err != nil { if err != nil {
return nil, err return nil, err
} }
+4 -2
View File
@@ -10,6 +10,8 @@ import (
"path" "path"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"gorm.io/gorm" "gorm.io/gorm"
@@ -156,7 +158,7 @@ func loadActivePagesProject(ctx context.Context, projectID uint, siteName string
if siteName == "" { if siteName == "" {
siteName = "pages" siteName = "pages"
} }
project, err := model.GetPagesProjectByID(ctx, projectID) project, err := repository.GetPagesProjectByID(ctx, projectID)
if err != nil { if err != nil {
if errorsIsNotFound(err) { if errorsIsNotFound(err) {
return nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目不存在", siteName) 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 { if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 {
return nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 项目没有激活部署", siteName) 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 err != nil {
if errorsIsNotFound(err) { if errorsIsNotFound(err) {
return nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 激活部署不存在", siteName) return nil, nil, fmt.Errorf("路由 %s Pages 配置无效: pages 激活部署不存在", siteName)
@@ -13,6 +13,8 @@ import (
"strconv" "strconv"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/oauth" "github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -48,7 +50,7 @@ func TestUploadDeploymentHandlerRecordsCurrentUserActor(t *testing.T) {
UploadDeploymentHandler(c) UploadDeploymentHandler(c)
assert.Equal(t, http.StatusOK, recorder.Code) 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.NoError(t, err)
require.Len(t, deployments, 1) require.Len(t, deployments, 1)
assert.Equal(t, "user:42", deployments[0].CreatedBy) assert.Equal(t, "user:42", deployments[0].CreatedBy)
@@ -83,7 +85,7 @@ func TestUploadDeploymentFromURLHandlerRecordsCurrentUserActor(t *testing.T) {
UploadDeploymentFromURLHandler(c) UploadDeploymentFromURLHandler(c)
assert.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String()) 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.NoError(t, err)
require.Len(t, deployments, 1) require.Len(t, deployments, 1)
assert.Equal(t, "user:77", deployments[0].CreatedBy) assert.Equal(t, "user:77", deployments[0].CreatedBy)
+30 -46
View File
@@ -15,10 +15,9 @@ import (
"strings" "strings"
"time" "time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm" "gorm.io/gorm"
"gorm.io/gorm/clause"
) )
const ( const (
@@ -120,7 +119,7 @@ type remoteSourceConfig struct {
// GetSource returns the current persisted source or a manual discriminator. // GetSource returns the current persisted source or a manual discriminator.
func GetSource(ctx context.Context, projectID uint) (*SourceView, error) { 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 return nil, err
} }
source, runtime, err := loadSourceByProject(ctx, projectID) source, runtime, err := loadSourceByProject(ctx, projectID)
@@ -156,7 +155,7 @@ func UpdateSourceAs(
changed := false changed := false
var persistedSource model.PagesProjectSource 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 var err error
switch strings.TrimSpace(input.SourceType) { switch strings.TrimSpace(input.SourceType) {
case PagesSourceTypeRemoteURL: case PagesSourceTypeRemoteURL:
@@ -169,7 +168,12 @@ func UpdateSourceAs(
if err != nil || !changed || strings.TrimSpace(input.SourceType) != PagesSourceTypeGitHubRelease { if err != nil || !changed || strings.TrimSpace(input.SourceType) != PagesSourceTypeGitHubRelease {
return err 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 { if err != nil {
return nil, err return nil, err
@@ -193,8 +197,7 @@ func UpdateSourceAs(
} }
func updateRemoteSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) (bool, error) { func updateRemoteSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) (bool, error) {
var project model.PagesProject if _, err := repository.LockPagesProjectByIDTx(tx, projectID); err != nil {
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return false, err return false, err
} }
existing, hasExisting, err := loadProjectSourceForUpdate(tx, projectID) 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) { func loadProjectSourceForUpdate(tx *gorm.DB, projectID uint) (*model.PagesProjectSource, bool, error) {
var source model.PagesProjectSource source, err := repository.LockPagesProjectSourceByProjectIDTx(tx, projectID)
err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", projectID).
First(&source).Error
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return &source, false, nil return &model.PagesProjectSource{}, false, nil
} }
if err != nil { if err != nil {
return nil, false, err return nil, false, err
} }
return &source, true, nil return source, true, nil
} }
func buildRemoteSourceConfig( func buildRemoteSourceConfig(
@@ -256,13 +256,13 @@ func createRemoteSourceTx(tx *gorm.DB, projectID uint, config remoteSourceConfig
ConfigVersion: 1, ConfigVersion: 1,
SourceIdentity: config.Identity, SourceIdentity: config.Identity,
} }
if err := tx.Create(source).Error; err != nil { if err := repository.CreatePagesProjectSourceTx(tx, source); err != nil {
return err return err
} }
return tx.Create(&model.PagesProjectSourceRuntime{ return repository.CreatePagesProjectSourceRuntimeTx(tx, &model.PagesProjectSourceRuntime{
SourceID: source.ID, SourceID: source.ID,
SyncStatus: pagesSourceStatusIdle, SyncStatus: pagesSourceStatusIdle,
}).Error })
} }
func updateExistingRemoteSourceTx( func updateExistingRemoteSourceTx(
@@ -273,14 +273,12 @@ func updateExistingRemoteSourceTx(
if !remoteSourceConfigChanged(existing, config) { if !remoteSourceConfigChanged(existing, config) {
return false, nil return false, nil
} }
var runtime model.PagesProjectSourceRuntime runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, existing.ID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
Where("source_id = ?", existing.ID).
First(&runtime).Error; err != nil {
return false, err return false, err
} }
identityChanged := existing.SourceIdentity != config.Identity 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, "source_type": PagesSourceTypeRemoteURL,
"remote_url": config.URL, "remote_url": config.URL,
"allow_insecure": config.AllowInsecure, "allow_insecure": config.AllowInsecure,
@@ -292,10 +290,10 @@ func updateExistingRemoteSourceTx(
"check_interval_minutes": 0, "check_interval_minutes": 0,
sourceColumnConfigVersion: existing.ConfigVersion + 1, sourceColumnConfigVersion: existing.ConfigVersion + 1,
"source_identity": config.Identity, "source_identity": config.Identity,
}).Error; err != nil { }); err != nil {
return false, err return false, err
} }
return true, resetRuntimeAfterSourceUpdate(tx, &runtime, identityChanged) return true, resetRuntimeAfterSourceUpdate(tx, runtime, identityChanged)
} }
func remoteSourceConfigChanged(existing *model.PagesProjectSource, config remoteSourceConfig) bool { 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. // DeleteSource idempotently switches a project back to manual mode.
func DeleteSource(ctx context.Context, projectID uint) (*SourceView, error) { func DeleteSource(ctx context.Context, projectID uint) (*SourceView, error) {
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 := repository.LockPagesProjectByIDTx(tx, projectID); err != nil {
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return err return err
} }
var source model.PagesProjectSource source, err := repository.LockPagesProjectSourceByProjectIDTx(tx, projectID)
err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", projectID).
First(&source).Error
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil return nil
} }
if err != nil { if err != nil {
return err return err
} }
var runtime model.PagesProjectSourceRuntime if _, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID); err != nil &&
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). !errors.Is(err, gorm.ErrRecordNotFound) {
Where("source_id = ?", source.ID).
First(&runtime).Error; err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return err 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 err
} }
return tx.Delete(&source).Error return repository.DeletePagesProjectSourceTx(tx, source)
}) })
if err != nil { if err != nil {
return nil, err 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) { func loadSourceByProject(ctx context.Context, projectID uint) (*model.PagesProjectSource, *model.PagesProjectSourceRuntime, error) {
var source model.PagesProjectSource return repository.GetPagesProjectSourceAndRuntimeByProjectID(ctx, projectID)
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
} }
func buildSourceView(source *model.PagesProjectSource, runtime *model.PagesProjectSourceRuntime) (*SourceView, error) { func buildSourceView(source *model.PagesProjectSource, runtime *model.PagesProjectSourceRuntime) (*SourceView, error) {
@@ -515,7 +499,7 @@ func resetRuntimeAfterSourceUpdate(tx *gorm.DB, runtime *model.PagesProjectSourc
} else { } else {
updates[sourceRuntimeColumnSyncStatus] = normalizedSourceRuntimeStatus(runtime) updates[sourceRuntimeColumnSyncStatus] = normalizedSourceRuntimeStatus(runtime)
} }
return tx.Model(runtime).Updates(updates).Error return repository.UpdatePagesProjectSourceRuntimeTx(tx, runtime, updates)
} }
func normalizedSourceRuntimeStatus(runtime *model.PagesProjectSourceRuntime) string { func normalizedSourceRuntimeStatus(runtime *model.PagesProjectSourceRuntime) string {
@@ -11,7 +11,6 @@ import (
"time" "time"
"github.com/Rain-kl/Wavelet/internal/apps/upload" "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/model"
"github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
@@ -62,7 +61,7 @@ func ReconcilePagesOrphanUploads(
} }
cutoff := now.UTC().Add(-pagesOrphanUploadIsolation) cutoff := now.UTC().Add(-pagesOrphanUploadIsolation)
systemUser := repository.GetSystemUser(ctx) systemUser := repository.GetSystemUser(ctx)
candidates, err := model.ListPagesOrphanUploadCandidates(ctx, model.PagesOrphanUploadCandidateQuery{ candidates, err := repository.ListPagesOrphanUploadCandidates(ctx, model.PagesOrphanUploadCandidateQuery{
SystemUserID: systemUser.ID, SystemUserID: systemUser.ID,
UploadType: upload.ReservedPagesDeploymentType, UploadType: upload.ReservedPagesDeploymentType,
Marker: pagesIngestMarkerV2, Marker: pagesIngestMarkerV2,
@@ -130,7 +129,7 @@ func reconcilePagesOrphanUploadCandidate(
outcome := pagesOrphanCleanupSkipped outcome := pagesOrphanCleanupSkipped
uploadLocked := false 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) scopeOutcome, proceed, err := lockPagesOrphanCleanupScope(ctx, tx, candidate.ID, marker)
if err != nil { if err != nil {
return err return err
@@ -179,14 +178,13 @@ func lockPagesOrphanCleanupScope(
return pagesOrphanCleanupSkipped, true, nil return pagesOrphanCleanupSkipped, true, nil
} }
var source model.PagesProjectSource source, err := repository.LockPagesProjectSourceByIDTx(tx, *marker.SourceID)
sourceExists, err := lockOptionalPagesCleanupRecord(tx, &source, "id = ?", *marker.SourceID) if errors.Is(err, gorm.ErrRecordNotFound) {
return pagesOrphanCleanupSkipped, true, nil
}
if err != nil { if err != nil {
return pagesOrphanCleanupSkipped, false, err return pagesOrphanCleanupSkipped, false, err
} }
if !sourceExists {
return pagesOrphanCleanupSkipped, true, nil
}
if source.ProjectID != marker.ProjectID { if source.ProjectID != marker.ProjectID {
logger.WarnF(ctx, logger.WarnF(ctx,
"[PagesSource] orphan upload source ownership mismatch: upload_id=%d project_id=%d source_id=%d source_project_id=%d", "[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 return pagesOrphanCleanupInvalidMarker, false, nil
} }
var runtime model.PagesProjectSourceRuntime runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID)
runtimeExists, err := lockOptionalPagesCleanupRecord(tx, &runtime, "source_id = ?", source.ID) if err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
if err != nil {
return pagesOrphanCleanupSkipped, false, err return pagesOrphanCleanupSkipped, false, err
} }
// Read the real clock only after obtaining the runtime row lock. The scanner // 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. // snapshot time is only an isolation cutoff and may be stale after lock wait.
leaseCheckedAt := time.Now() 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 pagesOrphanCleanupLeaseBusy, false, nil
} }
return pagesOrphanCleanupSkipped, true, nil return pagesOrphanCleanupSkipped, true, nil
+49 -64
View File
@@ -11,10 +11,9 @@ import (
"strings" "strings"
"time" "time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm" "gorm.io/gorm"
"gorm.io/gorm/clause"
) )
const ( const (
@@ -89,24 +88,16 @@ func acquireSourceLease(
now := time.Now() now := time.Now()
expiresAt := now.Add(leaseDuration) expiresAt := now.Add(leaseDuration)
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). rows, err := repository.TryAcquirePagesSourceRuntimeLease(ctx, sourceID, expectedConfigVersion, now, map[string]any{
Where("source_id = ?", sourceID). sourceRuntimeColumnLeaseToken: token,
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). sourceRuntimeColumnLeaseExpiresAt: expiresAt,
Where( sourceRuntimeColumnSyncStatus: status,
"EXISTS (SELECT 1 FROM of_pages_project_sources source WHERE source.id = ? AND source.config_version = ?)", sourceRuntimeColumnLastError: "",
sourceID, })
expectedConfigVersion, if err != nil {
). return nil, sourceLeaseStale, err
Updates(map[string]any{
sourceRuntimeColumnLeaseToken: token,
sourceRuntimeColumnLeaseExpiresAt: expiresAt,
sourceRuntimeColumnSyncStatus: status,
sourceRuntimeColumnLastError: "",
})
if result.Error != nil {
return nil, sourceLeaseStale, result.Error
} }
if result.RowsAffected == 0 { if rows == 0 {
outcome, inspectErr := inspectSourceLeaseMiss(ctx, sourceID, expectedConfigVersion, now) outcome, inspectErr := inspectSourceLeaseMiss(ctx, sourceID, expectedConfigVersion, now)
return nil, outcome, inspectErr return nil, outcome, inspectErr
} }
@@ -129,34 +120,30 @@ func loadSourceExecutionSnapshot(
token string, token string,
) (*sourceExecutionSnapshot, error) { ) (*sourceExecutionSnapshot, error) {
var snapshot sourceExecutionSnapshot var snapshot sourceExecutionSnapshot
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error {
var source model.PagesProjectSource source, err := repository.GetPagesProjectSourceByIDTx(tx, sourceID)
if err := tx.Where("id = ?", sourceID).First(&source).Error; err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return errSourceLeaseSnapshotStale return errSourceLeaseSnapshotStale
} }
return err return err
} }
var project model.PagesProject project, err := repository.LockPagesProjectByIDTx(tx, source.ProjectID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
First(&project, source.ProjectID).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return errSourceLeaseSnapshotStale return errSourceLeaseSnapshotStale
} }
return err return err
} }
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). source, err = repository.LockPagesProjectSourceByIDTx(tx, sourceID)
Where("id = ?", sourceID). if err != nil {
First(&source).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return errSourceLeaseSnapshotStale return errSourceLeaseSnapshotStale
} }
return err return err
} }
var runtime model.PagesProjectSourceRuntime runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
Where("source_id = ?", source.ID).
First(&runtime).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return errSourceLeaseSnapshotStale return errSourceLeaseSnapshotStale
} }
@@ -207,8 +194,8 @@ func inspectSourceLeaseMiss(
expectedConfigVersion int, expectedConfigVersion int,
now time.Time, now time.Time,
) (sourceLeaseOutcome, error) { ) (sourceLeaseOutcome, error) {
var source model.PagesProjectSource source, err := repository.GetPagesProjectSourceByID(ctx, sourceID)
if err := db.DB(ctx).Where("id = ?", sourceID).First(&source).Error; err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return sourceLeaseStale, nil return sourceLeaseStale, nil
} }
@@ -217,8 +204,8 @@ func inspectSourceLeaseMiss(
if source.ConfigVersion != expectedConfigVersion { if source.ConfigVersion != expectedConfigVersion {
return sourceLeaseStale, nil return sourceLeaseStale, nil
} }
var runtime model.PagesProjectSourceRuntime runtime, err := repository.GetPagesProjectSourceRuntimeBySourceID(ctx, sourceID)
if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return sourceLeaseStale, nil return sourceLeaseStale, nil
} }
@@ -255,13 +242,11 @@ func renewSourceLease(ctx context.Context, snapshot *sourceExecutionSnapshot, du
} }
now := time.Now() now := time.Now()
expiresAt := now.Add(duration) expiresAt := now.Add(duration)
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). rows, err := repository.RenewPagesSourceRuntimeLease(ctx, snapshot.SourceID, snapshot.LeaseToken, now, expiresAt)
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now). if err != nil {
Updates(map[string]any{sourceRuntimeColumnLeaseExpiresAt: expiresAt}) return false, err
if result.Error != nil {
return false, result.Error
} }
if result.RowsAffected == 0 { if rows == 0 {
return false, nil return false, nil
} }
snapshot.LeaseExpiresAt = expiresAt snapshot.LeaseExpiresAt = expiresAt
@@ -274,14 +259,19 @@ func failSourceLease(ctx context.Context, snapshot *sourceExecutionSnapshot, mes
} }
message = safeSourceRuntimeError(message) message = safeSourceRuntimeError(message)
now := time.Now() now := time.Now()
return db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). _, err := repository.UpdatePagesSourceRuntimeByActiveLease(
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now). ctx,
Updates(map[string]any{ snapshot.SourceID,
snapshot.LeaseToken,
now,
map[string]any{
sourceRuntimeColumnSyncStatus: pagesSourceStatusFailed, sourceRuntimeColumnSyncStatus: pagesSourceStatusFailed,
sourceRuntimeColumnLastError: message, sourceRuntimeColumnLastError: message,
sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil, sourceRuntimeColumnLeaseExpiresAt: nil,
}).Error },
)
return err
} }
func safeSourceRuntimeError(message string) string { func safeSourceRuntimeError(message string) string {
@@ -296,8 +286,8 @@ func safeSourceRuntimeError(message string) string {
} }
func sourceLeaseIsBusy(ctx context.Context, sourceID uint) (bool, error) { func sourceLeaseIsBusy(ctx context.Context, sourceID uint) (bool, error) {
var runtime model.PagesProjectSourceRuntime runtime, err := repository.GetPagesProjectSourceRuntimeBySourceID(ctx, sourceID)
if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { if err != nil {
return false, err return false, err
} }
return runtime.LeaseExpiresAt != nil && runtime.LeaseExpiresAt.After(time.Now()), nil return runtime.LeaseExpiresAt != nil && runtime.LeaseExpiresAt.After(time.Now()), nil
@@ -326,35 +316,30 @@ func recoverExpiredSourceLease(
sourceRuntimeColumnLeaseExpiresAt: nil, sourceRuntimeColumnLeaseExpiresAt: nil,
sourceRuntimeColumnNextCheckAt: nextCheckAt, sourceRuntimeColumnNextCheckAt: nextCheckAt,
} }
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). rows, err := repository.RecoverExpiredPagesSourceRuntimeLease(
Where("source_id = ?", sourceID). ctx, sourceID, token, expiresAt, status, now, updates,
Where("lease_token = ?", token). )
Where("lease_expires_at = ? AND lease_expires_at <= ?", expiresAt, now). if err != nil {
Where("sync_status = ?", status). return false, err
Updates(updates)
if result.Error != nil {
return false, result.Error
} }
return result.RowsAffected == 1, nil return rows == 1, nil
} }
// fenceAndNormalizeRuntime invalidates in-flight work while preserving safe // fenceAndNormalizeRuntime invalidates in-flight work while preserving safe
// seen/applied cursors. The caller must already hold the source row lock. // seen/applied cursors. The caller must already hold the source row lock.
func fenceAndNormalizeRuntime(tx *gorm.DB, sourceID uint) error { func fenceAndNormalizeRuntime(tx *gorm.DB, sourceID uint) error {
var runtime model.PagesProjectSourceRuntime runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, sourceID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
Where("source_id = ?", sourceID).
First(&runtime).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil return nil
} }
return err return err
} }
return tx.Model(&runtime).Updates(map[string]any{ return repository.UpdatePagesProjectSourceRuntimeTx(tx, runtime, map[string]any{
sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil, sourceRuntimeColumnLeaseExpiresAt: nil,
sourceRuntimeColumnSyncStatus: normalizedSourceRuntimeStatus(&runtime), sourceRuntimeColumnSyncStatus: normalizedSourceRuntimeStatus(runtime),
}).Error })
} }
func sourceHasSameReleaseReplacement(runtime *model.PagesProjectSourceRuntime) bool { func sourceHasSameReleaseReplacement(runtime *model.PagesProjectSourceRuntime) bool {
@@ -10,6 +10,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
) )
@@ -286,7 +288,7 @@ func TestSourceConfigAndProjectContentChangesFenceLease(t *testing.T) {
if renewed { if renewed {
t.Error("renewSourceLease(after content update) = true, want false") 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 { if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
} }
+36 -73
View File
@@ -11,11 +11,10 @@ import (
"fmt" "fmt"
"time" "time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
) )
const ( const (
@@ -65,20 +64,6 @@ type pagesSourceProviderBackoff struct {
RetryAt string `json:"retry_at"` 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 ( var (
pagesSourceScanNow = time.Now pagesSourceScanNow = time.Now
reconcilePagesSourceOrphans = ReconcilePagesOrphanUploads reconcilePagesSourceOrphans = ReconcilePagesOrphanUploads
@@ -173,17 +158,11 @@ func recoverExpiredPagesSourceLeases(
now time.Time, now time.Time,
summary *pagesSourceScanSummary, summary *pagesSourceScanSummary,
) error { ) error {
var candidates []expiredSourceLeaseCandidate candidates, err := repository.ListExpiredPagesSourceLeaseCandidates(
err := db.DB(ctx). ctx,
Table("of_pages_project_source_runtime AS runtime"). now,
Select(`runtime.source_id, runtime.lease_token, runtime.lease_expires_at, []string{pagesSourceStatusChecking, pagesSourceStatusSyncing},
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
if err != nil { if err != nil {
return err return err
} }
@@ -227,27 +206,18 @@ func scanDueGitHubSources(
now time.Time, now time.Time,
summary *pagesSourceScanSummary, summary *pagesSourceScanSummary,
) error { ) error {
dueQuery := func() *gorm.DB { dueCount, err := repository.CountDueGitHubPagesSourceChecks(
return db.DB(ctx). ctx, now, PagesSourceTypeGitHubRelease, githubReleaseSelectorLatest,
Table("of_pages_project_source_runtime AS runtime"). )
Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id"). if err != nil {
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 {
return err return err
} }
summary.DueSources = int(dueCount) summary.DueSources = int(dueCount)
var candidates []dueGitHubSourceCandidate candidates, err := repository.ListDueGitHubPagesSourceChecks(
if err := dueQuery(). ctx, now, PagesSourceTypeGitHubRelease, githubReleaseSelectorLatest, pagesSourceScanBatchSize,
Select("source.id AS source_id, source.config_version"). )
Order("runtime.next_check_at ASC"). if err != nil {
Order("source.id ASC").
Limit(pagesSourceScanBatchSize).
Scan(&candidates).Error; err != nil {
return err return err
} }
summary.SelectedSources = len(candidates) summary.SelectedSources = len(candidates)
@@ -261,8 +231,10 @@ func scanDueGitHubSources(
for _, candidate := range candidates { for _, candidate := range candidates {
scanOneDueGitHubSource(ctx, candidate, summary) scanOneDueGitHubSource(ctx, candidate, summary)
} }
var remainingDue int64 remainingDue, err := repository.CountDueGitHubPagesSourceChecks(
if err := dueQuery().Count(&remainingDue).Error; err != nil { ctx, now, PagesSourceTypeGitHubRelease, githubReleaseSelectorLatest,
)
if err != nil {
return err return err
} }
summary.Backlog = int(remainingDue) summary.Backlog = int(remainingDue)
@@ -272,7 +244,7 @@ func scanDueGitHubSources(
func scanOneDueGitHubSource( func scanOneDueGitHubSource(
ctx context.Context, ctx context.Context,
candidate dueGitHubSourceCandidate, candidate model.PagesDueGitHubSourceCandidate,
summary *pagesSourceScanSummary, summary *pagesSourceScanSummary,
) { ) {
snapshot, outcome, err := acquireSourceLease( snapshot, outcome, err := acquireSourceLease(
@@ -343,11 +315,8 @@ func recordPagesSourceProviderBackoff(
} }
retryAt := domainError.retryAt retryAt := domainError.retryAt
var runtime model.PagesProjectSourceRuntime runtime, err := repository.GetPagesProjectSourceRuntimeBySourceID(ctx, sourceID)
if err := db.DB(ctx). if err != nil {
Select("next_check_at").
Where("source_id = ?", sourceID).
First(&runtime).Error; err != nil {
logger.WarnF(ctx, "[PagesSourceScan] load provider backoff deadline failed: source_id=%d error=%v", sourceID, err) logger.WarnF(ctx, "[PagesSourceScan] load provider backoff deadline failed: source_id=%d error=%v", sourceID, err)
} else if runtime.NextCheckAt != nil { } else if runtime.NextCheckAt != nil {
retryAt = runtime.NextCheckAt retryAt = runtime.NextCheckAt
@@ -448,29 +417,23 @@ func recordPagesSourceAutoDispatchFailure(
if retryAt != nil && retryAt.After(next) { if retryAt != nil && retryAt.After(next) {
next = retryAt.In(now.Location()) next = retryAt.In(now.Location())
} }
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). rows, err := repository.RecordPagesSourceAutoDispatchFailure(
Where("source_id = ?", snapshot.SourceID). ctx,
Where("sync_status = ? AND last_seen_revision = ?", pagesSourceStatusUpdateAvailable, revision). snapshot.SourceID,
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). snapshot.SourceConfigVersion,
Where(`EXISTS ( PagesSourceTypeGitHubRelease,
SELECT 1 FROM of_pages_project_sources AS source githubReleaseSelectorLatest,
WHERE source.id = ? AND source.config_version = ? revision,
AND source.source_type = ? AND source.release_selector = ? pagesSourceStatusUpdateAvailable,
AND source.auto_update_enabled = ? now,
)`, map[string]any{
snapshot.SourceID,
snapshot.SourceConfigVersion,
PagesSourceTypeGitHubRelease,
githubReleaseSelectorLatest,
true,
).
Updates(map[string]any{
sourceRuntimeColumnSyncStatus: pagesSourceStatusUpdateAvailable, sourceRuntimeColumnSyncStatus: pagesSourceStatusUpdateAvailable,
sourceRuntimeColumnLastError: errPagesSourceTaskDispatchFailed, sourceRuntimeColumnLastError: errPagesSourceTaskDispatchFailed,
sourceRuntimeColumnNextCheckAt: &next, 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
} }
@@ -16,6 +16,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/integration/githubrelease" "github.com/Rain-kl/Wavelet/internal/integration/githubrelease"
"github.com/Rain-kl/Wavelet/internal/model" "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 { if err != nil || synced == nil || synced.Deployment == nil || synced.Stale {
t.Fatalf("syncGitHubSourceWithTrigger() = %+v, %v; want active deployment", synced, err) 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 { if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v", synced.Deployment.ID, err) t.Fatalf("GetPagesDeploymentByID(%d) error = %v", synced.Deployment.ID, err)
} }
+23 -35
View File
@@ -16,9 +16,9 @@ import (
"unicode/utf8" "unicode/utf8"
"github.com/Rain-kl/Wavelet/internal/apps/upload" "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/infra/task"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm" "gorm.io/gorm"
"gorm.io/gorm/clause" "gorm.io/gorm/clause"
@@ -375,14 +375,7 @@ func findSourceDeployment(
sourceIdentity string, sourceIdentity string,
revision string, revision string,
) (*model.PagesDeployment, error) { ) (*model.PagesDeployment, error) {
var deployment model.PagesDeployment return repository.GetPagesDeploymentBySourceRevision(ctx, projectID, sourceIdentity, revision)
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
} }
func commitSourceDeploymentWithTrigger( func commitSourceDeploymentWithTrigger(
@@ -408,7 +401,7 @@ func commitSourceDeploymentWithTrigger(
var committed model.PagesDeployment var committed model.PagesDeployment
reused := false reused := false
ingestReferenced := 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) state, err := lockSourceCommitState(tx, snapshot)
if err != nil { if err != nil {
return err return err
@@ -445,18 +438,15 @@ func commitSourceDeploymentWithTrigger(
func lockSourceCommitState(tx *gorm.DB, snapshot *sourceExecutionSnapshot) (*sourceCommitState, error) { func lockSourceCommitState(tx *gorm.DB, snapshot *sourceExecutionSnapshot) (*sourceCommitState, error) {
state := &sourceCommitState{} state := &sourceCommitState{}
var project model.PagesProject project, err := repository.LockPagesProjectByIDTx(tx, snapshot.ProjectID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
First(&project, snapshot.ProjectID).Error; err != nil {
return nil, sourceFenceRecordError(err) return nil, sourceFenceRecordError(err)
} }
if project.ContentConfigVersion != snapshot.ContentConfigVersion { if project.ContentConfigVersion != snapshot.ContentConfigVersion {
return nil, errSourceFinalFence return nil, errSourceFinalFence
} }
var source model.PagesProjectSource source, err := repository.LockPagesProjectSourceByIDAndProjectIDTx(tx, snapshot.SourceID, snapshot.ProjectID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
Where("id = ? AND project_id = ?", snapshot.SourceID, snapshot.ProjectID).
First(&source).Error; err != nil {
return nil, sourceFenceRecordError(err) return nil, sourceFenceRecordError(err)
} }
if source.ConfigVersion != snapshot.SourceConfigVersion || if source.ConfigVersion != snapshot.SourceConfigVersion ||
@@ -464,15 +454,13 @@ func lockSourceCommitState(tx *gorm.DB, snapshot *sourceExecutionSnapshot) (*sou
source.SourceType != snapshot.SourceType { source.SourceType != snapshot.SourceType {
return nil, errSourceFinalFence return nil, errSourceFinalFence
} }
var runtime model.PagesProjectSourceRuntime runtime, err := repository.LockPagesProjectSourceRuntimeBySourceIDTx(tx, source.ID)
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}). if err != nil {
Where("source_id = ?", source.ID).
First(&runtime).Error; err != nil {
return nil, sourceFenceRecordError(err) return nil, sourceFenceRecordError(err)
} }
state.Project = &project state.Project = project
state.Source = &source state.Source = source
state.Runtime = &runtime state.Runtime = runtime
if err := refreshSourceCommitLease(state, snapshot); err != nil { if err := refreshSourceCommitLease(state, snapshot); err != nil {
return nil, err return nil, err
} }
@@ -675,13 +663,12 @@ func activateSourceDeploymentTx(
} }
nextCheckAt = &next nextCheckAt = &next
} }
result := tx.Model(&model.PagesProjectSourceRuntime{}). rows, err := repository.UpdatePagesSourceRuntimeByActiveLeaseTx(
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", tx,
state.Runtime.SourceID, state.Runtime.SourceID,
state.Runtime.LeaseToken, state.Runtime.LeaseToken,
finishedAt, finishedAt,
). map[string]any{
Updates(map[string]any{
"last_seen_revision": revision, "last_seen_revision": revision,
"last_seen_detail": detailJSON, "last_seen_detail": detailJSON,
"last_applied_revision": revision, "last_applied_revision": revision,
@@ -693,11 +680,12 @@ func activateSourceDeploymentTx(
"next_check_at": nextCheckAt, "next_check_at": nextCheckAt,
sourceRuntimeColumnLeaseToken: "", sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil, sourceRuntimeColumnLeaseExpiresAt: nil,
}) },
if result.Error != nil { )
return result.Error if err != nil {
return err
} }
if result.RowsAffected != 1 { if rows != 1 {
return errSourceFinalFence return errSourceFinalFence
} }
return nil return nil
@@ -17,6 +17,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/upload" "github.com/Rain-kl/Wavelet/internal/apps/upload"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
@@ -70,7 +72,7 @@ func mustCreateActiveManualDeployment(
if _, err := ActivateDeploymentAs(ctx, projectID, view.ID, "user:1"); err != nil { 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) 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 { if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", view.ID, err) 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 { if first == nil || first.Stale || first.Reused || first.Deployment == nil {
t.Fatalf("syncRemoteSource(first) = %+v, want new active deployment", first) 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 { if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
} }
if storedProject.ActiveDeploymentID == nil || *storedProject.ActiveDeploymentID != first.Deployment.ID { if storedProject.ActiveDeploymentID == nil || *storedProject.ActiveDeploymentID != first.Deployment.ID {
t.Fatalf("project ActiveDeploymentID = %v, want %d", 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 { if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", first.Deployment.ID, err) 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 { if outcome == nil || !outcome.Stale || outcome.Deployment != nil {
t.Fatalf("syncRemoteSource(final fence) = %+v, want stale outcome without deployment", outcome) 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 { if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) 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 { if err != nil || first == nil || first.Deployment == nil {
t.Fatalf("syncRemoteSource(seed) = (%+v, %v), want deployment", first, err) 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 { if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", first.Deployment.ID, err) t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", first.Deployment.ID, err)
} }
@@ -403,14 +405,14 @@ func TestCommitSourceDeploymentRechecksLeaseAfterUploadLocks(t *testing.T) {
if nowCalls != 2 { if nowCalls != 2 {
t.Fatalf("sourceCommitNow calls = %d, want runtime-lock and post-upload checks", nowCalls) 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 { if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
} }
if storedProject.ActiveDeploymentID != nil { if storedProject.ActiveDeploymentID != nil {
t.Errorf("ActiveDeploymentID after expiry recheck = %v, want nil", storedProject.ActiveDeploymentID) 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 { if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) after expiry error = %v, want nil", deployment.ID, err) t.Fatalf("GetPagesDeploymentByID(%d) after expiry error = %v, want nil", deployment.ID, err)
} }
@@ -479,7 +481,7 @@ func assertPagesSyncFailureState(
wantDeploymentCount int64, wantDeploymentCount int64,
) { ) {
t.Helper() t.Helper()
project, err := model.GetPagesProjectByID(ctx, projectID) project, err := repository.GetPagesProjectByID(ctx, projectID)
if err != nil { if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", projectID, err) t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", projectID, err)
} }
@@ -589,7 +591,7 @@ func TestCommitSourceDeploymentRejectsDeletedTargetUpload(t *testing.T) {
if !errors.Is(err, errSourceFinalFence) { if !errors.Is(err, errSourceFinalFence) {
t.Errorf("commitSourceDeployment(deleted upload) error = %v, want %v", 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 { if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
} }
@@ -13,9 +13,9 @@ import (
"strconv" "strconv"
"strings" "strings"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm" "gorm.io/gorm"
) )
@@ -135,8 +135,8 @@ func (h *SourceActionHandler) Execute(ctx context.Context, payload []byte) (*tas
return nil, task.PermanentError(errPagesSourceActionInvalid) return nil, task.PermanentError(errPagesSourceActionInvalid)
} }
var source model.PagesProjectSource source, err := repository.GetPagesProjectSourceByID(ctx, input.SourceID)
if err := db.DB(ctx).Where("id = ?", input.SourceID).First(&source).Error; err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
task.AppendLog(ctx, "[resolve] 来源已不存在,本次任务跳过") task.AppendLog(ctx, "[resolve] 来源已不存在,本次任务跳过")
return &task.TaskResult{Message: errPagesSourceActionStale}, nil return &task.TaskResult{Message: errPagesSourceActionStale}, nil
@@ -179,7 +179,7 @@ func (h *SourceActionHandler) Execute(ctx context.Context, payload []byte) (*tas
if input.Action == sourceActionCheck { if input.Action == sourceActionCheck {
return executeGitHubCheckAction(ctx, snapshot) 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) { func executeGitHubCheckAction(ctx context.Context, snapshot *sourceExecutionSnapshot) (*task.TaskResult, error) {
@@ -313,14 +313,14 @@ func dispatchSourceActionByProject(
return nil, errors.New(errPagesSourceActionInvalid) return nil, errors.New(errPagesSourceActionInvalid)
} }
var source model.PagesProjectSource source, err := repository.GetPagesProjectSourceByProjectID(ctx, projectID)
if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New(errPagesSourceNotFound) return nil, errors.New(errPagesSourceNotFound)
} }
return nil, err 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 return nil, err
} }
busy, err := sourceLeaseIsBusy(ctx, source.ID) busy, err := sourceLeaseIsBusy(ctx, source.ID)
@@ -330,7 +330,7 @@ func dispatchSourceActionByProject(
if busy { if busy {
return nil, errors.New(errPagesSourceActionBusy) 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( 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) 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) return nil, errors.New(errPagesSourceTaskDispatchFailed)
} }
execution, err := model.GetTaskExecutionByTaskID(ctx, taskID) execution, err := repository.GetTaskExecutionByTaskID(ctx, taskID)
if err != nil { if err != nil {
logger.ErrorF(ctx, "[PagesSource] load dispatched execution failed: source_id=%d task_id=%s error=%v", source.ID, taskID, err) 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) return nil, errors.New(errPagesSourceTaskDispatchFailed)
+4 -2
View File
@@ -10,6 +10,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
) )
@@ -67,7 +69,7 @@ func mustCreatePagesSourceProject(t *testing.T, ctx context.Context, slug string
if err != nil { if err != nil {
t.Fatalf("CreateProject(%q) error = %v, want nil", slug, err) 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 { if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", view.ID, err) 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 { if sourceCount != 0 || runtimeCount != 0 || deploymentCount != 1 {
t.Errorf("DeleteSource counts = source:%d runtime:%d deployment:%d, want 0, 0, 1", sourceCount, runtimeCount, deploymentCount) 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 { if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err) t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
} }
@@ -11,7 +11,6 @@ import (
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
) )
type proxyRouteJSONFields struct { type proxyRouteJSONFields struct {
@@ -141,38 +140,3 @@ func applyProxyRouteUpstreamType(ctx context.Context, route *model.ProxyRoute, u
} }
return nil 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
}
+10 -8
View File
@@ -17,6 +17,8 @@ import (
"strings" "strings"
"unicode" "unicode"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm" "gorm.io/gorm"
) )
@@ -195,7 +197,7 @@ func getOrCreateOriginByAddress(ctx context.Context, address string) (*model.Ori
if err := validateOriginAddress(normalizedAddress); err != nil { if err := validateOriginAddress(normalizedAddress); err != nil {
return nil, err return nil, err
} }
existing, err := model.GetOriginByAddress(ctx, normalizedAddress) existing, err := repository.GetOriginByAddress(ctx, normalizedAddress)
if err == nil { if err == nil {
return existing, nil return existing, nil
} }
@@ -207,9 +209,9 @@ func getOrCreateOriginByAddress(ctx context.Context, address string) (*model.Ori
Address: normalizedAddress, Address: normalizedAddress,
Remark: "", Remark: "",
} }
if err := model.CreateOriginRecord(ctx, origin); err != nil { if err := repository.CreateOriginRecord(ctx, origin); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return model.GetOriginByAddress(ctx, normalizedAddress) return repository.GetOriginByAddress(ctx, normalizedAddress)
} }
return nil, err 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) { 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) { 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) { 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) { func parseLeafCertificate(certPEM string) (*x509.Certificate, error) {
@@ -270,7 +272,7 @@ func loadProxyRouteZoneDomains(ctx context.Context, ids []uint) ([]model.ZoneDom
} }
seen[id] = struct{}{} seen[id] = struct{}{}
} }
domains, err := model.ListZoneDomainsByIDs(ctx, ids) domains, err := repository.ListZoneDomainsByIDs(ctx, ids)
if err != nil { if err != nil {
return nil, errors.New(errProxyRouteZoneDomainNotFound) 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 { func validateProxyRouteSiteNameUniqueness(ctx context.Context, route *model.ProxyRoute, siteName string) error {
routes, err := model.ListProxyRoutes(ctx) routes, err := repository.ListProxyRoutes(ctx)
if err != nil { if err != nil {
return err return err
} }
+35 -26
View File
@@ -10,10 +10,9 @@ import (
"strings" "strings"
"time" "time"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm" "gorm.io/gorm"
"gorm.io/gorm/clause"
) )
// CustomHeaderInput 自定义响应头。 // CustomHeaderInput 自定义响应头。
@@ -103,7 +102,7 @@ type ZoneDomainView struct {
// ListProxyRoutes 列出全部代理规则。 // ListProxyRoutes 列出全部代理规则。
func ListProxyRoutes(ctx context.Context) ([]*View, error) { func ListProxyRoutes(ctx context.Context) ([]*View, error) {
routes, err := model.ListProxyRoutes(ctx) routes, err := repository.ListProxyRoutes(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -112,7 +111,7 @@ func ListProxyRoutes(ctx context.Context) ([]*View, error) {
// GetProxyRoute 获取代理规则详情。 // GetProxyRoute 获取代理规则详情。
func GetProxyRoute(ctx context.Context, id uint) (*View, error) { func GetProxyRoute(ctx context.Context, id uint) (*View, error) {
route, err := model.GetProxyRouteByID(ctx, id) route, err := repository.GetProxyRouteByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -125,17 +124,17 @@ func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
if err != nil { if err != nil {
return nil, err 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 { if err := lockPagesProjectsForRouteMutation(tx, 0, route); err != nil {
return err return err
} }
if err := tx.Create(route).Error; err != nil { if err := repository.CreateProxyRouteRecordTx(tx, route); err != nil {
return err return err
} }
return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs) return repository.ReplaceZoneDomainRouteBindingsTx(tx, route.ID, input.ZoneDomainIDs)
}); err != nil { }); err != nil {
if isUniqueConstraintError(err) { if mapped := mapProxyRoutePersistError(err); mapped != nil {
return nil, errors.New(errProxyRouteIdentityExists) return nil, mapped
} }
return nil, err return nil, err
} }
@@ -144,7 +143,7 @@ func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
// UpdateProxyRoute 更新代理规则。 // UpdateProxyRoute 更新代理规则。
func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -153,23 +152,39 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error)
if err != nil { if err != nil {
return nil, err 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 { if err := lockPagesProjectsForRouteMutation(tx, previousPagesProjectID, route); err != nil {
return err return err
} }
if err := updateProxyRouteRecord(tx, route); err != nil { if err := repository.UpdateProxyRouteRecordTx(tx, route); err != nil {
return err return err
} }
return replaceZoneDomainRouteBindings(tx, route.ID, input.ZoneDomainIDs) return repository.ReplaceZoneDomainRouteBindingsTx(tx, route.ID, input.ZoneDomainIDs)
}); err != nil { }); err != nil {
if isUniqueConstraintError(err) { if mapped := mapProxyRoutePersistError(err); mapped != nil {
return nil, errors.New(errProxyRouteIdentityExists) return nil, mapped
} }
return nil, err return nil, err
} }
return buildProxyRouteView(ctx, route) 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 { func pagesProjectIDForRoute(route *model.ProxyRoute) uint {
if route == nil || route.UpstreamType != proxyRouteUpstreamTypePages || route.PagesProjectID == nil { if route == nil || route.UpstreamType != proxyRouteUpstreamTypePages || route.PagesProjectID == nil {
return 0 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] }) sort.Slice(projectIDs, func(i int, j int) bool { return projectIDs[i] < projectIDs[j] })
for _, projectID := range projectIDs { for _, projectID := range projectIDs {
var project model.PagesProject project, err := repository.LockPagesProjectByIDTx(tx, projectID)
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&project, projectID).Error
if errors.Is(err, gorm.ErrRecordNotFound) && projectID != nextProjectID { if errors.Is(err, gorm.ErrRecordNotFound) && projectID != nextProjectID {
continue continue
} }
@@ -201,7 +215,7 @@ func lockPagesProjectsForRouteMutation(tx *gorm.DB, previousProjectID uint, rout
return err return err
} }
if projectID == nextProjectID { if projectID == nextProjectID {
if err := validateLockedPagesRouteProject(&project); err != nil { if err := validateLockedPagesRouteProject(project); err != nil {
return err return err
} }
} }
@@ -224,15 +238,10 @@ func validateLockedPagesRouteProject(project *model.PagesProject) error {
// DeleteProxyRoute 删除代理规则。 // DeleteProxyRoute 删除代理规则。
func DeleteProxyRoute(ctx context.Context, id uint) error { 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 err
} }
return db.DB(ctx).Transaction(func(tx *gorm.DB) error { return repository.DeleteProxyRouteAndUnbind(ctx, id)
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
})
} }
func buildProxyRoute(ctx context.Context, route *model.ProxyRoute, input Input) (*model.ProxyRoute, []model.ZoneDomain, error) { 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 { if route == nil {
return nil, errors.New("proxy route is 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 { if err != nil {
return nil, err return nil, err
} }
@@ -9,6 +9,7 @@ import (
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/glebarez/sqlite" "github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "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) require.NoError(t, db.DB(ctx).Delete(&model.PagesProject{}, project.ID).Error)
route := &model.ProxyRoute{UpstreamType: proxyRouteUpstreamTypePages, PagesProjectID: &project.ID} 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) return lockPagesProjectsForRouteMutation(tx, 0, route)
}) })
require.EqualError(t, err, errProxyRoutePagesNotFound) require.EqualError(t, err, errProxyRoutePagesNotFound)
@@ -128,7 +129,7 @@ func TestRouteCanMoveAwayFromAlreadyMissingPagesProject(t *testing.T) {
missingProjectID := uint(404) missingProjectID := uint(404)
route := &model.ProxyRoute{UpstreamType: "direct"} 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) return lockPagesProjectsForRouteMutation(tx, missingProjectID, route)
}) })
require.NoError(t, err) require.NoError(t, err)
@@ -8,7 +8,8 @@ import (
"errors" "errors"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm" "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) { 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 { if err != nil {
return "", nil, errors.New(errProxyRouteOriginNotFound) return "", nil, errors.New(errProxyRouteOriginNotFound)
} }
@@ -67,7 +68,7 @@ func resolveLegacyOriginInput(ctx context.Context, originURL string) (string, *u
if err != nil { if err != nil {
return "", nil, err return "", nil, err
} }
origin, findErr := model.GetOriginByAddress(ctx, address) origin, findErr := repository.GetOriginByAddress(ctx, address)
if findErr == nil { if findErr == nil {
return originURL, &origin.ID, nil return originURL, &origin.ID, nil
} }
+2 -2
View File
@@ -11,8 +11,8 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/openflare/agent" "github.com/Rain-kl/Wavelet/internal/apps/openflare/agent"
ofgeoip "github.com/Rain-kl/Wavelet/internal/apps/openflare/geoip" 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/model"
"github.com/Rain-kl/Wavelet/internal/repository"
) )
const nodeStatusOnline = "online" const nodeStatusOnline = "online"
@@ -88,7 +88,7 @@ func Heartbeat(ctx context.Context, node *model.OpenFlareNode, payload Heartbeat
node.LastSeenAt = &lastSeen node.LastSeenAt = &lastSeen
node.Status = nodeStatusOnline 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) return nil, fmt.Errorf("update relay heartbeat: %w", err)
} }
if err := reconcileRelayHealthEvents(ctx, node.NodeID, payload.RelayStatus, now); err != nil { if err := reconcileRelayHealthEvents(ctx, node.NodeID, payload.RelayStatus, now); err != nil {
+8 -6
View File
@@ -9,6 +9,8 @@ import (
"testing" "testing"
"time" "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/agent"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
@@ -36,7 +38,7 @@ func setupRelayTestDB(t *testing.T) func() {
db.SetDB(sqliteDB) db.SetDB(sqliteDB)
agent.ResetAuthCacheForTest() agent.ResetAuthCacheForTest()
resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore())
return func() { return func() {
resetObservabilityStore() resetObservabilityStore()
@@ -107,17 +109,17 @@ func TestHeartbeatPayloadBindingAndFrpsObservationInsert(t *testing.T) {
assert.Equal(t, "v0.1.0", stored.Version) assert.Equal(t, "v0.1.0", stored.Version)
assert.Equal(t, "0.61.0", stored.ExtVersion) 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) require.NoError(t, err)
assert.Equal(t, "relay-runtime", profile.Hostname) assert.Equal(t, "relay-runtime", profile.Hostname)
assert.Equal(t, "Ubuntu", profile.OSName) 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.NoError(t, err)
require.Len(t, snapshots, 1) require.Len(t, snapshots, 1)
assert.Equal(t, 12.5, snapshots[0].CPUUsagePercent) 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.NoError(t, err)
require.Len(t, frpsObs, 1) require.Len(t, frpsObs, 1)
assert.Equal(t, 7, frpsObs[0].FrpsConnections) assert.Equal(t, 7, frpsObs[0].FrpsConnections)
@@ -153,7 +155,7 @@ func TestHeartbeatRelayReconcilesFrpsUnhealthyEvent(t *testing.T) {
}) })
require.NoError(t, err) 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.NoError(t, err)
require.Len(t, events, 1) require.Len(t, events, 1)
assert.Equal(t, relayFrpsUnhealthyEventType, events[0].EventType) assert.Equal(t, relayFrpsUnhealthyEventType, events[0].EventType)
@@ -166,7 +168,7 @@ func TestHeartbeatRelayReconcilesFrpsUnhealthyEvent(t *testing.T) {
}) })
require.NoError(t, err) 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.NoError(t, err)
require.Len(t, events, 1) require.Len(t, events, 1)
assert.Equal(t, "resolved", events[0].Status) assert.Equal(t, "resolved", events[0].Status)
+3 -1
View File
@@ -8,6 +8,8 @@ import (
"errors" "errors"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/internal/shared/response"
"github.com/gin-gonic/gin" "github.com/gin-gonic/gin"
@@ -38,7 +40,7 @@ func authenticateAccessToken(ctx context.Context, token string) (*model.OpenFlar
if token == "" { if token == "" {
return nil, errors.New("missing agent token") return nil, errors.New("missing agent token")
} }
node, err := model.GetOpenFlareNodeByAccessToken(ctx, token) node, err := repository.GetOpenFlareNodeByAccessToken(ctx, token)
if err != nil { if err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) { if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New("invalid agent token") return nil, errors.New("invalid agent token")
@@ -9,6 +9,8 @@ import (
"net/http/httptest" "net/http/httptest"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/repository"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/shared/response" "github.com/Rain-kl/Wavelet/internal/shared/response"
@@ -44,7 +46,7 @@ func seedRelayNode(t *testing.T, nodeType, accessToken string) *model.OpenFlareN
NodeType: nodeType, NodeType: nodeType,
AccessToken: accessToken, AccessToken: accessToken,
} }
require.NoError(t, model.CreateOpenFlareNode(ctx, node)) require.NoError(t, repository.CreateOpenFlareNode(ctx, node))
return node return node
} }
+4 -10
View File
@@ -7,11 +7,11 @@ import (
"context" "context"
"time" "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/agent"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"go.uber.org/zap" "go.uber.org/zap"
"gorm.io/gorm"
) )
const relayFrpsUnhealthyEventType = "frps_unhealthy" const relayFrpsUnhealthyEventType = "frps_unhealthy"
@@ -35,13 +35,7 @@ func reconcileRelayHealthEvents(ctx context.Context, nodeID string, relayStatus
}, },
}) })
} }
conn := db.DB(ctx) return agent.ReconcileScopedNodeHealthEvents(ctx, nodeID, events, reportedAt, managedTypes)
if conn == nil {
return nil
}
return conn.Transaction(func(tx *gorm.DB) error {
return agent.ReconcileScopedNodeHealthEvents(tx, nodeID, events, reportedAt, managedTypes)
})
} }
func persistRelayHeartbeatObservability(ctx context.Context, nodeID string, payload HeartbeatPayload, reportedAt time.Time) { 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, FrpsClientCount: payload.FrpsClientCount,
FrpsProxies: agent.MarshalJSON(payload.FrpsProxies), 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)) zap.L().Error("persist relay frps observation failed", zap.String("node_id", nodeID), zap.Error(err))
} }
} }
@@ -193,15 +193,15 @@ func deleteAllObservabilityRows(ctx context.Context, target string) (int64, stri
) )
switch target { switch target {
case DatabaseCleanupTargetAccessLogs: case DatabaseCleanupTargetAccessLogs:
deleted, err = model.DeleteAllOpenFlareAccessLogs(ctx) deleted, err = repository.DeleteAllOpenFlareAccessLogs(ctx)
case DatabaseCleanupTargetMetricSnapshots: case DatabaseCleanupTargetMetricSnapshots:
deleted, err = model.DeleteAllOpenFlareMetricSnapshots(ctx) deleted, err = repository.DeleteAllOpenFlareMetricSnapshots(ctx)
case DatabaseCleanupTargetEdgeHealth: case DatabaseCleanupTargetEdgeHealth:
deleted, err = model.DeleteAllOpenFlareEdgeHealth(ctx) deleted, err = repository.DeleteAllOpenFlareEdgeHealth(ctx)
case DatabaseCleanupTargetObsFrps: case DatabaseCleanupTargetObsFrps:
deleted, err = model.DeleteAllOpenFlareNodeObservationFrps(ctx) deleted, err = repository.DeleteAllOpenFlareNodeObservationFrps(ctx)
case DatabaseCleanupTargetObsFrpc: case DatabaseCleanupTargetObsFrpc:
deleted, err = model.DeleteAllOpenFlareNodeObservationFrpc(ctx) deleted, err = repository.DeleteAllOpenFlareNodeObservationFrpc(ctx)
default: default:
return 0, "", errors.New("unsupported cleanup target") return 0, "", errors.New("unsupported cleanup target")
} }
@@ -226,15 +226,15 @@ func materializeObservabilityTableTTL(ctx context.Context, target string) (int64
) )
switch target { switch target {
case DatabaseCleanupTargetAccessLogs: case DatabaseCleanupTargetAccessLogs:
eligible, err = model.DeleteOpenFlareAccessLogsBefore(ctx, cutoff) eligible, err = repository.DeleteOpenFlareAccessLogsBefore(ctx, cutoff)
case DatabaseCleanupTargetMetricSnapshots: case DatabaseCleanupTargetMetricSnapshots:
eligible, err = model.DeleteOpenFlareMetricSnapshotsBefore(ctx, cutoff) eligible, err = repository.DeleteOpenFlareMetricSnapshotsBefore(ctx, cutoff)
case DatabaseCleanupTargetEdgeHealth: case DatabaseCleanupTargetEdgeHealth:
eligible, err = model.DeleteOpenFlareEdgeHealthBefore(ctx, cutoff) eligible, err = repository.DeleteOpenFlareEdgeHealthBefore(ctx, cutoff)
case DatabaseCleanupTargetObsFrps: case DatabaseCleanupTargetObsFrps:
eligible, err = model.DeleteOpenFlareNodeObservationFrpsBefore(ctx, cutoff) eligible, err = repository.DeleteOpenFlareNodeObservationFrpsBefore(ctx, cutoff)
case DatabaseCleanupTargetObsFrpc: case DatabaseCleanupTargetObsFrpc:
eligible, err = model.DeleteOpenFlareNodeObservationFrpcBefore(ctx, cutoff) eligible, err = repository.DeleteOpenFlareNodeObservationFrpcBefore(ctx, cutoff)
default: default:
return 0, "", errors.New("unsupported cleanup target") return 0, "", errors.New("unsupported cleanup target")
} }
@@ -27,8 +27,8 @@ func setupDatabaseCleanupTestDB(t *testing.T) context.Context {
require.NoError(t, err) require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{})) require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{}))
db.SetDB(sqliteDB) db.SetDB(sqliteDB)
resetAccessLogStore := model.SetAccessLogStoreForTest(model.NewMemoryAccessLogStore()) resetAccessLogStore := repository.SetAccessLogStoreForTest(repository.NewMemoryAccessLogStore())
resetObservabilityStore := model.SetObservabilityStoreForTest(model.NewMemoryObservabilityStore()) resetObservabilityStore := repository.SetObservabilityStoreForTest(repository.NewMemoryObservabilityStore())
t.Cleanup(func() { t.Cleanup(func() {
resetObservabilityStore() resetObservabilityStore()
resetAccessLogStore() resetAccessLogStore()
@@ -69,12 +69,12 @@ func TestCleanupDatabaseObservabilityMaterializeDoesNotClaimHardDelete(t *testin
now := time.Now().UTC() now := time.Now().UTC()
// One row past metric table TTL (30d), one still inside the window. // 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", NodeID: "node-a",
CapturedAt: now.Add(-40 * 24 * time.Hour), CapturedAt: now.Add(-40 * 24 * time.Hour),
CPUUsagePercent: 10, CPUUsagePercent: 10,
})) }))
require.NoError(t, model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{
NodeID: "node-a", NodeID: "node-a",
CapturedAt: now.Add(-12 * time.Hour), CapturedAt: now.Add(-12 * time.Hour),
CPUUsagePercent: 20, CPUUsagePercent: 20,
@@ -96,7 +96,7 @@ func TestCleanupDatabaseObservabilityMaterializeDoesNotClaimHardDelete(t *testin
assert.True(t, result.Cutoff.Before(now.Add(-29*24*time.Hour))) 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. // 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.NoError(t, err)
require.Len(t, rows, 1) require.Len(t, rows, 1)
assert.Equal(t, float64(20), rows[0].CPUUsagePercent) assert.Equal(t, float64(20), rows[0].CPUUsagePercent)
@@ -106,7 +106,7 @@ func TestCleanupDatabaseObservabilityDeletesAllRowsWhenRetentionMissing(t *testi
ctx := setupDatabaseCleanupTestDB(t) ctx := setupDatabaseCleanupTestDB(t)
now := time.Now().UTC() now := time.Now().UTC()
require.NoError(t, model.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{ require.NoError(t, repository.InsertOpenFlareAccessLogsBatch(ctx, []*model.OpenFlareAccessLog{
{ {
NodeID: "node-a", NodeID: "node-a",
LoggedAt: now.Add(-3 * time.Hour), 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.DeletedCount)
assert.Equal(t, int64(2), result.EligibleCount) 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) require.NoError(t, err)
assert.Empty(t, rows) assert.Empty(t, rows)
} }
@@ -144,7 +144,7 @@ func TestRunDatabaseAutoCleanupOnceClampsRetentionToTableTTL(t *testing.T) {
now := time.Now().UTC() now := time.Now().UTC()
// Access logs TTL=90d, metrics TTL=30d. Config retention=1 must clamp, not reject. // 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", NodeID: "node-a",
LoggedAt: now.Add(-100 * 24 * time.Hour), LoggedAt: now.Add(-100 * 24 * time.Hour),
RemoteAddr: "203.0.113.10", RemoteAddr: "203.0.113.10",
@@ -152,12 +152,12 @@ func TestRunDatabaseAutoCleanupOnceClampsRetentionToTableTTL(t *testing.T) {
Path: "/access", Path: "/access",
StatusCode: 200, StatusCode: 200,
}})) }}))
require.NoError(t, model.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{ require.NoError(t, repository.InsertOpenFlareMetricSnapshot(ctx, &model.OpenFlareMetricSnapshot{
NodeID: "node-a", NodeID: "node-a",
CapturedAt: now.Add(-40 * 24 * time.Hour), CapturedAt: now.Add(-40 * 24 * time.Hour),
CPUUsagePercent: 10, CPUUsagePercent: 10,
})) }))
require.NoError(t, model.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{ require.NoError(t, repository.InsertOpenFlareEdgeHealth(ctx, &model.OpenFlareEdgeHealth{
NodeID: "node-a", NodeID: "node-a",
CapturedAt: now.Add(-40 * 24 * time.Hour), CapturedAt: now.Add(-40 * 24 * time.Hour),
Status: "healthy", Status: "healthy",
@@ -181,15 +181,15 @@ func TestRunDatabaseAutoCleanupOnceClampsRetentionToTableTTL(t *testing.T) {
assert.GreaterOrEqual(t, *result.RetentionDays, result.TableTTLDays) 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) require.NoError(t, err)
assert.Empty(t, accessLogs) assert.Empty(t, accessLogs)
metricSnapshots, err := model.ListOpenFlareMetricSnapshotsSince(ctx, "", time.Time{}, 0) metricSnapshots, err := repository.ListOpenFlareMetricSnapshotsSince(ctx, "", time.Time{}, 0)
require.NoError(t, err) require.NoError(t, err)
assert.Empty(t, metricSnapshots) assert.Empty(t, metricSnapshots)
edgeHealth, err := model.ListOpenFlareEdgeHealth(ctx, "", time.Time{}, 0) edgeHealth, err := repository.ListOpenFlareEdgeHealth(ctx, "", time.Time{}, 0)
require.NoError(t, err) require.NoError(t, err)
assert.Empty(t, edgeHealth) assert.Empty(t, edgeHealth)
} }
+3 -2
View File
@@ -7,8 +7,9 @@ import (
"context" "context"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls" "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger" "github.com/Rain-kl/Wavelet/pkg/logger"
) )
@@ -16,7 +17,7 @@ import (
func RunSSLRenewJob(ctx context.Context) error { func RunSSLRenewJob(ctx context.Context) error {
logger.InfoF(ctx, "[OpenFlareTasks] SSL renew job started") logger.InfoF(ctx, "[OpenFlareTasks] SSL renew job started")
certificates, err := model.ListTLSCertificates(ctx) certificates, err := repository.ListTLSCertificates(ctx)
if err != nil { if err != nil {
logger.ErrorF(ctx, "[OpenFlareTasks] list certificates failed: %v", err) logger.ErrorF(ctx, "[OpenFlareTasks] list certificates failed: %v", err)
return err return err
@@ -8,6 +8,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls" "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
"github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/config"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
@@ -70,16 +72,16 @@ func TestRunSSLRenewJobTriggersDueCertificates(t *testing.T) {
KeyPEM: " ", KeyPEM: " ",
NotAfter: now.Add(30 * 24 * time.Hour), NotAfter: now.Add(30 * 24 * time.Hour),
} }
require.NoError(t, model.CreateTLSCertificateRecord(ctx, due)) require.NoError(t, repository.CreateTLSCertificateRecord(ctx, due))
require.NoError(t, model.CreateTLSCertificateRecord(ctx, fresh)) require.NoError(t, repository.CreateTLSCertificateRecord(ctx, fresh))
require.NoError(t, RunSSLRenewJob(ctx)) require.NoError(t, RunSSLRenewJob(ctx))
renewed, err := model.GetTLSCertificateByID(ctx, due.ID) renewed, err := repository.GetTLSCertificateByID(ctx, due.ID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "applying", renewed.ApplyStatus) assert.Equal(t, "applying", renewed.ApplyStatus)
unchanged, err := model.GetTLSCertificateByID(ctx, fresh.ID) unchanged, err := repository.GetTLSCertificateByID(ctx, fresh.ID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "ready", unchanged.ApplyStatus) assert.Equal(t, "ready", unchanged.ApplyStatus)
} }
@@ -9,6 +9,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -81,7 +83,7 @@ func TestRenewCertificateSetsApplying(t *testing.T) {
CertPEM: " ", CertPEM: " ",
KeyPEM: " ", 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 { restore := SetObtainCertificateFuncForTest(func(ctx context.Context, c *model.TLSCertificate) error {
return nil return nil
@@ -106,19 +108,19 @@ func TestConvertCertificateToACMEPreservesUploadOnFailure(t *testing.T) {
}) })
require.NoError(t, err) require.NoError(t, err)
stored, err := model.GetTLSCertificateByID(ctx, cert.ID) stored, err := repository.GetTLSCertificateByID(ctx, cert.ID)
require.NoError(t, err) require.NoError(t, err)
originalStoredCertPEM := stored.CertPEM originalStoredCertPEM := stored.CertPEM
originalStoredKeyPEM := stored.KeyPEM originalStoredKeyPEM := stored.KeyPEM
stored.ApplyStatus = "applying" stored.ApplyStatus = "applying"
stored.PrimaryDomain = "manual.example.com" 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") err = updateCertError(ctx, stored, "dns challenge failed")
require.Error(t, err) require.Error(t, err)
finalCert, err := model.GetTLSCertificateByID(ctx, cert.ID) finalCert, err := repository.GetTLSCertificateByID(ctx, cert.ID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, "upload", finalCert.Provider) assert.Equal(t, "upload", finalCert.Provider)
assert.Equal(t, "error", finalCert.ApplyStatus) assert.Equal(t, "error", finalCert.ApplyStatus)
@@ -141,14 +143,14 @@ func TestConvertCertificateToACMERejectsInvalidStates(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
cert.Provider = "acme" 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"}) _, err = ConvertCertificateToACME(ctx, cert.ID, ApplyInput{Name: "manual-cert"})
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), "only uploaded") assert.Contains(t, err.Error(), "only uploaded")
cert.Provider = "upload" cert.Provider = "upload"
cert.ApplyStatus = "applying" 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"}) _, err = ConvertCertificateToACME(ctx, cert.ID, ApplyInput{Name: "manual-cert"})
require.Error(t, err) require.Error(t, err)
assert.Contains(t, err.Error(), "already applying") assert.Contains(t, err.Error(), "already applying")
+28 -26
View File
@@ -12,6 +12,8 @@ import (
"mime/multipart" "mime/multipart"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
) )
@@ -71,17 +73,17 @@ type DNSAccountInput struct {
// ListCertificates 列出全部证书(不含 PEM)。 // ListCertificates 列出全部证书(不含 PEM)。
func ListCertificates(ctx context.Context) ([]model.TLSCertificate, error) { func ListCertificates(ctx context.Context) ([]model.TLSCertificate, error) {
return model.ListTLSCertificates(ctx) return repository.ListTLSCertificates(ctx)
} }
// GetCertificate 获取证书详情(不含 PEM)。 // GetCertificate 获取证书详情(不含 PEM)。
func GetCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) { func GetCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) {
return model.GetTLSCertificateByID(ctx, id) return repository.GetTLSCertificateByID(ctx, id)
} }
// GetCertificateContent 获取证书 PEM 内容。 // GetCertificateContent 获取证书 PEM 内容。
func GetCertificateContent(ctx context.Context, id uint) (*CertificateContent, error) { func GetCertificateContent(ctx context.Context, id uint) (*CertificateContent, error) {
certificate, err := model.GetTLSCertificateByID(ctx, id) certificate, err := repository.GetTLSCertificateByID(ctx, id)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -117,7 +119,7 @@ func CreateCertificate(ctx context.Context, input CertificateInput) (*model.TLSC
if err != nil { if err != nil {
return nil, err return nil, err
} }
if err = model.CreateTLSCertificateRecord(ctx, certificate); err != nil { if err = repository.CreateTLSCertificateRecord(ctx, certificate); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errCertificateNameExists) return nil, errors.New(errCertificateNameExists)
} }
@@ -149,7 +151,7 @@ func CreateCertificateFromFiles(ctx context.Context, name string, certFile *mult
// UpdateCertificate 更新上传证书。 // UpdateCertificate 更新上传证书。
func UpdateCertificate(ctx context.Context, id uint, input CertificateInput) (*model.TLSCertificate, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -157,7 +159,7 @@ func UpdateCertificate(ctx context.Context, id uint, input CertificateInput) (*m
if err != nil { if err != nil {
return nil, err return nil, err
} }
if err = model.SaveTLSCertificate(ctx, certificate); err != nil { if err = repository.SaveTLSCertificate(ctx, certificate); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errCertificateNameExists) return nil, errors.New(errCertificateNameExists)
} }
@@ -171,10 +173,10 @@ func DeleteCertificate(ctx context.Context, id uint) error {
if err := ensureCertificateNotReferenced(ctx, id); err != nil { if err := ensureCertificateNotReferenced(ctx, id); err != nil {
return err return err
} }
if _, err := model.GetTLSCertificateByID(ctx, id); err != nil { if _, err := repository.GetTLSCertificateByID(ctx, id); err != nil {
return err return err
} }
return model.DeleteTLSCertificateRecord(ctx, id) return repository.DeleteTLSCertificateRecord(ctx, id)
} }
// ApplyCertificate 申请 ACME 证书。 // ApplyCertificate 申请 ACME 证书。
@@ -188,7 +190,7 @@ func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertific
if cert.Name == "" { if cert.Name == "" {
return nil, errors.New(errCertificateNameRequired) return nil, errors.New(errCertificateNameRequired)
} }
if err := model.CreateTLSCertificateRecord(ctx, cert); err != nil { if err := repository.CreateTLSCertificateRecord(ctx, cert); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errCertificateNameExists) return nil, errors.New(errCertificateNameExists)
} }
@@ -205,7 +207,7 @@ func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertific
// UpdateACMECertificate 更新 ACME 证书配置。 // UpdateACMECertificate 更新 ACME 证书配置。
func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -216,7 +218,7 @@ func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*mod
if cert.Name == "" { if cert.Name == "" {
return nil, errors.New(errCertificateNameRequired) return nil, errors.New(errCertificateNameRequired)
} }
if err := model.SaveTLSCertificate(ctx, cert); err != nil { if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errCertificateNameExists) return nil, errors.New(errCertificateNameExists)
} }
@@ -233,7 +235,7 @@ func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*mod
// ConvertCertificateToACME 将上传证书转为 ACME 管理。 // ConvertCertificateToACME 将上传证书转为 ACME 管理。
func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*model.TLSCertificate, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -248,7 +250,7 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*
return nil, errors.New(errCertificateNameRequired) return nil, errors.New(errCertificateNameRequired)
} }
cert.ApplyMessage = "" cert.ApplyMessage = ""
if err := model.SaveTLSCertificate(ctx, cert); err != nil { if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
if isUniqueConstraintError(err) { if isUniqueConstraintError(err) {
return nil, errors.New(errCertificateNameExists) 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 { if err := obtainTLSCertificate(asyncCtx, c); err != nil {
return return
} }
latest, err := model.GetTLSCertificateByID(asyncCtx, c.ID) latest, err := repository.GetTLSCertificateByID(asyncCtx, c.ID)
if err != nil { if err != nil {
return return
} }
latest.Provider = tlsProviderACME latest.Provider = tlsProviderACME
latest.ApplyStatus = tlsApplyStatusReady latest.ApplyStatus = tlsApplyStatusReady
latest.ApplyMessage = "" latest.ApplyMessage = ""
_ = model.SaveTLSCertificate(asyncCtx, latest) _ = repository.SaveTLSCertificate(asyncCtx, latest)
}(cert) }(cert)
return sanitizeCertificateForResponse(cert), nil return sanitizeCertificateForResponse(cert), nil
@@ -275,7 +277,7 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (*
// RenewCertificate 续期 ACME 证书。 // RenewCertificate 续期 ACME 证书。
func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -295,7 +297,7 @@ func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, erro
cert.ApplyStatus = tlsApplyStatusApplying cert.ApplyStatus = tlsApplyStatusApplying
cert.ApplyMessage = "" cert.ApplyMessage = ""
if err := model.SaveTLSCertificate(ctx, cert); err != nil { if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
return nil, err return nil, err
} }
return sanitizeCertificateForResponse(cert), nil return sanitizeCertificateForResponse(cert), nil
@@ -303,7 +305,7 @@ func RenewCertificate(ctx context.Context, id uint) (*model.TLSCertificate, erro
// ListDNSAccounts 列出 DNS 账号。 // ListDNSAccounts 列出 DNS 账号。
func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) { func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) {
return model.ListDNSAccounts(ctx) return repository.ListDNSAccounts(ctx)
} }
// CreateDNSAccount 创建 DNS 账号。 // CreateDNSAccount 创建 DNS 账号。
@@ -320,7 +322,7 @@ func CreateDNSAccount(ctx context.Context, input DNSAccountInput) (*model.DNSAcc
if account.Name == "" || account.Type == "" || authorization == "" { if account.Name == "" || account.Type == "" || authorization == "" {
return nil, errors.New("DNS 账号参数不完整") 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 nil, err
} }
return sanitizeDNSAccountForResponse(account), nil return sanitizeDNSAccountForResponse(account), nil
@@ -328,7 +330,7 @@ func CreateDNSAccount(ctx context.Context, input DNSAccountInput) (*model.DNSAcc
// UpdateDNSAccount 更新 DNS 账号。 // UpdateDNSAccount 更新 DNS 账号。
func UpdateDNSAccount(ctx context.Context, id uint, input DNSAccountInput) (*model.DNSAccount, error) { 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 { if err != nil {
return nil, err return nil, err
} }
@@ -342,7 +344,7 @@ func UpdateDNSAccount(ctx context.Context, id uint, input DNSAccountInput) (*mod
if account.Name == "" || account.Type == "" || authorization == "" { if account.Name == "" || account.Type == "" || authorization == "" {
return nil, errors.New("DNS 账号参数不完整") 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 nil, err
} }
return sanitizeDNSAccountForResponse(account), nil return sanitizeDNSAccountForResponse(account), nil
@@ -350,22 +352,22 @@ func UpdateDNSAccount(ctx context.Context, id uint, input DNSAccountInput) (*mod
// DeleteDNSAccount 删除 DNS 账号。 // DeleteDNSAccount 删除 DNS 账号。
func DeleteDNSAccount(ctx context.Context, id uint) error { 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 return err
} }
count, err := model.CountTLSCertificatesByDNSAccountID(ctx, id) count, err := repository.CountTLSCertificatesByDNSAccountID(ctx, id)
if err != nil { if err != nil {
return err return err
} }
if count > 0 { if count > 0 {
return errors.New(errDNSAccountInUse) return errors.New(errDNSAccountInUse)
} }
return model.DeleteDNSAccountRecord(ctx, id) return repository.DeleteDNSAccountRecord(ctx, id)
} }
// GetDefaultAcmeAccount 获取默认 ACME 账号。 // GetDefaultAcmeAccount 获取默认 ACME 账号。
func GetDefaultAcmeAccount(ctx context.Context) (*model.AcmeAccount, error) { func GetDefaultAcmeAccount(ctx context.Context) (*model.AcmeAccount, error) {
account, err := model.GetDefaultAcmeAccount(ctx) account, err := repository.GetDefaultAcmeAccount(ctx)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -430,7 +432,7 @@ func fillAcmeCertificateFields(cert *model.TLSCertificate, input ApplyInput) {
} }
func ensureCertificateNotReferenced(ctx context.Context, id uint) error { func ensureCertificateNotReferenced(ctx context.Context, id uint) error {
count, err := model.CountZoneDomainsByCertificateID(ctx, id) count, err := repository.CountZoneDomainsByCertificateID(ctx, id)
if err != nil { if err != nil {
return err return err
} }
+3 -1
View File
@@ -16,6 +16,8 @@ import (
"testing" "testing"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/infra/config" "github.com/Rain-kl/Wavelet/internal/infra/config"
db "github.com/Rain-kl/Wavelet/internal/infra/persistence" db "github.com/Rain-kl/Wavelet/internal/infra/persistence"
"github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/infra/task"
@@ -123,7 +125,7 @@ func TestCreateCertificateEncryptsPrivateKey(t *testing.T) {
}) })
require.NoError(t, err) require.NoError(t, err)
stored, err := model.GetTLSCertificateByID(ctx, certificate.ID) stored, err := repository.GetTLSCertificateByID(ctx, certificate.ID)
require.NoError(t, err) require.NoError(t, err)
assert.NotEqual(t, keyPEM, stored.KeyPEM) assert.NotEqual(t, keyPEM, stored.KeyPEM)
assert.Contains(t, stored.KeyPEM, sensitiveValuePrefix) assert.Contains(t, stored.KeyPEM, sensitiveValuePrefix)
+5 -3
View File
@@ -9,6 +9,8 @@ import (
"strings" "strings"
"time" "time"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls/acme" "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls/acme"
"github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model" "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 { func obtainCertificate(ctx context.Context, cert *model.TLSCertificate) error {
task.AppendLog(ctx, "【续签任务】开始续签,设置申请状态为 applying...") task.AppendLog(ctx, "【续签任务】开始续签,设置申请状态为 applying...")
cert.ApplyStatus = tlsApplyStatusApplying cert.ApplyStatus = tlsApplyStatusApplying
if err := model.SaveTLSCertificate(ctx, cert); err != nil { if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
return err return err
} }
@@ -46,7 +48,7 @@ func obtainCertificate(ctx context.Context, cert *model.TLSCertificate) error {
} }
task.AppendLog(ctx, "【续签任务】正在解析 DNS 账户信息 (ID=%d)...", cert.DNSAccountID) 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 { if err != nil {
return updateCertError(ctx, cert, fmt.Sprintf("Failed to get DNS account: %v", err)) 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 { func updateCertError(ctx context.Context, cert *model.TLSCertificate, message string) error {
cert.ApplyStatus = "error" cert.ApplyStatus = "error"
cert.ApplyMessage = message cert.ApplyMessage = message
if err := model.SaveTLSCertificate(ctx, cert); err != nil { if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
return err return err
} }
return fmt.Errorf("%s", message) return fmt.Errorf("%s", message)
@@ -7,21 +7,23 @@ import (
"context" "context"
"fmt" "fmt"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls/acme" "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls/acme"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
) )
func resolveAcmeAccount(ctx context.Context, cert *model.TLSCertificate) (*model.AcmeAccount, error) { 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 { if err == nil {
return acmeAccount, nil return acmeAccount, nil
} }
acmeAccount, err = model.GetDefaultAcmeAccount(ctx) acmeAccount, err = repository.GetDefaultAcmeAccount(ctx)
if err != nil { if err != nil {
return nil, fmt.Errorf("failed to get ACME account: %w", err) return nil, fmt.Errorf("failed to get ACME account: %w", err)
} }
cert.AcmeAccountID = acmeAccount.ID cert.AcmeAccountID = acmeAccount.ID
if err := model.SaveTLSCertificate(ctx, cert); err != nil { if err := repository.SaveTLSCertificate(ctx, cert); err != nil {
return nil, err return nil, err
} }
return acmeAccount, nil return acmeAccount, nil
@@ -49,14 +51,14 @@ func persistAcmeAccountUpdates(
acmeAccount.URL = newAccountURL acmeAccount.URL = newAccountURL
} }
if acmeAccount.ID == 0 { 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) 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) return fmt.Errorf("failed to save ACME account: %w", dbErr)
} }
cert.AcmeAccountID = acmeAccount.ID 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 { 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.NotAfter = result.NotAfter
cert.ApplyStatus = tlsApplyStatusReady cert.ApplyStatus = tlsApplyStatusReady
cert.ApplyMessage = "" cert.ApplyMessage = ""
return model.SaveTLSCertificate(ctx, cert) return repository.SaveTLSCertificate(ctx, cert)
} }
+3 -2
View File
@@ -9,8 +9,9 @@ import (
"errors" "errors"
"fmt" "fmt"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/infra/task" "github.com/Rain-kl/Wavelet/internal/infra/task"
"github.com/Rain-kl/Wavelet/internal/model"
) )
const ( const (
@@ -77,7 +78,7 @@ func (h *SSLSingleRenewHandler) Execute(ctx context.Context, payload []byte) (*t
task.AppendLog(ctx, "开始续期证书,ID: %d", req.ID) task.AppendLog(ctx, "开始续期证书,ID: %d", req.ID)
cert, err := model.GetTLSCertificateByID(ctx, req.ID) cert, err := repository.GetTLSCertificateByID(ctx, req.ID)
if err != nil { if err != nil {
task.AppendLog(ctx, "获取证书记录失败 ID=%d: %v", req.ID, err) task.AppendLog(ctx, "获取证书记录失败 ID=%d: %v", req.ID, err)
return nil, fmt.Errorf("获取证书记录失败: %w", err) return nil, fmt.Errorf("获取证书记录失败: %w", err)
+7 -5
View File
@@ -9,6 +9,8 @@ import (
"errors" "errors"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -91,7 +93,7 @@ func TestSSLSingleRenewHandler_Execute(t *testing.T) {
Provider: "custom", Provider: "custom",
PrimaryDomain: "example.com", PrimaryDomain: "example.com",
} }
err := model.CreateTLSCertificateRecord(ctx, cert) err := repository.CreateTLSCertificateRecord(ctx, cert)
require.NoError(t, err) require.NoError(t, err)
payload, err := json.Marshal(SSLSingleRenewPayload{ID: cert.ID}) payload, err := json.Marshal(SSLSingleRenewPayload{ID: cert.ID})
@@ -108,13 +110,13 @@ func TestSSLSingleRenewHandler_Execute(t *testing.T) {
Provider: tlsProviderACME, Provider: tlsProviderACME,
PrimaryDomain: "success.example.com", PrimaryDomain: "success.example.com",
} }
err := model.CreateTLSCertificateRecord(ctx, cert) err := repository.CreateTLSCertificateRecord(ctx, cert)
require.NoError(t, err) require.NoError(t, err)
// Mock obtainCertificate to succeed // Mock obtainCertificate to succeed
restore := SetObtainCertificateFuncForTest(func(ctx context.Context, c *model.TLSCertificate) error { restore := SetObtainCertificateFuncForTest(func(ctx context.Context, c *model.TLSCertificate) error {
c.ApplyStatus = tlsApplyStatusReady c.ApplyStatus = tlsApplyStatusReady
return model.SaveTLSCertificate(ctx, c) return repository.SaveTLSCertificate(ctx, c)
}) })
defer restore() defer restore()
@@ -125,7 +127,7 @@ func TestSSLSingleRenewHandler_Execute(t *testing.T) {
require.NoError(t, err) require.NoError(t, err)
assert.Contains(t, res.Message, "续签成功") assert.Contains(t, res.Message, "续签成功")
updated, err := model.GetTLSCertificateByID(ctx, cert.ID) updated, err := repository.GetTLSCertificateByID(ctx, cert.ID)
require.NoError(t, err) require.NoError(t, err)
assert.Equal(t, tlsApplyStatusReady, updated.ApplyStatus) assert.Equal(t, tlsApplyStatusReady, updated.ApplyStatus)
}) })
@@ -136,7 +138,7 @@ func TestSSLSingleRenewHandler_Execute(t *testing.T) {
Provider: tlsProviderACME, Provider: tlsProviderACME,
PrimaryDomain: "fail.example.com", PrimaryDomain: "fail.example.com",
} }
err := model.CreateTLSCertificateRecord(ctx, cert) err := repository.CreateTLSCertificateRecord(ctx, cert)
require.NoError(t, err) require.NoError(t, err)
// Mock obtainCertificate to fail // Mock obtainCertificate to fail
+2 -2
View File
@@ -100,7 +100,7 @@ func SyncToUptimeKuma(ctx context.Context) error {
"scope", config.MonitorScope, "scope", config.MonitorScope,
) )
allRoutes, err := model.ListProxyRoutes(ctx) allRoutes, err := repository.ListProxyRoutes(ctx)
if err != nil { if err != nil {
return fmt.Errorf("failed to list local proxy routes: %w", err) 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 { if route == nil {
return "", fmt.Errorf("proxy route is 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 { if err != nil {
return "", err return "", err
} }
@@ -217,9 +217,9 @@ func TestSyncToUptimeKumaSuccess(t *testing.T) {
EnableHTTPS: false, EnableHTTPS: false,
} }
require.NoError(t, model.CreateProxyRouteRecord(ctx, routeA)) require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeA))
require.NoError(t, model.CreateProxyRouteRecord(ctx, routeB)) require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeB))
require.NoError(t, model.CreateProxyRouteRecord(ctx, routeC)) require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeC))
createRouteZoneDomain(t, ctx, routeA, "site-a.com") createRouteZoneDomain(t, ctx, routeA, "site-a.com")
createRouteZoneDomain(t, ctx, routeB, "site-b.com") createRouteZoneDomain(t, ctx, routeB, "site-b.com")
createRouteZoneDomain(t, ctx, routeC, "site-c.com") createRouteZoneDomain(t, ctx, routeC, "site-c.com")
@@ -320,8 +320,8 @@ func TestSyncToUptimeKumaSelectedScope(t *testing.T) {
EnableHTTPS: false, EnableHTTPS: false,
} }
require.NoError(t, model.CreateProxyRouteRecord(ctx, routeA)) require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeA))
require.NoError(t, model.CreateProxyRouteRecord(ctx, routeB)) require.NoError(t, repository.CreateProxyRouteRecord(ctx, routeB))
createRouteZoneDomain(t, ctx, routeA, "site-a.com") createRouteZoneDomain(t, ctx, routeA, "site-a.com")
createRouteZoneDomain(t, ctx, routeB, "site-b.com") createRouteZoneDomain(t, ctx, routeB, "site-b.com")
+7 -5
View File
@@ -17,6 +17,8 @@ import (
"strings" "strings"
"time" "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/agent"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket" "github.com/Rain-kl/Wavelet/internal/apps/openflare/websocket"
"github.com/Rain-kl/Wavelet/internal/model" "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. // SyncDueWAFIPGroups syncs all enabled automatic/subscription IP groups that are due.
func SyncDueWAFIPGroups(ctx context.Context) error { func SyncDueWAFIPGroups(ctx context.Context) error {
now := time.Now().UTC() now := time.Now().UTC()
groups, err := model.ListDueOpenFlareWAFIPGroups(ctx, now) groups, err := repository.ListDueOpenFlareWAFIPGroups(ctx, now)
if err != nil { if err != nil {
return err return err
} }
@@ -176,7 +178,7 @@ func syncIPGroupSubscription(ctx context.Context, group *model.OpenFlareWAFIPGro
group.NextSyncAt = &nextSyncAt group.NextSyncAt = &nextSyncAt
group.LastSyncStatus = "success" group.LastSyncStatus = "success"
group.LastSyncMessage = fmt.Sprintf("同步成功,共 %d 条 IP/IP 段", len(ips)) 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 return nil, err
} }
broadcastIPGroupToAgents(ctx, group.ID) broadcastIPGroupToAgents(ctx, group.ID)
@@ -259,7 +261,7 @@ func syncIPGroupAutomatic(ctx context.Context, group *model.OpenFlareWAFIPGroup,
group.NextSyncAt = &nextSyncAt group.NextSyncAt = &nextSyncAt
group.LastSyncStatus = "success" group.LastSyncStatus = "success"
group.LastSyncMessage = fmt.Sprintf("自动规则执行成功,共命中 %d 个 IP,当前生效 %d 个 IP", len(ips), len(finalIPs)) 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 return nil, err
} }
broadcastIPGroupToAgents(ctx, group.ID) broadcastIPGroupToAgents(ctx, group.ID)
@@ -283,7 +285,7 @@ func recordIPGroupSyncFailure(ctx context.Context, group *model.OpenFlareWAFIPGr
group.NextSyncAt = &nextSyncAt group.NextSyncAt = &nextSyncAt
group.LastSyncStatus = "failed" group.LastSyncStatus = "failed"
group.LastSyncMessage = syncErr.Error() group.LastSyncMessage = syncErr.Error()
_ = model.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group) _ = repository.UpdateOpenFlareWAFIPGroupSyncResult(ctx, group)
} }
func evaluateParsedIPGroupAutoConfig(ctx context.Context, config ipGroupAutoConfig, now time.Time) ([]string, error) { 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 { if lookback <= 0 {
lookback = defaultWAFIPGroupAutoLookbackDur lookback = defaultWAFIPGroupAutoLookbackDur
} }
aggregates, err := model.ListOpenFlareAccessLogWAFIPAggregates(ctx, model.OpenFlareAccessLogQuery{ aggregates, err := repository.ListOpenFlareAccessLogWAFIPAggregates(ctx, model.OpenFlareAccessLogQuery{
Since: now.Add(-lookback), Since: now.Add(-lookback),
Until: now, Until: now,
}) })

Some files were not shown because too many files have changed in this diff Show More