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

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