diff --git a/.agents/skills/new-api/SKILL.md b/.agents/skills/new-api/SKILL.md index a3451a53..9d9c8893 100644 --- a/.agents/skills/new-api/SKILL.md +++ b/.agents/skills/new-api/SKILL.md @@ -13,102 +13,38 @@ description: "Wavelet 项目专用:当新增或修改自定义业务 API、新 Wavelet 后端路由采用了**严格的框架层与业务层隔离机制**。请牢记以下开发原则: -1. **禁止修改框架级路由文件**: - - 以下文件属于系统框架/平台级接口,**禁止为了添加自定义业务接口而进行任何修改**: - - `internal/router/router.go`(核心入口委派) - - `internal/router/root/default.go`(公开文件服务、robots.txt、Swagger 及 /api/health 路由) - - `internal/router/root/frontend.go`(前端静态服务) - - `internal/router/v1/v1.go`(V1 分发层协调器) - - `internal/router/v1/admin.go`(框架管理员端管理接口) - - `internal/router/v1/user.go`(框架普通用户端基础接口、OAuth及公开接口) -2. **仅允许在 `custom.go` 中注册业务接口**: - - 所有的自定义/业务相关接口注册,有且仅有以下两个合法的承载点: - - [internal/router/root/custom.go](file:///Users/ryan/DEV/Go/Wavelet/internal/router/root/custom.go)(用于挂载到根路径的特殊业务接口) - - [internal/router/v1/custom.go](file:///Users/ryan/DEV/Go/Wavelet/internal/router/v1/custom.go)(用于挂载在 API V1 下的标准自定义业务接口) +### 插件目录标准结构 (`backend/openflare/plugins//` 或 `backend/plugins/domain//`) ---- - -## 路由归属判定表 (Where should I register my new API?) - -根据接口的**访问路径特征**和**访问身份/限制条件**,决定将新开发的 API 挂载至何处: - -| 目标 API 路径特征 | 访问身份/条件限制 | 对应的路由注册入口 | 是否允许修改 | -| :--- | :--- | :--- | :--- | -| **`/my-custom-path`** (挂载在根路径下的特殊业务接口) | 自定义控制 | `root/custom.go` 中的 `RegisterCustomRootRoutes` | **允许修改 (业务自定义入口)** | -| **`/api/v1/custom/...`** (API v1 下的定制业务接口) | 自定义控制 | `v1/custom.go` 中的 `RegisterCustomRoutes` | **允许修改 (业务自定义入口)** | -| **`/api/v1/admin/...`** (系统管理员管理端接口) | 需要管理员登录 (`admin.LoginAdminRequired()`) | `v1/admin.go` | **禁止修改 (仅限系统框架路由)** | -| **`/api/v1/user/...`** (框架普通用户基础接口) | 需要普通用户登录 (`oauth.LoginRequired()`) | `v1/user.go` | **禁止修改 (仅限系统框架路由)** | -| **`/api/v1/public/...`** (Captcha、Config 等系统公开接口) | 所有人 (无条件 / 公开) | `v1/user.go` | **禁止修改 (仅限系统框架路由)** | -| **`GET /f/:id`**, **`GET /robots.txt`**, **`GET /api/health`** (系统级默认及公开接口) | 所有人 (无条件 / 公开) | `root/default.go` | **禁止修改 (仅限系统框架路由)** | - ---- - -## 两个自定义路由包的用法与区别 (Root Custom vs V1 Custom) - -### 1. 根路径自定义包:`root/custom.go` - -* **适用场景**:适用于需要**直接挂载在主域名根路径下**的特殊自定义业务接口(如第三方 Webhook 回调、特定的短链接重定向、外部数据接口等,不需要 `/api/v1` 前缀)。 -* **用法示例**: - 在 [root/custom.go](file:///Users/ryan/DEV/Go/Wavelet/internal/router/root/custom.go) 中实现: - ```go - package root - - import ( - "OpenFlare/internal/apps/custom" - "github.com/gin-gonic/gin" - ) - - // RegisterCustomRootRoutes registers custom business routes that belong to the root path. - func RegisterCustomRootRoutes(r *gin.Engine) { - // 挂载到根路径下,如 GET /my-custom-webhook - r.GET("/my-custom-webhook", custom.HandleRootWebhook) - } - ``` - *(注:该函数已由 `root.go` 自动加载,你无需修改任何其他核心文件。)* - -### 2. V1 API 自定义包:`v1/custom.go` - -* **适用场景**:适用于普通的**自定义业务 API**,需要规范挂载在标准 API V1 路径下(即自动带有 `/api/v1/custom/...` 前缀,可选择性配置用户/管理员登录中间件)。 -* **用法示例**: - 在 [v1/custom.go](file:///Users/ryan/DEV/Go/Wavelet/internal/router/v1/custom.go) 中实现: - ```go - package v1 - - import ( - "OpenFlare/internal/apps/custom" - "github.com/gin-gonic/gin" - ) - - // RegisterCustomRoutes registers standard custom API routes under /api/v1. - func RegisterCustomRoutes(apiV1Router *gin.RouterGroup) { - customRouter := apiV1Router.Group("/custom") - { - // 挂载到 /api/v1/custom 下,例如:POST /api/v1/custom/action - customRouter.POST("/action", custom.DoActionHandler) - } - } - ``` - *(注:该函数已由 `v1/v1.go` 自动加载,你无需修改任何其他核心文件。)* - ---- - -## 建议创建/修改的文件结构 (Recommended Directory Structure) - -当新增一套定制的业务接口(例如名为 `custom` 的业务模块)时,建议采用以下标准文件结构: +所有标准插件与下游定制插件,**统一以 `backend/downstream/plugins/custom_example` 为基准模板**,严格采用物理子包隔离的分层架构: ```text -internal/ -├── router/ -│ ├── root/ -│ │ └── custom.go # [修改] 若为根路径 API,在此处注册,将路由委派给 apps/custom -│ └── v1/ -│ └── custom.go # [修改] 若为 v1 API,在此处注册,将路由委派给 apps/custom -└── apps/ - └── custom/ - ├── routers.go # [新建] HTTP Handlers (Gin),负责参数绑定、校验与响应 - ├── logics.go # [新建] 业务逻辑层:承载模块内闭环的纯 Go 业务逻辑,不依赖 gin.Context - └── errs.go # [新建] 存放模块特有的业务错误常量定义(可选) +backend/openflare/plugins// (或 backend/plugins/domain//) +├── plugin.go # 插件根入口:实现 core.Plugin,装配各子包并向 Cordis 注册 +│ +├── consts/ # package consts:常量、配置键名与错误码定义 +│ └── consts.go +│ +├── controller/ # package controller:HTTP 控制器与路由声明 (参数绑定、会话获取、信封响应) +│ └── hello/ # 业务分组/实体子包 +│ └── hello.go # 接口处理 Handler(直接以业务命名,禁止 controller_hello.go) +│ +├── service/ # package service:业务逻辑层(用例编排、事务控制、事件发布) +│ └── order.go # 订单业务用例实现(纯 Go 逻辑,禁止依赖 *gin.Context) +│ +├── dao/ # package dao:数据访问持久化层 DAL (GORM CRUD、SQL 转义防注入) +│ └── order.go # 订单数据访问实现(直接以业务命名,禁止 dao_order.go) +│ +├── model/ # package model:纯数据实体与 DTO(无外部依赖) +│ ├── entity/ # 数据库映射实体 (TableName() 带插件专属前缀) +│ │ └── order.go +│ └── do/ # 请求 Request DTO 与响应 Response DTO、领域对象 +│ └── order.go +│ +└── migrations/ # 专属嵌入式 Goose SQL 双方言迁移脚本 (//go:embed) + ├── postgres/ # PostgreSQL 迁移脚本 + └── sqlite/ # SQLite 迁移脚本 ``` +> ⚠️ **严禁**:严禁在根目录平铺 `handlers_*.go`、`service_*.go`、`dao_*.go` 等前缀文件,子包内文件直接按业务实体命名。严格约束 `controller -> service -> dao -> model` 单向依赖。 --- diff --git a/.agents/skills/push-notification/SKILL.md b/.agents/skills/push-notification/SKILL.md index d7dc7006..31361e6c 100644 --- a/.agents/skills/push-notification/SKILL.md +++ b/.agents/skills/push-notification/SKILL.md @@ -14,10 +14,9 @@ description: "Wavelet 项目专用:当需要开发或接入新的系统通知 Wavelet 的消息推送机制采用了**元数据驱动 + 统一触发器 + 异步任务派发**的解耦设计,其分层及职责划分如下: | 目录/包名 | 职责定位 | 包含内容与设计细节 | -| :--- | :--- | :--- | -| **`pkg/push/`** | 推送基础设施层 | 静态定义、不依赖系统数据库和任何框架。定义了统一接口 `Pusher`、单例 `PusherPool` 和多实现(Lark, Webhook, Email 等),提供配置验证及发送功能。 | -| **`internal/apps/admin/push/`** | 通知服务与后台任务层 | 包含以下核心文件:
1. [events.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/events.go):定义通知事件的结构模型(`NotificationMessage`, `EventMetadata`)、内置事件的动态注册中心(`BuiltInEvents` 及 `RegisterBuiltInEvent` 函数)以及统一触发器类 `EventTrigger`(包括其底层的派发引擎逻辑)。
2. [tasks.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/tasks.go):定义 Asynq 后台异步发送任务、处理器 `PushHandler` 及其校验逻辑,并记录推送历史审计。
3. [routers.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/routers.go):管理端接口,负责获取事件配置列表和更新配置。 | -| **`internal/apps/admin/push/custom_events/`** | 自定义通知事件包 | 事件元数据定义与 push 侧处理逻辑;**一个 Go 文件代表一个事件**。在 [register.go](file:///Users/ryan/DEV/Go/Wavelet/internal/apps/admin/push/custom_events/register.go) 统一装配,禁止 `init()` 副作用。 | +| **`backend/plugins/domain/msg_gateway/push/`** | 推送基础设施层 | 静态定义、不依赖系统数据库和任何框架。定义了统一接口 `Pusher` 和多实现(Lark, Webhook, Email 等),提供配置验证及发送功能。 | +| **`internal/apps/admin/push/`** | 通知服务与后台任务层 | 包含以下核心文件:
1. `events.go`:定义通知事件的结构模型(`NotificationMessage`, `EventMetadata`)、内置事件的动态注册中心(`BuiltInEvents` 及 `RegisterBuiltInEvent` 函数)以及统一触发器类 `EventTrigger`(包括其底层的派发引擎逻辑)。
2. `tasks.go`:定义 Asynq 后台异步发送任务、处理器 `PushHandler` 及其校验逻辑,并记录推送历史审计。
3. `routers.go`:管理端接口,负责获取事件配置列表和更新配置。 | +| **`internal/apps/admin/push/custom_events/`** | 自定义通知事件包 | 事件元数据定义与 push 侧处理逻辑;**一个 Go 文件代表一个事件**。在 `register.go` 统一装配,禁止 `init()` 副作用。 | | **`internal/listener/`** | 域事件分发层 | 核心域发射事件(如 `EmitAdminLoggedIn`),push 在 bootstrap 阶段通过 `OnAdminLoggedIn` 订阅,避免 auth/user 直接依赖 push。 | | **`internal/platform/bootstrap/`** | 应用装配根 | `RegisterPushDomainEvents()` 调用 `custom_events.Register()`;`Init` 中执行 `SyncEvents` 将内置事件元数据同步到数据库。 | | **数据库审计表** | 状态与历史审计 | `w_push_events` 存放每个通知事件的启用状态、启用渠道、发送目标和自定义渲染模板。
`w_push_histories` 存放消息发送记录用于审计。 | diff --git a/.github/workflows/build-image-openflare-agent.yml b/.github/workflows/build-image-openflare-agent.yml index 7397bd31..984ee335 100644 --- a/.github/workflows/build-image-openflare-agent.yml +++ b/.github/workflows/build-image-openflare-agent.yml @@ -18,7 +18,7 @@ permissions: env: IMAGE_NAME: openflare-agent - DOCKERFILE: docker/Dockerfile.agent + DOCKERFILE: manifest/docker/Dockerfile.agent jobs: build: diff --git a/.github/workflows/build-image-openflare-relay.yml b/.github/workflows/build-image-openflare-relay.yml index 289685dc..c479e65e 100644 --- a/.github/workflows/build-image-openflare-relay.yml +++ b/.github/workflows/build-image-openflare-relay.yml @@ -18,7 +18,7 @@ permissions: env: IMAGE_NAME: openflare-relay - DOCKERFILE: docker/Dockerfile.relay + DOCKERFILE: manifest/docker/Dockerfile.relay jobs: build: diff --git a/.github/workflows/build-image-openflare.yml b/.github/workflows/build-image-openflare.yml index 27d4e818..9eb21c05 100644 --- a/.github/workflows/build-image-openflare.yml +++ b/.github/workflows/build-image-openflare.yml @@ -24,7 +24,7 @@ permissions: env: IMAGE_NAME: openflare - DOCKERFILE: docker/Dockerfile + DOCKERFILE: manifest/docker/Dockerfile jobs: # Resolve version / registries once. No checkout: triggers alone determine the tag. diff --git a/.github/workflows/build-image-openflared.yml b/.github/workflows/build-image-openflared.yml index 753b75b6..af0f2345 100644 --- a/.github/workflows/build-image-openflared.yml +++ b/.github/workflows/build-image-openflared.yml @@ -18,7 +18,7 @@ permissions: env: IMAGE_NAME: openflared - DOCKERFILE: docker/Dockerfile.flared + DOCKERFILE: manifest/docker/Dockerfile.flared jobs: build: diff --git a/.github/workflows/build-release.yml b/.github/workflows/build-release.yml index dacde310..ed561a96 100644 --- a/.github/workflows/build-release.yml +++ b/.github/workflows/build-release.yml @@ -26,7 +26,7 @@ env: LICENSE README.md README_zh.md - config.example.yaml + manifest/config/config.default.yaml DEPLOYMENT_zh.md permissions: diff --git a/.gitignore b/.gitignore index cba47938..ddb1171a 100644 --- a/.gitignore +++ b/.gitignore @@ -11,6 +11,7 @@ # config config.yaml +manifest/config/config.yaml .env .env.* !.env.example @@ -94,3 +95,8 @@ backend/plugins/**/dist/ backend/core/**/dist/ backend/pkg/**/uploads/ /backend/openflare/plugins/server/upload/filesrv/uploads/ +/backend/plugins/domain/upload/filesrv/uploads/ +/backend/plugins/domain/upload/task/uploads/ +/backend/data/ +/backend/plugins/drivers/driver_http/dist/ +/backend/uploads/ diff --git a/AGENTS.md b/AGENTS.md index 2d637411..43ccf683 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -94,6 +94,41 @@ Strong success criteria let you loop independently. Weak criteria ("make it work - 上游暂缺而确属通用能力时,可先在本仓库实现并登记到 `backend/openflare/upstream-patches.md` (merge 上游后请确认补丁仍在),回流 Wavelet 后删除登记并重新 merge。 +### Cordis 架构核心防线与分层规范 +- **微内核 (`backend/core/`)**: + - 上下文总线(`Context`)、泛型依赖注入(`Container`)、生命周期编排(`Lifecycle`)、扩展点定义(`extpoints/`)与领域事件总线(`EventBus`)。 + - **严禁**包含任何具体业务逻辑,**严禁** import `gin`、`gorm`、`asynq` 等具体运行时依赖。 +- **服务契约 (`backend/core/contracts/`)**: + - 跨插件通信的统一公开 Go Interface(如 `AuthService`、`UserService`、`CacheService`、`DBService`、`StorageService`)与公共 DTO。 + - **严禁**包含任何具体业务实现或 SQL 操作。 +- **自包含插件 (`backend/plugins/`)**: + - 所有业务功能与驱动实现均以插件形式存在(`backend/plugins/drivers/`、`backend/plugins/infra/`、`backend/plugins/domain/` 或下游 `backend/openflare/plugins/`)。 + - 每个插件实现 `core.Plugin`(`Name() string` 与 `Apply(ctx *core.Context) error`)。 + - **统一插件分层架构与标准模板**: + - **开发模板唯一基准**:所有插件统一以 `backend/downstream/plugins/custom_example` 为基准模板构建。 + - **物理子包隔离规范**:统一采用物理子包结构(`plugin.go`, `consts/`, `controller/`, `service/`, `dao/`, `model/` [含 `entity/`, `do/`], `migrations/` [含 `postgres/`, `sqlite/`])。**严禁在根包平铺 `handlers_*`、`service_*`、`dao_*` 等前缀文件**,子包内文件直接按业务实体命名(如 `hello.go`, `user.go`),严格约束 `controller -> service -> dao -> model` 单向依赖。 +- **插件通信与依赖隔离**: + - **严禁跨包 import internal/私有实现**:插件之间严禁直接 import 对方具体实现包代码。 + - **单向服务契约调用**:调用方仅面向 `backend/core/contracts` 编程,在 `Apply` 中通过 `core.Provide[contracts.XxxService](ctx, svc)` 注册服务,通过 `core.Inject[contracts.XxxService](ctx)` 或 `ctx.Using(func(svc contracts.XxxService) { ... })` 声明式解析。 + - **事件总线广播**:状态联动与解耦通信统一通过强类型事件 `ctx.Events().Emit()` 广播,由感兴趣的插件通过 `ctx.Events().On()` 订阅,消除双向依赖与循环引用。 +- **扩展点自包含注册**: + - **HTTP 路由与白名单机制**: + - 插件自包含在 `Apply` 中通过 `ctx.Router().Group(...)` 挂载路由与中间件,禁止跨插件散落注册。 + - **白名单机制**:`driver_http` 与微内核扩展点提供路由白名单支持(`ctx.Router().RegisterWhitelist(patterns...)`),支持精确路径与通配符(如 `/api/v1/oauth/*`)。 + - **所有权主动声明**:认证域(`auth` 插件)与各业务插件必须在 `Apply` 中主动注册其公开/免鉴权接口(如 `/api/v1/user/login`、`/api/v1/oauth/callback`、`/api/v1/cap/*` 等)。 + - **鉴权中间件放行防线**:`auth` 提供的登录鉴权中间件(`LoginRequired`)必须先执行白名单匹配并自动放行,彻底杜绝免鉴权接口被全局或组级鉴权中间件误拦截(返回 401 Unauthorized)。 + - **异步与定时任务**:插件自包含在 `Apply` 中通过 `ctx.Task().Register(...)` 与 `ctx.Schedule().RegisterCron(...)` 声明。 + - **静态启动配置**:插件自包含在 `Apply` 中通过 `ctx.Config().Bind("", &cfg)` 读取**自己声明**的配置,字段以 tag 表达来源:`config`(yaml 路径)、`env`(覆盖变量名)、`default`、`autoEnable`(该变量存在即置真)、`secret`(导出脱敏)。需要在 `Apply` 之前被门禁求值的键,必须在 `DeclareConfig()` 中提前声明并实现 `core.ConfigGatedPlugin`。新增基础设施 key 保持顶层命名(`redis.*`),插件私有配置归 `plugins..*`。**严禁**再造全局配置单例或在 `backend/pkg/` 读取配置。 + - **动态设置**:插件自包含在 `Apply` 中通过 `ctx.Settings().Register(core.SettingSchema{...})` 声明可热更新的管理台设置模式(与上面的静态启动配置分属两层)。 + - **数据迁移**:插件自包含在内部维护 `migrations/*.sql`,通过 `//go:embed` 打包并在 `Apply` 中通过 `ctx.Migrations().Register(pluginID, embedFS)` 注入。 +- **表单一所有者原则 (Single Owner Principle)**: + - 每张数据表有且仅由一个所有者插件声明与维护(表名使用插件前缀如 `w_order_*`)。 + - 严禁插件 B 跨过所有者插件 A 直接 DDL/DML 旁路读写表 A,必须调用插件 A 暴露的 `contracts` 接口或订阅事件。 +- **平台服务复用**: + - 文件摄取统一使用 `upload.Ingest` / `contracts.StorageService`,禁止绕过存储域直接操作底层 Bucket 或直写文件表。 + - 业务缓存统一使用 `ctx.Cache()`(`contracts.CacheService`)或标准缓存框架,禁止自研不带失效广播的本地 map。 + - 数据库操作通过 `ctx.DB()`(`contracts.DBService`)获取受事务与 Trace 保护的连接。 + - 禁止删除 `frontend/node_modules`。 - `backend/pkg/util/` 保持纯净:禁止导入 Gin、GORM、sessions 等 HTTP/Web/DB 框架(会话选项在 `backend/openflare/plugins/server/oauth/session.go`)。 - 测试临时目录只用 `t.TempDir()`,禁止硬编码相对路径写源码树。 diff --git a/Makefile b/Makefile index 5f35bf6f..25d93fa3 100644 --- a/Makefile +++ b/Makefile @@ -40,7 +40,7 @@ build-embedded: code-check: @echo "==> Architecture guards..." - @command -v rg >/dev/null 2>&1 || { echo 'error: rg (ripgrep) is required for architecture guards' >&2; exit 1; } + scripts/check_cordis_architecture.sh @if rg -n 'db\.DB\(|db\.Redis' backend/openflare/plugins/server/kernel/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; \ @@ -108,7 +108,7 @@ cross-build: (version=$(or $(VERSION),dev))..." @mkdir -p bin docker build \ - --file docker/Dockerfile.cross \ + --file manifest/docker/Dockerfile.cross \ --target export \ --build-arg VERSION=$(or $(VERSION),dev) \ --build-arg BUILD_DATE="$(shell date -u +'%Y-%m-%dT%H:%M:%SZ')" \ diff --git a/README_zh.md b/README_zh.md index 0ce1e87c..ed0b6658 100644 --- a/README_zh.md +++ b/README_zh.md @@ -95,10 +95,10 @@ cd refreshing ### 2. 配置环境 ```bash -cp config.example.yaml config.yaml +cp manifest/config/config.default.yaml manifest/config/config.yaml ``` -编辑 `config.yaml`,配置数据库和 Redis。OIDC 认证源统一在管理后台的系统设置页面运行时配置。 +编辑 `manifest/config/config.yaml`,配置数据库和 Redis。OIDC 认证源统一在管理后台的系统设置页面运行时配置。 ### 3. 初始化数据库 @@ -156,7 +156,7 @@ pnpm dev ## ⚙️ 配置说明 -主要配置项(完整说明请参考 `config.example.yaml`): +主要配置项(完整说明请参考 `manifest/config/config.default.yaml`): | 配置项 | 说明 | 示例 | |--------|------|------| @@ -211,9 +211,8 @@ pnpm format ``` wavelet/ ├── main.go # 程序入口(委托给 internal/cmd) -├── config.example.yaml # 配置模板 ├── Makefile # 常用命令(swagger、tidy、license、cross-build) -├── docker/ # Docker 镜像构建文件(集成/前端/后端) +├── manifest/ # 项目清单与编排:docker 镜像构建、deploy (k8s)、config 配置(默认/覆盖) ├── docs/ # Swagger 自动生成文档 ├── frontend/ # Next.js 前端应用 │ ├── app/ # App Router 页面 diff --git a/backend/cmd/app.go b/backend/cmd/app.go index 07d5b7ba..2408a118 100644 --- a/backend/cmd/app.go +++ b/backend/cmd/app.go @@ -10,8 +10,7 @@ import ( "Wavelet/openflare/plugins/server/migrate" "Wavelet/plugins/domain/admin" "Wavelet/plugins/domain/auth" - "Wavelet/plugins/domain/cap" - "Wavelet/plugins/domain/message_gateway" + "Wavelet/plugins/domain/msg_gateway" "Wavelet/plugins/domain/risk_control" "Wavelet/plugins/domain/system" "Wavelet/plugins/domain/upload" @@ -104,15 +103,14 @@ func newOpenFlareApp(profile core.Profile, opts ...core.AppOption) *core.App { driver_inproc_cron.New(), ) - // 3. Register all 8 domain business plugins (admin first to ensure schema and base config tables exist) + // 3. Register all 7 domain business plugins (admin first to ensure schema and base config tables exist) app.Use( admin.New(), user.New(), auth.New(), - message_gateway.New(), + msg_gateway.New(), risk_control.New(), upload.New(), - cap.New(), system.New(), ) diff --git a/backend/cmd/app_test.go b/backend/cmd/app_test.go index 23ab2278..38c4bbf2 100644 --- a/backend/cmd/app_test.go +++ b/backend/cmd/app_test.go @@ -36,7 +36,7 @@ func TestNewOpenFlareAppRegistersServerAndWaveletUser(t *testing.T) { for _, p := range app.Plugins() { names[p.Name()] = true } - for _, n := range []string{"user", "auth", "cap", "admin", "server"} { + for _, n := range []string{"user", "auth", "admin", "server"} { if !names[n] { t.Errorf("missing plugin %s", n) } diff --git a/backend/core/appctx.go b/backend/core/appctx.go new file mode 100644 index 00000000..1c1e2201 --- /dev/null +++ b/backend/core/appctx.go @@ -0,0 +1,43 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package core + +import "context" + +type appContextKey struct{} + +// WithAppContext attaches the micro-kernel Context to a standard context.Context +// so request and worker handlers can Inject services without package-level setters. +func WithAppContext(ctx context.Context, app *Context) context.Context { + if ctx == nil { + ctx = context.Background() + } + if app == nil { + return ctx + } + return context.WithValue(ctx, appContextKey{}, app.Root()) +} + +// AppContext extracts the micro-kernel Context from ctx, if present. +func AppContext(ctx context.Context) *Context { + if ctx == nil { + return nil + } + if c, ok := ctx.(*Context); ok { + return c + } + app, _ := ctx.Value(appContextKey{}).(*Context) + return app +} + +// InjectFrom resolves T from ctx when it carries a micro-kernel Context +// (*Context itself, or a value attached by WithAppContext). +func InjectFrom[T any](ctx context.Context) (T, error) { + var zero T + app := AppContext(ctx) + if app == nil { + return zero, ErrNilContext + } + return Inject[T](app) +} diff --git a/backend/core/container.go b/backend/core/container.go index ed783572..a910a314 100644 --- a/backend/core/container.go +++ b/backend/core/container.go @@ -13,18 +13,20 @@ import ( // Container manages service registration and resolution using Go reflection and generics. type Container struct { - mu sync.RWMutex - parent *Container - services map[reflect.Type]any - listeners map[reflect.Type][]func(any) + mu sync.RWMutex + parent *Container + services map[reflect.Type]any + interfaceCache map[reflect.Type]any + listeners map[reflect.Type][]func(any) } // NewContainer creates a new IoC container instance with an optional parent container. func NewContainer(parent *Container) *Container { return &Container{ - parent: parent, - services: make(map[reflect.Type]any), - listeners: make(map[reflect.Type][]func(any)), + parent: parent, + services: make(map[reflect.Type]any), + interfaceCache: make(map[reflect.Type]any), + listeners: make(map[reflect.Type][]func(any)), } } @@ -45,6 +47,7 @@ func (c *Container) remove(targetType reflect.Type) { c.mu.Lock() defer c.mu.Unlock() delete(c.services, targetType) + c.interfaceCache = make(map[reflect.Type]any) } // Provide registers a typed service implementation into the Context hierarchy's root IoC container. @@ -88,6 +91,7 @@ func ProvideScoped[T any](ctx *Context, service T) { func (c *Container) provide(targetType reflect.Type, service any) { c.mu.Lock() c.services[targetType] = service + c.interfaceCache = make(map[reflect.Type]any) // Collect any matching listeners to invoke outside the lock var callbacks []func(any) @@ -131,17 +135,14 @@ func (c *Container) resolve(targetType reflect.Type) (any, error) { c.mu.RUnlock() return val, nil } + c.mu.RUnlock() - // 2. Interface assignment scan + // 2. Interface assignment scan & cache if targetType.Kind() == reflect.Interface { - for _, val := range c.services { - if reflect.TypeOf(val).Implements(targetType) { - c.mu.RUnlock() - return val, nil - } + if val, found := c.resolveInterface(targetType); found { + return val, nil } } - c.mu.RUnlock() // 3. Fallback to parent container if c.parent != nil { @@ -151,6 +152,35 @@ func (c *Container) resolve(targetType reflect.Type) (any, error) { return nil, fmt.Errorf("%w: %v", ErrServiceNotFound, targetType) } +func (c *Container) resolveInterface(targetType reflect.Type) (any, bool) { + c.mu.RLock() + if val, ok := c.interfaceCache[targetType]; ok { + c.mu.RUnlock() + return val, true + } + + var matched any + for _, val := range c.services { + if reflect.TypeOf(val).Implements(targetType) { + matched = val + break + } + } + c.mu.RUnlock() + + if matched == nil { + return nil, false + } + + c.mu.Lock() + if c.interfaceCache == nil { + c.interfaceCache = make(map[reflect.Type]any) + } + c.interfaceCache[targetType] = matched + c.mu.Unlock() + return matched, true +} + // MustInject resolves a service of type T or panics if the service is not found. func MustInject[T any](ctx *Context) T { s, err := Inject[T](ctx) @@ -201,13 +231,17 @@ func Using3[T1, T2, T3 any](ctx *Context, fn func(s1 T1, s2 T2, s3 T3)) error { // When registers a reactive hook that is called immediately if T is already provided, // or called as soon as T is provided in the future. +// +// Listeners are stored on the root container so they observe core.Provide, which +// always writes to the root. Registering on a Fiber child container would miss +// services provided by plugins that load later. func When[T any](ctx *Context, fn func(s T)) { if ctx == nil { panic("core: nil context provided to When") } targetType := reflect.TypeFor[T]() - c := ctx.Container() + c := ctx.Root().Container() // If already ready, execute immediately if s, err := Inject[T](ctx); err == nil { @@ -223,3 +257,9 @@ func When[T any](ctx *Context, fn func(s T)) { } }) } + +// Bind is When with a name that matches plugin wiring: fill a dependency as +// soon as the root container provides it. +func Bind[T any](ctx *Context, fn func(s T)) { + When(ctx, fn) +} diff --git a/backend/core/context_test.go b/backend/core/context_test.go index d6c66bc6..f7a0a3e4 100644 --- a/backend/core/context_test.go +++ b/backend/core/context_test.go @@ -313,6 +313,46 @@ func TestContextReactiveWhen(t *testing.T) { assert.True(t, immediateCalled) } +func TestWhenObservesProvideFromForkedFiberContext(t *testing.T) { + root := core.NewContext(context.Background()) + adminFiber := root.Fork() + lateFiber := root.Fork() + + var got atomic.Bool + core.When[SampleService](adminFiber, func(s SampleService) { + if s != nil { + got.Store(true) + } + }) + assert.False(t, got.Load()) + + core.Provide[SampleService](lateFiber, &sampleServiceImpl{}) + assert.True(t, got.Load(), "When on a Fiber child must observe Provide on the root") +} + +func TestBindIsWhen(t *testing.T) { + ctx := core.NewContext(context.Background()) + var called atomic.Bool + core.Bind[SampleService](ctx, func(s SampleService) { + called.Store(true) + }) + core.Provide[SampleService](ctx, &sampleServiceImpl{}) + assert.True(t, called.Load()) +} + +func TestInjectFromAppContext(t *testing.T) { + app := core.NewContext(context.Background()) + core.Provide[SampleService](app, &sampleServiceImpl{prefix: "Hi:"}) + + req := core.WithAppContext(context.Background(), app) + svc, err := core.InjectFrom[SampleService](req) + require.NoError(t, err) + assert.Equal(t, "Hi: Ada", svc.Greet("Ada")) + + _, err = core.InjectFrom[SampleService](context.Background()) + assert.ErrorIs(t, err, core.ErrNilContext) +} + func TestContextDisposerLifecycle(t *testing.T) { parent := core.NewContext(context.Background()) child := parent.Fork() @@ -531,10 +571,16 @@ func TestContext_ScopedExtpoints_RevertibleEffects(t *testing.T) { root := core.NewContext(context.Background()) child := root.Fork() - // Register route, task, schedule, setting, event on child + // Register route, task, schedule, setting, event, middleware, whitelist on child child.Router().GET("/test-route", func() {}) assert.Equal(t, 1, len(root.Router().Routes())) + child.Router().Use("scoped_middleware") + assert.Equal(t, 1, len(root.Router().Middlewares())) + + child.Router().RegisterWhitelist("/api/v1/scoped/*") + assert.True(t, root.Router().IsWhitelisted("/api/v1/scoped/test")) + child.Tasks().Register("test:task", func() {}) assert.Equal(t, 1, len(root.Tasks().Tasks())) @@ -553,8 +599,48 @@ func TestContext_ScopedExtpoints_RevertibleEffects(t *testing.T) { // All child effects should be cleanly revoked in LIFO order assert.Equal(t, 0, len(root.Router().Routes())) + assert.Equal(t, 0, len(root.Router().Middlewares())) + assert.False(t, root.Router().IsWhitelisted("/api/v1/scoped/test")) assert.Equal(t, 0, len(root.Tasks().Tasks())) assert.Equal(t, 0, len(root.Schedules().Schedules())) assert.Equal(t, 0, len(root.Settings().Schemas())) assert.Equal(t, 0, root.Events().Listeners("test:event")) } + +func TestContainer_InterfaceResolutionCache(t *testing.T) { + ctx := core.NewContext(context.Background()) + svc := &sampleServiceImpl{prefix: "Cached:"} + + core.Provide[SampleService](ctx, svc) + + // 1. Initial resolution populates interfaceCache + res1, err := core.Inject[SampleService](ctx) + require.NoError(t, err) + assert.Equal(t, "Cached: Alice", res1.Greet("Alice")) + + // 2. Subsequent resolutions hit interfaceCache + res2, err := core.Inject[SampleService](ctx) + require.NoError(t, err) + assert.Same(t, res1, res2) + + // 3. Concurrent lookups + var wg sync.WaitGroup + for i := 0; i < 20; i++ { + wg.Add(1) + go func() { + defer wg.Done() + r, e := core.Inject[SampleService](ctx) + assert.NoError(t, e) + assert.Equal(t, "Cached: Bob", r.Greet("Bob")) + }() + } + wg.Wait() + + // 4. Overriding/providing another service invalidates cache + svc2 := &sampleServiceImpl{prefix: "Updated:"} + core.Provide[SampleService](ctx, svc2) + + res3, err := core.Inject[SampleService](ctx) + require.NoError(t, err) + assert.Equal(t, "Updated: Alice", res3.Greet("Alice")) +} diff --git a/backend/core/contracts/config.go b/backend/core/contracts/config.go new file mode 100644 index 00000000..0a17928c --- /dev/null +++ b/backend/core/contracts/config.go @@ -0,0 +1,33 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package contracts + +import ( + "context" + "time" +) + +// SystemConfigDTO represents a system configuration key-value entry. +type SystemConfigDTO struct { + Key string `json:"key"` + Value string `json:"value"` + Type string `json:"type"` + Visibility int `json:"visibility"` + Description string `json:"description"` + UpdatedAt time.Time `json:"updated_at"` + CreatedAt time.Time `json:"created_at"` +} + +// SystemConfigService defines the unified contract for querying and mutating system configurations. +type SystemConfigService interface { + GetByKey(ctx context.Context, key string) (SystemConfigDTO, error) + ListByKeys(ctx context.Context, keys []string) (map[string]SystemConfigDTO, error) + ListVisible(ctx context.Context) ([]SystemConfigDTO, error) + ListByType(ctx context.Context, configType string) ([]SystemConfigDTO, error) + GetIntByKey(ctx context.Context, key string) (int, error) + GetBoolByKey(ctx context.Context, key string) (bool, error) + SaveOrUpdate(ctx context.Context, key, value string) error + InvalidateCache(ctx context.Context, key string) error + InvalidateAllCaches(ctx context.Context) error +} diff --git a/backend/core/contracts/config_public.go b/backend/core/contracts/config_public.go index 31cd92b2..57fd4d37 100644 --- a/backend/core/contracts/config_public.go +++ b/backend/core/contracts/config_public.go @@ -5,8 +5,10 @@ package contracts import "context" -// PublicConfigProvider supplies the payload for GET /api/v1/config/public -// when a downstream plugin replaces Wavelet's default {configs, app} JSON. +// PublicConfigProvider supplies GET /api/v1/config/public. +// The owner of w_system_configs (admin) must provide this. The payload is a +// flat key/value map of visibility=1 rows; the frontend reads keys such as +// cap_login_enabled directly off data. type PublicConfigProvider interface { - PublicConfig(ctx context.Context) (any, error) + PublicConfig(ctx context.Context) (map[string]string, error) } diff --git a/backend/core/contracts/events.go b/backend/core/contracts/events.go index 8de13c0a..43f52740 100644 --- a/backend/core/contracts/events.go +++ b/backend/core/contracts/events.go @@ -151,3 +151,8 @@ type UserDeletedEvent struct { CurrentUserID uint64 `json:"current_user_id,string"` TargetUserID uint64 `json:"target_user_id,string"` } + +// SystemCleanupEvent fires when a periodic system cleanup is triggered. +type SystemCleanupEvent struct { + TriggeredAt string `json:"triggered_at"` +} diff --git a/backend/core/contracts/limiter.go b/backend/core/contracts/limiter.go new file mode 100644 index 00000000..9d5b2718 --- /dev/null +++ b/backend/core/contracts/limiter.go @@ -0,0 +1,36 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package contracts defines unified service interfaces and DTOs for cross-plugin communication. +package contracts + +import ( + "context" + "time" +) + +// Rate specifies a rate limit of Limit events permitted within a Period. +type Rate struct { + Limit int `json:"limit"` + Period time.Duration `json:"period"` +} + +// RateLimitResult holds the outcome of a rate limit check. +type RateLimitResult struct { + Allowed bool `json:"allowed"` + Remaining int `json:"remaining"` + ResetAfter time.Duration `json:"reset_after"` + RetryAfter time.Duration `json:"retry_after"` +} + +// LimiterService defines the rate limiting service contract for cross-plugin communication. +type LimiterService interface { + // Allow checks whether 1 event for the given key is permitted under the specified rate. + Allow(ctx context.Context, key string, rate Rate) (*RateLimitResult, error) + + // AllowN checks whether n events for the given key are permitted under the specified rate. + AllowN(ctx context.Context, key string, rate Rate, n int) (*RateLimitResult, error) + + // Reset clears the rate limit state for the given key. + Reset(ctx context.Context, key string) error +} diff --git a/backend/core/contracts/push.go b/backend/core/contracts/push.go index 9769a144..37863342 100644 --- a/backend/core/contracts/push.go +++ b/backend/core/contracts/push.go @@ -6,6 +6,7 @@ package contracts import "context" +// PushNotificationTemplate defines notification message template payload. type PushNotificationTemplate struct { Title string Content string @@ -13,6 +14,7 @@ type PushNotificationTemplate struct { Ext map[string]any } +// PushEventMeta defines metadata for a system push event. type PushEventMeta struct { Key string Name string @@ -20,6 +22,7 @@ type PushEventMeta struct { DefaultTemplate PushNotificationTemplate } +// PushRegistry defines the interface for registering built-in events. type PushRegistry interface { RegisterBuiltInEvent(meta PushEventMeta) SyncEvents(ctx context.Context) error diff --git a/backend/core/contracts/task.go b/backend/core/contracts/task.go index 3fe1c01c..efd05de6 100644 --- a/backend/core/contracts/task.go +++ b/backend/core/contracts/task.go @@ -43,6 +43,12 @@ type TaskResultDTO struct { Detail any `json:"detail,omitempty"` } +// TaskHandler is the preferred background task handler. Drivers invoke Execute +// and persist Message/Detail onto the execution record. +type TaskHandler interface { + Execute(ctx context.Context, payload []byte) (*TaskResultDTO, error) +} + // TaskExecutionDTO represents a single task execution record. type TaskExecutionDTO struct { ID uint64 `json:"id,string"` @@ -65,6 +71,14 @@ type TaskExecutionDTO struct { UpdatedAt time.Time `json:"updated_at"` } +// Canonical triggered_by values persisted on task executions and shown in admin UI. +const ( + TaskTriggerSystem = "system" + TaskTriggerManual = "manual" + TaskTriggerRetry = "retry" + TaskTriggerSchedule = "schedule" +) + // TaskService defines the unified contract for dispatching and tracking background tasks. type TaskService interface { Dispatch(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error) @@ -76,4 +90,5 @@ type TaskService interface { AppendLog(ctx context.Context, format string, args ...any) ListExecutions(ctx context.Context, taskType, status string, page, pageSize int) ([]TaskExecutionDTO, int64, error) GetExecution(ctx context.Context, id uint64) (*TaskExecutionDTO, error) + GetExecutionByTaskID(ctx context.Context, taskID string) (*TaskExecutionDTO, error) } diff --git a/backend/core/contracts/upload.go b/backend/core/contracts/upload.go new file mode 100644 index 00000000..2d5b0a04 --- /dev/null +++ b/backend/core/contracts/upload.go @@ -0,0 +1,80 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package contracts + +import ( + "context" + "database/sql/driver" + "encoding/json" + "fmt" + "io" + "time" +) + +// UploadMetadataDTO represents upload metadata JSON. +type UploadMetadataDTO struct { + Width int `json:"width,omitempty"` + Height int `json:"height,omitempty"` + Duration float64 `json:"duration,omitempty"` + OriginalMime string `json:"original_mime,omitempty"` + UserAgent string `json:"user_agent,omitempty"` + ClientIP string `json:"client_ip,omitempty"` + Bucket string `json:"bucket,omitempty"` + Extra map[string]any `json:"extra,omitempty"` +} + +// Value implements the driver.Valuer interface for database serialization. +func (m UploadMetadataDTO) Value() (driver.Value, error) { + return json.Marshal(m) +} + +// Scan implements the sql.Scanner interface for database deserialization. +func (m *UploadMetadataDTO) Scan(value any) error { + if value == nil { + *m = UploadMetadataDTO{} + return nil + } + switch v := value.(type) { + case []byte: + return json.Unmarshal(v, m) + case string: + return json.Unmarshal([]byte(v), m) + default: + return fmt.Errorf("cannot scan type %T into UploadMetadataDTO", value) + } +} + +// UploadDTO represents an uploaded file record. +type UploadDTO struct { + ID uint64 `json:"id"` + UserID uint64 `json:"user_id"` + FileName string `json:"file_name"` + FilePath string `json:"file_path"` + MimeType string `json:"mime_type"` + Size int64 `json:"size"` + Hash string `json:"hash"` + Status string `json:"status"` + Type string `json:"type"` + Metadata UploadMetadataDTO `json:"metadata"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// OpenedUploadDTO encapsulates the retrieved object stream and its metadata. +type OpenedUploadDTO struct { + Upload UploadDTO + Body io.ReadCloser + ContentType string + ContentLength int64 +} + +// UploadService defines the unified contract for managed file uploads and media entities. +type UploadService interface { + GetByID(ctx context.Context, id uint64) (*UploadDTO, error) + OpenStoredUpload(ctx context.Context, id uint64) (*OpenedUploadDTO, error) + Remove(ctx context.Context, id uint64) error + RemoveOwned(ctx context.Context, id uint64, userID uint64) error + FindByHash(ctx context.Context, hash string, size int64) (*UploadDTO, error) + RebuildStats(ctx context.Context) error +} diff --git a/backend/core/events.go b/backend/core/events.go index d30be096..e6ee840c 100644 --- a/backend/core/events.go +++ b/backend/core/events.go @@ -283,7 +283,8 @@ func (b *EventBus) Waterfall(ctx context.Context, topic string, initialPayload a } // Parallel executes all subscribers of the topic concurrently in separate goroutines. -// It waits for all handlers to complete and collects any errors via errors.Join. +// It waits for all handlers to complete or returns immediately if ctx is cancelled/timed out, +// collecting any handler errors via errors.Join. // //nolint:contextcheck func (b *EventBus) Parallel(ctx context.Context, topic string, payload any) error { @@ -325,17 +326,25 @@ func (b *EventBus) Parallel(ctx context.Context, topic string, payload any) erro }(l) } - wg.Wait() - close(errCh) + done := make(chan struct{}) + go func() { + wg.Wait() + close(done) + }() - var errs []error - for err := range errCh { - if err != nil { - errs = append(errs, err) + select { + case <-done: + close(errCh) + var errs []error + for err := range errCh { + if err != nil { + errs = append(errs, err) + } } + return errors.Join(errs...) + case <-ctx.Done(): + return ctx.Err() } - - return errors.Join(errs...) } // Serial executes subscribers strictly in sequence. diff --git a/backend/core/events_test.go b/backend/core/events_test.go index 730661fc..d4801632 100644 --- a/backend/core/events_test.go +++ b/backend/core/events_test.go @@ -11,6 +11,7 @@ import ( "sync" "sync/atomic" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -368,6 +369,25 @@ func TestEventBusParallel(t *testing.T) { assert.Equal(t, int64(20), counter.Load()) } +func TestEventBusParallelContextTimeout(t *testing.T) { + bus := core.NewEventBus() + + bus.On("test:timeout", func(ctx context.Context) error { + select { + case <-time.After(200 * time.Millisecond): + return nil + case <-ctx.Done(): + return ctx.Err() + } + }) + + ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond) + defer cancel() + + err := bus.Parallel(ctx, "test:timeout", nil) + assert.ErrorIs(t, err, context.DeadlineExceeded) +} + func TestEventBusSerial(t *testing.T) { bus := core.NewEventBus() diff --git a/backend/core/extpoints/config_resolve.go b/backend/core/extpoints/config_resolve.go index 3d8c7c37..4b2e1f7b 100644 --- a/backend/core/extpoints/config_resolve.go +++ b/backend/core/extpoints/config_resolve.go @@ -181,7 +181,7 @@ func formatEntryValue(value any, secret bool) string { return "" } rv := reflect.ValueOf(value) - if rv.Kind() == reflect.Ptr { + if rv.Kind() == reflect.Pointer { if rv.IsNil() { return "" } diff --git a/backend/core/extpoints/config_value.go b/backend/core/extpoints/config_value.go index b1cb0b98..5d1da6b5 100644 --- a/backend/core/extpoints/config_value.go +++ b/backend/core/extpoints/config_value.go @@ -33,7 +33,7 @@ func convertValue(raw any, typ reflect.Type) (any, error) { return convertSlice(raw, typ) case reflect.Struct: return convertStruct(raw, typ) - case reflect.Ptr: + case reflect.Pointer: return convertPointer(raw, typ) default: return nil, fmt.Errorf("%w: %s is not a supported configuration type", ErrConfigType, typ) @@ -44,7 +44,7 @@ func convertValue(raw any, typ reflect.Type) (any, error) { // Nested pointers are rejected so configuration tags stay one level deep. func convertPointer(raw any, typ reflect.Type) (any, error) { elemType := typ.Elem() - if elemType.Kind() == reflect.Ptr { + if elemType.Kind() == reflect.Pointer { return nil, fmt.Errorf("%w: %s is not a supported configuration type", ErrConfigType, typ) } elem, err := convertValue(raw, elemType) diff --git a/backend/core/extpoints/extpoints_test.go b/backend/core/extpoints/extpoints_test.go index 6d51b999..1c558db0 100644 --- a/backend/core/extpoints/extpoints_test.go +++ b/backend/core/extpoints/extpoints_test.go @@ -80,12 +80,12 @@ func TestRouterExtension(t *testing.T) { for _, route := range routes { if route.Method == "GET" && route.Path == "/api/v1/orders" { foundOrderGet = true - assert.Equal(t, []any{mGlobal, mAPI, "api_extra_middleware"}, route.Middlewares) + assert.Equal(t, []any{mAPI, "api_extra_middleware"}, route.Middlewares) assert.Equal(t, []any{hList}, route.Handlers) } if route.Method == "PUT" && route.Path == "/api/v1/admin/users/:id" { foundUserPut = true - assert.Equal(t, []any{mGlobal, mAPI, "api_extra_middleware", mAdmin}, route.Middlewares) + assert.Equal(t, []any{mAPI, "api_extra_middleware", mAdmin}, route.Middlewares) assert.Equal(t, []any{hUserPut}, route.Handlers) } } @@ -93,6 +93,21 @@ func TestRouterExtension(t *testing.T) { assert.True(t, foundUserPut) } +func TestRouterGlobalMiddlewareIsNotSnapshottedOntoRoutes(t *testing.T) { + r := extpoints.NewRouterRegistry() + r.GET("/before", "handler") + r.Use("late_global") + r.GET("/after", "handler") + + for _, route := range r.Routes() { + if len(route.Middlewares) != 0 { + t.Errorf("route %s %s Middlewares = %v, want none (globals live on Router.Middlewares)", + route.Method, route.Path, route.Middlewares) + } + } + assert.Equal(t, []any{"late_global"}, r.Middlewares()) +} + func TestRouterWhitelist(t *testing.T) { r := extpoints.NewRouterRegistry() require.NotNil(t, r) @@ -221,10 +236,30 @@ func TestTaskExtension(t *testing.T) { assert.True(t, ok) assert.Equal(t, "order:cancel_timeout", task.Pattern) + byType, ok := tr.Get("cancel_timeout") + assert.True(t, ok, "Get should resolve admin type identifier") + assert.Equal(t, "order:cancel_timeout", byType.Pattern) + _, ok = tr.Get("unknown") assert.False(t, ok) } +func TestTaskRegisterRejectsNilHandler(t *testing.T) { + tr := extpoints.NewTaskRegistry() + assert.Panics(t, func() { + tr.Register("broken:task", nil) + }) +} + +func TestTaskRegisterRejectsDuplicateType(t *testing.T) { + tr := extpoints.NewTaskRegistry() + handler := func(ctx context.Context, payload []byte) error { return nil } + tr.Register("system:cleanup", handler, extpoints.WithTaskType("system_cleanup")) + assert.Panics(t, func() { + tr.Register("admin:system_cleanup", handler, extpoints.WithTaskType("system_cleanup")) + }) +} + func TestScheduleExtension(t *testing.T) { sr := extpoints.NewScheduleRegistry() require.NotNil(t, sr) @@ -344,7 +379,7 @@ func TestContextExtensionPointsIntegration(t *testing.T) { func TestExtensionPointsUnregister(t *testing.T) { ctx := core.NewContext(context.Background()) - // 1. Router unregister + // 1. Router unregister (routes, middlewares, whitelist) rd := ctx.Router().GET("/temp", "temp_handler") assert.Greater(t, rd.ID, uint64(0)) assert.Len(t, ctx.Router().Routes(), 1) @@ -356,6 +391,20 @@ func TestExtensionPointsUnregister(t *testing.T) { assert.True(t, ctx.Router().UnregisterByID(rd2.ID)) assert.Len(t, ctx.Router().Routes(), 0) + ctx.Router().Use("mw1") + assert.Len(t, ctx.Router().Middlewares(), 1) + if reg, ok := ctx.Router().(*extpoints.RouterRegistry); ok { + ids := reg.UseWithID("mw2") + assert.Len(t, ctx.Router().Middlewares(), 2) + assert.True(t, ctx.Router().UnregisterMiddlewareByID(ids[0])) + assert.Len(t, ctx.Router().Middlewares(), 1) + } + + ctx.Router().RegisterWhitelist("/api/v1/temp/*") + assert.True(t, ctx.Router().IsWhitelisted("/api/v1/temp/item")) + ctx.Router().UnregisterWhitelist("/api/v1/temp/*") + assert.False(t, ctx.Router().IsWhitelisted("/api/v1/temp/item")) + // 2. Task unregister ctx.Task().Register("temp:task", "handler") assert.Len(t, ctx.Task().Tasks(), 1) diff --git a/backend/core/extpoints/router.go b/backend/core/extpoints/router.go index 44258028..92bcda33 100644 --- a/backend/core/extpoints/router.go +++ b/backend/core/extpoints/router.go @@ -39,17 +39,26 @@ type RouterExtension interface { Middlewares() []any Unregister(method, path string) bool UnregisterByID(id uint64) bool + UnregisterMiddlewareByID(id uint64) bool RegisterWhitelist(patterns ...string) + UnregisterWhitelist(patterns ...string) Whitelist() []string IsWhitelisted(path string) bool } +// middlewareDefinition holds an assigned ID and handler for registered middleware. +type middlewareDefinition struct { + ID uint64 + Handler any +} + // RouterRegistry implements RouterExtension as the root route and middleware collector. type RouterRegistry struct { mu sync.RWMutex nextID uint64 + nextMWID uint64 routes []RouteDefinition - middlewares []any + middlewares []middlewareDefinition whitelist PathWhitelist } @@ -60,17 +69,48 @@ func NewRouterRegistry() *RouterRegistry { // Use registers global middlewares to the router. func (r *RouterRegistry) Use(middlewares ...any) { - r.mu.Lock() - defer r.mu.Unlock() - r.middlewares = append(r.middlewares, middlewares...) + r.UseWithID(middlewares...) } -// Middlewares returns a copy of registered root middlewares. +// UseWithID registers global middlewares to the router and returns their assigned IDs. +func (r *RouterRegistry) UseWithID(middlewares ...any) []uint64 { + r.mu.Lock() + defer r.mu.Unlock() + + ids := make([]uint64, 0, len(middlewares)) + for _, mw := range middlewares { + r.nextMWID++ + r.middlewares = append(r.middlewares, middlewareDefinition{ + ID: r.nextMWID, + Handler: mw, + }) + ids = append(ids, r.nextMWID) + } + return ids +} + +// UnregisterMiddlewareByID removes a registered global middleware by its unique ID. +func (r *RouterRegistry) UnregisterMiddlewareByID(id uint64) bool { + r.mu.Lock() + defer r.mu.Unlock() + + for i, mw := range r.middlewares { + if mw.ID == id { + r.middlewares = append(r.middlewares[:i], r.middlewares[i+1:]...) + return true + } + } + return false +} + +// Middlewares returns a copy of registered root middleware handlers. func (r *RouterRegistry) Middlewares() []any { r.mu.RLock() defer r.mu.RUnlock() res := make([]any, len(r.middlewares)) - copy(res, r.middlewares) + for i, mw := range r.middlewares { + res[i] = mw.Handler + } return res } @@ -95,11 +135,12 @@ func (r *RouterRegistry) addRoute(method, fullPath string, handlers ...any) Rout r.nextID++ rd := RouteDefinition{ - ID: r.nextID, - Method: strings.ToUpper(method), - Path: fullPath, - Handlers: handlers, - Middlewares: append([]any(nil), r.middlewares...), + ID: r.nextID, + Method: strings.ToUpper(method), + Path: fullPath, + Handlers: handlers, + // Global Router.Use middlewares are applied at HTTP Start from + // Router.Middlewares(), so late-registered plugins still wrap earlier routes. } r.routes = append(r.routes, rd) return rd @@ -195,6 +236,11 @@ func (r *RouterRegistry) RegisterWhitelist(patterns ...string) { r.whitelist.Add(patterns...) } +// UnregisterWhitelist removes path patterns from the whitelist. +func (r *RouterRegistry) UnregisterWhitelist(patterns ...string) { + r.whitelist.Remove(patterns...) +} + // Whitelist returns a copy of all registered whitelist path patterns. func (r *RouterRegistry) Whitelist() []string { return r.whitelist.Patterns() @@ -241,9 +287,7 @@ func (g *RouterGroup) addRoute(method, fullPath string, handlers ...any) RouteDe g.registry.mu.Lock() defer g.registry.mu.Unlock() - allMiddlewares := make([]any, 0, len(g.registry.middlewares)+len(g.middlewares)) - allMiddlewares = append(allMiddlewares, g.registry.middlewares...) - allMiddlewares = append(allMiddlewares, g.middlewares...) + allMiddlewares := append([]any(nil), g.middlewares...) g.registry.nextID++ rd := RouteDefinition{ @@ -268,6 +312,11 @@ func (g *RouterGroup) UnregisterByID(id uint64) bool { return g.registry.UnregisterByID(id) } +// UnregisterMiddlewareByID removes a middleware by ID via the root registry. +func (g *RouterGroup) UnregisterMiddlewareByID(id uint64) bool { + return g.registry.UnregisterMiddlewareByID(id) +} + // GET registers a GET route in this group. func (g *RouterGroup) GET(path string, handlers ...any) RouteDefinition { return g.Handle("GET", path, handlers...) @@ -332,6 +381,13 @@ func (g *RouterGroup) RegisterWhitelist(patterns ...string) { } } +// UnregisterWhitelist removes path patterns under this group prefix from the whitelist. +func (g *RouterGroup) UnregisterWhitelist(patterns ...string) { + for _, p := range patterns { + g.registry.UnregisterWhitelist(joinPaths(g.prefix, p)) + } +} + // Whitelist returns a copy of all registered whitelist path patterns. func (g *RouterGroup) Whitelist() []string { return g.registry.Whitelist() @@ -466,6 +522,27 @@ func (w *PathWhitelist) Replace(patterns ...string) { w.patterns = compiled } +// Remove removes matching patterns from the whitelist. +func (w *PathWhitelist) Remove(patterns ...string) { + if len(patterns) == 0 { + return + } + targets := make(map[string]struct{}, len(patterns)) + for _, p := range patterns { + targets[cleanPath(p)] = struct{}{} + } + + w.mu.Lock() + defer w.mu.Unlock() + filtered := w.patterns[:0] + for _, p := range w.patterns { + if _, remove := targets[p.raw]; !remove { + filtered = append(filtered, p) + } + } + w.patterns = filtered +} + // Match reports whether path matches any registered pattern. Equivalent to calling // MatchPathPattern for every pattern, except the path is normalised and split once. func (w *PathWhitelist) Match(path string) bool { diff --git a/backend/core/extpoints/router_raw_test.go b/backend/core/extpoints/router_raw_test.go index 2d4f62b0..31b0f38f 100644 --- a/backend/core/extpoints/router_raw_test.go +++ b/backend/core/extpoints/router_raw_test.go @@ -1,7 +1,12 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + package extpoints import "testing" +// TestHandleRawPreservesTrailingSlash 验证 HandleRaw 能表达 /x 与 /x/ 两条不同路由, +// 而 Handle 会归一化掉尾部斜杠(server 插件的 list 端点历史行为依赖这一点)。 func TestHandleRawPreservesTrailingSlash(t *testing.T) { r := &RouterRegistry{} g := r.Group("/api/v1/nodes") @@ -19,9 +24,29 @@ func TestHandleRawPreservesTrailingSlash(t *testing.T) { t.Errorf("HandleRaw(\"/\") path = %q, want %q", slashed.Path, "/api/v1/nodes/") } if slashed.ID == slashless.ID { - t.Error("HandleRaw must allocate its own route ID") + t.Error("HandleRaw must allocate its own route ID so scoped teardown can unregister both") } if got := len(r.Routes()); got != 2 { t.Errorf("registry routes = %d, want 2", got) } + if !r.UnregisterByID(slashed.ID) { + t.Error("UnregisterByID(HandleRaw route) = false, want true") + } + if got := len(r.Routes()); got != 1 { + t.Errorf("routes after unregister = %d, want 1", got) + } +} + +// TestRegistryHandleRawKeepsAbsolutePath 根注册表上 HandleRaw 只做绝对化处理。 +func TestRegistryHandleRawKeepsAbsolutePath(t *testing.T) { + r := &RouterRegistry{} + if got := r.HandleRaw("GET", "/health/").Path; got != "/health/" { + t.Errorf("path = %q, want %q", got, "/health/") + } + if got := r.HandleRaw("POST", "submit").Path; got != "/submit" { + t.Errorf("path = %q, want %q", got, "/submit") + } + if got := r.BasePath(); got != "" { + t.Errorf("registry BasePath() = %q, want empty", got) + } } diff --git a/backend/core/extpoints/task.go b/backend/core/extpoints/task.go index 1217ed64..3e22b1e3 100644 --- a/backend/core/extpoints/task.go +++ b/backend/core/extpoints/task.go @@ -5,6 +5,8 @@ package extpoints import ( "Wavelet/core/contracts" + "fmt" + "reflect" "sync" "time" ) @@ -230,10 +232,15 @@ func NewTaskRegistry() *TaskRegistry { } // Register registers a task pattern and its handler with optional configuration. +// A nil handler panics. A non-empty Type that is already used by another pattern panics. func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) { t.mu.Lock() defer t.mu.Unlock() + if isNilTaskHandler(handler) { + panic(fmt.Sprintf("extpoints: nil handler for task pattern %q", pattern)) + } + td := TaskDefinition{ Pattern: pattern, Handler: handler, @@ -245,8 +252,23 @@ func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) opt(&td) } } + if td.Type == "" { + td.Type = pattern + } - if _, exists := t.lookup[pattern]; exists { + for _, item := range t.tasks { + if item.Pattern == pattern { + continue + } + if item.Type == td.Type { + panic(fmt.Sprintf("extpoints: duplicate task type %q (patterns %q and %q)", td.Type, item.Pattern, pattern)) + } + } + + if existing, exists := t.lookup[pattern]; exists { + if existing.Type != "" && existing.Type != pattern { + delete(t.lookup, existing.Type) + } for i, item := range t.tasks { if item.Pattern == pattern { t.tasks[i] = td @@ -258,13 +280,44 @@ func (t *TaskRegistry) Register(pattern string, handler any, opts ...TaskOption) } t.lookup[pattern] = td + if td.Type != pattern { + t.lookup[td.Type] = td + } +} + +func isNilTaskHandler(handler any) bool { + if handler == nil { + return true + } + v := reflect.ValueOf(handler) + switch v.Kind() { + case reflect.Chan, reflect.Func, reflect.Map, reflect.Pointer, reflect.UnsafePointer, reflect.Interface, reflect.Slice: + return v.IsNil() + default: + return false + } } // Unregister removes a registered task definition by its pattern. func (t *TaskRegistry) Unregister(pattern string) bool { - return unregisterEntry(&t.mu, t.lookup, &t.tasks, pattern, func(item TaskDefinition) bool { - return item.Pattern == pattern - }) + t.mu.Lock() + defer t.mu.Unlock() + td, ok := t.lookup[pattern] + if !ok { + return false + } + delete(t.lookup, td.Pattern) + if td.Type != "" && td.Type != td.Pattern { + delete(t.lookup, td.Type) + } + filtered := t.tasks[:0] + for _, item := range t.tasks { + if item.Pattern != td.Pattern { + filtered = append(filtered, item) + } + } + t.tasks = filtered + return true } // Tasks returns a copy of all registered TaskDefinitions. @@ -276,7 +329,7 @@ func (t *TaskRegistry) Tasks() []TaskDefinition { return res } -// Get retrieves a task definition by its pattern. +// Get retrieves a task definition by its pattern or admin type identifier. func (t *TaskRegistry) Get(pattern string) (TaskDefinition, bool) { t.mu.RLock() defer t.mu.RUnlock() diff --git a/backend/core/scoped_extpoints.go b/backend/core/scoped_extpoints.go index f567544a..812709d6 100644 --- a/backend/core/scoped_extpoints.go +++ b/backend/core/scoped_extpoints.go @@ -22,6 +22,19 @@ func newScopedRouterExtension(ctx *Context, underlying extpoints.RouterExtension } func (s *scopedRouterExtension) Use(middlewares ...any) { + if len(middlewares) == 0 { + return + } + if reg, ok := s.underlying.(*extpoints.RouterRegistry); ok { + ids := reg.UseWithID(middlewares...) + s.ctx.OnDispose(func() error { + for _, id := range ids { + reg.UnregisterMiddlewareByID(id) + } + return nil + }) + return + } s.underlying.Use(middlewares...) } @@ -107,8 +120,22 @@ func (s *scopedRouterExtension) UnregisterByID(id uint64) bool { return s.underlying.UnregisterByID(id) } +func (s *scopedRouterExtension) UnregisterMiddlewareByID(id uint64) bool { + return s.underlying.UnregisterMiddlewareByID(id) +} + func (s *scopedRouterExtension) RegisterWhitelist(patterns ...string) { s.underlying.RegisterWhitelist(patterns...) + if len(patterns) > 0 { + s.ctx.OnDispose(func() error { + s.underlying.UnregisterWhitelist(patterns...) + return nil + }) + } +} + +func (s *scopedRouterExtension) UnregisterWhitelist(patterns ...string) { + s.underlying.UnregisterWhitelist(patterns...) } func (s *scopedRouterExtension) Whitelist() []string { diff --git a/backend/docs/docs.go b/backend/docs/docs.go index e90035dd..3a69627f 100644 --- a/backend/docs/docs.go +++ b/backend/docs/docs.go @@ -1023,7 +1023,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.ChannelDTO" + "$ref": "#/definitions/do.ChannelDTO" } } } @@ -1057,7 +1057,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.CreateChannelRequest" + "$ref": "#/definitions/do.CreateChannelRequest" } } ], @@ -1073,7 +1073,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.ChannelDTO" + "$ref": "#/definitions/do.ChannelDTO" } } } @@ -1118,7 +1118,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.Definition" + "$ref": "#/definitions/do.Definition" } } } @@ -1199,7 +1199,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.UpdateChannelRequest" + "$ref": "#/definitions/do.UpdateChannelRequest" } } ], @@ -1215,7 +1215,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.ChannelDTO" + "$ref": "#/definitions/do.ChannelDTO" } } } @@ -1312,7 +1312,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.PushChannel" + "$ref": "#/definitions/entity.PushChannel" } } } @@ -1346,7 +1346,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.CreatePushChannelRequest" + "$ref": "#/definitions/do.CreatePushChannelRequest" } } ], @@ -1362,7 +1362,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.PushChannel" + "$ref": "#/definitions/entity.PushChannel" } } } @@ -1422,7 +1422,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.TestPushChannelRequest" + "$ref": "#/definitions/do.TestPushChannelRequest" } } ], @@ -1469,7 +1469,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.UpdatePushChannelRequest" + "$ref": "#/definitions/do.UpdatePushChannelRequest" } } ], @@ -1485,7 +1485,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.PushChannel" + "$ref": "#/definitions/entity.PushChannel" } } } @@ -1557,7 +1557,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.PushEvent" + "$ref": "#/definitions/entity.PushEvent" } } } @@ -1591,7 +1591,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.CreatePushEventRequest" + "$ref": "#/definitions/do.CreatePushEventRequest" } } ], @@ -1607,7 +1607,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.PushEvent" + "$ref": "#/definitions/entity.PushEvent" } } } @@ -1674,7 +1674,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.UpdatePushEventRequest" + "$ref": "#/definitions/do.UpdatePushEventRequest" } } ], @@ -1866,7 +1866,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.TestPushRequest" + "$ref": "#/definitions/do.TestPushRequest" } } ], @@ -2675,7 +2675,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Schedule" + "$ref": "#/definitions/model.Schedule" } } } @@ -2737,7 +2737,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Schedule" + "$ref": "#/definitions/model.Schedule" } } } @@ -2819,7 +2819,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Schedule" + "$ref": "#/definitions/model.Schedule" } } } @@ -3008,7 +3008,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Template" + "$ref": "#/definitions/model.Template" } } } @@ -3146,7 +3146,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Template" + "$ref": "#/definitions/model.Template" } } } @@ -3226,7 +3226,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Template" + "$ref": "#/definitions/model.Template" } } } @@ -4928,7 +4928,7 @@ const docTemplate = `{ "name": "request", "in": "body", "schema": { - "$ref": "#/definitions/cap.challengeRequest" + "$ref": "#/definitions/dto.ChallengeRequest" } } ], @@ -4944,7 +4944,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/cap.ChallengeResponse" + "$ref": "#/definitions/dto.ChallengeResponse" } } } @@ -4977,7 +4977,7 @@ const docTemplate = `{ "name": "request", "in": "body", "schema": { - "$ref": "#/definitions/cap.challengeRequest" + "$ref": "#/definitions/dto.ChallengeRequest" } } ], @@ -4993,7 +4993,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/cap.ChallengeResponse" + "$ref": "#/definitions/dto.ChallengeResponse" } } } @@ -5029,7 +5029,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/cap.redeemRequest" + "$ref": "#/definitions/dto.RedeemRequest" } } ], @@ -5045,7 +5045,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/cap.RedeemResponse" + "$ref": "#/definitions/dto.RedeemResponse" } } } @@ -5069,7 +5069,7 @@ const docTemplate = `{ }, "/api/v1/config/public": { "get": { - "description": "返回系统配置表中 visibility 为 1 的配置键值集合", + "description": "返回系统配置表中 visibility 为 1 的扁平键值集合(如 cap_login_enabled)", "consumes": [ "application/json" ], @@ -13222,7 +13222,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.BindingDTO" + "$ref": "#/definitions/do.BindingDTO" } } } @@ -13262,7 +13262,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.BindRequest" + "$ref": "#/definitions/do.BindRequest" } } ], @@ -13278,7 +13278,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.BindingDTO" + "$ref": "#/definitions/do.BindingDTO" } } } @@ -13375,7 +13375,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.PublicChannelDTO" + "$ref": "#/definitions/do.PublicChannelDTO" } } } @@ -13412,7 +13412,7 @@ const docTemplate = `{ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/auth.CallbackRequest" + "$ref": "#/definitions/dto.CallbackRequest" } } ], @@ -13428,7 +13428,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.OAuthCallbackResult" + "$ref": "#/definitions/dto.OAuthCallbackResult" } } } @@ -13582,7 +13582,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.OAuthAuthorizeResponse" + "$ref": "#/definitions/dto.OAuthAuthorizeResponse" } } } @@ -13671,7 +13671,7 @@ const docTemplate = `{ "data": { "type": "array", "items": { - "$ref": "#/definitions/auth.AuthSourceView" + "$ref": "#/definitions/dto.AuthSourceView" } } } @@ -13709,7 +13709,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.BasicUserInfo" + "$ref": "#/definitions/dto.BasicUserInfo" } } } @@ -13762,7 +13762,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.OAuthAuthorizeResponse" + "$ref": "#/definitions/dto.OAuthAuthorizeResponse" } } } @@ -14397,7 +14397,7 @@ const docTemplate = `{ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.BasicUserInfo" + "$ref": "#/definitions/dto.BasicUserInfo" } } } @@ -14654,7 +14654,7 @@ const docTemplate = `{ }, "/api/v1/user/login": { "post": { - "description": "使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。", + "description": "使用用户名和密码登录系统,验证通过后建立 Session 并返回用户信息。", "consumes": [ "application/json" ], @@ -14664,7 +14664,7 @@ const docTemplate = `{ "tags": [ "user" ], - "summary": "用户密码登录", + "summary": "用户登录", "parameters": [ { "description": "登录请求参数", @@ -14689,6 +14689,12 @@ const docTemplate = `{ "$ref": "#/definitions/response.Any" } }, + "429": { + "description": "登录尝试过于频繁", + "schema": { + "$ref": "#/definitions/response.Any" + } + }, "500": { "description": "服务内部错误", "schema": { @@ -14824,6 +14830,12 @@ const docTemplate = `{ "$ref": "#/definitions/response.Any" } }, + "429": { + "description": "注册尝试过于频繁", + "schema": { + "$ref": "#/definitions/response.Any" + } + }, "500": { "description": "服务内部错误", "schema": { @@ -14877,6 +14889,17 @@ const docTemplate = `{ "user" ], "summary": "发送邮箱验证码", + "parameters": [ + { + "description": "目标邮箱", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/user.sendEmailCodeRequest" + } + } + ], "responses": { "200": { "description": "发送成功", @@ -14889,6 +14912,12 @@ const docTemplate = `{ "schema": { "$ref": "#/definitions/response.Any" } + }, + "500": { + "description": "发送失败", + "schema": { + "$ref": "#/definitions/response.Any" + } } } } @@ -15211,36 +15240,6 @@ const docTemplate = `{ } } }, - "Wavelet_plugins_domain_admin_model.Schedule": { - "type": "object", - "properties": { - "created_at": { - "type": "string" - }, - "cron": { - "type": "string" - }, - "id": { - "type": "string", - "example": "0" - }, - "is_active": { - "type": "boolean" - }, - "name": { - "type": "string" - }, - "payload": { - "type": "string" - }, - "task_type": { - "type": "string" - }, - "updated_at": { - "type": "string" - } - } - }, "Wavelet_plugins_domain_admin_model.SystemConfig": { "type": "object", "properties": { @@ -15327,41 +15326,6 @@ const docTemplate = `{ } } }, - "Wavelet_plugins_domain_admin_model.Template": { - "type": "object", - "properties": { - "content": { - "type": "string" - }, - "created_at": { - "type": "string" - }, - "description": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "is_system": { - "type": "boolean" - }, - "key": { - "type": "string" - }, - "name": { - "type": "string" - }, - "subject": { - "type": "string" - }, - "type": { - "type": "string" - }, - "updated_at": { - "type": "string" - } - } - }, "agent.ActiveConfigMeta": { "type": "object", "properties": { @@ -15669,179 +15633,6 @@ const docTemplate = `{ } } }, - "auth.AuthSourceView": { - "type": "object", - "properties": { - "client_secret_configured": { - "type": "boolean" - }, - "display_name": { - "type": "string" - }, - "icon_url": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "is_active": { - "type": "boolean" - }, - "name": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, - "auth.BasicUserInfo": { - "type": "object", - "properties": { - "avatar_url": { - "type": "string" - }, - "bio": { - "type": "string" - }, - "email": { - "type": "string" - }, - "gender": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "is_admin": { - "type": "boolean" - }, - "location": { - "type": "string" - }, - "need_change_password": { - "type": "boolean" - }, - "nickname": { - "type": "string" - }, - "phone": { - "type": "string" - }, - "username": { - "type": "string" - }, - "website": { - "type": "string" - } - } - }, - "auth.CallbackRequest": { - "type": "object", - "required": [ - "code", - "state" - ], - "properties": { - "code": { - "type": "string" - }, - "state": { - "type": "string" - } - } - }, - "auth.OAuthAuthorizeResponse": { - "type": "object", - "properties": { - "authorize_url": { - "type": "string" - } - } - }, - "auth.OAuthCallbackResult": { - "type": "object", - "properties": { - "status": { - "type": "string" - }, - "user": { - "$ref": "#/definitions/auth.BasicUserInfo" - } - } - }, - "cap.ChallengeResponse": { - "type": "object", - "properties": { - "challenge": { - "type": "object", - "properties": { - "c": { - "type": "integer" - }, - "d": { - "type": "integer" - }, - "s": { - "type": "integer" - } - } - }, - "expires": { - "description": "ms timestamp", - "type": "integer" - }, - "token": { - "type": "string" - } - } - }, - "cap.RedeemResponse": { - "type": "object", - "properties": { - "error": { - "type": "string" - }, - "expires": { - "type": "integer" - }, - "success": { - "type": "boolean" - }, - "token": { - "type": "string" - } - } - }, - "cap.challengeRequest": { - "type": "object", - "properties": { - "scope": { - "type": "string" - } - } - }, - "cap.redeemRequest": { - "type": "object", - "required": [ - "solutions", - "token" - ], - "properties": { - "scope": { - "type": "string" - }, - "solutions": { - "type": "array", - "items": { - "type": "integer" - } - }, - "token": { - "type": "string" - } - } - }, "cloudflare.AvailableDomain": { "type": "object", "properties": { @@ -16491,6 +16282,573 @@ const docTemplate = `{ } } }, + "do.BindRequest": { + "type": "object", + "properties": { + "channel_id": { + "type": "string" + }, + "code": { + "type": "string" + } + } + }, + "do.BindingDTO": { + "type": "object", + "properties": { + "channel_id": { + "type": "string", + "example": "0" + }, + "channel_name": { + "type": "string" + }, + "channel_type": { + "type": "string" + }, + "created_at": { + "type": "string" + }, + "id": { + "type": "string", + "example": "0" + }, + "platform_user_id": { + "type": "string" + }, + "user_id": { + "type": "string", + "example": "0" + } + } + }, + "do.ChannelDTO": { + "type": "object", + "properties": { + "credentials": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "extra": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "id": { + "type": "string", + "example": "0" + }, + "name": { + "type": "string" + }, + "owner_id": { + "type": "string", + "example": "0" + }, + "owner_scope": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "do.CreateChannelRequest": { + "type": "object", + "properties": { + "credentials": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "extra": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "name": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "do.CreatePushChannelRequest": { + "type": "object", + "required": [ + "name", + "type" + ], + "properties": { + "description": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "name": { + "type": "string" + }, + "other": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "do.CreatePushEventRequest": { + "type": "object", + "properties": { + "channels": { + "type": "array", + "items": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "event_key": { + "type": "string" + }, + "targets": { + "type": "array", + "items": { + "type": "string" + } + }, + "task_type": { + "type": "string" + }, + "template": { + "type": "string" + } + } + }, + "do.Definition": { + "type": "object", + "properties": { + "fields": { + "type": "array", + "items": { + "$ref": "#/definitions/do.Field" + } + }, + "type": { + "type": "string" + } + } + }, + "do.Field": { + "type": "object", + "properties": { + "key": { + "type": "string" + }, + "required": { + "type": "boolean" + }, + "type": { + "type": "string" + } + } + }, + "do.PublicChannelDTO": { + "type": "object", + "properties": { + "id": { + "type": "string", + "example": "0" + }, + "name": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "do.TestPushChannelRequest": { + "type": "object", + "properties": { + "name": { + "type": "string" + }, + "other": { + "type": "string" + }, + "target": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "do.TestPushRequest": { + "type": "object", + "required": [ + "config" + ], + "properties": { + "config": { + "$ref": "#/definitions/push.Config" + }, + "target": { + "type": "string" + } + } + }, + "do.UpdateChannelRequest": { + "type": "object", + "properties": { + "credentials": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "extra": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "name": { + "type": "string" + } + } + }, + "do.UpdatePushChannelRequest": { + "type": "object", + "required": [ + "type" + ], + "properties": { + "description": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "other": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "do.UpdatePushEventRequest": { + "type": "object", + "required": [ + "template" + ], + "properties": { + "channels": { + "type": "array", + "items": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "targets": { + "type": "array", + "items": { + "type": "string" + } + }, + "template": { + "type": "string" + } + } + }, + "dto.AuthSourceView": { + "type": "object", + "properties": { + "client_secret_configured": { + "type": "boolean" + }, + "display_name": { + "type": "string" + }, + "icon_url": { + "type": "string" + }, + "id": { + "type": "integer" + }, + "is_active": { + "type": "boolean" + }, + "name": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "dto.BasicUserInfo": { + "type": "object", + "properties": { + "avatar_url": { + "type": "string" + }, + "bio": { + "type": "string" + }, + "email": { + "type": "string" + }, + "gender": { + "type": "string" + }, + "id": { + "type": "string", + "example": "0" + }, + "is_admin": { + "type": "boolean" + }, + "location": { + "type": "string" + }, + "need_change_password": { + "type": "boolean" + }, + "nickname": { + "type": "string" + }, + "phone": { + "type": "string" + }, + "username": { + "type": "string" + }, + "website": { + "type": "string" + } + } + }, + "dto.CallbackRequest": { + "type": "object", + "required": [ + "code", + "state" + ], + "properties": { + "code": { + "type": "string" + }, + "state": { + "type": "string" + } + } + }, + "dto.ChallengeRequest": { + "type": "object", + "properties": { + "scope": { + "type": "string" + } + } + }, + "dto.ChallengeResponse": { + "type": "object", + "properties": { + "challenge": { + "type": "object", + "properties": { + "c": { + "type": "integer" + }, + "d": { + "type": "integer" + }, + "s": { + "type": "integer" + } + } + }, + "expires": { + "description": "ms timestamp", + "type": "integer" + }, + "token": { + "type": "string" + } + } + }, + "dto.OAuthAuthorizeResponse": { + "type": "object", + "properties": { + "authorize_url": { + "type": "string" + } + } + }, + "dto.OAuthCallbackResult": { + "type": "object", + "properties": { + "status": { + "type": "string" + }, + "user": { + "$ref": "#/definitions/dto.BasicUserInfo" + } + } + }, + "dto.RedeemRequest": { + "type": "object", + "required": [ + "solutions", + "token" + ], + "properties": { + "scope": { + "type": "string" + }, + "solutions": { + "type": "array", + "items": { + "type": "integer" + } + }, + "token": { + "type": "string" + } + } + }, + "dto.RedeemResponse": { + "type": "object", + "properties": { + "error": { + "type": "string" + }, + "expires": { + "type": "integer" + }, + "success": { + "type": "boolean" + }, + "token": { + "type": "string" + } + } + }, + "entity.PushChannel": { + "type": "object", + "properties": { + "created_at": { + "type": "string" + }, + "description": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "id": { + "type": "integer" + }, + "name": { + "type": "string" + }, + "other": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "updated_at": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "entity.PushEvent": { + "type": "object", + "properties": { + "channels": { + "type": "array", + "items": { + "type": "string" + } + }, + "created_at": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "event_key": { + "type": "string" + }, + "id": { + "type": "integer" + }, + "name": { + "type": "string" + }, + "targets": { + "type": "array", + "items": { + "type": "string" + } + }, + "task_type": { + "type": "string" + }, + "template": { + "type": "string" + }, + "updated_at": { + "type": "string" + } + } + }, "flared.ApplyLogPayload": { "type": "object", "properties": { @@ -16804,46 +17162,6 @@ const docTemplate = `{ } } }, - "model.BindRequest": { - "type": "object", - "properties": { - "channel_id": { - "type": "string" - }, - "code": { - "type": "string" - } - } - }, - "model.BindingDTO": { - "type": "object", - "properties": { - "channel_id": { - "type": "string", - "example": "0" - }, - "channel_name": { - "type": "string" - }, - "channel_type": { - "type": "string" - }, - "created_at": { - "type": "string" - }, - "id": { - "type": "string", - "example": "0" - }, - "platform_user_id": { - "type": "string" - }, - "user_id": { - "type": "string", - "example": "0" - } - } - }, "model.BrowserItem": { "type": "object", "properties": { @@ -16855,43 +17173,6 @@ const docTemplate = `{ } } }, - "model.ChannelDTO": { - "type": "object", - "properties": { - "credentials": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "extra": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "id": { - "type": "string", - "example": "0" - }, - "name": { - "type": "string" - }, - "owner_id": { - "type": "string", - "example": "0" - }, - "owner_scope": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, "model.ConfigVersion": { "type": "object", "properties": { @@ -16950,91 +17231,6 @@ const docTemplate = `{ } } }, - "model.CreateChannelRequest": { - "type": "object", - "properties": { - "credentials": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "extra": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "name": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, - "model.CreatePushChannelRequest": { - "type": "object", - "required": [ - "name", - "type" - ], - "properties": { - "description": { - "type": "string" - }, - "enabled": { - "type": "boolean" - }, - "name": { - "type": "string" - }, - "other": { - "type": "string" - }, - "token": { - "type": "string" - }, - "type": { - "type": "string" - }, - "url": { - "type": "string" - } - } - }, - "model.CreatePushEventRequest": { - "type": "object", - "properties": { - "channels": { - "type": "array", - "items": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "event_key": { - "type": "string" - }, - "targets": { - "type": "array", - "items": { - "type": "string" - } - }, - "task_type": { - "type": "string" - }, - "template": { - "type": "string" - } - } - }, "model.CreateScheduleRequest": { "type": "object", "required": [ @@ -17221,20 +17417,6 @@ const docTemplate = `{ } } }, - "model.Definition": { - "type": "object", - "properties": { - "fields": { - "type": "array", - "items": { - "$ref": "#/definitions/model.Field" - } - }, - "type": { - "type": "string" - } - } - }, "model.DispatchTaskRequest": { "type": "object", "required": [ @@ -17297,20 +17479,6 @@ const docTemplate = `{ } } }, - "model.Field": { - "type": "object", - "properties": { - "key": { - "type": "string" - }, - "required": { - "type": "boolean" - }, - "type": { - "type": "string" - } - } - }, "model.ListUsersResponse": { "type": "object", "properties": { @@ -17633,92 +17801,31 @@ const docTemplate = `{ } } }, - "model.PublicChannelDTO": { + "model.Schedule": { "type": "object", "properties": { + "created_at": { + "type": "string" + }, + "cron": { + "type": "string" + }, "id": { "type": "string", "example": "0" }, - "name": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, - "model.PushChannel": { - "type": "object", - "properties": { - "created_at": { - "type": "string" - }, - "description": { - "type": "string" - }, - "enabled": { + "is_active": { "type": "boolean" }, - "id": { - "type": "integer" - }, "name": { "type": "string" }, - "other": { + "payload": { "type": "string" }, - "token": { - "type": "string" - }, - "type": { - "type": "string" - }, - "updated_at": { - "type": "string" - }, - "url": { - "type": "string" - } - } - }, - "model.PushEvent": { - "type": "object", - "properties": { - "channels": { - "type": "array", - "items": { - "type": "string" - } - }, - "created_at": { - "type": "string" - }, - "enabled": { - "type": "boolean" - }, - "event_key": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "name": { - "type": "string" - }, - "targets": { - "type": "array", - "items": { - "type": "string" - } - }, "task_type": { "type": "string" }, - "template": { - "type": "string" - }, "updated_at": { "type": "string" } @@ -17893,39 +18000,37 @@ const docTemplate = `{ "TaskExecutionStatusFailed" ] }, - "model.TestPushChannelRequest": { + "model.Template": { "type": "object", "properties": { + "content": { + "type": "string" + }, + "created_at": { + "type": "string" + }, + "description": { + "type": "string" + }, + "id": { + "type": "integer" + }, + "is_system": { + "type": "boolean" + }, + "key": { + "type": "string" + }, "name": { "type": "string" }, - "other": { - "type": "string" - }, - "target": { - "type": "string" - }, - "token": { + "subject": { "type": "string" }, "type": { "type": "string" }, - "url": { - "type": "string" - } - } - }, - "model.TestPushRequest": { - "type": "object", - "required": [ - "config" - ], - "properties": { - "config": { - "$ref": "#/definitions/push.Config" - }, - "target": { + "updated_at": { "type": "string" } } @@ -18023,81 +18128,6 @@ const docTemplate = `{ } } }, - "model.UpdateChannelRequest": { - "type": "object", - "properties": { - "credentials": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "extra": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "name": { - "type": "string" - } - } - }, - "model.UpdatePushChannelRequest": { - "type": "object", - "required": [ - "type" - ], - "properties": { - "description": { - "type": "string" - }, - "enabled": { - "type": "boolean" - }, - "other": { - "type": "string" - }, - "token": { - "type": "string" - }, - "type": { - "type": "string" - }, - "url": { - "type": "string" - } - } - }, - "model.UpdatePushEventRequest": { - "type": "object", - "required": [ - "template" - ], - "properties": { - "channels": { - "type": "array", - "items": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "targets": { - "type": "array", - "items": { - "type": "string" - } - }, - "template": { - "type": "string" - } - } - }, "model.UpdateScheduleRequest": { "type": "object", "required": [ @@ -20524,6 +20554,10 @@ const docTemplate = `{ "description": "AppID 或 SMTP 用户名", "type": "string" }, + "other": { + "description": "附加配置 (如 ChatID / UserKey / 扩展 JSON)", + "type": "string" + }, "secret": { "description": "签名密钥或 SMTP 密码/Token", "type": "string" @@ -20892,6 +20926,17 @@ const docTemplate = `{ } } }, + "user.sendEmailCodeRequest": { + "type": "object", + "required": [ + "email" + ], + "properties": { + "email": { + "type": "string" + } + } + }, "user.updateProfileRequest": { "type": "object", "properties": { diff --git a/backend/docs/swagger.json b/backend/docs/swagger.json index 3bc39020..dafc71ac 100644 --- a/backend/docs/swagger.json +++ b/backend/docs/swagger.json @@ -1016,7 +1016,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.ChannelDTO" + "$ref": "#/definitions/do.ChannelDTO" } } } @@ -1050,7 +1050,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.CreateChannelRequest" + "$ref": "#/definitions/do.CreateChannelRequest" } } ], @@ -1066,7 +1066,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.ChannelDTO" + "$ref": "#/definitions/do.ChannelDTO" } } } @@ -1111,7 +1111,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.Definition" + "$ref": "#/definitions/do.Definition" } } } @@ -1192,7 +1192,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.UpdateChannelRequest" + "$ref": "#/definitions/do.UpdateChannelRequest" } } ], @@ -1208,7 +1208,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.ChannelDTO" + "$ref": "#/definitions/do.ChannelDTO" } } } @@ -1305,7 +1305,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.PushChannel" + "$ref": "#/definitions/entity.PushChannel" } } } @@ -1339,7 +1339,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.CreatePushChannelRequest" + "$ref": "#/definitions/do.CreatePushChannelRequest" } } ], @@ -1355,7 +1355,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.PushChannel" + "$ref": "#/definitions/entity.PushChannel" } } } @@ -1415,7 +1415,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.TestPushChannelRequest" + "$ref": "#/definitions/do.TestPushChannelRequest" } } ], @@ -1462,7 +1462,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.UpdatePushChannelRequest" + "$ref": "#/definitions/do.UpdatePushChannelRequest" } } ], @@ -1478,7 +1478,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.PushChannel" + "$ref": "#/definitions/entity.PushChannel" } } } @@ -1550,7 +1550,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.PushEvent" + "$ref": "#/definitions/entity.PushEvent" } } } @@ -1584,7 +1584,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.CreatePushEventRequest" + "$ref": "#/definitions/do.CreatePushEventRequest" } } ], @@ -1600,7 +1600,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.PushEvent" + "$ref": "#/definitions/entity.PushEvent" } } } @@ -1667,7 +1667,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.UpdatePushEventRequest" + "$ref": "#/definitions/do.UpdatePushEventRequest" } } ], @@ -1859,7 +1859,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.TestPushRequest" + "$ref": "#/definitions/do.TestPushRequest" } } ], @@ -2668,7 +2668,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Schedule" + "$ref": "#/definitions/model.Schedule" } } } @@ -2730,7 +2730,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Schedule" + "$ref": "#/definitions/model.Schedule" } } } @@ -2812,7 +2812,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Schedule" + "$ref": "#/definitions/model.Schedule" } } } @@ -3001,7 +3001,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Template" + "$ref": "#/definitions/model.Template" } } } @@ -3139,7 +3139,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Template" + "$ref": "#/definitions/model.Template" } } } @@ -3219,7 +3219,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Template" + "$ref": "#/definitions/model.Template" } } } @@ -4921,7 +4921,7 @@ "name": "request", "in": "body", "schema": { - "$ref": "#/definitions/cap.challengeRequest" + "$ref": "#/definitions/dto.ChallengeRequest" } } ], @@ -4937,7 +4937,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/cap.ChallengeResponse" + "$ref": "#/definitions/dto.ChallengeResponse" } } } @@ -4970,7 +4970,7 @@ "name": "request", "in": "body", "schema": { - "$ref": "#/definitions/cap.challengeRequest" + "$ref": "#/definitions/dto.ChallengeRequest" } } ], @@ -4986,7 +4986,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/cap.ChallengeResponse" + "$ref": "#/definitions/dto.ChallengeResponse" } } } @@ -5022,7 +5022,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/cap.redeemRequest" + "$ref": "#/definitions/dto.RedeemRequest" } } ], @@ -5038,7 +5038,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/cap.RedeemResponse" + "$ref": "#/definitions/dto.RedeemResponse" } } } @@ -5062,7 +5062,7 @@ }, "/api/v1/config/public": { "get": { - "description": "返回系统配置表中 visibility 为 1 的配置键值集合", + "description": "返回系统配置表中 visibility 为 1 的扁平键值集合(如 cap_login_enabled)", "consumes": [ "application/json" ], @@ -13215,7 +13215,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.BindingDTO" + "$ref": "#/definitions/do.BindingDTO" } } } @@ -13255,7 +13255,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.BindRequest" + "$ref": "#/definitions/do.BindRequest" } } ], @@ -13271,7 +13271,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.BindingDTO" + "$ref": "#/definitions/do.BindingDTO" } } } @@ -13368,7 +13368,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.PublicChannelDTO" + "$ref": "#/definitions/do.PublicChannelDTO" } } } @@ -13405,7 +13405,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/auth.CallbackRequest" + "$ref": "#/definitions/dto.CallbackRequest" } } ], @@ -13421,7 +13421,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.OAuthCallbackResult" + "$ref": "#/definitions/dto.OAuthCallbackResult" } } } @@ -13575,7 +13575,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.OAuthAuthorizeResponse" + "$ref": "#/definitions/dto.OAuthAuthorizeResponse" } } } @@ -13664,7 +13664,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/auth.AuthSourceView" + "$ref": "#/definitions/dto.AuthSourceView" } } } @@ -13702,7 +13702,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.BasicUserInfo" + "$ref": "#/definitions/dto.BasicUserInfo" } } } @@ -13755,7 +13755,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.OAuthAuthorizeResponse" + "$ref": "#/definitions/dto.OAuthAuthorizeResponse" } } } @@ -14390,7 +14390,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.BasicUserInfo" + "$ref": "#/definitions/dto.BasicUserInfo" } } } @@ -14647,7 +14647,7 @@ }, "/api/v1/user/login": { "post": { - "description": "使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。", + "description": "使用用户名和密码登录系统,验证通过后建立 Session 并返回用户信息。", "consumes": [ "application/json" ], @@ -14657,7 +14657,7 @@ "tags": [ "user" ], - "summary": "用户密码登录", + "summary": "用户登录", "parameters": [ { "description": "登录请求参数", @@ -14682,6 +14682,12 @@ "$ref": "#/definitions/response.Any" } }, + "429": { + "description": "登录尝试过于频繁", + "schema": { + "$ref": "#/definitions/response.Any" + } + }, "500": { "description": "服务内部错误", "schema": { @@ -14817,6 +14823,12 @@ "$ref": "#/definitions/response.Any" } }, + "429": { + "description": "注册尝试过于频繁", + "schema": { + "$ref": "#/definitions/response.Any" + } + }, "500": { "description": "服务内部错误", "schema": { @@ -14870,6 +14882,17 @@ "user" ], "summary": "发送邮箱验证码", + "parameters": [ + { + "description": "目标邮箱", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/user.sendEmailCodeRequest" + } + } + ], "responses": { "200": { "description": "发送成功", @@ -14882,6 +14905,12 @@ "schema": { "$ref": "#/definitions/response.Any" } + }, + "500": { + "description": "发送失败", + "schema": { + "$ref": "#/definitions/response.Any" + } } } } @@ -15204,36 +15233,6 @@ } } }, - "Wavelet_plugins_domain_admin_model.Schedule": { - "type": "object", - "properties": { - "created_at": { - "type": "string" - }, - "cron": { - "type": "string" - }, - "id": { - "type": "string", - "example": "0" - }, - "is_active": { - "type": "boolean" - }, - "name": { - "type": "string" - }, - "payload": { - "type": "string" - }, - "task_type": { - "type": "string" - }, - "updated_at": { - "type": "string" - } - } - }, "Wavelet_plugins_domain_admin_model.SystemConfig": { "type": "object", "properties": { @@ -15320,41 +15319,6 @@ } } }, - "Wavelet_plugins_domain_admin_model.Template": { - "type": "object", - "properties": { - "content": { - "type": "string" - }, - "created_at": { - "type": "string" - }, - "description": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "is_system": { - "type": "boolean" - }, - "key": { - "type": "string" - }, - "name": { - "type": "string" - }, - "subject": { - "type": "string" - }, - "type": { - "type": "string" - }, - "updated_at": { - "type": "string" - } - } - }, "agent.ActiveConfigMeta": { "type": "object", "properties": { @@ -15662,179 +15626,6 @@ } } }, - "auth.AuthSourceView": { - "type": "object", - "properties": { - "client_secret_configured": { - "type": "boolean" - }, - "display_name": { - "type": "string" - }, - "icon_url": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "is_active": { - "type": "boolean" - }, - "name": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, - "auth.BasicUserInfo": { - "type": "object", - "properties": { - "avatar_url": { - "type": "string" - }, - "bio": { - "type": "string" - }, - "email": { - "type": "string" - }, - "gender": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "is_admin": { - "type": "boolean" - }, - "location": { - "type": "string" - }, - "need_change_password": { - "type": "boolean" - }, - "nickname": { - "type": "string" - }, - "phone": { - "type": "string" - }, - "username": { - "type": "string" - }, - "website": { - "type": "string" - } - } - }, - "auth.CallbackRequest": { - "type": "object", - "required": [ - "code", - "state" - ], - "properties": { - "code": { - "type": "string" - }, - "state": { - "type": "string" - } - } - }, - "auth.OAuthAuthorizeResponse": { - "type": "object", - "properties": { - "authorize_url": { - "type": "string" - } - } - }, - "auth.OAuthCallbackResult": { - "type": "object", - "properties": { - "status": { - "type": "string" - }, - "user": { - "$ref": "#/definitions/auth.BasicUserInfo" - } - } - }, - "cap.ChallengeResponse": { - "type": "object", - "properties": { - "challenge": { - "type": "object", - "properties": { - "c": { - "type": "integer" - }, - "d": { - "type": "integer" - }, - "s": { - "type": "integer" - } - } - }, - "expires": { - "description": "ms timestamp", - "type": "integer" - }, - "token": { - "type": "string" - } - } - }, - "cap.RedeemResponse": { - "type": "object", - "properties": { - "error": { - "type": "string" - }, - "expires": { - "type": "integer" - }, - "success": { - "type": "boolean" - }, - "token": { - "type": "string" - } - } - }, - "cap.challengeRequest": { - "type": "object", - "properties": { - "scope": { - "type": "string" - } - } - }, - "cap.redeemRequest": { - "type": "object", - "required": [ - "solutions", - "token" - ], - "properties": { - "scope": { - "type": "string" - }, - "solutions": { - "type": "array", - "items": { - "type": "integer" - } - }, - "token": { - "type": "string" - } - } - }, "cloudflare.AvailableDomain": { "type": "object", "properties": { @@ -16484,6 +16275,573 @@ } } }, + "do.BindRequest": { + "type": "object", + "properties": { + "channel_id": { + "type": "string" + }, + "code": { + "type": "string" + } + } + }, + "do.BindingDTO": { + "type": "object", + "properties": { + "channel_id": { + "type": "string", + "example": "0" + }, + "channel_name": { + "type": "string" + }, + "channel_type": { + "type": "string" + }, + "created_at": { + "type": "string" + }, + "id": { + "type": "string", + "example": "0" + }, + "platform_user_id": { + "type": "string" + }, + "user_id": { + "type": "string", + "example": "0" + } + } + }, + "do.ChannelDTO": { + "type": "object", + "properties": { + "credentials": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "extra": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "id": { + "type": "string", + "example": "0" + }, + "name": { + "type": "string" + }, + "owner_id": { + "type": "string", + "example": "0" + }, + "owner_scope": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "do.CreateChannelRequest": { + "type": "object", + "properties": { + "credentials": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "extra": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "name": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "do.CreatePushChannelRequest": { + "type": "object", + "required": [ + "name", + "type" + ], + "properties": { + "description": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "name": { + "type": "string" + }, + "other": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "do.CreatePushEventRequest": { + "type": "object", + "properties": { + "channels": { + "type": "array", + "items": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "event_key": { + "type": "string" + }, + "targets": { + "type": "array", + "items": { + "type": "string" + } + }, + "task_type": { + "type": "string" + }, + "template": { + "type": "string" + } + } + }, + "do.Definition": { + "type": "object", + "properties": { + "fields": { + "type": "array", + "items": { + "$ref": "#/definitions/do.Field" + } + }, + "type": { + "type": "string" + } + } + }, + "do.Field": { + "type": "object", + "properties": { + "key": { + "type": "string" + }, + "required": { + "type": "boolean" + }, + "type": { + "type": "string" + } + } + }, + "do.PublicChannelDTO": { + "type": "object", + "properties": { + "id": { + "type": "string", + "example": "0" + }, + "name": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "do.TestPushChannelRequest": { + "type": "object", + "properties": { + "name": { + "type": "string" + }, + "other": { + "type": "string" + }, + "target": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "do.TestPushRequest": { + "type": "object", + "required": [ + "config" + ], + "properties": { + "config": { + "$ref": "#/definitions/push.Config" + }, + "target": { + "type": "string" + } + } + }, + "do.UpdateChannelRequest": { + "type": "object", + "properties": { + "credentials": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "extra": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "name": { + "type": "string" + } + } + }, + "do.UpdatePushChannelRequest": { + "type": "object", + "required": [ + "type" + ], + "properties": { + "description": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "other": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "do.UpdatePushEventRequest": { + "type": "object", + "required": [ + "template" + ], + "properties": { + "channels": { + "type": "array", + "items": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "targets": { + "type": "array", + "items": { + "type": "string" + } + }, + "template": { + "type": "string" + } + } + }, + "dto.AuthSourceView": { + "type": "object", + "properties": { + "client_secret_configured": { + "type": "boolean" + }, + "display_name": { + "type": "string" + }, + "icon_url": { + "type": "string" + }, + "id": { + "type": "integer" + }, + "is_active": { + "type": "boolean" + }, + "name": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "dto.BasicUserInfo": { + "type": "object", + "properties": { + "avatar_url": { + "type": "string" + }, + "bio": { + "type": "string" + }, + "email": { + "type": "string" + }, + "gender": { + "type": "string" + }, + "id": { + "type": "string", + "example": "0" + }, + "is_admin": { + "type": "boolean" + }, + "location": { + "type": "string" + }, + "need_change_password": { + "type": "boolean" + }, + "nickname": { + "type": "string" + }, + "phone": { + "type": "string" + }, + "username": { + "type": "string" + }, + "website": { + "type": "string" + } + } + }, + "dto.CallbackRequest": { + "type": "object", + "required": [ + "code", + "state" + ], + "properties": { + "code": { + "type": "string" + }, + "state": { + "type": "string" + } + } + }, + "dto.ChallengeRequest": { + "type": "object", + "properties": { + "scope": { + "type": "string" + } + } + }, + "dto.ChallengeResponse": { + "type": "object", + "properties": { + "challenge": { + "type": "object", + "properties": { + "c": { + "type": "integer" + }, + "d": { + "type": "integer" + }, + "s": { + "type": "integer" + } + } + }, + "expires": { + "description": "ms timestamp", + "type": "integer" + }, + "token": { + "type": "string" + } + } + }, + "dto.OAuthAuthorizeResponse": { + "type": "object", + "properties": { + "authorize_url": { + "type": "string" + } + } + }, + "dto.OAuthCallbackResult": { + "type": "object", + "properties": { + "status": { + "type": "string" + }, + "user": { + "$ref": "#/definitions/dto.BasicUserInfo" + } + } + }, + "dto.RedeemRequest": { + "type": "object", + "required": [ + "solutions", + "token" + ], + "properties": { + "scope": { + "type": "string" + }, + "solutions": { + "type": "array", + "items": { + "type": "integer" + } + }, + "token": { + "type": "string" + } + } + }, + "dto.RedeemResponse": { + "type": "object", + "properties": { + "error": { + "type": "string" + }, + "expires": { + "type": "integer" + }, + "success": { + "type": "boolean" + }, + "token": { + "type": "string" + } + } + }, + "entity.PushChannel": { + "type": "object", + "properties": { + "created_at": { + "type": "string" + }, + "description": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "id": { + "type": "integer" + }, + "name": { + "type": "string" + }, + "other": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "updated_at": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "entity.PushEvent": { + "type": "object", + "properties": { + "channels": { + "type": "array", + "items": { + "type": "string" + } + }, + "created_at": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "event_key": { + "type": "string" + }, + "id": { + "type": "integer" + }, + "name": { + "type": "string" + }, + "targets": { + "type": "array", + "items": { + "type": "string" + } + }, + "task_type": { + "type": "string" + }, + "template": { + "type": "string" + }, + "updated_at": { + "type": "string" + } + } + }, "flared.ApplyLogPayload": { "type": "object", "properties": { @@ -16797,46 +17155,6 @@ } } }, - "model.BindRequest": { - "type": "object", - "properties": { - "channel_id": { - "type": "string" - }, - "code": { - "type": "string" - } - } - }, - "model.BindingDTO": { - "type": "object", - "properties": { - "channel_id": { - "type": "string", - "example": "0" - }, - "channel_name": { - "type": "string" - }, - "channel_type": { - "type": "string" - }, - "created_at": { - "type": "string" - }, - "id": { - "type": "string", - "example": "0" - }, - "platform_user_id": { - "type": "string" - }, - "user_id": { - "type": "string", - "example": "0" - } - } - }, "model.BrowserItem": { "type": "object", "properties": { @@ -16848,43 +17166,6 @@ } } }, - "model.ChannelDTO": { - "type": "object", - "properties": { - "credentials": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "extra": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "id": { - "type": "string", - "example": "0" - }, - "name": { - "type": "string" - }, - "owner_id": { - "type": "string", - "example": "0" - }, - "owner_scope": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, "model.ConfigVersion": { "type": "object", "properties": { @@ -16943,91 +17224,6 @@ } } }, - "model.CreateChannelRequest": { - "type": "object", - "properties": { - "credentials": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "extra": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "name": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, - "model.CreatePushChannelRequest": { - "type": "object", - "required": [ - "name", - "type" - ], - "properties": { - "description": { - "type": "string" - }, - "enabled": { - "type": "boolean" - }, - "name": { - "type": "string" - }, - "other": { - "type": "string" - }, - "token": { - "type": "string" - }, - "type": { - "type": "string" - }, - "url": { - "type": "string" - } - } - }, - "model.CreatePushEventRequest": { - "type": "object", - "properties": { - "channels": { - "type": "array", - "items": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "event_key": { - "type": "string" - }, - "targets": { - "type": "array", - "items": { - "type": "string" - } - }, - "task_type": { - "type": "string" - }, - "template": { - "type": "string" - } - } - }, "model.CreateScheduleRequest": { "type": "object", "required": [ @@ -17214,20 +17410,6 @@ } } }, - "model.Definition": { - "type": "object", - "properties": { - "fields": { - "type": "array", - "items": { - "$ref": "#/definitions/model.Field" - } - }, - "type": { - "type": "string" - } - } - }, "model.DispatchTaskRequest": { "type": "object", "required": [ @@ -17290,20 +17472,6 @@ } } }, - "model.Field": { - "type": "object", - "properties": { - "key": { - "type": "string" - }, - "required": { - "type": "boolean" - }, - "type": { - "type": "string" - } - } - }, "model.ListUsersResponse": { "type": "object", "properties": { @@ -17626,92 +17794,31 @@ } } }, - "model.PublicChannelDTO": { + "model.Schedule": { "type": "object", "properties": { + "created_at": { + "type": "string" + }, + "cron": { + "type": "string" + }, "id": { "type": "string", "example": "0" }, - "name": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, - "model.PushChannel": { - "type": "object", - "properties": { - "created_at": { - "type": "string" - }, - "description": { - "type": "string" - }, - "enabled": { + "is_active": { "type": "boolean" }, - "id": { - "type": "integer" - }, "name": { "type": "string" }, - "other": { + "payload": { "type": "string" }, - "token": { - "type": "string" - }, - "type": { - "type": "string" - }, - "updated_at": { - "type": "string" - }, - "url": { - "type": "string" - } - } - }, - "model.PushEvent": { - "type": "object", - "properties": { - "channels": { - "type": "array", - "items": { - "type": "string" - } - }, - "created_at": { - "type": "string" - }, - "enabled": { - "type": "boolean" - }, - "event_key": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "name": { - "type": "string" - }, - "targets": { - "type": "array", - "items": { - "type": "string" - } - }, "task_type": { "type": "string" }, - "template": { - "type": "string" - }, "updated_at": { "type": "string" } @@ -17886,39 +17993,37 @@ "TaskExecutionStatusFailed" ] }, - "model.TestPushChannelRequest": { + "model.Template": { "type": "object", "properties": { + "content": { + "type": "string" + }, + "created_at": { + "type": "string" + }, + "description": { + "type": "string" + }, + "id": { + "type": "integer" + }, + "is_system": { + "type": "boolean" + }, + "key": { + "type": "string" + }, "name": { "type": "string" }, - "other": { - "type": "string" - }, - "target": { - "type": "string" - }, - "token": { + "subject": { "type": "string" }, "type": { "type": "string" }, - "url": { - "type": "string" - } - } - }, - "model.TestPushRequest": { - "type": "object", - "required": [ - "config" - ], - "properties": { - "config": { - "$ref": "#/definitions/push.Config" - }, - "target": { + "updated_at": { "type": "string" } } @@ -18016,81 +18121,6 @@ } } }, - "model.UpdateChannelRequest": { - "type": "object", - "properties": { - "credentials": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "extra": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "name": { - "type": "string" - } - } - }, - "model.UpdatePushChannelRequest": { - "type": "object", - "required": [ - "type" - ], - "properties": { - "description": { - "type": "string" - }, - "enabled": { - "type": "boolean" - }, - "other": { - "type": "string" - }, - "token": { - "type": "string" - }, - "type": { - "type": "string" - }, - "url": { - "type": "string" - } - } - }, - "model.UpdatePushEventRequest": { - "type": "object", - "required": [ - "template" - ], - "properties": { - "channels": { - "type": "array", - "items": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "targets": { - "type": "array", - "items": { - "type": "string" - } - }, - "template": { - "type": "string" - } - } - }, "model.UpdateScheduleRequest": { "type": "object", "required": [ @@ -20517,6 +20547,10 @@ "description": "AppID 或 SMTP 用户名", "type": "string" }, + "other": { + "description": "附加配置 (如 ChatID / UserKey / 扩展 JSON)", + "type": "string" + }, "secret": { "description": "签名密钥或 SMTP 密码/Token", "type": "string" @@ -20885,6 +20919,17 @@ } } }, + "user.sendEmailCodeRequest": { + "type": "object", + "required": [ + "email" + ], + "properties": { + "email": { + "type": "string" + } + } + }, "user.updateProfileRequest": { "type": "object", "properties": { diff --git a/backend/docs/swagger.yaml b/backend/docs/swagger.yaml index 18668c08..0c20ed10 100644 --- a/backend/docs/swagger.yaml +++ b/backend/docs/swagger.yaml @@ -156,26 +156,6 @@ definitions: type: type: string type: object - Wavelet_plugins_domain_admin_model.Schedule: - properties: - created_at: - type: string - cron: - type: string - id: - example: "0" - type: string - is_active: - type: boolean - name: - type: string - payload: - type: string - task_type: - type: string - updated_at: - type: string - type: object Wavelet_plugins_domain_admin_model.SystemConfig: properties: created_at: @@ -233,29 +213,6 @@ definitions: updated_at: type: string type: object - Wavelet_plugins_domain_admin_model.Template: - properties: - content: - type: string - created_at: - type: string - description: - type: string - id: - type: integer - is_system: - type: boolean - key: - type: string - name: - type: string - subject: - type: string - type: - type: string - updated_at: - type: string - type: object agent.ActiveConfigMeta: properties: checksum: @@ -456,119 +413,6 @@ definitions: totalPage: type: integer type: object - auth.AuthSourceView: - properties: - client_secret_configured: - type: boolean - display_name: - type: string - icon_url: - type: string - id: - type: integer - is_active: - type: boolean - name: - type: string - type: - type: string - type: object - auth.BasicUserInfo: - properties: - avatar_url: - type: string - bio: - type: string - email: - type: string - gender: - type: string - id: - type: integer - is_admin: - type: boolean - location: - type: string - need_change_password: - type: boolean - nickname: - type: string - phone: - type: string - username: - type: string - website: - type: string - type: object - auth.CallbackRequest: - properties: - code: - type: string - state: - type: string - required: - - code - - state - type: object - auth.OAuthAuthorizeResponse: - properties: - authorize_url: - type: string - type: object - auth.OAuthCallbackResult: - properties: - status: - type: string - user: - $ref: '#/definitions/auth.BasicUserInfo' - type: object - cap.ChallengeResponse: - properties: - challenge: - properties: - c: - type: integer - d: - type: integer - s: - type: integer - type: object - expires: - description: ms timestamp - type: integer - token: - type: string - type: object - cap.RedeemResponse: - properties: - error: - type: string - expires: - type: integer - success: - type: boolean - token: - type: string - type: object - cap.challengeRequest: - properties: - scope: - type: string - type: object - cap.redeemRequest: - properties: - scope: - type: string - solutions: - items: - type: integer - type: array - token: - type: string - required: - - solutions - - token - type: object cloudflare.AvailableDomain: properties: domain: @@ -996,6 +840,379 @@ definitions: ttl_minutes: type: integer type: object + do.BindRequest: + properties: + channel_id: + type: string + code: + type: string + type: object + do.BindingDTO: + properties: + channel_id: + example: "0" + type: string + channel_name: + type: string + channel_type: + type: string + created_at: + type: string + id: + example: "0" + type: string + platform_user_id: + type: string + user_id: + example: "0" + type: string + type: object + do.ChannelDTO: + properties: + credentials: + additionalProperties: + type: string + type: object + enabled: + type: boolean + extra: + additionalProperties: + type: string + type: object + id: + example: "0" + type: string + name: + type: string + owner_id: + example: "0" + type: string + owner_scope: + type: string + type: + type: string + type: object + do.CreateChannelRequest: + properties: + credentials: + additionalProperties: + type: string + type: object + enabled: + type: boolean + extra: + additionalProperties: + type: string + type: object + name: + type: string + type: + type: string + type: object + do.CreatePushChannelRequest: + properties: + description: + type: string + enabled: + type: boolean + name: + type: string + other: + type: string + token: + type: string + type: + type: string + url: + type: string + required: + - name + - type + type: object + do.CreatePushEventRequest: + properties: + channels: + items: + type: string + type: array + enabled: + type: boolean + event_key: + type: string + targets: + items: + type: string + type: array + task_type: + type: string + template: + type: string + type: object + do.Definition: + properties: + fields: + items: + $ref: '#/definitions/do.Field' + type: array + type: + type: string + type: object + do.Field: + properties: + key: + type: string + required: + type: boolean + type: + type: string + type: object + do.PublicChannelDTO: + properties: + id: + example: "0" + type: string + name: + type: string + type: + type: string + type: object + do.TestPushChannelRequest: + properties: + name: + type: string + other: + type: string + target: + type: string + token: + type: string + type: + type: string + url: + type: string + type: object + do.TestPushRequest: + properties: + config: + $ref: '#/definitions/push.Config' + target: + type: string + required: + - config + type: object + do.UpdateChannelRequest: + properties: + credentials: + additionalProperties: + type: string + type: object + enabled: + type: boolean + extra: + additionalProperties: + type: string + type: object + name: + type: string + type: object + do.UpdatePushChannelRequest: + properties: + description: + type: string + enabled: + type: boolean + other: + type: string + token: + type: string + type: + type: string + url: + type: string + required: + - type + type: object + do.UpdatePushEventRequest: + properties: + channels: + items: + type: string + type: array + enabled: + type: boolean + targets: + items: + type: string + type: array + template: + type: string + required: + - template + type: object + dto.AuthSourceView: + properties: + client_secret_configured: + type: boolean + display_name: + type: string + icon_url: + type: string + id: + type: integer + is_active: + type: boolean + name: + type: string + type: + type: string + type: object + dto.BasicUserInfo: + properties: + avatar_url: + type: string + bio: + type: string + email: + type: string + gender: + type: string + id: + example: "0" + type: string + is_admin: + type: boolean + location: + type: string + need_change_password: + type: boolean + nickname: + type: string + phone: + type: string + username: + type: string + website: + type: string + type: object + dto.CallbackRequest: + properties: + code: + type: string + state: + type: string + required: + - code + - state + type: object + dto.ChallengeRequest: + properties: + scope: + type: string + type: object + dto.ChallengeResponse: + properties: + challenge: + properties: + c: + type: integer + d: + type: integer + s: + type: integer + type: object + expires: + description: ms timestamp + type: integer + token: + type: string + type: object + dto.OAuthAuthorizeResponse: + properties: + authorize_url: + type: string + type: object + dto.OAuthCallbackResult: + properties: + status: + type: string + user: + $ref: '#/definitions/dto.BasicUserInfo' + type: object + dto.RedeemRequest: + properties: + scope: + type: string + solutions: + items: + type: integer + type: array + token: + type: string + required: + - solutions + - token + type: object + dto.RedeemResponse: + properties: + error: + type: string + expires: + type: integer + success: + type: boolean + token: + type: string + type: object + entity.PushChannel: + properties: + created_at: + type: string + description: + type: string + enabled: + type: boolean + id: + type: integer + name: + type: string + other: + type: string + token: + type: string + type: + type: string + updated_at: + type: string + url: + type: string + type: object + entity.PushEvent: + properties: + channels: + items: + type: string + type: array + created_at: + type: string + enabled: + type: boolean + event_key: + type: string + id: + type: integer + name: + type: string + targets: + items: + type: string + type: array + task_type: + type: string + template: + type: string + updated_at: + type: string + type: object flared.ApplyLogPayload: properties: checksum: @@ -1202,33 +1419,6 @@ definitions: url: type: string type: object - model.BindRequest: - properties: - channel_id: - type: string - code: - type: string - type: object - model.BindingDTO: - properties: - channel_id: - example: "0" - type: string - channel_name: - type: string - channel_type: - type: string - created_at: - type: string - id: - example: "0" - type: string - platform_user_id: - type: string - user_id: - example: "0" - type: string - type: object model.BrowserItem: properties: browser: @@ -1236,31 +1426,6 @@ definitions: count: type: integer type: object - model.ChannelDTO: - properties: - credentials: - additionalProperties: - type: string - type: object - enabled: - type: boolean - extra: - additionalProperties: - type: string - type: object - id: - example: "0" - type: string - name: - type: string - owner_id: - example: "0" - type: string - owner_scope: - type: string - type: - type: string - type: object model.ConfigVersion: properties: checksum: @@ -1299,62 +1464,6 @@ definitions: version: type: string type: object - model.CreateChannelRequest: - properties: - credentials: - additionalProperties: - type: string - type: object - enabled: - type: boolean - extra: - additionalProperties: - type: string - type: object - name: - type: string - type: - type: string - type: object - model.CreatePushChannelRequest: - properties: - description: - type: string - enabled: - type: boolean - name: - type: string - other: - type: string - token: - type: string - type: - type: string - url: - type: string - required: - - name - - type - type: object - model.CreatePushEventRequest: - properties: - channels: - items: - type: string - type: array - enabled: - type: boolean - event_key: - type: string - targets: - items: - type: string - type: array - task_type: - type: string - template: - type: string - type: object model.CreateScheduleRequest: properties: cron: @@ -1485,15 +1594,6 @@ definitions: version: type: string type: object - model.Definition: - properties: - fields: - items: - $ref: '#/definitions/model.Field' - type: array - type: - type: string - type: object model.DispatchTaskRequest: properties: end_time: @@ -1535,15 +1635,6 @@ definitions: description: '"select" 或 "exec"' type: string type: object - model.Field: - properties: - key: - type: string - required: - type: boolean - type: - type: string - type: object model.ListUsersResponse: properties: total: @@ -1756,63 +1847,23 @@ definitions: value: type: string type: object - model.PublicChannelDTO: + model.Schedule: properties: + created_at: + type: string + cron: + type: string id: example: "0" type: string - name: - type: string - type: - type: string - type: object - model.PushChannel: - properties: - created_at: - type: string - description: - type: string - enabled: + is_active: type: boolean - id: - type: integer name: type: string - other: + payload: type: string - token: - type: string - type: - type: string - updated_at: - type: string - url: - type: string - type: object - model.PushEvent: - properties: - channels: - items: - type: string - type: array - created_at: - type: string - enabled: - type: boolean - event_key: - type: string - id: - type: integer - name: - type: string - targets: - items: - type: string - type: array task_type: type: string - template: - type: string updated_at: type: string type: object @@ -1930,30 +1981,29 @@ definitions: - TaskExecutionStatusRunning - TaskExecutionStatusSucceeded - TaskExecutionStatusFailed - model.TestPushChannelRequest: + model.Template: properties: + content: + type: string + created_at: + type: string + description: + type: string + id: + type: integer + is_system: + type: boolean + key: + type: string name: type: string - other: - type: string - target: - type: string - token: + subject: type: string type: type: string - url: + updated_at: type: string type: object - model.TestPushRequest: - properties: - config: - $ref: '#/definitions/push.Config' - target: - type: string - required: - - config - type: object model.TestSMTPRequest: properties: smtp_host: @@ -2018,55 +2068,6 @@ definitions: - max_size_mb - ttl_minutes type: object - model.UpdateChannelRequest: - properties: - credentials: - additionalProperties: - type: string - type: object - enabled: - type: boolean - extra: - additionalProperties: - type: string - type: object - name: - type: string - type: object - model.UpdatePushChannelRequest: - properties: - description: - type: string - enabled: - type: boolean - other: - type: string - token: - type: string - type: - type: string - url: - type: string - required: - - type - type: object - model.UpdatePushEventRequest: - properties: - channels: - items: - type: string - type: array - enabled: - type: boolean - targets: - items: - type: string - type: array - template: - type: string - required: - - template - type: object model.UpdateScheduleRequest: properties: cron: @@ -3669,6 +3670,9 @@ definitions: key: description: AppID 或 SMTP 用户名 type: string + other: + description: 附加配置 (如 ChatID / UserKey / 扩展 JSON) + type: string secret: description: 签名密钥或 SMTP 密码/Token type: string @@ -3917,6 +3921,13 @@ definitions: - password - username type: object + user.sendEmailCodeRequest: + properties: + email: + type: string + required: + - email + type: object user.updateProfileRequest: properties: avatar_url: @@ -4884,7 +4895,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.ChannelDTO' + $ref: '#/definitions/do.ChannelDTO' type: array type: object security: @@ -4902,7 +4913,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.CreateChannelRequest' + $ref: '#/definitions/do.CreateChannelRequest' produces: - application/json responses: @@ -4913,7 +4924,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.ChannelDTO' + $ref: '#/definitions/do.ChannelDTO' type: object "400": description: Bad Request @@ -4964,7 +4975,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.UpdateChannelRequest' + $ref: '#/definitions/do.UpdateChannelRequest' produces: - application/json responses: @@ -4975,7 +4986,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.ChannelDTO' + $ref: '#/definitions/do.ChannelDTO' type: object "400": description: Bad Request @@ -5033,7 +5044,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.Definition' + $ref: '#/definitions/do.Definition' type: array type: object security: @@ -5055,7 +5066,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.PushChannel' + $ref: '#/definitions/entity.PushChannel' type: array type: object security: @@ -5073,7 +5084,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.CreatePushChannelRequest' + $ref: '#/definitions/do.CreatePushChannelRequest' produces: - application/json responses: @@ -5084,7 +5095,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.PushChannel' + $ref: '#/definitions/entity.PushChannel' type: object security: - SessionCookie: [] @@ -5129,7 +5140,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.UpdatePushChannelRequest' + $ref: '#/definitions/do.UpdatePushChannelRequest' produces: - application/json responses: @@ -5140,7 +5151,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.PushChannel' + $ref: '#/definitions/entity.PushChannel' type: object security: - SessionCookie: [] @@ -5173,7 +5184,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.TestPushChannelRequest' + $ref: '#/definitions/do.TestPushChannelRequest' produces: - application/json responses: @@ -5200,7 +5211,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.PushEvent' + $ref: '#/definitions/entity.PushEvent' type: array type: object security: @@ -5218,7 +5229,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.CreatePushEventRequest' + $ref: '#/definitions/do.CreatePushEventRequest' produces: - application/json responses: @@ -5229,7 +5240,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.PushEvent' + $ref: '#/definitions/entity.PushEvent' type: object security: - SessionCookie: [] @@ -5277,7 +5288,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.UpdatePushEventRequest' + $ref: '#/definitions/do.UpdatePushEventRequest' produces: - application/json responses: @@ -5379,7 +5390,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.TestPushRequest' + $ref: '#/definitions/do.TestPushRequest' produces: - application/json responses: @@ -5862,7 +5873,7 @@ paths: - properties: data: items: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Schedule' + $ref: '#/definitions/model.Schedule' type: array type: object "401": @@ -5899,7 +5910,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Schedule' + $ref: '#/definitions/model.Schedule' type: object "400": description: Cron 表达式无效、异步任务类型不存在或参数错误 @@ -5990,7 +6001,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Schedule' + $ref: '#/definitions/model.Schedule' type: object "400": description: Cron 表达式无效、参数错误 @@ -6061,7 +6072,7 @@ paths: - properties: data: items: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Template' + $ref: '#/definitions/model.Template' type: array type: object "401": @@ -6189,7 +6200,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Template' + $ref: '#/definitions/model.Template' type: object "401": description: 未登录 @@ -6238,7 +6249,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Template' + $ref: '#/definitions/model.Template' type: object "400": description: 参数错误 @@ -7221,7 +7232,7 @@ paths: in: body name: request schema: - $ref: '#/definitions/cap.challengeRequest' + $ref: '#/definitions/dto.ChallengeRequest' produces: - application/json responses: @@ -7232,7 +7243,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/cap.ChallengeResponse' + $ref: '#/definitions/dto.ChallengeResponse' type: object "500": description: 内部服务错误 @@ -7250,7 +7261,7 @@ paths: in: body name: request schema: - $ref: '#/definitions/cap.challengeRequest' + $ref: '#/definitions/dto.ChallengeRequest' produces: - application/json responses: @@ -7261,7 +7272,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/cap.ChallengeResponse' + $ref: '#/definitions/dto.ChallengeResponse' type: object "500": description: 内部服务错误 @@ -7281,7 +7292,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/cap.redeemRequest' + $ref: '#/definitions/dto.RedeemRequest' produces: - application/json responses: @@ -7292,7 +7303,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/cap.RedeemResponse' + $ref: '#/definitions/dto.RedeemResponse' type: object "400": description: 参数错误或核销失败 @@ -7309,7 +7320,7 @@ paths: get: consumes: - application/json - description: 返回系统配置表中 visibility 为 1 的配置键值集合 + description: 返回系统配置表中 visibility 为 1 的扁平键值集合(如 cap_login_enabled) produces: - application/json responses: @@ -12204,7 +12215,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.BindingDTO' + $ref: '#/definitions/do.BindingDTO' type: array type: object "401": @@ -12227,7 +12238,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.BindRequest' + $ref: '#/definitions/do.BindRequest' produces: - application/json responses: @@ -12238,7 +12249,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.BindingDTO' + $ref: '#/definitions/do.BindingDTO' type: object "400": description: Bad Request @@ -12296,7 +12307,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.PublicChannelDTO' + $ref: '#/definitions/do.PublicChannelDTO' type: array type: object "401": @@ -12331,7 +12342,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.OAuthAuthorizeResponse' + $ref: '#/definitions/dto.OAuthAuthorizeResponse' type: object "400": description: 认证源不存在或未启用 @@ -12355,7 +12366,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/auth.CallbackRequest' + $ref: '#/definitions/dto.CallbackRequest' produces: - application/json responses: @@ -12366,7 +12377,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.OAuthCallbackResult' + $ref: '#/definitions/dto.OAuthCallbackResult' type: object "400": description: state 无效、参数错误或认证源错误 @@ -12459,7 +12470,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.OAuthAuthorizeResponse' + $ref: '#/definitions/dto.OAuthAuthorizeResponse' type: object "400": description: 认证源不存在或未配置 @@ -12510,7 +12521,7 @@ paths: - properties: data: items: - $ref: '#/definitions/auth.AuthSourceView' + $ref: '#/definitions/dto.AuthSourceView' type: array type: object summary: 获取可用登录源 @@ -12529,7 +12540,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.BasicUserInfo' + $ref: '#/definitions/dto.BasicUserInfo' type: object "401": description: 未登录 @@ -12905,7 +12916,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.BasicUserInfo' + $ref: '#/definitions/dto.BasicUserInfo' type: object "401": description: 未登录 @@ -13063,7 +13074,7 @@ paths: post: consumes: - application/json - description: 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。 + description: 使用用户名和密码登录系统,验证通过后建立 Session 并返回用户信息。 parameters: - description: 登录请求参数 in: body @@ -13082,11 +13093,15 @@ paths: description: 用户名或密码错误 schema: $ref: '#/definitions/response.Any' + "429": + description: 登录尝试过于频繁 + schema: + $ref: '#/definitions/response.Any' "500": description: 服务内部错误 schema: $ref: '#/definitions/response.Any' - summary: 用户密码登录 + summary: 用户登录 tags: - user /api/v1/user/logout: @@ -13166,6 +13181,10 @@ paths: description: 参数错误、用户名已存在或注册已关闭 schema: $ref: '#/definitions/response.Any' + "429": + description: 注册尝试过于频繁 + schema: + $ref: '#/definitions/response.Any' "500": description: 服务内部错误 schema: @@ -13197,6 +13216,13 @@ paths: consumes: - application/json description: 向指定邮箱发送验证码(用于注册场景) + parameters: + - description: 目标邮箱 + in: body + name: request + required: true + schema: + $ref: '#/definitions/user.sendEmailCodeRequest' produces: - application/json responses: @@ -13208,6 +13234,10 @@ paths: description: 参数错误 schema: $ref: '#/definitions/response.Any' + "500": + description: 发送失败 + schema: + $ref: '#/definitions/response.Any' summary: 发送邮箱验证码 tags: - user diff --git a/backend/downstream/README.md b/backend/downstream/README.md index 2c6a0925..afceb751 100644 --- a/backend/downstream/README.md +++ b/backend/downstream/README.md @@ -1,77 +1,147 @@ -# Downstream Custom Plugins +# 下游插件开发指南 (Downstream Custom Plugins) -This directory is the designated location for downstream (deployment-specific) Cordis plugins. +本目录为 Wavelet 下游业务定制插件(Deployment-specific plugins)的专属开发目录。 -## Architecture +所有插件开发均以标准模板 [`custom_example`](./plugins/custom_example) 为基准进行构建。 -``` +--- + +## 目录结构 + +```text downstream/ ├── README.md └── plugins/ - └── custom_example/ # Example plugin — copy & rename to get started - └── plugin.go + └── custom_example/ # 标准插件开发基准模板(可直接复制并重命名开发新插件) + ├── plugin.go # 插件入口:实现 core.Plugin (Name & Apply) + ├── consts/ # 常量与错误码定义 + │ └── consts.go + ├── controller/ # 控制器层 (HTTP API 接口、参数绑定与信封响应) + │ └── hello/ + │ └── hello.go + ├── service/ # 业务逻辑层 (用例编排、事务控制与领域事件) + ├── dao/ # 数据访问层 (GORM CRUD、SQL 防注入与转义) + ├── model/ # 数据模型与实体定义 + │ ├── do/ # Domain Object 领域对象与 DTO + │ └── entity/ # 数据表映射实体 (TableName() 带专属前缀) + └── migrations/ # Goose SQL 双方言独立数据库迁移 + ├── postgres/ # PostgreSQL 迁移脚本 + └── sqlite/ # SQLite 迁移脚本 ``` -Downstream plugins follow the same `core.Plugin` contract as platform plugins: +--- -```go -type Plugin interface { - Name() string - Apply(ctx *core.Context) error -} +## 插件开发规范与分层职责 + +每个下游插件均需遵循标准分层架构(`controller -> service -> dao -> model`): + +1. **`plugin.go` (插件装配入口)**: + - 实现 `core.Plugin` 接口(`Name() string` 与 `Apply(ctx *core.Context) error`)。 + - 负责在 `Apply` 中注册路由组(`ctx.Router()`)、异步任务(`ctx.Task()`)、定时调度(`ctx.Schedule()`)与数据库迁移(`ctx.Migrations()`)。 + - 依赖注入统一使用 `core.Inject` 或 `ctx.Using` 获取平台服务(如 `contracts.AuthService`、`contracts.DBService`、`contracts.CacheService`)。 + +2. **`controller/` (控制器层 / Handler)**: + - 负责 HTTP API 请求参数绑定(`c.ShouldBindJSON` / `c.ShouldBindQuery`)、用户会话获取(`oauth.GetCurrentUser`)。 + - 调用 Service 层处理业务,严禁直接包含复杂业务逻辑或直接执行 SQL 操作。 + - 统一使用 `response.OK` 或 `response.Abort*` 返回标准信封响应。 + +3. **`service/` (业务逻辑层)**: + - 纯 Go 业务用例,方法入参首位统一为 `context.Context`。 + - 严禁依赖 `*gin.Context` 或 HTTP 相关对象,确保逻辑具备可移植性与可测试性。 + - 涉及数据修改可通过 `ctx.Events().Emit()` 发射强类型领域事件。 + +4. **`dao/` (数据访问层 / Repository)**: + - 负责底层数据库交互,通过 `contracts.DBService` 获取受保护的 GORM 数据库句柄。 + - 模糊查询必须调用 `pkg/util.EscapeLike` 并显式声明 `ESCAPE '\\'` 防注入。 + - 遵循**表单一所有者原则**:严禁跨过其他所有者插件直读或修改其他插件的数据表。 + +5. **`model/` (模型与 DTO)**: + - `model/entity/`:数据库表映射结构体,`TableName()` 必须带有插件专属前缀(如 `w_custom_*`)。 + - `model/do/`:业务领域对象、入参校验 Request DTO 与响应 Response DTO。 + +6. **`consts/` (常量与错误定义)**: + - 定义插件内部常量、配置键名及驼峰式(camelCase)错误标识字符串。 + +7. **`migrations/` (双方言 SQL 迁移)**: + - 包含 `postgres/` 与 `sqlite/` 双方言 Goose SQL 迁移文件,使用 `//go:embed` 打包并在 `Apply()` 中注册。 + +--- + +## 快速上手 + +### 1. 基于模板复制创建新插件 + +```bash +cp -r backend/downstream/plugins/custom_example backend/downstream/plugins/my_plugin ``` -## Rules - -1. **Naming**: Each plugin directory name becomes its import path and plugin ID (kebab-case recommended). -2. **Dependencies**: Downstream plugins may import `core/`, `core/contracts/`, `pkg/`, and `plugins/infra/` packages from the platform. They MUST NOT import domain plugin internal packages — use `core.Inject[contracts.XxxService](ctx)` instead. -3. **Registration**: Add your downstream plugin to `cmd/app.go` before the platform plugins or after, depending on which services it needs: - ```go - // newWaveletApp in cmd/app.go - app.Use( - database.New(), - cache.New(), - logger.New(), - storage.New(), - // ... platform domain plugins ... - custom_hello.New(), // your downstream plugin - driver_http.New(), - driver_asynq_worker.New(), - driver_asynq_cron.New(), - ) - ``` -4. **Migration**: If your plugin needs database tables, embed SQL files in a `migrations/` directory and register via `ctx.Migrations().Register(...)` in `Apply()`. - -## Quick Start +### 2. 实现插件入口 (`plugin.go`) ```go -package custom_example +package my_plugin import ( - "github.com/Rain-kl/Wavelet/core" - "github.com/Rain-kl/Wavelet/core/contracts" - "github.com/gin-gonic/gin" + "Wavelet/core" + "Wavelet/core/contracts" + "net/http" + + "github.com/gin-gonic/gin" ) type Plugin struct{} -func New() *Plugin { return &Plugin{} } +func New() *Plugin { + return &Plugin{} +} -func (p *Plugin) Name() string { return "custom_example" } +func (p *Plugin) Name() string { + return "my_plugin" +} func (p *Plugin) Apply(ctx *core.Context) error { - // Example: register a route that uses AuthService - var authSvc contracts.AuthService - if err := ctx.Using(func(svc contracts.AuthService) { authSvc = svc }); err != nil { - return err - } + // 通过容器解析认证服务 + var authSvc contracts.AuthService + if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil { + return err + } - g := ctx.Router().Group("/api/v1/custom", authSvc.RequireAuthMiddleware().(gin.HandlerFunc)) - g.GET("/hello", func(c *gin.Context) { - user, _ := authSvc.GetCurrentUser(c.Request.Context()) - c.JSON(200, gin.H{"message": "Hello " + user.Username}) - }) + // 注册带鉴权中间件的路由组 + g := ctx.Router().Group("/api/v1/my-plugin", authSvc.RequireAuthMiddleware().(gin.HandlerFunc)) + g.GET("/hello", func(c *gin.Context) { + user, err := authSvc.GetCurrentUser(c.Request.Context()) + if err != nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "unauthorized"}) + return + } + c.JSON(http.StatusOK, gin.H{"message": "Hello " + user.Username}) + }) - return nil + return nil } -``` \ No newline at end of file +``` + +### 3. 在应用装配入口注册插件 (`cmd/app.go`) + +在 `cmd/app.go` 中的 `newWaveletApp` 函数内注册你的新插件: + +```go +app.Use( + database.New(), + cache.New(), + logger.New(), + storage.New(), + // ... 官方平台插件 ... + my_plugin.New(), // 注册下游定制插件 + driver_http.New(), + driver_asynq_worker.New(), + driver_asynq_cron.New(), +) +``` + +--- + +## 严格红线与规范 (Guardrails) + +- **严禁跨插件直接 import**:下游插件可依赖 `core/`、`core/contracts/`、`pkg/`、`plugins/infra/`,**严禁直接 import `plugins/domain/*` 内部私有实现**,一律通过 `contracts` 接口或事件总线调用。 +- **禁止 GORM AutoMigrate**:数据表结构定义必须通过 `migrations/` 下嵌入的 Goose SQL 管理。 +- **Goroutine 并发安全**:严禁裸 `go func()`,后台并发任务统一使用 `backend/pkg/util.Go`。 diff --git a/backend/downstream/plugins/custom_example/consts/consts.go b/backend/downstream/plugins/custom_example/consts/consts.go new file mode 100644 index 00000000..9454d22f --- /dev/null +++ b/backend/downstream/plugins/custom_example/consts/consts.go @@ -0,0 +1,5 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package consts defines constants and error codes for custom_example plugin. +package consts diff --git a/backend/downstream/plugins/custom_example/controller/hello/hello.go b/backend/downstream/plugins/custom_example/controller/hello/hello.go new file mode 100644 index 00000000..0cf687a3 --- /dev/null +++ b/backend/downstream/plugins/custom_example/controller/hello/hello.go @@ -0,0 +1,5 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package hello provides HTTP API handlers for the custom_example plugin. +package hello diff --git a/backend/downstream/plugins/custom_example/dao/.gitkeep b/backend/downstream/plugins/custom_example/dao/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/backend/downstream/plugins/custom_example/migrations/.gitkeep b/backend/downstream/plugins/custom_example/migrations/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/backend/downstream/plugins/custom_example/migrations/postgres/.gitkeep b/backend/downstream/plugins/custom_example/migrations/postgres/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/backend/downstream/plugins/custom_example/migrations/sqlite/.gitkeep b/backend/downstream/plugins/custom_example/migrations/sqlite/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/backend/downstream/plugins/custom_example/model/.gitkeep b/backend/downstream/plugins/custom_example/model/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/backend/downstream/plugins/custom_example/model/do/.gitkeep b/backend/downstream/plugins/custom_example/model/do/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/backend/downstream/plugins/custom_example/model/entity/.gitkeep b/backend/downstream/plugins/custom_example/model/entity/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/backend/downstream/plugins/custom_example/service/.gitkeep b/backend/downstream/plugins/custom_example/service/.gitkeep new file mode 100644 index 00000000..e69de29b diff --git a/backend/go.mod b/backend/go.mod index d4d04a12..6191b934 100644 --- a/backend/go.mod +++ b/backend/go.mod @@ -1,14 +1,14 @@ module Wavelet -go 1.25.7 +go 1.25.10 require ( github.com/ClickHouse/clickhouse-go/v2 v2.48.0 github.com/alicebob/miniredis/v2 v2.38.0 github.com/aliyun/alibabacloud-oss-go-sdk-v2 v1.5.3 - github.com/aws/aws-sdk-go-v2 v1.43.4 - github.com/aws/aws-sdk-go-v2/config v1.32.35 - github.com/aws/aws-sdk-go-v2/credentials v1.19.34 + github.com/aws/aws-sdk-go-v2 v1.45.1 + github.com/aws/aws-sdk-go-v2/config v1.33.1 + github.com/aws/aws-sdk-go-v2/credentials v1.20.1 github.com/aws/aws-sdk-go-v2/service/s3 v1.106.5 github.com/bodgit/sevenzip v1.6.5 github.com/bwmarrin/snowflake v0.3.0 @@ -20,13 +20,13 @@ require ( github.com/gin-gonic/gin v1.12.0 github.com/glebarez/sqlite v1.11.0 github.com/go-acme/lego/v4 v4.35.2 - github.com/go-jose/go-jose/v4 v4.1.4 - github.com/google/go-cmp v0.7.0 + github.com/go-redis/redis_rate/v10 v10.0.1 github.com/google/uuid v1.6.0 github.com/gorilla/sessions v1.4.0 github.com/gorilla/websocket v1.5.3 github.com/hibiken/asynq v0.26.0 github.com/maypok86/otter/v2 v2.3.0 + github.com/nikoksr/notify v1.6.0 github.com/oschwald/maxminddb-golang v1.13.1 github.com/peterbourgon/diskv/v3 v3.0.1 github.com/pressly/goose/v3 v3.27.3 @@ -36,7 +36,7 @@ require ( github.com/shopspring/decimal v1.4.0 github.com/spf13/cobra v1.10.2 github.com/spf13/viper v1.21.0 - github.com/stretchr/testify v1.11.1 + github.com/stretchr/testify v1.12.1 github.com/studio-b12/gowebdav v0.13.0 github.com/swaggo/files v1.0.1 github.com/swaggo/gin-swagger v1.6.1 @@ -44,6 +44,7 @@ require ( github.com/tencent-connect/botgo v0.2.1 github.com/ulikunitz/xz v0.5.16 github.com/uptrace/opentelemetry-go-extra/otelzap v0.3.2 + github.com/wneessen/go-mail v0.8.1 github.com/yuin/gopher-lua v1.1.2 go.opentelemetry.io/contrib/instrumentation/github.com/gin-gonic/gin/otelgin v0.70.0 go.opentelemetry.io/contrib/instrumentation/net/http/otelhttp v0.70.0 @@ -52,20 +53,19 @@ require ( go.opentelemetry.io/otel/sdk v1.45.0 go.opentelemetry.io/otel/trace v1.45.0 go.uber.org/zap v1.28.0 - golang.org/x/crypto v0.54.0 + golang.org/x/crypto v0.55.0 golang.org/x/image v0.44.0 golang.org/x/mod v0.38.0 - golang.org/x/net v0.57.0 + golang.org/x/net v0.58.0 golang.org/x/oauth2 v0.36.0 golang.org/x/sync v0.22.0 gopkg.in/natefinch/lumberjack.v2 v2.2.1 gopkg.in/telebot.v4 v4.0.0-beta.10 gorm.io/driver/clickhouse v0.7.0 gorm.io/driver/postgres v1.6.2 - gorm.io/driver/sqlite v1.6.0 gorm.io/gorm v1.31.2 gorm.io/plugin/dbresolver v1.6.2 - gorm.io/plugin/opentelemetry v0.1.16 + gorm.io/plugin/opentelemetry v0.1.14 ) require ( @@ -74,62 +74,67 @@ require ( github.com/KyleBanks/depth v1.2.1 // indirect github.com/andybalholm/brotli v1.2.2 // indirect github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.16 // indirect - github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.35 // indirect - github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.35 // indirect - github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.35 // indirect - github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.36 // indirect - github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.15 // indirect + github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.19.1 // indirect + github.com/aws/aws-sdk-go-v2/internal/configsources v1.5.1 // indirect + github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.8.1 // indirect + github.com/aws/aws-sdk-go-v2/internal/v4a v1.5.1 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.19 // indirect github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.28 // indirect - github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.35 // indirect + github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.14.1 // indirect github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.36 // indirect - github.com/aws/aws-sdk-go-v2/service/signin v1.5.4 // indirect - github.com/aws/aws-sdk-go-v2/service/sso v1.33.4 // indirect - github.com/aws/aws-sdk-go-v2/service/ssooidc v1.38.4 // indirect - github.com/aws/aws-sdk-go-v2/service/sts v1.45.4 // indirect - github.com/aws/smithy-go v1.27.6 // indirect + github.com/aws/aws-sdk-go-v2/service/signin v1.7.1 // indirect + github.com/aws/aws-sdk-go-v2/service/sso v1.35.1 // indirect + github.com/aws/aws-sdk-go-v2/service/ssooidc v1.40.1 // indirect + github.com/aws/aws-sdk-go-v2/service/sts v1.47.1 // indirect + github.com/aws/smithy-go v1.28.1 // indirect + github.com/blinkbean/dingtalk v1.1.3 // indirect github.com/bodgit/plumbing v1.3.0 // indirect github.com/bodgit/windows v1.0.1 // indirect - github.com/boj/redistore v1.4.2 // indirect + github.com/boj/redistore v1.4.1 // indirect + github.com/bwmarrin/discordgo v0.29.0 // indirect github.com/bytedance/gopkg v0.1.4 // indirect github.com/bytedance/sonic v1.15.2 // indirect github.com/bytedance/sonic/loader v0.5.2 // indirect github.com/cenkalti/backoff/v5 v5.0.3 // indirect github.com/cespare/xxhash/v2 v2.3.0 // indirect github.com/cloudwego/base64x v0.1.7 // indirect - github.com/davecgh/go-spew v1.1.2-0.20180830191138-d8f796af33cc // indirect github.com/dustin/go-humanize v1.0.1 // indirect github.com/felixge/httpsnoop v1.1.0 // indirect - github.com/fsnotify/fsnotify v1.10.1 // indirect + github.com/fsnotify/fsnotify v1.9.0 // indirect github.com/gabriel-vasile/mimetype v1.4.15 // indirect github.com/gin-contrib/sse v1.1.1 // indirect - github.com/glebarez/go-sqlite v1.23.0 // indirect + github.com/glebarez/go-sqlite v1.21.2 // indirect github.com/go-faster/city v1.0.1 // indirect - github.com/go-faster/errors v0.8.0 // indirect + github.com/go-faster/errors v0.7.1 // indirect + github.com/go-jose/go-jose/v4 v4.1.4 // indirect + github.com/go-lark/lark v1.16.0 // indirect github.com/go-logr/logr v1.4.4 // indirect github.com/go-logr/stdr v1.2.2 // indirect - github.com/go-openapi/jsonpointer v1.0.0 // indirect - github.com/go-openapi/jsonreference v1.0.0 // indirect - github.com/go-openapi/spec v0.22.9 // indirect - github.com/go-openapi/swag/conv v0.28.0 // indirect - github.com/go-openapi/swag/jsonutils v0.28.0 // indirect - github.com/go-openapi/swag/loading v0.28.0 // indirect - github.com/go-openapi/swag/pools v0.28.0 // indirect - github.com/go-openapi/swag/stringutils v0.28.0 // indirect - github.com/go-openapi/swag/typeutils v0.28.0 // indirect - github.com/go-openapi/swag/yamlutils v0.28.0 // indirect + github.com/go-openapi/jsonpointer v0.22.1 // indirect + github.com/go-openapi/jsonreference v0.21.2 // indirect + github.com/go-openapi/spec v0.22.0 // indirect + github.com/go-openapi/swag/conv v0.25.1 // indirect + github.com/go-openapi/swag/jsonname v0.25.1 // indirect + github.com/go-openapi/swag/jsonutils v0.25.1 // indirect + github.com/go-openapi/swag/loading v0.25.1 // indirect + github.com/go-openapi/swag/stringutils v0.25.1 // indirect + github.com/go-openapi/swag/typeutils v0.25.1 // indirect + github.com/go-openapi/swag/yamlutils v0.25.1 // indirect github.com/go-playground/locales v0.14.1 // indirect github.com/go-playground/universal-translator v0.18.1 // indirect github.com/go-playground/validator/v10 v10.30.3 // indirect github.com/go-resty/resty/v2 v2.17.2 // indirect github.com/go-sql-driver/mysql v1.10.0 // indirect + github.com/go-telegram-bot-api/telegram-bot-api v4.6.4+incompatible // indirect github.com/go-viper/mapstructure/v2 v2.5.0 // indirect github.com/goccy/go-json v0.10.6 // indirect github.com/goccy/go-yaml v1.19.2 // indirect github.com/gomodule/redigo v1.9.3 // indirect - github.com/google/btree v1.1.3 // indirect + github.com/google/btree v1.0.0 // indirect github.com/gorilla/context v1.1.2 // indirect github.com/gorilla/securecookie v1.1.2 // indirect - github.com/grpc-ecosystem/grpc-gateway/v2 v2.30.0 // indirect + github.com/gregdel/pushover v1.4.0 // indirect + github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 // indirect github.com/hashicorp/go-version v1.9.0 // indirect github.com/hashicorp/golang-lru/v2 v2.0.7 // indirect github.com/inconshreveable/mousetrap v1.1.0 // indirect @@ -140,11 +145,10 @@ require ( github.com/jinzhu/inflection v1.0.0 // indirect github.com/jinzhu/now v1.1.5 // indirect github.com/json-iterator/go v1.1.13-0.20220915233716-71ac16282d12 // indirect - github.com/klauspost/compress v1.19.2 // indirect + github.com/klauspost/compress v1.19.1 // indirect github.com/klauspost/cpuid/v2 v2.4.0 // indirect github.com/leodido/go-urn v1.5.0 // indirect github.com/mattn/go-isatty v0.0.24 // indirect - github.com/mattn/go-sqlite3 v1.14.22 // indirect github.com/mfridman/interpolate v0.0.2 // indirect github.com/miekg/dns v1.1.72 // indirect github.com/modern-go/concurrent v0.0.0-20180306012644-bacd9c7ef1dd // indirect @@ -152,8 +156,7 @@ require ( github.com/ncruces/go-strftime v1.0.0 // indirect github.com/paulmach/orb v0.13.0 // indirect github.com/pelletier/go-toml/v2 v2.4.3 // indirect - github.com/pierrec/lz4/v4 v4.1.28 // indirect - github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 // indirect + github.com/pierrec/lz4/v4 v4.1.27 // indirect github.com/quic-go/qpack v0.6.0 // indirect github.com/quic-go/quic-go v0.61.0 // indirect github.com/redis/go-redis/extra/rediscmd/v9 v9.22.0 // indirect @@ -161,21 +164,24 @@ require ( github.com/sagikazarmark/locafero v0.12.0 // indirect github.com/segmentio/asm v1.2.1 // indirect github.com/sethvargo/go-retry v0.4.0 // indirect + github.com/slack-go/slack v0.29.0 // indirect github.com/spf13/afero v1.15.0 // indirect github.com/spf13/cast v1.10.0 // indirect github.com/spf13/pflag v1.0.10 // indirect github.com/stangelandcl/ppmd v0.1.1 // indirect + github.com/stretchr/objx v0.5.3 // indirect github.com/subosito/gotenv v1.6.0 // indirect - github.com/tidwall/gjson v1.9.3 // indirect - github.com/tidwall/match v1.1.1 // indirect - github.com/tidwall/pretty v1.2.0 // indirect + github.com/technoweenie/multipartstreamer v1.0.1 // indirect + github.com/tidwall/gjson v1.19.0 // indirect + github.com/tidwall/match v1.2.0 // indirect + github.com/tidwall/pretty v1.2.1 // indirect github.com/twitchyliquid64/golang-asm v0.15.1 // indirect - github.com/ugorji/go/codec v1.3.2 // indirect + github.com/ugorji/go/codec v1.3.1 // indirect github.com/uptrace/opentelemetry-go-extra/otelutil v0.3.2 // indirect go.mongodb.org/mongo-driver/v2 v2.8.0 // indirect go.opentelemetry.io/auto/sdk v1.2.1 // indirect go.opentelemetry.io/otel/exporters/otlp/otlptrace v1.45.0 // indirect - go.opentelemetry.io/otel/log v0.6.0 // indirect + go.opentelemetry.io/otel/log v0.12.2 // indirect go.opentelemetry.io/otel/metric v1.45.0 // indirect go.opentelemetry.io/proto/otlp v1.11.0 // indirect go.uber.org/atomic v1.11.0 // indirect @@ -184,19 +190,16 @@ require ( go4.org v0.0.0-20260112195520-a5071408f32f // indirect golang.org/x/arch v0.29.0 // indirect golang.org/x/sys v0.47.0 // indirect - golang.org/x/text v0.40.0 // indirect + golang.org/x/text v0.41.0 // indirect golang.org/x/time v0.15.0 // indirect golang.org/x/tools v0.48.0 // indirect google.golang.org/genproto/googleapis/api v0.0.0-20260803160001-6ac0973c030d // indirect - google.golang.org/genproto/googleapis/rpc v0.0.0-20260803160001-6ac0973c030d // indirect - google.golang.org/grpc v1.83.0 // indirect - google.golang.org/protobuf v1.36.11 // indirect - gopkg.in/yaml.v3 v3.0.1 // indirect + google.golang.org/genproto/googleapis/rpc v0.0.0-20260819154853-08b0e4226688 // indirect + google.golang.org/grpc v1.83.2 // indirect + google.golang.org/protobuf v1.36.12 // indirect gorm.io/driver/mysql v1.6.0 // indirect - modernc.org/libc v1.74.4 // indirect + modernc.org/libc v1.74.3 // indirect modernc.org/mathutil v1.7.1 // indirect modernc.org/memory v1.11.0 // indirect - modernc.org/sqlite v1.56.0 // indirect + modernc.org/sqlite v1.54.0 // indirect ) - -exclude github.com/gomodule/redigo v2.0.0+incompatible diff --git a/backend/go.sum b/backend/go.sum index 574eb7ec..44dad23c 100644 --- a/backend/go.sum +++ b/backend/go.sum @@ -84,58 +84,62 @@ github.com/armon/go-metrics v0.0.0-20180917152333-f0300d1749da/go.mod h1:Q73ZrmV github.com/armon/go-metrics v0.3.10/go.mod h1:4O98XIr/9W0sxpJ8UaYkvjk10Iff7SnFrb4QAOwNTFc= github.com/armon/go-radix v0.0.0-20180808171621-7fddfc383310/go.mod h1:ufUuZ+zHj4x4TnLV4JWEpy2hxWSpsRywHrMgIH9cCH8= github.com/armon/go-radix v1.0.0/go.mod h1:ufUuZ+zHj4x4TnLV4JWEpy2hxWSpsRywHrMgIH9cCH8= -github.com/aws/aws-sdk-go-v2 v1.43.4 h1:b9FTvbRwy+JCsfp2Wp6wV/KbOx3Aj7nkoFb2cRX0IhE= -github.com/aws/aws-sdk-go-v2 v1.43.4/go.mod h1:70vwSy16txshwG+g55WkpgPKDIByzHI8ccBsOteo3bQ= +github.com/aws/aws-sdk-go-v2 v1.45.1 h1:iIoG3NaLhV6UZpPXyPXlDj2I9oS8tV/nMcMnITCC6Ks= +github.com/aws/aws-sdk-go-v2 v1.45.1/go.mod h1:bttEH6JqnUL8LepvDVfdrds/fZ5bCIxzpe3abyUrhDU= github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.16 h1:aiuaKlDweRC5qExJondpWjOgyzMHpofpwspGXUtwn4c= github.com/aws/aws-sdk-go-v2/aws/protocol/eventstream v1.7.16/go.mod h1:nG/LOlmox9BDe9HvQnXWzgcK8uKbgBMZ/Hp5pVt/21I= -github.com/aws/aws-sdk-go-v2/config v1.32.35 h1:UEzXuET8E42lxBPijuACu/tEK7v5lFPlk0Q+GT5WD9E= -github.com/aws/aws-sdk-go-v2/config v1.32.35/go.mod h1:KaMtJpFa2JlL2BStjjHQVwQpzZEmw+ND/EgVrfFoo2g= -github.com/aws/aws-sdk-go-v2/credentials v1.19.34 h1:y6GkSmcv5myd1ngrYbGmiLlwQqB6TQhOuN/tbSSuWDY= -github.com/aws/aws-sdk-go-v2/credentials v1.19.34/go.mod h1:w3dTcnDVoQIewjo7JG45hduAToikiIFLC4FIO7fndvw= -github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.35 h1:+S7kbJoLDDQ5tE+lHrUBgMkzC8NLgsaioS2F3dVoFAE= -github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.18.35/go.mod h1:Ak7xXviIARfFdNUJ9Etb0bdVDt/KAvKjMGJVLWXDzik= -github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.35 h1:kzVuGlatQtYinwBJEEyLAbggepCoavosiaHHX9+fD+c= -github.com/aws/aws-sdk-go-v2/internal/configsources v1.4.35/go.mod h1:0yLx0yEI+SfqeJMPvOtIEFoZbiQYXMGszBueiutQyaI= -github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.35 h1:WK6CjihTuLisCjSKKbildJ79sGZZgbBz3iNa7VsKIhU= -github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.7.35/go.mod h1:KYleN57luLoe97R7vTnx8PMcVrr9gAcRECtOjl91DNg= -github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.36 h1:jbGY4CXLzZElOXgGsexlC3Hi+3YM0rSmk4opFXKqg/k= -github.com/aws/aws-sdk-go-v2/internal/v4a v1.4.36/go.mod h1:uBu/9aKsS/UQGc72RAt3y54kjgYQxmhut8ZD2dXCDNE= -github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.15 h1:JJLBQxwY+AFwuPAi5ivGc1ChnTdUt4cXMv7e76m2c/Y= -github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.15/go.mod h1:lQknBIe78MVL0cQOQDlag8KGflMbMEVFx9mB6O8ENvk= +github.com/aws/aws-sdk-go-v2/config v1.33.1 h1:bq9jze1hQ5YTCLoVxNnbp0T7rglrlOE7N9YsHqjGkEw= +github.com/aws/aws-sdk-go-v2/config v1.33.1/go.mod h1:2A3HQwG4zaL5Tm80rc6RZj8LmWWv4WYT5v8raSz/L7A= +github.com/aws/aws-sdk-go-v2/credentials v1.20.1 h1:Z8GRNEx0u9sDkZOq4PUnN8mjGwbUQGRzMSXpvt3d8xQ= +github.com/aws/aws-sdk-go-v2/credentials v1.20.1/go.mod h1:uBIK00kFo95dnemqfFMTWx0X8YRqsh6ecIoCjjOkZqM= +github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.19.1 h1:YIEBqcqRnpi4Pfv0YHImtgi6czGCwKHANC7SwmUAVD0= +github.com/aws/aws-sdk-go-v2/feature/ec2/imds v1.19.1/go.mod h1:imEf0oufgAo8KAkCHhrOdqGEC0YWx1PPBQH82shSxGw= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.5.1 h1:pc138gM1CW+XPc60rEwUlwwuwWFQK16CI1T7v1F9Oec= +github.com/aws/aws-sdk-go-v2/internal/configsources v1.5.1/go.mod h1:1+koxpPIbfBdfzP6vojm5/zTpTQ/micYwlxIiNB3TxI= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.8.1 h1:K0JsbZQj+1h208Ro1zHeA4l7bMp0NvRffHQ91q8Ol1s= +github.com/aws/aws-sdk-go-v2/internal/endpoints/v2 v2.8.1/go.mod h1:W3/vL6EtCIatICGy9ab29QhMuae+cOKPWcMxv02CO+Q= +github.com/aws/aws-sdk-go-v2/internal/v4a v1.5.1 h1:yhw5KD1phVyP9vijxOUzDfEtJx+bt+L63k+VfuiYFAA= +github.com/aws/aws-sdk-go-v2/internal/v4a v1.5.1/go.mod h1:ZW2e0d7DYlRxlS9hEiMXE47gTdX5KRN4byUiNbUpG+Q= +github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.19 h1:bAdDl/HkGCcGPoe25ToSHEw23VIxt6CT5fLcg111BKg= +github.com/aws/aws-sdk-go-v2/service/internal/accept-encoding v1.13.19/go.mod h1:KaUzbLxv4CeSxh6ZCl9B4m7CuFenS8kUEaDs+f/DQr4= github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.28 h1:Q1TF1J9jVD+vFo0LzNnmNdQ9EAt52TS+MQlq9Ir+Yxo= github.com/aws/aws-sdk-go-v2/service/internal/checksum v1.9.28/go.mod h1:4KqXXC/p1hrotmouDFbrRoWaLy962b9PMUReCG6+uWo= -github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.35 h1:BBEElKh4a+rKshvjrfpajTe9CbpZvrbb4Jkg2PB7RzA= -github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.13.35/go.mod h1:zaZk983w//8beSruBVec/mr4CmDwgZitW/qzGhAAX0g= +github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.14.1 h1:RmmWQPREQdk9U+PfqeHW3MqZaBaNK7TpV9W3RY+b+7g= +github.com/aws/aws-sdk-go-v2/service/internal/presigned-url v1.14.1/go.mod h1:0A3W4F+68ZnNk5XcNL/e9HFMwnP8RlEicFfy6eOEDyw= github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.36 h1:EUIwBoN+q7UmhAejxgD27APiRjh1vwCFo53gSqdT0BM= github.com/aws/aws-sdk-go-v2/service/internal/s3shared v1.19.36/go.mod h1:6u00gmlTGR6W0b2k9NBrld7MnOEmf1Spqx0VVt6AqyE= github.com/aws/aws-sdk-go-v2/service/s3 v1.106.5 h1:HpN6GgZ3T8pSvRp81ZsgumNjlvRsa+9M0ZL2o6W4uLY= github.com/aws/aws-sdk-go-v2/service/s3 v1.106.5/go.mod h1:5FTZoQxhmLEiCAtYVk6V+t0iS/B5yGZVLZ3Wq5FDJZI= -github.com/aws/aws-sdk-go-v2/service/signin v1.5.4 h1:cOJELVNrq5Q3Udry2GLuHUM7MhwpeaQRdYaoa6GI/yI= -github.com/aws/aws-sdk-go-v2/service/signin v1.5.4/go.mod h1:f4LxzKBtaTxD7xh3PiVg3CE1tchQemfmghaJr+NbK2c= -github.com/aws/aws-sdk-go-v2/service/sso v1.33.4 h1:AMW7a7S8iQaHjBYZdU3PCq4GKRPijTPRAc7e6XtEThY= -github.com/aws/aws-sdk-go-v2/service/sso v1.33.4/go.mod h1:QQNsFV1DVXoXcZt18FS8lI8rtUrlDyAuWZLQ5shunv4= -github.com/aws/aws-sdk-go-v2/service/ssooidc v1.38.4 h1:AsbZcJAQPRmHDJG8K1N0pof/1zPWjVT8TFlTWuGLSvo= -github.com/aws/aws-sdk-go-v2/service/ssooidc v1.38.4/go.mod h1:6imqztH0//t0mKbl6yWl7swSEl7F/w32oAmqB3vP1ag= -github.com/aws/aws-sdk-go-v2/service/sts v1.45.4 h1:w/AryDYMjSUANSQ2uoZxJovUsMTwWJNTv3IMex30Y+4= -github.com/aws/aws-sdk-go-v2/service/sts v1.45.4/go.mod h1:WeBiAa67azG7Su9Vf+ChGDBLiAozJCXzdjXiPBUwtbc= -github.com/aws/smithy-go v1.27.6 h1:0zjT8jgK3jbrTT7JJ3EE6JsMhX8JTrZ+f1sEndYDXrA= -github.com/aws/smithy-go v1.27.6/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc= +github.com/aws/aws-sdk-go-v2/service/signin v1.7.1 h1:mdMtSVKdQ3+mzBh+l0ogrFYZVQUCg6pJZOirA2ARsYE= +github.com/aws/aws-sdk-go-v2/service/signin v1.7.1/go.mod h1:9IqUlsJDbUPcg6cgx3WEzXdjrbWzLDQrak0aaSqlTcI= +github.com/aws/aws-sdk-go-v2/service/sso v1.35.1 h1:B6WFn91tobD6gG4724ONHaqrpKsoETGnv98LHe/yIGM= +github.com/aws/aws-sdk-go-v2/service/sso v1.35.1/go.mod h1:tWuiVBUtPBr8/rgRiYS8Uf85sHcAN+G7XS3D3CEoUh8= +github.com/aws/aws-sdk-go-v2/service/ssooidc v1.40.1 h1:6yeYCWFvgbI2TI3K6jr9LtBNhXgJ7g4xqD+DEiaDDmM= +github.com/aws/aws-sdk-go-v2/service/ssooidc v1.40.1/go.mod h1:naFe83jSMuYkH+QjQPX8n1MLhBkeCFM5Lsnh5m5wz3c= +github.com/aws/aws-sdk-go-v2/service/sts v1.47.1 h1:Sv2xPnRHlThSUtVujYuUBPI/Il8si6UPHXL8DMiB/F0= +github.com/aws/aws-sdk-go-v2/service/sts v1.47.1/go.mod h1:mKo/CzaCz8qytGW70NG4vIIGAx1HXTlb5lHNkC5k3lk= +github.com/aws/smithy-go v1.28.1 h1:R/nXH00c8qcfCzQVELtRw+eLQWtzv+VAIEFJ1/xxXlQ= +github.com/aws/smithy-go v1.28.1/go.mod h1:YE2RhdIuDbA5E5bTdciG9KrW3+TiEONeUWCqxX9i1Fc= github.com/beorn7/perks v0.0.0-20180321164747-3a771d992973/go.mod h1:Dwedo/Wpr24TaqPxmxbtue+5NUziq4I4S80YR8gNf3Q= github.com/beorn7/perks v1.0.0/go.mod h1:KWe93zE9D1o94FZ5RNwFwVgaQK1VOXiVxmqh+CedLV8= github.com/beorn7/perks v1.0.1/go.mod h1:G2ZrVWU2WbWT9wwq4/hrbKbnv/1ERSJQ0ibhJ6rlkpw= github.com/bgentry/speakeasy v0.1.0/go.mod h1:+zsyZBPWlz7T6j88CTgSN5bM796AkVf0kBD4zp0CCIs= +github.com/blinkbean/dingtalk v1.1.3 h1:MbidFZYom7DTFHD/YIs+eaI7kRy52kmWE/sy0xjo6E4= +github.com/blinkbean/dingtalk v1.1.3/go.mod h1:9BaLuGSBqY3vT5hstValh48DbsKO7vaHaJnG9pXwbto= github.com/bodgit/plumbing v1.3.0 h1:pf9Itz1JOQgn7vEOE7v7nlEfBykYqvUYioC61TwWCFU= github.com/bodgit/plumbing v1.3.0/go.mod h1:JOTb4XiRu5xfnmdnDJo6GmSbSbtSyufrsyZFByMtKEs= github.com/bodgit/sevenzip v1.6.5 h1:7H7BxgmeX0j6UX42lH+KXQ92WgMQJ49DoocFdfHbCng= github.com/bodgit/sevenzip v1.6.5/go.mod h1:GhuB6Lq1xCpP1sps+horjZ8lgiKPJcy2zUX3prla9wc= github.com/bodgit/windows v1.0.1 h1:tF7K6KOluPYygXa3Z2594zxlkbKPAOvqr97etrGNIz4= github.com/bodgit/windows v1.0.1/go.mod h1:a6JLwrB4KrTR5hBpp8FI9/9W9jJfeQ2h4XDXU74ZCdM= -github.com/boj/redistore v1.4.2 h1:44FVJnBTdzDV9VpaByCOaQs0ND8hzABD2xBHcAIbX9s= -github.com/boj/redistore v1.4.2/go.mod h1:jjh65GXAH+5lj29pPRnQHdRNOt/lmP0LOaK+fiG2Fu8= +github.com/boj/redistore v1.4.1 h1:lP9ZZWqKMq2RIqexlZX1w1ODSnegL+puxGIujkU5tIw= +github.com/boj/redistore v1.4.1/go.mod h1:c0Tvw6aMjslog4jHIAcNv6EtJM849YoOAhMY7JBbWpI= github.com/bsm/ginkgo/v2 v2.12.0 h1:Ny8MWAHyOepLGlLKYmXG4IEkioBysk6GpaRTLC8zwWs= github.com/bsm/ginkgo/v2 v2.12.0/go.mod h1:SwYbGRRDovPVboqFv0tPTcG1sN61LM1Z4ARdbAV9g4c= github.com/bsm/gomega v1.27.10 h1:yeMWxP2pV2fG3FgAODIY8EiRE3dy0aeFYt4l7wh6yKA= github.com/bsm/gomega v1.27.10/go.mod h1:JyEr/xRbxbtgWNi8tIEVPUYZ5Dzef52k01W3YH0H+O0= +github.com/bwmarrin/discordgo v0.29.0 h1:FmWeXFaKUwrcL3Cx65c20bTRW+vOb6k8AnaP+EgjDno= +github.com/bwmarrin/discordgo v0.29.0/go.mod h1:NJZpH+1AfhIcyQsPeuBKsUtYrRnjkyu0kIVMCHkZtRY= github.com/bwmarrin/snowflake v0.3.0 h1:xm67bEhkKh6ij1790JB83OujPR5CzNe8QuQqAgISZN0= github.com/bwmarrin/snowflake v0.3.0/go.mod h1:NdZxfVWX+oR6y2K0o6qAYv6gIOP9rjG0/E9WsDpxqwE= github.com/bytedance/gopkg v0.1.4 h1:oZnQwnX82KAIWb7033bEwtxvTqXcYMxDBaQxo5JJHWM= @@ -214,8 +218,8 @@ github.com/frankban/quicktest v1.14.6/go.mod h1:4ptaffx2x8+WTWXmUCuVU6aPUX1/Mz7z github.com/fsnotify/fsnotify v1.4.7/go.mod h1:jwhsz4b93w/PPRr/qN1Yymfu8t87LnFCMoQvtojpjFo= github.com/fsnotify/fsnotify v1.4.9/go.mod h1:znqG4EE+3YCdAaPaxE2ZRY/06pZUdp0tY4IgpuI1SZQ= github.com/fsnotify/fsnotify v1.5.4/go.mod h1:OVB6XrOHzAwXMpEM7uPOzcehqUV2UqJxmVXmkdnm1bU= -github.com/fsnotify/fsnotify v1.10.1 h1:b0/UzAf9yR5rhf3RPm9gf3ehBPpf0oZKIjtpKrx59Ho= -github.com/fsnotify/fsnotify v1.10.1/go.mod h1:TLheqan6HD6GBK6PrDWyDPBaEV8LspOxvPSjC+bVfgo= +github.com/fsnotify/fsnotify v1.9.0 h1:2Ml+OJNzbYCTzsxtv8vKSFD9PbJjmhYF14k/jKC7S9k= +github.com/fsnotify/fsnotify v1.9.0/go.mod h1:8jBTzvmWwFyi3Pb8djgCCO5IBqzKJ/Jwo8TRcHyHii0= github.com/gabriel-vasile/mimetype v1.4.15 h1:05iP/CYtZ/w455R/KZM6rZ5ieAdh99UPtd+d3YzLmaI= github.com/gabriel-vasile/mimetype v1.4.15/go.mod h1:azpTcoLcDZRNgFou5j+APrqQx9HqVPWa6ijYQIIVswQ= github.com/ghodss/yaml v1.0.0/go.mod h1:4dBDuWmgqj2HViK6kFavaiC9ZROes6MMH2rRYeMEF04= @@ -227,16 +231,16 @@ github.com/gin-contrib/sse v1.1.1 h1:uGYpNwTacv5R68bSGMapo62iLTRa9l5zxGCps4hK6ko github.com/gin-contrib/sse v1.1.1/go.mod h1:QXzuVkA0YO7o/gun03UI1Q+FTI8ZV/n5t03kIQAI89s= github.com/gin-gonic/gin v1.12.0 h1:b3YAbrZtnf8N//yjKeU2+MQsh2mY5htkZidOM7O0wG8= github.com/gin-gonic/gin v1.12.0/go.mod h1:VxccKfsSllpKshkBWgVgRniFFAzFb9csfngsqANjnLc= -github.com/glebarez/go-sqlite v1.23.0 h1:FyhIq4jqmgphQAUlY79zPldYGwISEZikaDfhiGWkkaI= -github.com/glebarez/go-sqlite v1.23.0/go.mod h1:IIYrOH3L0rHY3jb4IXOHoWdklNajSGUN2eJcvK8WrnI= +github.com/glebarez/go-sqlite v1.21.2 h1:3a6LFC4sKahUunAmynQKLZceZCOzUthkRkEAl9gAXWo= +github.com/glebarez/go-sqlite v1.21.2/go.mod h1:sfxdZyhQjTM2Wry3gVYWaW072Ri1WMdWJi0k6+3382k= github.com/glebarez/sqlite v1.11.0 h1:wSG0irqzP6VurnMEpFGer5Li19RpIRi2qvQz++w0GMw= github.com/glebarez/sqlite v1.11.0/go.mod h1:h8/o8j5wiAsqSPoWELDUdJXhjAhsVliSn7bWZjOhrgQ= github.com/go-acme/lego/v4 v4.35.2 h1:uVQg+KC/yj9R2g7Q9W5wDqhvQvxV5SMu5eqFVoN5xZU= github.com/go-acme/lego/v4 v4.35.2/go.mod h1:pX2jN5n8OphMGY1IaMjYm5DAEzguBaKRt8AvJAgJXpc= github.com/go-faster/city v1.0.1 h1:4WAxSZ3V2Ws4QRDrscLEDcibJY8uf41H6AhXDrNDcGw= github.com/go-faster/city v1.0.1/go.mod h1:jKcUJId49qdW3L1qKHH/3wPeUstCVpVSXTM6vO3VcTw= -github.com/go-faster/errors v0.8.0 h1:9T9eJrM+72dFk7n4DfhuaDDe6cyuFCSW2oNUkN77Yqc= -github.com/go-faster/errors v0.8.0/go.mod h1:5ySTjWFiphBs07IKuiL69nxdfd5+fzh1u7FPGZP2quo= +github.com/go-faster/errors v0.7.1 h1:MkJTnDoEdi9pDabt1dpWf7AA8/BaSYZqibYyhZ20AYg= +github.com/go-faster/errors v0.7.1/go.mod h1:5ySTjWFiphBs07IKuiL69nxdfd5+fzh1u7FPGZP2quo= github.com/go-gl/glfw v0.0.0-20190409004039-e6da0acd62b1/go.mod h1:vR7hzQXu2zJy9AVAgeJqvqgH9Q5CA+iKCZ2gyEVpxRU= github.com/go-gl/glfw/v3.3/glfw v0.0.0-20191125211704-12ad95a8df72/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8= github.com/go-gl/glfw/v3.3/glfw v0.0.0-20200222043503-6f7a984d4dc4/go.mod h1:tQ2UAYgL5IevRw8kRxooKSPJfGvJ9fJQFa0TUsXzTg8= @@ -245,6 +249,8 @@ github.com/go-jose/go-jose/v4 v4.1.4/go.mod h1:x4oUasVrzR7071A4TnHLGSPpNOm2a21K9 github.com/go-kit/kit v0.8.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as= github.com/go-kit/kit v0.9.0/go.mod h1:xBxKIO96dXMWWy0MnWVtmwkA9/13aqxPnvrjFYMA2as= github.com/go-kit/log v0.1.0/go.mod h1:zbhenjAZHb184qTLMA9ZjW7ThYL0H2mk7Q6pNt4vbaY= +github.com/go-lark/lark v1.16.0 h1:U6BwkLM9wrZedSM7cIiMofganr8PCvJN+M75w2lf2Gg= +github.com/go-lark/lark v1.16.0/go.mod h1:6ltbSztPZRT6IaO9ZIQyVaY5pVp/KeMizDYtfZkU+vM= github.com/go-logfmt/logfmt v0.3.0/go.mod h1:Qt1PoO58o5twSAckw1HlFXLmHsOX5/0LbT9GBnD5lWE= github.com/go-logfmt/logfmt v0.4.0/go.mod h1:3RMwSq7FuexP4Kalkev3ejPJsZTpXXBr9+V4qmtdjCk= github.com/go-logfmt/logfmt v0.5.0/go.mod h1:wCYkCAKZfumFQihp8CzCvQ3paCTfi41vtzG1KdI/P7A= @@ -253,33 +259,29 @@ github.com/go-logr/logr v1.4.4 h1:tG4xh9yMsRCAiodLVTxyrkzSZ9+o0L1Kg/+cPVcbP/8= github.com/go-logr/logr v1.4.4/go.mod h1:9T104GzyrTigFIr8wt5mBrctHMim0Nb2HLGrmQ40KvY= github.com/go-logr/stdr v1.2.2 h1:hSWxHoqTgW2S2qGc0LTAI563KZ5YKYRhT3MFKZMbjag= github.com/go-logr/stdr v1.2.2/go.mod h1:mMo/vtBO5dYbehREoey6XUKy/eSumjCCveDpRre4VKE= -github.com/go-openapi/jsonpointer v1.0.0 h1:kR9tHqY0CtZaOPVFm622dPVNhrvYpwr4uCxgL3h1H8s= -github.com/go-openapi/jsonpointer v1.0.0/go.mod h1:Z3rw7dWu1p9IgitXCFamSlA5lmDiklEB6vkaxcNZW5Y= -github.com/go-openapi/jsonreference v1.0.0 h1:jlmTr6torcd1YgDQvSfNmRtKzYDO4FGBkrAdlAVWnpY= -github.com/go-openapi/jsonreference v1.0.0/go.mod h1:jtwdyGbJk0Xhe5Y+rwtglQP6Sb1WZST4rT32LWB+sv0= -github.com/go-openapi/spec v0.22.9 h1:/vKIFDcGKp0ktZWGbym/tJEWbk6/XOEmAVU0kqKMH+w= -github.com/go-openapi/spec v0.22.9/go.mod h1:b/mNUYIOQOyIiUzUzXEE8xzyZqf93KvM9hQGP91yfl0= -github.com/go-openapi/swag v0.28.0 h1:xkgbOSKj6DZziNpyqRRAOt3GJGtgjgsd2RoyT30VWuw= -github.com/go-openapi/swag/conv v0.28.0 h1:GtqqbyFe7vR5Y7ehxG9W6/OvrSFdf1OLeTGp40TqxH8= -github.com/go-openapi/swag/conv v0.28.0/go.mod h1:mbUE+mzctnhxi864m0Q07SpN8OowD9JhxmxuYvZZD/k= -github.com/go-openapi/swag/jsonutils v0.28.0 h1:YIch6FwO7RXzeAnbO8Tu7dWBZeUEH+4nA0HXltVTnv4= -github.com/go-openapi/swag/jsonutils v0.28.0/go.mod h1:CYM3WlTUcagR2ZoHdz54di/cbBqt82tuxuXgAjxw+mg= -github.com/go-openapi/swag/jsonutils/fixtures_test v0.28.0 h1:qV+VVUAx5Oro8WjVWpZeql7YReTKhT4smR4zhcOQZr0= -github.com/go-openapi/swag/jsonutils/fixtures_test v0.28.0/go.mod h1:mofwUWx70wvskwESqRJ//k/9kURmCgyJl5m5Ppoh5kY= -github.com/go-openapi/swag/loading v0.28.0 h1:td8QZdZC9MIYGGSnSPKShKiK22I2tU5UQvuUhIBPRLU= -github.com/go-openapi/swag/loading v0.28.0/go.mod h1:rXB0QiQX5mMveXEA7ouM4KiiM9jVJe4K6BVbwhD1M4k= -github.com/go-openapi/swag/pools v0.28.0 h1:HPMZWSAfce3rdVTFcjFiCIBtDg9h4x2QlRrHipwhxeU= -github.com/go-openapi/swag/pools v0.28.0/go.mod h1:kVQefhSK5RWuRe7BXsL8htgBPAMpN7HDGpGEknqugeE= -github.com/go-openapi/swag/stringutils v0.28.0 h1:ixsc9iYgDPubHL/8nSkbnryEHpD2VRlBMLKpQyPXcDU= -github.com/go-openapi/swag/stringutils v0.28.0/go.mod h1:lzRN95CxXmA03XcDWHLOb6nOMcxCqR5rGY0lOgsfRoM= -github.com/go-openapi/swag/typeutils v0.28.0 h1:nRBKSBXjDgf01VDPB3fWeD9nQuhCOVeIYAkUx2tbkyY= -github.com/go-openapi/swag/typeutils v0.28.0/go.mod h1:Srm0xFNRZ1Y+vCxJclo5qzx8aj+1pAKda/YfFPrG0dQ= -github.com/go-openapi/swag/yamlutils v0.28.0 h1:TV3JXH6DS46KUroDtMLAYHGkdWf5VDq3wVWFirmzROY= -github.com/go-openapi/swag/yamlutils v0.28.0/go.mod h1:x0q/yndZHEgk9Rx3DyDqzFUmHy55KTvIZldvF2dTJXs= -github.com/go-openapi/testify/enable/yaml/v2 v2.6.0 h1:gGHwAJ0R/5jU8BEGDbfRNR3hL68dAVi84WuOApp29B0= -github.com/go-openapi/testify/enable/yaml/v2 v2.6.0/go.mod h1:tY+St1SGq4NFl0QIqdTY4aEdbChAHxhyB77XQi9iJCo= -github.com/go-openapi/testify/v2 v2.6.0 h1:5PKH2HE7YJ/LuRPQGvSxBRlFXNQhSetBLlGAgUEu3ug= -github.com/go-openapi/testify/v2 v2.6.0/go.mod h1:SgsVHtfooshd0tublTtJ50FPKhujf47YRqauXXOUxfw= +github.com/go-openapi/jsonpointer v0.22.1 h1:sHYI1He3b9NqJ4wXLoJDKmUmHkWy/L7rtEo92JUxBNk= +github.com/go-openapi/jsonpointer v0.22.1/go.mod h1:pQT9OsLkfz1yWoMgYFy4x3U5GY5nUlsOn1qSBH5MkCM= +github.com/go-openapi/jsonreference v0.21.2 h1:Wxjda4M/BBQllegefXrY/9aq1fxBA8sI5M/lFU6tSWU= +github.com/go-openapi/jsonreference v0.21.2/go.mod h1:pp3PEjIsJ9CZDGCNOyXIQxsNuroxm8FAJ/+quA0yKzQ= +github.com/go-openapi/spec v0.22.0 h1:xT/EsX4frL3U09QviRIZXvkh80yibxQmtoEvyqug0Tw= +github.com/go-openapi/spec v0.22.0/go.mod h1:K0FhKxkez8YNS94XzF8YKEMULbFrRw4m15i2YUht4L0= +github.com/go-openapi/swag v0.19.15 h1:D2NRCBzS9/pEY3gP9Nl8aDqGUcPFrwG2p+CNFrLyrCM= +github.com/go-openapi/swag/conv v0.25.1 h1:+9o8YUg6QuqqBM5X6rYL/p1dpWeZRhoIt9x7CCP+he0= +github.com/go-openapi/swag/conv v0.25.1/go.mod h1:Z1mFEGPfyIKPu0806khI3zF+/EUXde+fdeksUl2NiDs= +github.com/go-openapi/swag/jsonname v0.25.1 h1:Sgx+qbwa4ej6AomWC6pEfXrA6uP2RkaNjA9BR8a1RJU= +github.com/go-openapi/swag/jsonname v0.25.1/go.mod h1:71Tekow6UOLBD3wS7XhdT98g5J5GR13NOTQ9/6Q11Zo= +github.com/go-openapi/swag/jsonutils v0.25.1 h1:AihLHaD0brrkJoMqEZOBNzTLnk81Kg9cWr+SPtxtgl8= +github.com/go-openapi/swag/jsonutils v0.25.1/go.mod h1:JpEkAjxQXpiaHmRO04N1zE4qbUEg3b7Udll7AMGTNOo= +github.com/go-openapi/swag/jsonutils/fixtures_test v0.25.1 h1:DSQGcdB6G0N9c/KhtpYc71PzzGEIc/fZ1no35x4/XBY= +github.com/go-openapi/swag/jsonutils/fixtures_test v0.25.1/go.mod h1:kjmweouyPwRUEYMSrbAidoLMGeJ5p6zdHi9BgZiqmsg= +github.com/go-openapi/swag/loading v0.25.1 h1:6OruqzjWoJyanZOim58iG2vj934TysYVptyaoXS24kw= +github.com/go-openapi/swag/loading v0.25.1/go.mod h1:xoIe2EG32NOYYbqxvXgPzne989bWvSNoWoyQVWEZicc= +github.com/go-openapi/swag/stringutils v0.25.1 h1:Xasqgjvk30eUe8VKdmyzKtjkVjeiXx1Iz0zDfMNpPbw= +github.com/go-openapi/swag/stringutils v0.25.1/go.mod h1:JLdSAq5169HaiDUbTvArA2yQxmgn4D6h4A+4HqVvAYg= +github.com/go-openapi/swag/typeutils v0.25.1 h1:rD/9HsEQieewNt6/k+JBwkxuAHktFtH3I3ysiFZqukA= +github.com/go-openapi/swag/typeutils v0.25.1/go.mod h1:9McMC/oCdS4BKwk2shEB7x17P6HmMmA6dQRtAkSnNb8= +github.com/go-openapi/swag/yamlutils v0.25.1 h1:mry5ez8joJwzvMbaTGLhw8pXUnhDK91oSJLDPF1bmGk= +github.com/go-openapi/swag/yamlutils v0.25.1/go.mod h1:cm9ywbzncy3y6uPm/97ysW8+wZ09qsks+9RS8fLWKqg= github.com/go-playground/assert/v2 v2.0.1/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4= github.com/go-playground/assert/v2 v2.2.0 h1:JvknZsQTYeFEAhQwI4qEt9cyV5ONwRHC+lYKSsYSR8s= github.com/go-playground/assert/v2 v2.2.0/go.mod h1:VDjEfimB/XKnb+ZQfWdccd7VUvScMdVu0Titje2rxJ4= @@ -293,6 +295,8 @@ github.com/go-playground/validator/v10 v10.4.1/go.mod h1:nlOn6nFhuKACm19sB/8EGNn github.com/go-playground/validator/v10 v10.30.3 h1:4MU6YkEwx7GbcPJOZxrtbu+QfF3pJLJuaYTeAH0DYy8= github.com/go-playground/validator/v10 v10.30.3/go.mod h1:4Axh7oCNGcoGkqLoE4YWt6n20mcEIsPRlB7vPk3lpyc= github.com/go-redis/redis/v8 v8.11.4/go.mod h1:2Z2wHZXdQpCDXEGzqMockDpNyYvi2l4Pxt6RJr792+w= +github.com/go-redis/redis_rate/v10 v10.0.1 h1:calPxi7tVlxojKunJwQ72kwfozdy25RjA0bCj1h0MUo= +github.com/go-redis/redis_rate/v10 v10.0.1/go.mod h1:EMiuO9+cjRkR7UvdvwMO7vbgqJkltQHtwbdIQvaBKIU= github.com/go-resty/resty/v2 v2.6.0/go.mod h1:PwvJS6hvaPkjtjNg9ph+VrSD92bi5Zq73w/BIH7cC3Q= github.com/go-resty/resty/v2 v2.17.2 h1:FQW5oHYcIlkCNrMD2lloGScxcHJ0gkjshV3qcQAyHQk= github.com/go-resty/resty/v2 v2.17.2/go.mod h1:kCKZ3wWmwJaNc7S29BRtUhJwy7iqmn+2mLtQrOyQlVA= @@ -300,6 +304,10 @@ github.com/go-sql-driver/mysql v1.10.0 h1:Q+1LV8DkHJvSYAdR83XzuhDaTykuDx0l6fkXxo github.com/go-sql-driver/mysql v1.10.0/go.mod h1:M+cqaI7+xxXGG9swrdeUIoPG3Y3KCkF0pZej+SK+nWk= github.com/go-stack/stack v1.8.0/go.mod h1:v0f6uXyyMGvRgIKkXu+yp6POWl0qKG85gN/melR3HDY= github.com/go-task/slim-sprig v0.0.0-20210107165309-348f09dbbbc0/go.mod h1:fyg7847qk6SyHyPtNmDHnmrv/HOrqktSC+C9fM+CJOE= +github.com/go-telegram-bot-api/telegram-bot-api v4.6.4+incompatible h1:2cauKuaELYAEARXRkq2LrJ0yDDv1rW7+wrTEdVL3uaU= +github.com/go-telegram-bot-api/telegram-bot-api v4.6.4+incompatible/go.mod h1:qf9acutJ8cwBUhm1bqgz6Bei9/C/c93FPDljKWwsOgM= +github.com/go-test/deep v1.1.1 h1:0r/53hagsehfO4bzD2Pgr/+RgHqhmf+k1Bpse2cTu1U= +github.com/go-test/deep v1.1.1/go.mod h1:5C2ZWiW0ErCdrYzpqxLbTX7MG14M9iiw8DgHncVwcsE= github.com/go-viper/mapstructure/v2 v2.5.0 h1:vM5IJoUAy3d7zRSVtIwQgBj7BiWtMPfmPEgAXnvj1Ro= github.com/go-viper/mapstructure/v2 v2.5.0/go.mod h1:oJDH3BJKyqBA2TXFhDsKDGDTlndYOZ6rGS0BRZIxGhM= github.com/goccy/go-json v0.10.6 h1:p8HrPJzOakx/mn/bQtjgNjdTcN+/S6FcG2CTtQOrHVU= @@ -347,9 +355,8 @@ github.com/golang/snappy v0.0.3/go.mod h1:/XxbfmMg8lxefKM7IXC3fBNl/7bRcc72aCRzEW github.com/gomodule/redigo v1.9.3 h1:dNPSXeXv6HCq2jdyWfjgmhBdqnR6PRO3m/G05nvpPC8= github.com/gomodule/redigo v1.9.3/go.mod h1:KsU3hiK/Ay8U42qpaJk+kuNa3C+spxapWpM+ywhcgtw= github.com/google/btree v0.0.0-20180813153112-4030bb1f1f0c/go.mod h1:lNA+9X1NB3Zf8V7Ke586lFgjr2dZNuvo3lPJSGZ5JPQ= +github.com/google/btree v1.0.0 h1:0udJVsspx3VBr5FwtLhQQtuAsVc79tTq0ocGIPAU6qo= github.com/google/btree v1.0.0/go.mod h1:lNA+9X1NB3Zf8V7Ke586lFgjr2dZNuvo3lPJSGZ5JPQ= -github.com/google/btree v1.1.3 h1:CVpQJjYgC4VbzxeGVHfvZrv1ctoYCAI8vbl07Fcxlyg= -github.com/google/btree v1.1.3/go.mod h1:qOPhT0dTNdNzV6Z/lhRX0YXUafgPLFUh+gZMl761Gm4= github.com/google/go-cmp v0.2.0/go.mod h1:oXzfMopK8JAjlY9xF4vHSVASa0yLyX7SntLO5aqRK0M= github.com/google/go-cmp v0.3.0/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU= github.com/google/go-cmp v0.3.1/go.mod h1:8QqcDgzrUqlUb/G2PQTWiueGozuR1884gddMywk6iLU= @@ -389,8 +396,8 @@ github.com/google/pprof v0.0.0-20210226084205-cbba55b83ad5/go.mod h1:kpwsk12EmLe github.com/google/pprof v0.0.0-20210601050228-01bbb1931b22/go.mod h1:kpwsk12EmLew5upagYY7GY0pfYCcupk39gWOCRROcvE= github.com/google/pprof v0.0.0-20210609004039-a478d1d731e9/go.mod h1:kpwsk12EmLew5upagYY7GY0pfYCcupk39gWOCRROcvE= github.com/google/pprof v0.0.0-20210720184732-4bb14d4b1be1/go.mod h1:kpwsk12EmLew5upagYY7GY0pfYCcupk39gWOCRROcvE= -github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3 h1:LMLX+LgTNWpfvCBdFebv6EsYotImrt/Ppc5cXIriCSo= -github.com/google/pprof v0.0.0-20260802141513-ef3492d7dac3/go.mod h1:jl5iWTm0/hd5PjEYEOuwAJ57L/CibdZfrqZ5XA5GrCk= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e h1:ijClszYn+mADRFY17kjQEVQ1XRhq2/JR1M3sGqeJoxs= +github.com/google/pprof v0.0.0-20250317173921-a4b03ec1a45e/go.mod h1:boTsfXsheKC2y+lKOCMpSfarhxDeIzfZG1jqGcPl3cA= github.com/google/renameio v0.1.0/go.mod h1:KWCgfxg9yswjAJkECMjeO8J8rahYeXnNhOm40UhjYkI= github.com/google/uuid v1.1.2/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= github.com/google/uuid v1.3.0/go.mod h1:TIyPZe4MgqvfeYDBFedMoGGpEw/LqOeaOT+nhxU+yHo= @@ -413,10 +420,12 @@ github.com/gorilla/sessions v1.4.0/go.mod h1:FLWm50oby91+hl7p/wRxDth9bWSuk0qVL2e github.com/gorilla/websocket v1.4.2/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= github.com/gorilla/websocket v1.5.3 h1:saDtZ6Pbx/0u+bgYQ3q96pZgCzfhKXGPqt7kZ72aNNg= github.com/gorilla/websocket v1.5.3/go.mod h1:YR8l580nyteQvAITg2hZ9XVh4b55+EU/adAjf1fMHhE= +github.com/gregdel/pushover v1.4.0 h1:P77WAJ2zPG+b0mEsmMjWGrPMuvhkh9k3v7OviwsoveE= +github.com/gregdel/pushover v1.4.0/go.mod h1:EcaO66Nn1StkpEm1iKtBTV3d2A16SoMsVER1PthX7to= github.com/grpc-ecosystem/go-grpc-prometheus v1.2.0/go.mod h1:8NvIoxWQoOIhqOTXgfV/d3M/q6VIi02HzZEHgUlZvzk= github.com/grpc-ecosystem/grpc-gateway v1.16.0/go.mod h1:BDjrQk3hbvj6Nolgz8mAMFbcEtjT1g+wF4CSlocrBnw= -github.com/grpc-ecosystem/grpc-gateway/v2 v2.30.0 h1:/Tnpcb2E0Pz/tN9s3bfEY2Q8ePCEX9iuS+cneUwncnw= -github.com/grpc-ecosystem/grpc-gateway/v2 v2.30.0/go.mod h1:zOBXOsUaBSjKgmH4OGzV1esUpR3oUSCPYVd2cUBjKYY= +github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0 h1:5VipnvEpbqr2gA2VbM+nYVbkIF28c5ZQfqCBQ5g2xfk= +github.com/grpc-ecosystem/grpc-gateway/v2 v2.29.0/go.mod h1:Hyl3n6Twe1hvtd9XUXDec4pTvgMSEixRuQKPTMH2bNs= github.com/hashicorp/consul/api v1.12.0/go.mod h1:6pVBMo0ebnYdt2S3H87XhekM/HHrUoTD2XXb/VrZVy0= github.com/hashicorp/consul/sdk v0.8.0/go.mod h1:GBvyrGALthsZObzUGsfgHZQDXjg4lOjagTIwIR1vPms= github.com/hashicorp/errwrap v1.0.0/go.mod h1:YH+1FKiLXxHSkmPseP+kNlulaMuP3n2brvKWEqk/Jc4= @@ -468,6 +477,11 @@ github.com/jinzhu/inflection v1.0.0 h1:K317FqzuhWc8YvSVlFMCCUb36O/S9MCKRDI7QkRKD github.com/jinzhu/inflection v1.0.0/go.mod h1:h+uFLlag+Qp1Va5pdKtLDYj+kHp5pxUVkryuEj+Srlc= github.com/jinzhu/now v1.1.5 h1:/o9tlHleP7gOFmsnYNz3RGnqzefHA47wQpKrrdTIwXQ= github.com/jinzhu/now v1.1.5/go.mod h1:d3SSVoowX0Lcu0IBviAWJpolVfI5UJVZZ7cO71lE/z8= +github.com/joho/godotenv v1.3.0/go.mod h1:7hK45KPybAkOC6peb+G5yklZfMxEjkZhHbwpqxOKXbg= +github.com/joho/godotenv v1.5.1 h1:7eLL/+HRGLY0ldzfGMeQkb7vMd0as4CfYvUVzLqw0N0= +github.com/joho/godotenv v1.5.1/go.mod h1:f4LDr5Voq0i2e/R5DDNOoa2zzDfwtkZa6DnEwAbqwq4= +github.com/jordan-wright/email v4.0.1-0.20210109023952-943e75fe5223+incompatible h1:jdpOPRN1zP63Td1hDQbZW73xKmzDvZHzVdNYxhnTMDA= +github.com/jordan-wright/email v4.0.1-0.20210109023952-943e75fe5223+incompatible/go.mod h1:1c7szIrayyPPB/987hsnvNzLushdWf4o/79s3P08L8A= github.com/jpillora/backoff v1.0.0/go.mod h1:J/6gKK9jxlEcS3zixgDgUAsiuZ7yrSoa/FX5e0EB2j4= github.com/json-iterator/go v1.1.6/go.mod h1:+SdeFBvtyEkXs7REEP0seUULqWtbJapLOCVDaaPEHmU= github.com/json-iterator/go v1.1.9/go.mod h1:KdQUCv79m/52Kvf8AW2vK1V8akMuk1QjK/uOdHXbAo4= @@ -482,8 +496,8 @@ github.com/julienschmidt/httprouter v1.2.0/go.mod h1:SYymIcj16QtmaHHD7aYtjjsJG7V github.com/julienschmidt/httprouter v1.3.0/go.mod h1:JR6WtHb+2LUe8TCKY3cZOxFyyO8IZAc4RVcycCCAKdM= github.com/kisielk/errcheck v1.5.0/go.mod h1:pFxgyoBC7bSaBwPgfKdkLd5X25qrDl4LWUI2bnpBCr8= github.com/kisielk/gotool v1.0.0/go.mod h1:XhKaO+MFFWcvkIS/tQcRk01m1F5IRFswLeQ+oQHNcck= -github.com/klauspost/compress v1.19.2 h1:hMRETovs/pu/dVWN7zIT1PGG8t509MwT6bO7XSi26R8= -github.com/klauspost/compress v1.19.2/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= +github.com/klauspost/compress v1.19.1 h1:VsB4HPswih7mmZ8WleSFQ75c/Ui1M4trX5oAsJnhSlk= +github.com/klauspost/compress v1.19.1/go.mod h1:cwPg85FWrGar70rWktvGQj8/hthj3wpl0PGDogxkrSQ= github.com/klauspost/cpuid/v2 v2.4.0 h1:S6Hrbc7+ywsr0r+RLapfGBHfyefhCTwEh3A0tV913Dw= github.com/klauspost/cpuid/v2 v2.4.0/go.mod h1:19jmZ9mjzoF//ddRSUsv0zfBTJWh3QJh9FNxZTMrGxU= github.com/konsorten/go-windows-terminal-sequences v1.0.1/go.mod h1:T0+1ngSBFLxvqU3pZ+m/2kptfBszLMUkC4ZK/EgS/cQ= @@ -548,6 +562,8 @@ github.com/mwitkow/go-conntrack v0.0.0-20161129095857-cc309e4a2223/go.mod h1:qRW github.com/mwitkow/go-conntrack v0.0.0-20190716064945-2f068394615f/go.mod h1:qRWi+5nqEBWmkhHvq77mSJWrCKwh8bxhgT7d/eI7P4U= github.com/ncruces/go-strftime v1.0.0 h1:HMFp8mLCTPp341M/ZnA4qaf7ZlsbTc+miZjCLOFAw7w= github.com/ncruces/go-strftime v1.0.0/go.mod h1:Fwc5htZGVVkseilnfgOVb9mKy6w1naJmn9CehxcKcls= +github.com/nikoksr/notify v1.6.0 h1:H9pyvNyo4/47vn7uPTBbb+ld6Y/KmpXo9Ghf8cGqAlU= +github.com/nikoksr/notify v1.6.0/go.mod h1:GBrx8S2GI0ZtXdobxNmXJjPJN4P6qXbjPW78zlLSg6s= github.com/nxadm/tail v1.4.4/go.mod h1:kenIhsEOeOJmVchQTgglprH7qJGnHDVpk1VPCcaMI8A= github.com/nxadm/tail v1.4.8/go.mod h1:+ncqLTQzXmGhMZNUePPaPqPvBxHAIsmXswZKocGu+AU= github.com/onsi/ginkgo v1.6.0/go.mod h1:lLunBs/Ym6LB5Z9jYTR76FiuTmxDTDusOGeTQH+WWjE= @@ -568,16 +584,14 @@ github.com/pelletier/go-toml/v2 v2.4.3 h1:GTRvJQutkOSftxIFD5xw9aepkYNuPWmVJpffdD github.com/pelletier/go-toml/v2 v2.4.3/go.mod h1:2gIqNv+qfxSVS7cM2xJQKtLSTLUE9V8t9Stt+h56mCY= github.com/peterbourgon/diskv/v3 v3.0.1 h1:x06SQA46+PKIUftmEujdwSEpIx8kR+M9eLYsUxeYveU= github.com/peterbourgon/diskv/v3 v3.0.1/go.mod h1:kJ5Ny7vLdARGU3WUuy6uzO6T0nb/2gWcT1JiBvRmb5o= -github.com/pierrec/lz4/v4 v4.1.28 h1:pPEPwRJ4kybBTfGt28q7lQsRJQHhC08axprdLD5Ppio= -github.com/pierrec/lz4/v4 v4.1.28/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4= +github.com/pierrec/lz4/v4 v4.1.27 h1:+PhzhWDrjRj89TH2sw43nE3+4+W8lSxIuQadEHZyjUk= +github.com/pierrec/lz4/v4 v4.1.27/go.mod h1:EoQMVJgeeEOMsCqCzqFm2O0cJvljX2nGZjcRIPL34O4= github.com/pkg/diff v0.0.0-20210226163009-20ebb0f2a09e/go.mod h1:pJLUxLENpZxwdsKMEsNbx1VGcRFpLqf3715MtcvvzbA= github.com/pkg/errors v0.8.0/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pkg/errors v0.8.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pkg/errors v0.9.1/go.mod h1:bwawxfHBFNV+L2hUp1rHADufV3IMtnDRdf1r5NINEl0= github.com/pkg/sftp v1.13.1/go.mod h1:3HaPG6Dq1ILlpPZRO0HVMrsydcdLt6HRDccSgb87qRg= github.com/pmezard/go-difflib v1.0.0/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= -github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2 h1:Jamvg5psRIccs7FGNTlIRMkT8wgtp5eCXdBlqhYGL6U= -github.com/pmezard/go-difflib v1.0.1-0.20181226105442-5d4384ee4fb2/go.mod h1:iKH77koFhYxTK1pcRnkKkqfTogsbg7gZNVY4sRDYZ/4= github.com/posener/complete v1.1.1/go.mod h1:em0nMJCgc9GFtwrmVmEMR/ZL6WyhyjMBndrE9hABlRI= github.com/posener/complete v1.2.3/go.mod h1:WZIdtGGp+qx0sLrYKtIRAruyNpv6hFCicSgv7Sy7s/s= github.com/pressly/goose/v3 v3.27.3 h1:pIglVHjw99r4e/hDHHwbl9vfOsDMqUokfkXo6+n/RxA= @@ -637,6 +651,8 @@ github.com/shopspring/decimal v1.4.0/go.mod h1:gawqmDU56v4yIKSwfBSFip1HdCCXN8/+D github.com/sirupsen/logrus v1.2.0/go.mod h1:LxeOpSwHxABJmUn/MG1IvRgCAasNZTLOkJPxbbu5VWo= github.com/sirupsen/logrus v1.4.2/go.mod h1:tLMulIdttU9McNUspp0xgXVQah82FyeX6MwdIuYE2rE= github.com/sirupsen/logrus v1.6.0/go.mod h1:7uNnSEd1DgxDLC74fIahvMZmmYsHGZGEOFrfsX/uA88= +github.com/slack-go/slack v0.29.0 h1:ohhMNgp9DmPKiLhH/pNZV4NxhOXKgNy0SH8FzVHNerI= +github.com/slack-go/slack v0.29.0/go.mod h1:UEe+jmo9WLlwHB04qsOrTDvqM7Aa4rQL3O5wF3n0hx4= github.com/spaolacci/murmur3 v0.0.0-20180118202830-f09979ecbc72/go.mod h1:JwIasOWyU6f++ZhiEuf87xNszmSA2myDM2Kzu9HwQUA= github.com/spf13/afero v1.8.2/go.mod h1:CtAatgMJh6bJEIs48Ay/FOnkljP3WeGUG0MC1RfAqwo= github.com/spf13/afero v1.15.0 h1:b/YBCLWAJdFWJTN9cLhiXXcD7mzKn9Dm86dNnfyQw1I= @@ -675,8 +691,8 @@ github.com/stretchr/testify v1.8.0/go.mod h1:yNjHg4UonilssWZ8iaSj1OCr/vHnekPRkoO github.com/stretchr/testify v1.8.1/go.mod h1:w2LPCIKwWwSfY2zedu0+kehJoqGctiVI29o6fzry7u4= github.com/stretchr/testify v1.8.4/go.mod h1:sz/lmYIOXD/1dqDmKjjqLyZ2RngseejIcXlSw2iwfAo= github.com/stretchr/testify v1.10.0/go.mod h1:r2ic/lqez/lEtzL7wO/rwa5dbSLXVDPFyf8C91i36aY= -github.com/stretchr/testify v1.11.1 h1:7s2iGBzp5EwR7/aIZr8ao5+dra3wiQyKjjFuvgVKu7U= -github.com/stretchr/testify v1.11.1/go.mod h1:wZwfW3scLgRK+23gO65QZefKpKQRnfz6sD981Nm4B6U= +github.com/stretchr/testify v1.12.1 h1:EuwCh5fleGS7H32xRwO3wRGT7DxrDhLAT6FF8MpWDWE= +github.com/stretchr/testify v1.12.1/go.mod h1:MDEgiDPPsNp5cuIrHPPCyornHKgEVbtFUmoNlxoYthg= github.com/studio-b12/gowebdav v0.13.0 h1:OcwSg6IQHOFNdYHn3bPOHwSE8looG8N56Y5xTT1asqQ= github.com/studio-b12/gowebdav v0.13.0/go.mod h1:bHA7t77X/QFExdeAnDzK6vKM34kEZAcE1OX4MfiwjkE= github.com/subosito/gotenv v1.4.1/go.mod h1:ayKnFf/c6rvx/2iiLrJUk1e6plDbT3edrFNGqEflhK0= @@ -688,25 +704,32 @@ github.com/swaggo/gin-swagger v1.6.1 h1:Ri06G4gc9N4t4k8hekMigJ9zKTFSlqj/9paAQCQs github.com/swaggo/gin-swagger v1.6.1/go.mod h1:LQ+hJStHakCWRiK/YNYtJOu4mR2FP+pxLnILT/qNiTw= github.com/swaggo/swag v1.16.6 h1:qBNcx53ZaX+M5dxVyTrgQ0PJ/ACK+NzhwcbieTt+9yI= github.com/swaggo/swag v1.16.6/go.mod h1:ngP2etMK5a0P3QBizic5MEwpRmluJZPHjXcMoj4Xesg= +github.com/technoweenie/multipartstreamer v1.0.1 h1:XRztA5MXiR1TIRHxH2uNxXxaIkKQDeX7m2XsSOlQEnM= +github.com/technoweenie/multipartstreamer v1.0.1/go.mod h1:jNVxdtShOxzAsukZwTSw6MDx5eUJoiEBsSvzDU9uzog= github.com/tencent-connect/botgo v0.2.1 h1:+BrTt9Zh+awL28GWC4g5Na3nQaGRWb0N5IctS8WqBCk= github.com/tencent-connect/botgo v0.2.1/go.mod h1:oO1sG9ybhXNickvt+CVym5khwQ+uKhTR+IhTqEfOVsI= -github.com/tidwall/gjson v1.9.3 h1:hqzS9wAHMO+KVBBkLxYdkEeeFHuqr95GfClRLKlgK0E= github.com/tidwall/gjson v1.9.3/go.mod h1:/wbyibRr2FHMks5tjHJ5F8dMZh3AcwJEMf5vlfC0lxk= -github.com/tidwall/match v1.1.1 h1:+Ho715JplO36QYgwN9PGYNhgZvoUSc9X2c80KVTi+GA= +github.com/tidwall/gjson v1.19.0 h1:xwxm7n691Uf3u5OFjzngavjGTh55KX5q/9w9xHW88JU= +github.com/tidwall/gjson v1.19.0/go.mod h1:V37/opeE/JbLUOfH0QTXiNez2l0RUjYUhpT4szFQAfc= github.com/tidwall/match v1.1.1/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= -github.com/tidwall/pretty v1.2.0 h1:RWIZEg2iJ8/g6fDDYzMpobmaoGh5OLl4AXtGUGPcqCs= +github.com/tidwall/match v1.2.0 h1:0pt8FlkOwjN2fPt4bIl4BoNxb98gGHN2ObFEDkrfZnM= +github.com/tidwall/match v1.2.0/go.mod h1:eRSPERbgtNPcGhD8UCthc6PmLEQXEWd3PRB5JTxsfmM= github.com/tidwall/pretty v1.2.0/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= +github.com/tidwall/pretty v1.2.1 h1:qjsOFOWWQl+N3RsoF5/ssm1pHmJJwhjlSbZ51I6wMl4= +github.com/tidwall/pretty v1.2.1/go.mod h1:ITEVvHYasfjBbM0u2Pg8T2nJnzm8xPwvNhhsoaGGjNU= github.com/tv42/httpunix v0.0.0-20150427012821-b75d8614f926/go.mod h1:9ESjWnEqriFuLhtthL60Sar/7RFoluCcXsuvEwTV5KM= github.com/twitchyliquid64/golang-asm v0.15.1 h1:SU5vSMR7hnwNxj24w34ZyCi/FmDZTkS4MhqMhdFk5YI= github.com/twitchyliquid64/golang-asm v0.15.1/go.mod h1:a1lVb/DtPvCB8fslRZhAngC2+aY1QWCk3Cedj/Gdt08= -github.com/ugorji/go/codec v1.3.2 h1:zkEASHHyEClGeURfgNT9PJZVfAbs9oEX9QXggwWNJbc= -github.com/ugorji/go/codec v1.3.2/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4= +github.com/ugorji/go/codec v1.3.1 h1:waO7eEiFDwidsBN6agj1vJQ4AG7lh2yqXyOXqhgQuyY= +github.com/ugorji/go/codec v1.3.1/go.mod h1:pRBVtBSKl77K30Bv8R2P+cLSGaTtex6fsA2Wjqmfxj4= github.com/ulikunitz/xz v0.5.16 h1:ld6NyySjx5lowVKwJvMRLnW5nxKX/xnpSiFYZ/Lxur0= github.com/ulikunitz/xz v0.5.16/go.mod h1:H9Rt/W6/Qj27PGauhQc6nfCDy7vHpzsOThBSaYDoEhw= github.com/uptrace/opentelemetry-go-extra/otelutil v0.3.2 h1:3/aHKUq7qaFMWxyQV0W2ryNgg8x8rVeKVA20KJUkfS0= github.com/uptrace/opentelemetry-go-extra/otelutil v0.3.2/go.mod h1:Zit4b8AQXaXvA68+nzmbyDzqiyFRISyw1JiD5JqUBjw= github.com/uptrace/opentelemetry-go-extra/otelzap v0.3.2 h1:cj/Z6FKTTYBnstI0Lni9PA+k2foounKIPUmj1LBwNiQ= github.com/uptrace/opentelemetry-go-extra/otelzap v0.3.2/go.mod h1:LDaXk90gKEC2nC7JH3Lpnhfu+2V7o/TsqomJJmqA39o= +github.com/wneessen/go-mail v0.8.1 h1:tVcncj02/QySVFw3zr/kXOzZcuFQqBNT6K+Rbgm/pcM= +github.com/wneessen/go-mail v0.8.1/go.mod h1:dWZ61zadzCIyvB4y1/YzC5O7MrbbzBfPkARmbosdf8w= github.com/xyproto/randomstring v1.0.5 h1:YtlWPoRdgMu3NZtP45drfy1GKoojuR7hmRcnhZqKjWU= github.com/xyproto/randomstring v1.0.5/go.mod h1:rgmS5DeNXLivK7YprL0pY+lTuhNQW3iGxZ18UQApw/E= github.com/yuin/goldmark v1.1.25/go.mod h1:3hX8gzYuyVAZsxl0MRgGTJEmQBFcNTphYh9decYSb74= @@ -748,8 +771,8 @@ go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.45.0 h1:fG5MC go.opentelemetry.io/otel/exporters/otlp/otlptrace/otlptracegrpc v1.45.0/go.mod h1:BmAYTn+3ysbRe+IU2msxmf5Rx3g6DHvex+tWI3LdhYI= go.opentelemetry.io/otel/exporters/stdout/stdouttrace v1.45.0 h1:lsA/S1bxgdbyFGkTj+3meEdJ6ADVU7QoFstV6MXgE68= go.opentelemetry.io/otel/exporters/stdout/stdouttrace v1.45.0/go.mod h1:L7u+MirGoB1bjeLH66+xDykF4RC8C3RN7lIFpBiewUo= -go.opentelemetry.io/otel/log v0.6.0 h1:nH66tr+dmEgW5y+F9LanGJUBYPrRgP4g2EkmPE3LeK8= -go.opentelemetry.io/otel/log v0.6.0/go.mod h1:KdySypjQHhP069JX0z/t26VHwa8vSwzgaKmXtIB3fJM= +go.opentelemetry.io/otel/log v0.12.2 h1:yob9JVHn2ZY24byZeaXpTVoPS6l+UrrxmxmPKohXTwc= +go.opentelemetry.io/otel/log v0.12.2/go.mod h1:ShIItIxSYxufUMt+1H5a2wbckGli3/iCfuEbVZi/98E= go.opentelemetry.io/otel/metric v1.45.0 h1:7Eg1uH7CJ5cXv9is6tnBe1FI6rj1nwUdbFypRm3br/M= go.opentelemetry.io/otel/metric v1.45.0/go.mod h1:HAPbm1nd3p1PmFH7v2dR+6BjXxw+Lq4a2+pndMAm08s= go.opentelemetry.io/otel/sdk v1.45.0 h1:4VVSMgQ83dUgW2aoX5f6JgLvHwIvzcuLnF9lUdCSpCw= @@ -793,8 +816,8 @@ golang.org/x/crypto v0.0.0-20210921155107-089bfa567519/go.mod h1:GvvjBRRGRdwPK5y golang.org/x/crypto v0.0.0-20211108221036-ceb1ce70b4fa/go.mod h1:GvvjBRRGRdwPK5ydBHafDWAxML/pGHZbMvKqRZ5+Abc= golang.org/x/crypto v0.0.0-20220411220226-7b82a4e95df4/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4= golang.org/x/crypto v0.16.0/go.mod h1:gCAAfMLgwOJRpTjQ2zCCt2OcSfYMTeZVSRtQlPC7Nq4= -golang.org/x/crypto v0.54.0 h1:YLIA59K4fiNzHzjnZt2tUJQjQtUWfWbeHBqKtk3eScw= -golang.org/x/crypto v0.54.0/go.mod h1:KWL8ny2AZdGR2cWmzeHrp2azQPGogOv+HeQaVEXC2dk= +golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M= +golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis= golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= golang.org/x/exp v0.0.0-20190306152737-a1d7652674e8/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= golang.org/x/exp v0.0.0-20190510132918-efd6b22b2522/go.mod h1:ZjyILWgesfNpC6sMxTJOJm9Kp84zZh5NQWvqDGG3Qr8= @@ -891,8 +914,8 @@ golang.org/x/net v0.6.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= golang.org/x/net v0.7.0/go.mod h1:2Tu9+aMcznHK/AK1HMvgo6xiTLG5rD5rZLDS+rp2Bjs= golang.org/x/net v0.10.0/go.mod h1:0qNGK6F8kojg2nk9dLZ2mShWaEBan6FAoqfSigmmuDg= golang.org/x/net v0.19.0/go.mod h1:CfAk/cbD4CthTvqiEl8NpboMuiuOYsAr/7NOjZJtv1U= -golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= -golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= +golang.org/x/net v0.58.0 h1:ynWG7rqYi4ccpTEuPZ2QGWHktVEM9DMCj9yzDE0Q7To= +golang.org/x/net v0.58.0/go.mod h1:YwCddHnFlT7eLQqVprV19OnhLGtc5xOKgE0RyqgfWAU= golang.org/x/oauth2 v0.0.0-20180821212333-d2e6202438be/go.mod h1:N/0e6XlmueqKjAGxoOufVs8QHGRruUQn6yWY3a++T0U= golang.org/x/oauth2 v0.0.0-20190226205417-e64efc72b421/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw= golang.org/x/oauth2 v0.0.0-20190604053449-0f29369cfe45/go.mod h1:gOpvHmFTYa4IltrdGE7lF6nIHvwfUNPOp7c8zoXwtLw= @@ -1038,8 +1061,8 @@ golang.org/x/text v0.3.7/go.mod h1:u+2+/6zg+i71rQMx5EYifcz6MCKuco9NR6JIITiCfzQ= golang.org/x/text v0.7.0/go.mod h1:mrYo+phRRbMaCq/xk9113O4dZlRixOauAjOtrjsXDZ8= golang.org/x/text v0.9.0/go.mod h1:e1OnstbJyHTd6l/uOt8jFFHp6TRDWZR/bV3emEE/zU8= golang.org/x/text v0.14.0/go.mod h1:18ZOQIKpY8NJVqYksKHtTdi31H5itFRjB5/qKTNYzSU= -golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= -golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= golang.org/x/time v0.0.0-20181108054448-85acf8d2951c/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ= golang.org/x/time v0.0.0-20190308202827-9d24e82272b4/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ= golang.org/x/time v0.0.0-20191024005414-555d28b269f0/go.mod h1:tRJNPiyCQ0inRvYxbN9jk5I+vvW/OXSQhTDSoE431IQ= @@ -1239,8 +1262,8 @@ google.golang.org/genproto v0.0.0-20220505152158-f39f71e6c8f3/go.mod h1:RAyBrSAP google.golang.org/genproto v0.0.0-20220519153652-3a47de7e79bd/go.mod h1:RAyBrSAP7Fh3Nc84ghnVLDPuV51xc9agzmm4Ph6i0Q4= google.golang.org/genproto/googleapis/api v0.0.0-20260803160001-6ac0973c030d h1:FarXi840EJWSHYTN3ERkADbPWjl307+FGrA22KAVjjc= google.golang.org/genproto/googleapis/api v0.0.0-20260803160001-6ac0973c030d/go.mod h1:K/+WGbmBY7aNW1HDw1fJnKYo10i0DkAX6pows00dLig= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260803160001-6ac0973c030d h1:IL4hdHzcUv2l/gcg98/Rj3FbtE6axwqslOW8SW0C+S0= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260803160001-6ac0973c030d/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260819154853-08b0e4226688 h1:cYNAzI2sUwhmCcoj9TxvihSrqsxt6uIkj3rDRhSDmW4= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260819154853-08b0e4226688/go.mod h1:DjtHYE8FKJLivXcBEjGwndXfIC23G0VpXiXKqG179uA= google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c= google.golang.org/grpc v1.20.1/go.mod h1:10oTOabMzJvdu6/UiuZezV6QK5dSlG84ov/aaiqXj38= google.golang.org/grpc v1.21.1/go.mod h1:oYelfM1adQP15Ek0mdvEgi9Df8B9CZIaU1084ijfRaM= @@ -1271,8 +1294,8 @@ google.golang.org/grpc v1.44.0/go.mod h1:k+4IHHFw41K8+bbowsex27ge2rCb65oeWqe4jJ5 google.golang.org/grpc v1.45.0/go.mod h1:lN7owxKUQEqMfSyQikvvk5tf/6zMPsrK+ONuO11+0rQ= google.golang.org/grpc v1.46.0/go.mod h1:vN9eftEi1UMyUsIF80+uQXhHjbXYbm0uXoFCACuMGWk= google.golang.org/grpc v1.46.2/go.mod h1:vN9eftEi1UMyUsIF80+uQXhHjbXYbm0uXoFCACuMGWk= -google.golang.org/grpc v1.83.0 h1:JeNZEKJFbQxArAMl+hiytHauacDNqJUllNfmIMmpqnQ= -google.golang.org/grpc v1.83.0/go.mod h1:kDyl6SKsiHKt0uylY5gtn5cEjkrIOhQOGDgIc4JGwzQ= +google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU= +google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8= google.golang.org/grpc/cmd/protoc-gen-go-grpc v1.1.0/go.mod h1:6Kw0yEErY5E/yWrBtf03jp27GLLJujG4z/JK95pnjjw= google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8= google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0= @@ -1288,13 +1311,12 @@ google.golang.org/protobuf v1.26.0-rc.1/go.mod h1:jlhhOSvTdKEhbULTjvd4ARK9grFBp0 google.golang.org/protobuf v1.26.0/go.mod h1:9q0QmTI4eRPtz6boOQmLYwt+qCgq0jsYwAQnmE0givc= google.golang.org/protobuf v1.27.1/go.mod h1:9q0QmTI4eRPtz6boOQmLYwt+qCgq0jsYwAQnmE0givc= google.golang.org/protobuf v1.28.0/go.mod h1:HV8QOd/L58Z+nl8r43ehVNZIU/HEI6OcFqwMG9pJV4I= -google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= -google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= +google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/alecthomas/kingpin.v2 v2.2.6/go.mod h1:FMv+mEhP44yOT+4EoQTLFTRgOQ1FBLkstjWtayDeSgw= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= -gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c h1:Hei/4ADfdWqJk1ZMxUNpqntNwaWcugrBjAiHlqqRiVk= gopkg.in/check.v1 v1.0.0-20201130134442-10cb98267c6c/go.mod h1:JHkPIbrfpd72SG/EVd6muEfDQjcINNoR0C8j2r3qZ4Q= gopkg.in/errgo.v2 v2.1.0/go.mod h1:hNsd1EY+bozCKY1Ytp96fpM3vjJbqLJn88ws8XvfDNI= gopkg.in/fsnotify.v1 v1.4.7/go.mod h1:Tz8NjZHkW78fSQdbUxIjBTcgA1z1m8ZHf0WmKUhAMys= @@ -1314,7 +1336,6 @@ gopkg.in/yaml.v2 v2.3.0/go.mod h1:hI93XBmqTisBFMUTm0b8Fm+jr3Dg1NNxqwp+5A1VGuI= gopkg.in/yaml.v2 v2.4.0/go.mod h1:RDklbk79AGWmwhnvt/jBztapEOGDOx6ZbXqjP6csGnQ= gopkg.in/yaml.v3 v3.0.0-20200313102051-9f266ea9e77c/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gopkg.in/yaml.v3 v3.0.0-20210107192922-496545a6307b/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= -gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA= gopkg.in/yaml.v3 v3.0.1/go.mod h1:K4uyk7z7BCEPqu6E+C64Yfv1cQ7kz7rIZviUmN+EgEM= gorm.io/driver/clickhouse v0.7.0 h1:BCrqvgONayvZRgtuA6hdya+eAW5P2QVagV3OlEp1vtA= gorm.io/driver/clickhouse v0.7.0/go.mod h1:TmNo0wcVTsD4BBObiRnCahUgHJHjBIwuRejHwYt3JRs= @@ -1328,8 +1349,8 @@ gorm.io/gorm v1.31.2 h1:3o8FXNo9v9S858gil+3LlZA1LkCOzgb4g5BL64FgaCo= gorm.io/gorm v1.31.2/go.mod h1:XyQVbO2k6YkOis7C2437jSit3SsDK72s7n7rsSHd+Gs= gorm.io/plugin/dbresolver v1.6.2 h1:F4b85TenghUeITqe3+epPSUtHH7RIk3fXr5l83DF8Pc= gorm.io/plugin/dbresolver v1.6.2/go.mod h1:tctw63jdrOezFR9HmrKnPkmig3m5Edem9fdxk9bQSzM= -gorm.io/plugin/opentelemetry v0.1.16 h1:Kypj2YYAliJqkIczDZDde6P6sFMhKSlG5IpngMFQGpc= -gorm.io/plugin/opentelemetry v0.1.16/go.mod h1:P3RmTeZXT+9n0F1ccUqR5uuTvEXDxF8k2UpO7mTIB2Y= +gorm.io/plugin/opentelemetry v0.1.14 h1:xivP39t/0JgcceDl+BLwVAJHihjFEUj0ZocMSBwZ7ZY= +gorm.io/plugin/opentelemetry v0.1.14/go.mod h1:ZAp4v5vU1CCcK9Oo8/va5rl6NStrzpSU+a70evd+W/g= honnef.co/go/tools v0.0.0-20190102054323-c2f93a96b099/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= honnef.co/go/tools v0.0.0-20190106161140-3f1c8253044a/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= honnef.co/go/tools v0.0.0-20190418001031-e561f6794a2a/go.mod h1:rf3lG4BRIbNafJWhAfAdb/ePZxsR/4RtNHQocxwk9r4= @@ -1349,8 +1370,8 @@ modernc.org/gc/v3 v3.1.4 h1:2g65LGVSmFQrXeITAw97x7hCRvZFcyE1uDP+7Vng7JI= modernc.org/gc/v3 v3.1.4/go.mod h1:HFK/6AGESC7Ex+EZJhJ2Gni6cTaYpSMmU/cT9RmlfYY= modernc.org/goabi0 v0.2.0 h1:HvEowk7LxcPd0eq6mVOAEMai46V+i7Jrj13t4AzuNks= modernc.org/goabi0 v0.2.0/go.mod h1:CEFRnnJhKvWT1c1JTI3Avm+tgOWbkOu5oPA8eH8LnMI= -modernc.org/libc v1.74.4 h1:fX1Omw4o2/1C2iRkkIsrQTasJQldLhRmuPreXLoWs9k= -modernc.org/libc v1.74.4/go.mod h1:eeQAS9W3sZeKYMFubydxJpII9ybHWshk+7or7bLG9co= +modernc.org/libc v1.74.3 h1:a4J+Z8aVaxPyjyxRAdJzw246PqpcFGvVPnfT/AuM5Ws= +modernc.org/libc v1.74.3/go.mod h1:4H7h/MJ8wnjL8RAbp9v3OXgnk22X7MouHIhDbvP3gj4= modernc.org/mathutil v1.7.1 h1:GCZVGXdaN8gTqB1Mf/usp1Y/hSqgI2vAGGP4jZMCxOU= modernc.org/mathutil v1.7.1/go.mod h1:4p5IwJITfppl0G4sUEDtCr4DthTaT47/N3aT6MhfgJg= modernc.org/memory v1.11.0 h1:o4QC8aMQzmcwCK3t3Ux/ZHmwFPzE6hf2Y5LbkRs+hbI= @@ -1359,8 +1380,8 @@ modernc.org/opt v0.2.0 h1:tGyef5ApycA7FSEOMraay9SaTk5zmbx7Tu+cJs4QKZg= modernc.org/opt v0.2.0/go.mod h1:03fq9lsNfvkYSfxrfUhZCWPk1lm4cq4N+Bh//bEtgns= modernc.org/sortutil v1.2.1 h1:+xyoGf15mM3NMlPDnFqrteY07klSFxLElE2PVuWIJ7w= modernc.org/sortutil v1.2.1/go.mod h1:7ZI3a3REbai7gzCLcotuw9AC4VZVpYMjDzETGsSMqJE= -modernc.org/sqlite v1.56.0 h1:/D8e2RfFqoy/Zc6PuC76U28zFwmI/sYx1Kjm4yEn9e0= -modernc.org/sqlite v1.56.0/go.mod h1:yCJ2cmAaIkHQ25oXWrF8H4O1lIfPYPR26yCEDj2P3pQ= +modernc.org/sqlite v1.54.0 h1:JCxR4qwkJvOaqAoYcgDoO25Nc+ROg6EJ2LfBVzdrgog= +modernc.org/sqlite v1.54.0/go.mod h1:4ntCLuNmnH8+GNqjka1wNg7KJd5/Hi5FYp8K+XQ7GZw= modernc.org/strutil v1.2.1 h1:UneZBkQA+DX2Rp35KcM69cSsNES9ly8mQWD71HKlOA0= modernc.org/strutil v1.2.1/go.mod h1:EHkiggD70koQxjVdSBM3JKM7k6L0FbGE5eymy9i3B9A= modernc.org/token v1.1.0 h1:Xl7Ap9dKaEs5kLoOQeQmPWevfnk/DM5qcLcYlA8ys6Y= diff --git a/backend/openflare/plugins/flared/flared/runner.go b/backend/openflare/plugins/flared/flared/runner.go index 572fa076..1512e397 100644 --- a/backend/openflare/plugins/flared/flared/runner.go +++ b/backend/openflare/plugins/flared/flared/runner.go @@ -15,6 +15,7 @@ import ( "Wavelet/openflare/plugins/flared/sync" "Wavelet/openflare/plugins/flared/wsclient" edgerunner "Wavelet/openflare/share/edge/runner" + "Wavelet/pkg/util" ) // Runner is the top-level orchestrator for the flared agent. It wires together @@ -31,8 +32,8 @@ type Runner struct { // Run starts all background services and enters the WebSocket reconnect loop. // It blocks until ctx is cancelled or an unrecoverable error occurs. func (r *Runner) Run(ctx context.Context) error { - go r.HeartbeatService.Run(ctx) - go r.SyncService.Run(ctx) + util.Go(func() { r.HeartbeatService.Run(ctx) }) + util.Go(func() { r.SyncService.Run(ctx) }) return edgerunner.RunWSReconnectLoop(ctx, edgerunner.WSReconnectConfig{ ComponentName: "flared", diff --git a/backend/openflare/plugins/relay/frps/manager.go b/backend/openflare/plugins/relay/frps/manager.go index a65f81aa..121aade7 100644 --- a/backend/openflare/plugins/relay/frps/manager.go +++ b/backend/openflare/plugins/relay/frps/manager.go @@ -21,6 +21,7 @@ import ( "time" service "Wavelet/openflare/share/protocol" + "Wavelet/pkg/util" ) const ( @@ -143,7 +144,7 @@ func (m *Manager) UpdateConfig(ctx context.Context, cfg *service.RelayConfig) { m.lastError = err.Error() return } - go m.supervise(ctx, generation) + util.Go(func() { m.supervise(ctx, generation) }) } return } @@ -167,7 +168,7 @@ func (m *Manager) UpdateConfig(ctx context.Context, cfg *service.RelayConfig) { return } - go m.supervise(ctx, generation) + util.Go(func() { m.supervise(ctx, generation) }) } func (m *Manager) renderConfig(cfg *service.RelayConfig) error { diff --git a/backend/openflare/plugins/relay/relay/runner.go b/backend/openflare/plugins/relay/relay/runner.go index 6ed93ac5..80349626 100644 --- a/backend/openflare/plugins/relay/relay/runner.go +++ b/backend/openflare/plugins/relay/relay/runner.go @@ -17,6 +17,7 @@ import ( "Wavelet/openflare/plugins/relay/wsclient" edgerunner "Wavelet/openflare/share/edge/runner" service "Wavelet/openflare/share/protocol" + "Wavelet/pkg/util" ) // Runner manages the relay process. @@ -31,7 +32,7 @@ type Runner struct { // Run starts the relay process by initiating the heartbeat and WS reconnection loop. func (r *Runner) Run(ctx context.Context) error { - go r.HeartbeatService.Run(ctx) + util.Go(func() { r.HeartbeatService.Run(ctx) }) return edgerunner.RunWSReconnectLoop(ctx, edgerunner.WSReconnectConfig{ ComponentName: "relay", diff --git a/backend/openflare/plugins/server/domain/cloudflare/reconcile_test.go b/backend/openflare/plugins/server/domain/cloudflare/reconcile_test.go index b4545ef0..7a21d841 100644 --- a/backend/openflare/plugins/server/domain/cloudflare/reconcile_test.go +++ b/backend/openflare/plugins/server/domain/cloudflare/reconcile_test.go @@ -11,7 +11,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/credential" "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/repository" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "gorm.io/gorm" @@ -60,8 +59,8 @@ func setupCloudflareLogicDB(t *testing.T) (context.Context, uint) { ); err != nil { t.Fatalf("AutoMigrate() error = %v", err) } - db.SetDB(conn) - t.Cleanup(func() { db.SetDB(nil) }) + repository.SetDBForTest(conn) + t.Cleanup(func() { repository.SetDBForTest(nil) }) ctx := context.Background() sealed, err := credential.Seal(`{"api_token":"test-token"}`) if err != nil { @@ -143,11 +142,11 @@ func TestCreateMemberCopiesGroupDefaultProxied(t *testing.T) { t.Fatalf("GetCFPointingGroup() error = %v", err) } zone := model.Zone{Domain: "example.net"} - if err := db.DB(ctx).Create(&zone).Error; err != nil { + if err := repository.DB(ctx).Create(&zone).Error; err != nil { t.Fatalf("Create(zone) error = %v", err) } domain := model.ZoneDomain{ZoneID: zone.ID, Domain: "www.example.net"} - if err := db.DB(ctx).Create(&domain).Error; err != nil { + if err := repository.DB(ctx).Create(&domain).Error; err != nil { t.Fatalf("Create(domain) error = %v", err) } restore := SetDispatchTaskForTest(func(context.Context, string, []byte, string) (string, error) { return "task-1", nil }) diff --git a/backend/openflare/plugins/server/domain/cloudflare/routers_test.go b/backend/openflare/plugins/server/domain/cloudflare/routers_test.go index 165b3efc..e2123c35 100644 --- a/backend/openflare/plugins/server/domain/cloudflare/routers_test.go +++ b/backend/openflare/plugins/server/domain/cloudflare/routers_test.go @@ -11,7 +11,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/pkg/response" - db "Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" ) @@ -58,7 +57,7 @@ func TestGetGroupWithOrphanedMemberHealsAndSucceeds(t *testing.T) { } // Simulate orphaned member by deleting the ZoneDomain directly - if err := db.DB(ctx).Exec("DELETE FROM of_zone_domains WHERE id = ?", member.ZoneDomainID).Error; err != nil { + if err := repository.DB(ctx).Exec("DELETE FROM of_zone_domains WHERE id = ?", member.ZoneDomainID).Error; err != nil { t.Fatalf("DELETE FROM of_zone_domains error = %v", err) } diff --git a/backend/openflare/plugins/server/domain/dashboard/logics_test.go b/backend/openflare/plugins/server/domain/dashboard/logics_test.go index 67622787..0e747c1a 100644 --- a/backend/openflare/plugins/server/domain/dashboard/logics_test.go +++ b/backend/openflare/plugins/server/domain/dashboard/logics_test.go @@ -12,7 +12,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/testhelper" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -28,10 +27,10 @@ func setupDashboardTestDB(t *testing.T) func() { require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareNode{})) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) testhelper.SetupLogStoresForTest(t) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) } } @@ -43,7 +42,7 @@ func TestGetOverviewStructure(t *testing.T) { now := time.Now().UTC() lastSeen := now.Add(-15 * time.Second) // within default 60s offline threshold - require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + require.NoError(t, repository.DB(ctx).Create(&model.OpenFlareNode{ NodeID: "node-dashboard-1", Name: "Edge 1", IP: "10.0.0.1", @@ -52,7 +51,7 @@ func TestGetOverviewStructure(t *testing.T) { CurrentVersion: "v1.0.0", LastSeenAt: &lastSeen, }).Error) - require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + require.NoError(t, repository.DB(ctx).Create(&model.OpenFlareNode{ NodeID: "node-dashboard-2", Name: "Edge 2", IP: "10.0.0.2", diff --git a/backend/openflare/plugins/server/domain/fleet/agent/middleware_test.go b/backend/openflare/plugins/server/domain/fleet/agent/middleware_test.go index 0539f5a1..b479e0e0 100644 --- a/backend/openflare/plugins/server/domain/fleet/agent/middleware_test.go +++ b/backend/openflare/plugins/server/domain/fleet/agent/middleware_test.go @@ -15,7 +15,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/testhelper" "Wavelet/pkg/response" - db "Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" @@ -36,11 +35,11 @@ func setupAgentAuthTestDB(t *testing.T) func() { &model.SystemConfig{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) tokenCache.reset() return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) tokenCache.reset() } } @@ -51,7 +50,7 @@ func TestAuthenticateAccessToken(t *testing.T) { ctx := context.Background() now := time.Now() - require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + require.NoError(t, repository.DB(ctx).Create(&model.OpenFlareNode{ NodeID: "node-auth-1", Name: "edge", AccessToken: "valid-agent-token", @@ -98,7 +97,7 @@ func TestAgentAuthMiddleware(t *testing.T) { ctx := context.Background() now := time.Now() - require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + require.NoError(t, repository.DB(ctx).Create(&model.OpenFlareNode{ NodeID: "node-mw-1", Name: "edge", AccessToken: "middleware-token", @@ -145,7 +144,7 @@ func TestAgentRegisterAuthMiddleware(t *testing.T) { ctx := context.Background() now := time.Now() - require.NoError(t, db.DB(ctx).Create(&model.OpenFlareNode{ + require.NoError(t, repository.DB(ctx).Create(&model.OpenFlareNode{ NodeID: "node-register-1", Name: "edge", AccessToken: "existing-node-token", diff --git a/backend/openflare/plugins/server/domain/fleet/agent/waf_ip_group_test.go b/backend/openflare/plugins/server/domain/fleet/agent/waf_ip_group_test.go index 52570f6a..e68b82a2 100644 --- a/backend/openflare/plugins/server/domain/fleet/agent/waf_ip_group_test.go +++ b/backend/openflare/plugins/server/domain/fleet/agent/waf_ip_group_test.go @@ -14,7 +14,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/share/protocol" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -34,9 +33,9 @@ func setupWAFIPGroupTestDB(t *testing.T) func() { &model.ConfigVersion{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) } } @@ -60,7 +59,7 @@ func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID snapshotJSON, err := json.Marshal(snapshot) require.NoError(t, err) - require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{ + require.NoError(t, repository.DB(ctx).Create(&model.ConfigVersion{ Version: "20260618-001", SnapshotJSON: string(snapshotJSON), Checksum: "test-checksum", @@ -108,7 +107,7 @@ func seedActiveConfigWithWAFGraphIPGroup(t *testing.T, ctx context.Context, ipGr snapshotJSON, err := json.Marshal(snapshot) require.NoError(t, err) - require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{ + require.NoError(t, repository.DB(ctx).Create(&model.ConfigVersion{ Version: "20260713-graph-001", SnapshotJSON: string(snapshotJSON), Checksum: "graph-test-checksum", @@ -142,7 +141,7 @@ func TestChangedWAFIPGroupsForAgentRejectsMalformedIPMatchConfig(t *testing.T) { defer cleanup() ctx := context.Background() - require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{ + require.NoError(t, repository.DB(ctx).Create(&model.ConfigVersion{ Version: "20260713-malformed-001", SnapshotJSON: `{"waf":{"rule_groups":[{"id":7,"graph":{"entry":"match","nodes":{` + `"match":{"type":"ip_match","config":{"ip_group_ids":"not-an-array"}}}}}],"bindings":[]}}`, diff --git a/backend/openflare/plugins/server/domain/fleet/async_tasks_test.go b/backend/openflare/plugins/server/domain/fleet/async_tasks_test.go index 100776bc..1ab8f0ef 100644 --- a/backend/openflare/plugins/server/domain/fleet/async_tasks_test.go +++ b/backend/openflare/plugins/server/domain/fleet/async_tasks_test.go @@ -9,7 +9,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/repository" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -23,8 +22,8 @@ func TestUptimeKumaSyncHandlerSkipsWhenDisabled(t *testing.T) { }) require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{})) - db.SetDB(sqliteDB) - t.Cleanup(func() { db.SetDB(nil) }) + repository.SetDBForTest(sqliteDB) + t.Cleanup(func() { repository.SetDBForTest(nil) }) ctx := context.Background() require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyUptimeKumaEnabled, "false")) diff --git a/backend/openflare/plugins/server/domain/fleet/flared/middleware_test.go b/backend/openflare/plugins/server/domain/fleet/flared/middleware_test.go index 2206fecb..bcea0711 100644 --- a/backend/openflare/plugins/server/domain/fleet/flared/middleware_test.go +++ b/backend/openflare/plugins/server/domain/fleet/flared/middleware_test.go @@ -13,7 +13,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/pkg/response" - db "Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" @@ -30,10 +29,10 @@ func setupFlaredMiddlewareTestDB(t *testing.T) func() { }) require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareNode{})) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) } } diff --git a/backend/openflare/plugins/server/domain/fleet/flared/observability_test.go b/backend/openflare/plugins/server/domain/fleet/flared/observability_test.go index 741a99db..121d4cfd 100644 --- a/backend/openflare/plugins/server/domain/fleet/flared/observability_test.go +++ b/backend/openflare/plugins/server/domain/fleet/flared/observability_test.go @@ -11,7 +11,6 @@ import ( "Wavelet/openflare/plugins/server/domain/fleet/agent" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -33,11 +32,11 @@ func setupFlaredObservabilityTestDB(t *testing.T) func() { &model.ConfigVersion{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) agent.ResetAuthCacheForTest() return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) agent.ResetAuthCacheForTest() } } @@ -54,7 +53,7 @@ func TestHeartbeatFlaredEmitsHealthEventOnUnhealthy(t *testing.T) { Status: "pending", NodeType: "tunnel_client", } - require.NoError(t, db.DB(ctx).Create(node).Error) + require.NoError(t, repository.DB(ctx).Create(node).Error) _, err := Heartbeat(ctx, node, HeartbeatPayload{ ClientVersion: "v0.2.0", diff --git a/backend/openflare/plugins/server/domain/fleet/integration/agent_protocol_test.go b/backend/openflare/plugins/server/domain/fleet/integration/agent_protocol_test.go index 740c8da4..9ff66037 100644 --- a/backend/openflare/plugins/server/domain/fleet/integration/agent_protocol_test.go +++ b/backend/openflare/plugins/server/domain/fleet/integration/agent_protocol_test.go @@ -14,7 +14,6 @@ import ( ofnode "Wavelet/openflare/plugins/server/domain/fleet/node" "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/testhelper" - db "Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" @@ -39,7 +38,7 @@ func setupProtocolTestEnv(t *testing.T) (*gin.Engine, func()) { &model.ConfigVersion{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) agent.ResetAuthCacheForTest() testhelper.SetupLogStoresForTest(t) @@ -47,7 +46,7 @@ func setupProtocolTestEnv(t *testing.T) (*gin.Engine, func()) { mountOpenFlareTestRoutes(engine) cleanup := func() { - db.SetDB(nil) + repository.SetDBForTest(nil) agent.ResetAuthCacheForTest() } return engine, cleanup diff --git a/backend/openflare/plugins/server/domain/fleet/integration/core_chain_test.go b/backend/openflare/plugins/server/domain/fleet/integration/core_chain_test.go index db2e403a..4dba4c5f 100644 --- a/backend/openflare/plugins/server/domain/fleet/integration/core_chain_test.go +++ b/backend/openflare/plugins/server/domain/fleet/integration/core_chain_test.go @@ -11,8 +11,8 @@ import ( "Wavelet/openflare/plugins/server/domain/fleet/agent" "Wavelet/openflare/plugins/server/kernel/model" + "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/testhelper" - db "Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" @@ -55,7 +55,7 @@ func setupCoreChainTest(t *testing.T) (*gin.Engine, adminSeed, func()) { &model.ZoneDomain{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) agent.ResetAuthCacheForTest() seed, err := seedAdminWithAccessToken(sqliteDB) @@ -65,7 +65,7 @@ func setupCoreChainTest(t *testing.T) (*gin.Engine, adminSeed, func()) { mountOpenFlareTestRoutes(engine) cleanup := func() { - db.SetDB(nil) + repository.SetDBForTest(nil) agent.ResetAuthCacheForTest() } @@ -144,12 +144,12 @@ func TestCoreChainMigrationFlow(t *testing.T) { t.Run("create proxy route linked to origin", func(t *testing.T) { // Create Zone and ZoneDomain directly in the DB zone := model.Zone{Domain: "example.com"} - require.NoError(t, db.DB(context.Background()).Create(&zone).Error) + require.NoError(t, repository.DB(context.Background()).Create(&zone).Error) zoneDomain := model.ZoneDomain{ ZoneID: zone.ID, Domain: "core-chain.example.com", } - require.NoError(t, db.DB(context.Background()).Create(&zoneDomain).Error) + require.NoError(t, repository.DB(context.Background()).Create(&zoneDomain).Error) rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/proxy-routes/"), map[string]any{ "site_name": "core-chain-site", diff --git a/backend/openflare/plugins/server/domain/fleet/integration/security_test.go b/backend/openflare/plugins/server/domain/fleet/integration/security_test.go index 980686b9..c9ca61c7 100644 --- a/backend/openflare/plugins/server/domain/fleet/integration/security_test.go +++ b/backend/openflare/plugins/server/domain/fleet/integration/security_test.go @@ -17,9 +17,9 @@ import ( "time" "Wavelet/openflare/plugins/server/kernel/model" + "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/runtimeconfig" "Wavelet/openflare/plugins/server/kernel/testhelper" - db "Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" @@ -51,7 +51,7 @@ func setupSecurityTest(t *testing.T) (*gin.Engine, adminSeed, func()) { &model.SystemConfig{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) seed, err := seedAdminWithAccessToken(sqliteDB) require.NoError(t, err) @@ -64,7 +64,7 @@ func setupSecurityTest(t *testing.T) (*gin.Engine, adminSeed, func()) { cleanup := func() { runtimeconfig.Set(previous) - db.SetDB(nil) + repository.SetDBForTest(nil) } return engine, seed, cleanup @@ -205,12 +205,12 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) { t.Run("create proxy route for WAF binding", func(t *testing.T) { // Create Zone and ZoneDomain directly in the DB routeZone := model.Zone{Domain: "example-route.com"} - require.NoError(t, db.DB(context.Background()).Create(&routeZone).Error) + require.NoError(t, repository.DB(context.Background()).Create(&routeZone).Error) routeZoneDomain := model.ZoneDomain{ ZoneID: routeZone.ID, Domain: "route.example-route.com", } - require.NoError(t, db.DB(context.Background()).Create(&routeZoneDomain).Error) + require.NoError(t, repository.DB(context.Background()).Create(&routeZoneDomain).Error) rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/proxy-routes/"), map[string]any{ "site_name": "security-site", diff --git a/backend/openflare/plugins/server/domain/fleet/node/logics_test.go b/backend/openflare/plugins/server/domain/fleet/node/logics_test.go index edd7adee..233e7e85 100644 --- a/backend/openflare/plugins/server/domain/fleet/node/logics_test.go +++ b/backend/openflare/plugins/server/domain/fleet/node/logics_test.go @@ -15,7 +15,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/testhelper" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -40,13 +39,15 @@ func setupNodeTestDB(t *testing.T) func() { &model.OpenFlareNode{}, &model.SystemConfig{}, &model.OpenFlareApplyLog{}, + &model.OpenFlareNodeSystemProfile{}, + &model.OpenFlareHealthEvent{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) testhelper.SetupLogStoresForTest(t) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) } } @@ -174,7 +175,7 @@ func TestListNodesWithApplyLogMetadata(t *testing.T) { require.NoError(t, err) applyAt := time.Now().UTC().Truncate(time.Second) - require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{ + require.NoError(t, repository.DB(ctx).Create(&model.OpenFlareApplyLog{ NodeID: created.NodeID, Version: "20260618-001", Result: "success", @@ -280,7 +281,7 @@ func TestRequestOpenrestyRestart(t *testing.T) { func seedActiveConfigVersion(t *testing.T, ctx context.Context) { t.Helper() - conn := db.DB(ctx) + conn := repository.DB(ctx) require.NotNil(t, conn) require.NoError(t, conn.AutoMigrate(&model.ConfigVersion{})) require.NoError(t, conn.Create(&model.ConfigVersion{ @@ -311,7 +312,7 @@ func TestRequestForceSyncRequiresActiveConfig(t *testing.T) { cleanup := setupNodeTestDB(t) defer cleanup() ctx := context.Background() - conn := db.DB(ctx) + conn := repository.DB(ctx) require.NotNil(t, conn) require.NoError(t, conn.AutoMigrate(&model.ConfigVersion{})) diff --git a/backend/openflare/plugins/server/domain/fleet/relay/logics_test.go b/backend/openflare/plugins/server/domain/fleet/relay/logics_test.go index 239287a6..0b2c6e57 100644 --- a/backend/openflare/plugins/server/domain/fleet/relay/logics_test.go +++ b/backend/openflare/plugins/server/domain/fleet/relay/logics_test.go @@ -14,7 +14,6 @@ import ( "Wavelet/openflare/plugins/server/domain/fleet/agent" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -38,12 +37,12 @@ func setupRelayTestDB(t *testing.T) func() { &model.OpenFlareNodeObservationFrps{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) agent.ResetAuthCacheForTest() testhelper.SetupLogStoresForTest(t) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) agent.ResetAuthCacheForTest() } } @@ -63,7 +62,7 @@ func TestHeartbeatPayloadBindingAndFrpsObservationInsert(t *testing.T) { NodeType: "tunnel_relay", RelayStatus: "unknown", } - require.NoError(t, db.DB(ctx).Create(node).Error) + require.NoError(t, repository.DB(ctx).Create(node).Error) proxies := []ProxyStat{ { @@ -103,7 +102,7 @@ func TestHeartbeatPayloadBindingAndFrpsObservationInsert(t *testing.T) { require.NoError(t, err) var stored model.OpenFlareNode - require.NoError(t, db.DB(ctx).Where("node_id = ?", node.NodeID).First(&stored).Error) + require.NoError(t, repository.DB(ctx).Where("node_id = ?", node.NodeID).First(&stored).Error) assert.Equal(t, "online", stored.Status) assert.Equal(t, "healthy", stored.RelayStatus) assert.Equal(t, "203.0.113.9", stored.IP) @@ -147,7 +146,7 @@ func TestHeartbeatRelayReconcilesFrpsUnhealthyEvent(t *testing.T) { NodeType: "tunnel_relay", RelayStatus: "healthy", } - require.NoError(t, db.DB(ctx).Create(node).Error) + require.NoError(t, repository.DB(ctx).Create(node).Error) _, err := Heartbeat(ctx, node, HeartbeatPayload{ Version: "v0.1.0", diff --git a/backend/openflare/plugins/server/domain/fleet/relay/middleware_test.go b/backend/openflare/plugins/server/domain/fleet/relay/middleware_test.go index cb850582..47f227c2 100644 --- a/backend/openflare/plugins/server/domain/fleet/relay/middleware_test.go +++ b/backend/openflare/plugins/server/domain/fleet/relay/middleware_test.go @@ -13,7 +13,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/pkg/response" - db "Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" @@ -30,10 +29,10 @@ func setupRelayMiddlewareTestDB(t *testing.T) func() { }) require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareNode{})) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) } } diff --git a/backend/openflare/plugins/server/domain/fleet/websocket/agent_hub.go b/backend/openflare/plugins/server/domain/fleet/websocket/agent_hub.go index 9fe4ff64..1f6a9539 100644 --- a/backend/openflare/plugins/server/domain/fleet/websocket/agent_hub.go +++ b/backend/openflare/plugins/server/domain/fleet/websocket/agent_hub.go @@ -12,6 +12,8 @@ import ( "time" "github.com/gin-gonic/gin" + + "Wavelet/pkg/util" ) const ( @@ -64,7 +66,7 @@ func ServeAgent(c *gin.Context, nodeID string, onStatus AgentStatusHandler) { slog.Debug("agent ws connected", "node_id", nodeID, "remote", client.remoteAddr) - go client.writePump() + util.Go(client.writePump) client.readPump() } diff --git a/backend/openflare/plugins/server/domain/fleet/websocket/flared_hub.go b/backend/openflare/plugins/server/domain/fleet/websocket/flared_hub.go index 926a52e7..3572783d 100644 --- a/backend/openflare/plugins/server/domain/fleet/websocket/flared_hub.go +++ b/backend/openflare/plugins/server/domain/fleet/websocket/flared_hub.go @@ -8,6 +8,8 @@ import ( "sync" "github.com/gin-gonic/gin" + + "Wavelet/pkg/util" ) const ( @@ -51,7 +53,7 @@ func ServeFlared(c *gin.Context, nodeID string) { slog.Debug("flared ws connected", "node_id", nodeID, "remote", c.Request.RemoteAddr) - go client.writePump() + util.Go(client.writePump) client.readPump() } diff --git a/backend/openflare/plugins/server/domain/fleet/websocket/relay_hub.go b/backend/openflare/plugins/server/domain/fleet/websocket/relay_hub.go index 6fd1bd05..bbe9af59 100644 --- a/backend/openflare/plugins/server/domain/fleet/websocket/relay_hub.go +++ b/backend/openflare/plugins/server/domain/fleet/websocket/relay_hub.go @@ -8,6 +8,8 @@ import ( "sync" "github.com/gin-gonic/gin" + + "Wavelet/pkg/util" ) // RelayWSConnectedLastSeenValue is the sentinel last_seen_at value when relay WS is connected. @@ -45,7 +47,7 @@ func ServeRelay(c *gin.Context, nodeID string) { slog.Debug("relay ws connected", "node_id", nodeID, "remote", c.Request.RemoteAddr) - go client.writePump() + util.Go(client.writePump) client.readPump() } diff --git a/backend/openflare/plugins/server/domain/observability/chwriter/live_ch_test.go b/backend/openflare/plugins/server/domain/observability/chwriter/live_ch_test.go index 19f24b94..f49035a1 100644 --- a/backend/openflare/plugins/server/domain/observability/chwriter/live_ch_test.go +++ b/backend/openflare/plugins/server/domain/observability/chwriter/live_ch_test.go @@ -14,7 +14,6 @@ import ( "Wavelet/openflare/plugins/server/domain/observability/chwriter" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // Run with Docker ClickHouse + config.yaml: diff --git a/backend/openflare/plugins/server/domain/observability/log_db_switch_test.go b/backend/openflare/plugins/server/domain/observability/log_db_switch_test.go index b3ac3a2c..5d9ed98b 100644 --- a/backend/openflare/plugins/server/domain/observability/log_db_switch_test.go +++ b/backend/openflare/plugins/server/domain/observability/log_db_switch_test.go @@ -23,7 +23,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/repository/logstore" "Wavelet/openflare/plugins/server/kernel/runtimeconfig" - db "Wavelet/plugins/infra/database" ) var logDBSwitchDBSeq int64 @@ -57,13 +56,13 @@ func TestCopyAccessLogsPreservesIDs(t *testing.T) { srcDB := newLogDBSwitchDB(t) dstDB := newLogDBSwitchDB(t) - db.SetDB(srcDB) + repository.SetDBForTest(srcDB) src, err := logstore.Active(ctx) // 无 reader 时按 seed 规则解析为 sqlite require.NoError(t, err) - db.SetDB(dstDB) + repository.SetDBForTest(dstDB) dst, err := logstore.BuildForMigration(ctx, "sqlite") require.NoError(t, err) - t.Cleanup(func() { db.SetDB(nil) }) + t.Cleanup(func() { repository.SetDBForTest(nil) }) now := time.Now().UTC() rows := []analyticsmodel.NodeAccessLog{ @@ -103,13 +102,13 @@ func TestCopyUserAccessLogsPreservesIDs(t *testing.T) { srcDB := newLogDBSwitchDB(t) dstDB := newLogDBSwitchDB(t) - db.SetDB(srcDB) + repository.SetDBForTest(srcDB) src, err := logstore.Active(ctx) require.NoError(t, err) - db.SetDB(dstDB) + repository.SetDBForTest(dstDB) dst, err := logstore.BuildForMigration(ctx, "sqlite") require.NoError(t, err) - t.Cleanup(func() { db.SetDB(nil) }) + t.Cleanup(func() { repository.SetDBForTest(nil) }) now := time.Now().UTC() rows := []analyticsmodel.UserAccessLog{ @@ -142,8 +141,8 @@ func TestClearTargetLogTablesClearsUserAccessLogs(t *testing.T) { ctx := context.Background() dstDB := newLogDBSwitchDB(t) - db.SetDB(dstDB) - t.Cleanup(func() { db.SetDB(nil) }) + repository.SetDBForTest(dstDB) + t.Cleanup(func() { repository.SetDBForTest(nil) }) dst, err := logstore.BuildForMigration(ctx, "sqlite") require.NoError(t, err) @@ -168,8 +167,8 @@ func TestClearTargetLogTablesDuringMigration(t *testing.T) { defer logstore.ResetForTest() gdb := newLogDBSwitchDB(t) - db.SetDB(gdb) - t.Cleanup(func() { db.SetDB(nil) }) + repository.SetDBForTest(gdb) + t.Cleanup(func() { repository.SetDBForTest(nil) }) ctx := context.Background() logstore.SetConfigReader(func(ctx context.Context, key string) (string, error) { @@ -205,8 +204,8 @@ func TestValidateSwitch(t *testing.T) { }) gdb := newLogDBSwitchDB(t) - db.SetDB(gdb) - t.Cleanup(func() { db.SetDB(nil) }) + repository.SetDBForTest(gdb) + t.Cleanup(func() { repository.SetDBForTest(nil) }) ctx := context.Background() setLogDB := func(v string) { require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyLogDatabase, v)) @@ -288,8 +287,8 @@ func TestExecuteFailureClearsMigrationFlag(t *testing.T) { defer logstore.ResetForTest() gdb := newLogDBSwitchDB(t) - db.SetDB(gdb) - t.Cleanup(func() { db.SetDB(nil) }) + repository.SetDBForTest(gdb) + t.Cleanup(func() { repository.SetDBForTest(nil) }) ctx := context.Background() // FRESH DB:log_db_migration 行不存在(不预置),log_database 预置为 sqlite。 @@ -323,8 +322,8 @@ func TestSetMigrationFlagObservableThroughCache(t *testing.T) { defer logstore.ResetForTest() gdb := newLogDBSwitchDB(t) - db.SetDB(gdb) - t.Cleanup(func() { db.SetDB(nil) }) + repository.SetDBForTest(gdb) + t.Cleanup(func() { repository.SetDBForTest(nil) }) ctx := context.Background() // 按 bootstrap 同款注入 repository 读取,走 RAM 缓存路径。 @@ -353,8 +352,8 @@ func TestFlipLogDatabaseRefreshesCachedConfig(t *testing.T) { defer logstore.ResetForTest() gdb := newLogDBSwitchDB(t) - db.SetDB(gdb) - t.Cleanup(func() { db.SetDB(nil) }) + repository.SetDBForTest(gdb) + t.Cleanup(func() { repository.SetDBForTest(nil) }) ctx := context.Background() logstore.SetConfigReader(func(ctx context.Context, key string) (string, error) { diff --git a/backend/openflare/plugins/server/domain/option/logics_test.go b/backend/openflare/plugins/server/domain/option/logics_test.go index 2bb06ce2..9c003ce2 100644 --- a/backend/openflare/plugins/server/domain/option/logics_test.go +++ b/backend/openflare/plugins/server/domain/option/logics_test.go @@ -9,7 +9,7 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/repository" - db "Wavelet/plugins/infra/database" + "Wavelet/openflare/plugins/server/kernel/testhelper" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -26,7 +26,8 @@ func setupOptionTestDB(t *testing.T) func() { require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{})) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) + repository.SetSystemConfigService(testhelper.NewMockSystemConfigService(sqliteDB)) // 预填充一些业务配置用于测试 seedConfigs := []model.SystemConfig{ @@ -38,14 +39,15 @@ func setupOptionTestDB(t *testing.T) func() { } return func() { - db.SetDB(nil) + repository.SetSystemConfigService(nil) + repository.SetDBForTest(nil) } } // setTestConfig 设置测试配置的辅助函数 func setTestConfig(t *testing.T, ctx context.Context, key, value string) { t.Helper() - require.NoError(t, db.DB(ctx).Model(&model.SystemConfig{}).Where("key = ?", key).Update("value", value).Error) + require.NoError(t, repository.DB(ctx).Model(&model.SystemConfig{}).Where("key = ?", key).Update("value", value).Error) } func TestListOptionsFiltersSecretKeys(t *testing.T) { @@ -89,7 +91,7 @@ func TestUpdateOpenRestyOptionPersistsToSystemConfig(t *testing.T) { defer cleanup() ctx := context.Background() - require.NoError(t, db.DB(ctx).Create(&model.SystemConfig{ + require.NoError(t, repository.DB(ctx).Create(&model.SystemConfig{ Key: model.ConfigKeyOpenRestyEventsUse, Value: "epoll", Type: "business", diff --git a/backend/openflare/plugins/server/domain/option/uptimekuma/client.go b/backend/openflare/plugins/server/domain/option/uptimekuma/client.go index 48d6bcfc..aeb7c7da 100644 --- a/backend/openflare/plugins/server/domain/option/uptimekuma/client.go +++ b/backend/openflare/plugins/server/domain/option/uptimekuma/client.go @@ -16,6 +16,8 @@ import ( "strings" "sync" "time" + + "Wavelet/pkg/util" ) const emitAckTimeout = 10 * time.Second @@ -137,7 +139,7 @@ func (c *SocketIOClient) Connect() error { _ = respConnect.Body.Close() slog.Debug("Namespace connected successfully to Uptime Kuma", "sid", c.sid) - go c.pollLoop() + util.Go(c.pollLoop) return nil } diff --git a/backend/openflare/plugins/server/domain/option/uptimekuma/sync_test.go b/backend/openflare/plugins/server/domain/option/uptimekuma/sync_test.go index 08e40344..6ed3622e 100644 --- a/backend/openflare/plugins/server/domain/option/uptimekuma/sync_test.go +++ b/backend/openflare/plugins/server/domain/option/uptimekuma/sync_test.go @@ -17,7 +17,7 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/repository" - db "Wavelet/plugins/infra/database" + "Wavelet/openflare/plugins/server/kernel/testhelper" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -125,17 +125,19 @@ func setupSyncTestDB(t *testing.T) func() { require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.ProxyRoute{}, &model.Zone{}, &model.ZoneDomain{}, &model.SystemConfig{})) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) + repository.SetSystemConfigService(testhelper.NewMockSystemConfigService(sqliteDB)) return func() { - db.SetDB(nil) + repository.SetSystemConfigService(nil) + repository.SetDBForTest(nil) } } func createRouteZoneDomain(t *testing.T, ctx context.Context, route *model.ProxyRoute, domain string) { t.Helper() zone := &model.Zone{Domain: fmt.Sprintf("zone-%d.example", route.ID)} - require.NoError(t, db.DB(ctx).Create(zone).Error) - require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{ + require.NoError(t, repository.DB(ctx).Create(zone).Error) + require.NoError(t, repository.DB(ctx).Create(&model.ZoneDomain{ ZoneID: zone.ID, ProxyRouteID: &route.ID, Domain: domain, @@ -166,7 +168,7 @@ func backupUptimeKumaConfig(ctx context.Context) func() { return func() { // 恢复所有配置 for key, value := range oldValues { - _ = db.DB(ctx).Model(&model.SystemConfig{}).Where("key = ?", key).Update("value", value).Error + _ = repository.DB(ctx).Model(&model.SystemConfig{}).Where("key = ?", key).Update("value", value).Error } } } @@ -197,7 +199,7 @@ func TestSyncToUptimeKumaSuccess(t *testing.T) { restore := backupUptimeKumaConfig(ctx) defer restore() - require.NoError(t, db.DB(ctx).Where("1 = 1").Delete(&model.ProxyRoute{}).Error) + require.NoError(t, repository.DB(ctx).Where("1 = 1").Delete(&model.ProxyRoute{}).Error) routeA := &model.ProxyRoute{ SiteName: "site-a", @@ -306,7 +308,7 @@ func TestSyncToUptimeKumaSelectedScope(t *testing.T) { restore := backupUptimeKumaConfig(ctx) defer restore() - require.NoError(t, db.DB(ctx).Where("1 = 1").Delete(&model.ProxyRoute{}).Error) + require.NoError(t, repository.DB(ctx).Where("1 = 1").Delete(&model.ProxyRoute{}).Error) routeA := &model.ProxyRoute{ SiteName: "site-a", diff --git a/backend/openflare/plugins/server/domain/pages/github_source_identity_test.go b/backend/openflare/plugins/server/domain/pages/github_source_identity_test.go index b8685a40..56e90ca2 100644 --- a/backend/openflare/plugins/server/domain/pages/github_source_identity_test.go +++ b/backend/openflare/plugins/server/domain/pages/github_source_identity_test.go @@ -9,7 +9,7 @@ import ( "time" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" + "Wavelet/openflare/plugins/server/kernel/repository" ) func TestGitHubSourceIdentityLengthPrefixesFieldsAndResetsRuntime(t *testing.T) { @@ -57,7 +57,7 @@ func TestGitHubSourceIdentityLengthPrefixesFieldsAndResetsRuntime(t *testing.T) syncedAt := time.Now().Add(-30 * time.Second) nextCheckAt := time.Now().Add(time.Hour) leaseExpiresAt := time.Now().Add(time.Minute) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", firstSource.ID). Updates(map[string]any{ "etag": `"old-etag"`, diff --git a/backend/openflare/plugins/server/domain/pages/github_source_test.go b/backend/openflare/plugins/server/domain/pages/github_source_test.go index 69e98a8d..625238a7 100644 --- a/backend/openflare/plugins/server/domain/pages/github_source_test.go +++ b/backend/openflare/plugins/server/domain/pages/github_source_test.go @@ -20,7 +20,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/task" "Wavelet/openflare/share/githubrelease" - db "Wavelet/plugins/infra/database" "github.com/hibiken/asynq" "gorm.io/gorm" @@ -62,7 +61,7 @@ func mustConfigureGitHubSourceWithoutDispatch( if err := validateGitHubSourceInput(input); err != nil { t.Fatalf("validateGitHubSourceInput(%+v) error = %v, want nil", input, err) } - if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := repository.DB(ctx).Transaction(func(tx *gorm.DB) error { _, err := updateGitHubSourceTx(tx, projectID, input) return err }); err != nil { @@ -78,11 +77,11 @@ func mustLoadPagesSource( ) (*model.PagesProjectSource, *model.PagesProjectSourceRuntime) { t.Helper() var source model.PagesProjectSource - if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { + if err := repository.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { t.Fatalf("load source for project %d error = %v, want nil", projectID, err) } var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { + if err := repository.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { t.Fatalf("load runtime for source %d error = %v, want nil", source.ID, err) } return &source, &runtime @@ -136,14 +135,14 @@ func TestGitHubSourceValidationNormalizationAndProviderSwitch(t *testing.T) { } var taskCount int64 - if err := db.DB(ctx).Model(&model.TaskExecution{}).Count(&taskCount).Error; err != nil { + if err := repository.DB(ctx).Model(&model.TaskExecution{}).Count(&taskCount).Error; err != nil { t.Fatalf("count initial checks error = %v, want nil", err) } if _, err := UpdateSourceAs(ctx, project.ID, input, "user:42"); err != nil { t.Fatalf("UpdateSourceAs(GitHub no-op) error = %v, want nil", err) } var noOpTaskCount int64 - if err := db.DB(ctx).Model(&model.TaskExecution{}).Count(&noOpTaskCount).Error; err != nil { + if err := repository.DB(ctx).Model(&model.TaskExecution{}).Count(&noOpTaskCount).Error; err != nil { t.Fatalf("count no-op checks error = %v, want nil", err) } if noOpTaskCount != taskCount { @@ -275,7 +274,7 @@ func TestInitialCheckFailureUsesExactConfigFence(t *testing.T) { RepositoryURL: "https://github.com/a/b", }) staleVersion := source.ConfigVersion - if err := db.DB(ctx).Model(source).Update("config_version", staleVersion+1).Error; err != nil { + if err := repository.DB(ctx).Model(source).Update("config_version", staleVersion+1).Error; err != nil { t.Fatalf("increment source config version error = %v, want nil", err) } markInitialCheckDispatchFailed(ctx, source.ID, staleVersion) @@ -294,7 +293,7 @@ func TestGitHubCheckUsesETagAndDetectsSameReleaseReplacement(t *testing.T) { }) appliedRevision := strings.Repeat("a", 64) appliedDetail := `{"provider":"github","release_id":"100","asset_id":"1","tag":"release/v1","asset_name":"dist.zip"}` - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{ + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{ "etag": `"old-etag"`, "last_applied_revision": appliedRevision, "last_applied_detail": appliedDetail, @@ -375,7 +374,7 @@ func TestGitHubCheckNotModifiedRefreshesRuntimeWithoutDeployment(t *testing.T) { t.Errorf("304 runtime = %+v, want refreshed timestamps/etag and released lease", runtime) } var deployments int64 - if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&deployments).Error; err != nil { + if err := repository.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&deployments).Error; err != nil { t.Fatalf("count deployments error = %v, want nil", err) } if deployments != 0 { @@ -391,7 +390,7 @@ func TestGitHubTargetMismatchPreservesAttentionAndExpeditesRecheck(t *testing.T) RepositoryURL: "https://github.com/a/b", }) appliedRevision := strings.Repeat("a", 64) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{ + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{ "last_applied_revision": appliedRevision, "last_applied_detail": `{"provider":"github","release_id":"100","asset_id":"1","tag":"v1","asset_name":"dist.zip"}`, }).Error; err != nil { @@ -430,7 +429,7 @@ func TestGitHubTargetMismatchPreservesAttentionAndExpeditesRecheck(t *testing.T) t.Errorf("target mismatch NextCheckAt = %v, want server deadline >= %v", runtime.NextCheckAt, retryAt) } var deployments int64 - if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&deployments).Error; err != nil { + if err := repository.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&deployments).Error; err != nil { t.Fatalf("count mismatch deployments error = %v, want nil", err) } if deployments != 0 { @@ -454,7 +453,7 @@ func TestGitHubCheckLostLeaseReturnsStaleWithoutOverwritingRuntime(t *testing.T) }, }) snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionCheck) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{ + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{ "lease_token": "new-owner", "lease_expires_at": time.Now().Add(time.Minute), "sync_status": pagesSourceStatusSyncing, @@ -741,7 +740,7 @@ func TestGitHubSyncActivatesExactConfirmedReplacement(t *testing.T) { if err != nil { t.Fatalf("buildGitHubSourceTarget() error = %v, want nil", err) } - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{ + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{ "last_seen_revision": target.Revision, "last_seen_detail": target.DetailJSON, "last_applied_revision": strings.Repeat("a", 64), diff --git a/backend/openflare/plugins/server/domain/pages/helpers.go b/backend/openflare/plugins/server/domain/pages/helpers.go index fda6f34d..41479777 100644 --- a/backend/openflare/plugins/server/domain/pages/helpers.go +++ b/backend/openflare/plugins/server/domain/pages/helpers.go @@ -320,6 +320,7 @@ func removePagesUploadIfUnreferenced(ctx context.Context, projectID uint, upload if uploadID == 0 { return nil } + shouldRemove := false err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { if projectID != 0 { if _, projectErr := repository.LockPagesProjectByIDTx(tx, projectID); projectErr != nil && @@ -350,13 +351,13 @@ func removePagesUploadIfUnreferenced(ctx context.Context, projectID uint, upload if references > 0 { return nil } - _, err := ofupload.RemoveLockedTx(tx, &uploadRecord) - return err + shouldRemove = true + return nil }) - // Always invalidate after transaction completion, including idempotent no-op, - // so a prior post-commit cache interruption can heal on retry. - ofupload.InvalidateUploadMetaCache(ctx, uploadID) - return err + if err != nil || !shouldRemove { + return err + } + return ofupload.Remove(ctx, uploadID) } func inspectPagesPackage(packagePath string, format pagesarchive.Format, rootDir string, entryFile string, limits pagesLimits) (*deploymentManifest, error) { diff --git a/backend/openflare/plugins/server/domain/pages/logics.go b/backend/openflare/plugins/server/domain/pages/logics.go index 1f1a2380..42f3105e 100644 --- a/backend/openflare/plugins/server/domain/pages/logics.go +++ b/backend/openflare/plugins/server/domain/pages/logics.go @@ -1020,7 +1020,8 @@ func hydrateLegacyDeploymentUpload( if deployment.UploadID > 0 { uploadRecord, err := ofupload.GetActiveUpload(ctx, deployment.UploadID) if err == nil { - return &uploadRecord, nil + record := model.FromUploadDTO(uploadRecord) + return &record, nil } } @@ -1066,7 +1067,8 @@ func hydrateLegacyDeploymentUpload( } deployment.UploadID = winnerUploadID deployment.ArtifactPath = "" - return &winner, nil + winnerRecord := model.FromUploadDTO(winner) + return &winnerRecord, nil } func attachLegacyDeploymentUpload( diff --git a/backend/openflare/plugins/server/domain/pages/logics_test.go b/backend/openflare/plugins/server/domain/pages/logics_test.go index 78f73c1a..72b20a1e 100644 --- a/backend/openflare/plugins/server/domain/pages/logics_test.go +++ b/backend/openflare/plugins/server/domain/pages/logics_test.go @@ -23,8 +23,6 @@ import ( oftask "Wavelet/openflare/plugins/server/kernel/task" "Wavelet/openflare/plugins/server/kernel/testhelper" "Wavelet/pkg/idgen" - uploadshared "Wavelet/plugins/domain/upload/shared" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -39,6 +37,9 @@ func setupPagesTestDB(t *testing.T) func() { DisableForeignKeyConstraintWhenMigrating: true, }) require.NoError(t, err) + if sqlDB, err := sqliteDB.DB(); err == nil { + sqlDB.SetMaxOpenConns(1) + } require.NoError(t, sqliteDB.AutoMigrate( &model.User{}, &model.Upload{}, @@ -74,30 +75,33 @@ func setupPagesTestDB(t *testing.T) func() { }, }).Error) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) + repository.SetSystemConfigService(testhelper.NewMockSystemConfigService(sqliteDB)) require.NoError(t, idgen.Init(1)) - oftask.SetService(&testhelper.NoopTaskService{}) - mockStorage := uploadshared.NewMockStorageService() - uploadshared.SetDBService(db.NewService(sqliteDB)) - uploadshared.SetStorageService(mockStorage) + noopTask := &testhelper.NoopTaskService{} + oftask.SetService(noopTask) + repository.SetTaskService(noopTask) + mockStorage := testhelper.NewMockStorageService() ofupload.SetStorage(mockStorage) + ofupload.SetUploadService(testhelper.NewMockUploadService(sqliteDB)) _ = repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyPagesMaxPackageSizeMB) _ = repository.InvalidateSystemConfigCache(context.Background(), model.ConfigKeyPagesMaxHistoryCount) return func() { ofupload.SetStorage(nil) - uploadshared.ResetServices() - db.SetDB(nil) + ofupload.SetUploadService(nil) + repository.SetTaskService(nil) + oftask.SetService(nil) + repository.SetSystemConfigService(nil) + repository.SetDBForTest(nil) } } func setupPagesStorageMock(t *testing.T) (restore func(), disable func()) { t.Helper() - mock := uploadshared.NewMockStorageService() - uploadshared.SetStorageService(mock) + mock := testhelper.NewMockStorageService() ofupload.SetStorage(mock) restore = func() { ofupload.SetStorage(nil) - uploadshared.ResetServices() } disable = restore return restore, disable @@ -283,10 +287,10 @@ func TestUploadDeploymentStoresPackageInUploadFramework(t *testing.T) { assert.Empty(t, storedDeployment.ArtifactPath) var uploadCount int64 - require.NoError(t, db.DB(ctx).Model(&model.Upload{}).Count(&uploadCount).Error) + require.NoError(t, repository.DB(ctx).Model(&model.Upload{}).Count(&uploadCount).Error) assert.Equal(t, int64(1), uploadCount) var uploadRecord model.Upload - require.NoError(t, db.DB(ctx).First(&uploadRecord, storedDeployment.UploadID).Error) + require.NoError(t, repository.DB(ctx).First(&uploadRecord, storedDeployment.UploadID).Error) assert.Equal(t, ofupload.ReservedPagesDeploymentType, uploadRecord.Type) assert.Equal(t, pagesIngestMarkerV2, uploadRecord.Metadata.Extra[pagesIngestMarkerKey]) assert.Equal(t, fmt.Sprint(project.ID), uploadRecord.Metadata.Extra[pagesProjectIDMetadataKey]) @@ -323,8 +327,8 @@ func TestOpenDeploymentPackageHydratesLegacyArtifactPath(t *testing.T) { TotalSize: 10, CreatedBy: "test", } - require.NoError(t, db.DB(ctx).Create(deployment).Error) - require.NoError(t, db.DB(ctx).Create(&model.PagesDeploymentFile{ + require.NoError(t, repository.DB(ctx).Create(deployment).Error) + require.NoError(t, repository.DB(ctx).Create(&model.PagesDeploymentFile{ DeploymentID: deployment.ID, Path: "index.html", Size: 6, @@ -334,7 +338,7 @@ func TestOpenDeploymentPackageHydratesLegacyArtifactPath(t *testing.T) { _, err = ActivateDeployment(ctx, project.ID, deployment.ID) require.NoError(t, err) - require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{ + require.NoError(t, repository.DB(ctx).Create(&model.ConfigVersion{ Version: "v2026-legacy", SnapshotJSON: fmt.Sprintf(`{"routes":[{"upstream_type":"pages","pages_deployment":{"deployment_id":%d}}]}`, deployment.ID), MainConfig: "", @@ -363,7 +367,7 @@ func TestOpenDeploymentPackageHydratesLegacyArtifactPath(t *testing.T) { assert.Empty(t, storedDeployment.ArtifactPath) var uploadCount int64 - require.NoError(t, db.DB(ctx).Model(&model.Upload{}).Count(&uploadCount).Error) + require.NoError(t, repository.DB(ctx).Model(&model.Upload{}).Count(&uploadCount).Error) assert.Equal(t, int64(1), uploadCount) packageObj2, err := OpenDeploymentPackage(ctx, deployment.ID) @@ -400,7 +404,7 @@ func TestOpenDeploymentPackageRequiresActiveConfigSnapshot(t *testing.T) { require.Error(t, err) assert.Contains(t, err.Error(), "激活配置") - require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{ + require.NoError(t, repository.DB(ctx).Create(&model.ConfigVersion{ Version: "v2026-001", SnapshotJSON: fmt.Sprintf(`{"routes":[{"upstream_type":"pages","pages_deployment":{"deployment_id":%d}}]}`, deployment.ID), MainConfig: "", @@ -454,7 +458,7 @@ func TestProjectLatestRejectsWhenNotOnActiveConfigOrNotActive(t *testing.T) { _, _, err = GetProjectLatestPackageHash(ctx, project.ID) require.Error(t, err) - require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{ + require.NoError(t, repository.DB(ctx).Create(&model.ConfigVersion{ Version: "v-gate", SnapshotJSON: fmt.Sprintf( `{"routes":[{"upstream_type":"pages","pages_project_id":%d,"pages_deployment":{"project_id":%d,"deployment_id":%d}}]}`, @@ -476,7 +480,7 @@ func TestProjectLatestRejectsWhenNotOnActiveConfigOrNotActive(t *testing.T) { assert.NotEmpty(t, hash) // Disabled project rejects. - require.NoError(t, db.DB(ctx).Model(&model.PagesProject{}).Where("id = ?", project.ID).Update("enabled", false).Error) + require.NoError(t, repository.DB(ctx).Model(&model.PagesProject{}).Where("id = ?", project.ID).Update("enabled", false).Error) _, _, err = GetProjectLatestPackageHash(ctx, project.ID) require.Error(t, err) } @@ -572,7 +576,7 @@ func TestPruneProjectDeploymentHistory(t *testing.T) { defer disableStorage() ctx := context.Background() - require.NoError(t, db.DB(ctx).Model(&model.SystemConfig{}). + require.NoError(t, repository.DB(ctx).Model(&model.SystemConfig{}). Where("key = ?", model.ConfigKeyPagesMaxHistoryCount). Update("value", "2").Error) require.NoError(t, repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyPagesMaxHistoryCount)) @@ -634,7 +638,7 @@ func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) { defer disableStorage() ctx := context.Background() - require.NoError(t, db.DB(ctx).Model(&model.SystemConfig{}). + require.NoError(t, repository.DB(ctx).Model(&model.SystemConfig{}). Where("key = ?", model.ConfigKeyPagesMaxHistoryCount). Update("value", "1").Error) require.NoError(t, repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyPagesMaxHistoryCount)) @@ -671,7 +675,7 @@ func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) { assert.True(t, kept[newCandidate.ID]) assert.False(t, kept[oldCandidate.ID]) var removedUpload model.Upload - require.NoError(t, db.DB(ctx).First(&removedUpload, oldCandidate.UploadID).Error) + require.NoError(t, repository.DB(ctx).First(&removedUpload, oldCandidate.UploadID).Error) assert.Equal(t, model.UploadStatusDeleted, removedUpload.Status) _, err = ActivateDeployment(ctx, project.ID, newCandidate.ID) @@ -741,12 +745,12 @@ func TestDeleteDeploymentAndProjectSoftDeleteUnreferencedArtifacts(t *testing.T) require.NoError(t, DeleteDeployment(ctx, project.ID, second.ID)) var secondUpload model.Upload - require.NoError(t, db.DB(ctx).First(&secondUpload, second.UploadID).Error) + require.NoError(t, repository.DB(ctx).First(&secondUpload, second.UploadID).Error) assert.Equal(t, model.UploadStatusDeleted, secondUpload.Status) require.NoError(t, DeleteProject(ctx, project.ID)) var firstUpload model.Upload - require.NoError(t, db.DB(ctx).First(&firstUpload, first.UploadID).Error) + require.NoError(t, repository.DB(ctx).First(&firstUpload, first.UploadID).Error) assert.Equal(t, model.UploadStatusDeleted, firstUpload.Status) _, err = repository.GetPagesProjectByID(ctx, project.ID) assert.Error(t, err) diff --git a/backend/openflare/plugins/server/domain/pages/package_metadata.go b/backend/openflare/plugins/server/domain/pages/package_metadata.go index 7c04bebe..49eba9fa 100644 --- a/backend/openflare/plugins/server/domain/pages/package_metadata.go +++ b/backend/openflare/plugins/server/domain/pages/package_metadata.go @@ -58,7 +58,7 @@ func GetProjectLatestPackageMetadata(ctx context.Context, projectID uint) (*Proj return &ProjectLatestPackageMetadata{ DeploymentID: deployment.ID, Hash: hash, - PackageSize: uploadRecord.FileSize, + PackageSize: uploadRecord.Size, FileCount: deployment.FileCount, TotalSize: deployment.TotalSize, }, nil diff --git a/backend/openflare/plugins/server/domain/pages/package_metadata_test.go b/backend/openflare/plugins/server/domain/pages/package_metadata_test.go index d18fdb3c..09aad456 100644 --- a/backend/openflare/plugins/server/domain/pages/package_metadata_test.go +++ b/backend/openflare/plugins/server/domain/pages/package_metadata_test.go @@ -11,7 +11,7 @@ import ( "testing" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" + "Wavelet/openflare/plugins/server/kernel/repository" ) func TestGetProjectLatestPackageMetadataPreservesZeroTotalSize(t *testing.T) { @@ -46,7 +46,7 @@ func TestGetProjectLatestPackageMetadataPreservesZeroTotalSize(t *testing.T) { if _, err := ActivateDeployment(ctx, project.ID, deployment.ID); err != nil { t.Fatalf("ActivateDeployment() error = %v", err) } - if err := db.DB(ctx).Create(&model.ConfigVersion{ + if err := repository.DB(ctx).Create(&model.ConfigVersion{ Version: "v-package-metadata", SnapshotJSON: fmt.Sprintf( `{"routes":[{"upstream_type":"pages","pages_project_id":%d}]}`, diff --git a/backend/openflare/plugins/server/domain/pages/rebind_test.go b/backend/openflare/plugins/server/domain/pages/rebind_test.go index 20acf1e0..90f71b10 100644 --- a/backend/openflare/plugins/server/domain/pages/rebind_test.go +++ b/backend/openflare/plugins/server/domain/pages/rebind_test.go @@ -9,7 +9,7 @@ import ( "testing" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" + "Wavelet/openflare/plugins/server/kernel/repository" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -36,7 +36,7 @@ func TestRebindSnapshotPagesToCurrentActive(t *testing.T) { Status: model.PagesDeploymentStatusUploaded, FileCount: 1, } - require.NoError(t, db.DB(ctx).Create(old).Error) + require.NoError(t, repository.DB(ctx).Create(old).Error) active := &model.PagesDeployment{ ProjectID: project.ID, DeploymentNumber: 2, @@ -44,8 +44,8 @@ func TestRebindSnapshotPagesToCurrentActive(t *testing.T) { Status: model.PagesDeploymentStatusActive, FileCount: 1, } - require.NoError(t, db.DB(ctx).Create(active).Error) - require.NoError(t, db.DB(ctx).Model(&model.PagesProject{}). + require.NoError(t, repository.DB(ctx).Create(active).Error) + require.NoError(t, repository.DB(ctx).Model(&model.PagesProject{}). Where("id = ?", project.ID). Update("active_deployment_id", active.ID).Error) diff --git a/backend/openflare/plugins/server/domain/pages/routers_source_test.go b/backend/openflare/plugins/server/domain/pages/routers_source_test.go index 47f99332..78cbec6d 100644 --- a/backend/openflare/plugins/server/domain/pages/routers_source_test.go +++ b/backend/openflare/plugins/server/domain/pages/routers_source_test.go @@ -15,8 +15,8 @@ import ( "Wavelet/core/contracts" "Wavelet/openflare/plugins/server/kernel/model" + "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/testhelper" - db "Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" ) @@ -117,7 +117,7 @@ func TestPagesSourceHandlersReturnStableActionErrors(t *testing.T) { false, ) future := time.Now().Add(time.Minute) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", busySource.ID). Updates(map[string]any{ "sync_status": pagesSourceStatusSyncing, @@ -181,7 +181,7 @@ func TestSyncSourceHandlerAcceptsEmptyBodyAndEmptyObject(t *testing.T) { } var executions []model.TaskExecution - if err := db.DB(ctx).Where("task_type = ?", TaskTypePagesSourceAction).Order("id asc").Find(&executions).Error; err != nil { + if err := repository.DB(ctx).Where("task_type = ?", TaskTypePagesSourceAction).Order("id asc").Find(&executions).Error; err != nil { t.Fatalf("list Pages source task executions error = %v, want nil", err) } if got, want := len(executions), 2; got != want { diff --git a/backend/openflare/plugins/server/domain/pages/source_orphan_cleanup.go b/backend/openflare/plugins/server/domain/pages/source_orphan_cleanup.go index 65bde366..22cd41ef 100644 --- a/backend/openflare/plugins/server/domain/pages/source_orphan_cleanup.go +++ b/backend/openflare/plugins/server/domain/pages/source_orphan_cleanup.go @@ -131,7 +131,6 @@ func reconcilePagesOrphanUploadCandidate( } outcome := pagesOrphanCleanupSkipped - uploadLocked := false err := repository.WithPagesTx(ctx, func(tx *gorm.DB) error { scopeOutcome, proceed, err := lockPagesOrphanCleanupScope(ctx, tx, candidate.ID, marker) if err != nil { @@ -141,7 +140,7 @@ func reconcilePagesOrphanUploadCandidate( outcome = scopeOutcome return nil } - lockedOutcome, locked, err := reconcileLockedPagesOrphanUpload( + lockedOutcome, _, err := reconcileLockedPagesOrphanUpload( ctx, tx, candidate.ID, @@ -153,16 +152,15 @@ func reconcilePagesOrphanUploadCandidate( return err } outcome = lockedOutcome - uploadLocked = locked return nil }) if err != nil { return pagesOrphanCleanupSkipped, err } - if uploadLocked { - // Also heal a prior post-commit cache invalidation interruption when the - // status transition was an idempotent no-op. - ofupload.InvalidateUploadMetaCache(ctx, candidate.ID) + if outcome == pagesOrphanCleanupReconciled { + if err := ofupload.Remove(ctx, candidate.ID); err != nil { + return pagesOrphanCleanupSkipped, err + } } return outcome, nil } @@ -252,14 +250,7 @@ func reconcileLockedPagesOrphanUpload( return pagesOrphanCleanupReferenced, true, nil } - transitioned, err := ofupload.RemoveLockedTx(tx, &lockedUpload) - if err != nil { - return pagesOrphanCleanupSkipped, true, err - } - if transitioned { - return pagesOrphanCleanupReconciled, true, nil - } - return pagesOrphanCleanupSkipped, true, nil + return pagesOrphanCleanupReconciled, true, nil } func lockOptionalPagesCleanupRecord( diff --git a/backend/openflare/plugins/server/domain/pages/source_orphan_cleanup_test.go b/backend/openflare/plugins/server/domain/pages/source_orphan_cleanup_test.go index ccd98cb1..f957be81 100644 --- a/backend/openflare/plugins/server/domain/pages/source_orphan_cleanup_test.go +++ b/backend/openflare/plugins/server/domain/pages/source_orphan_cleanup_test.go @@ -12,7 +12,7 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/ofupload" - db "Wavelet/plugins/infra/database" + "Wavelet/openflare/plugins/server/kernel/repository" "gorm.io/gorm" "gorm.io/gorm/clause" @@ -164,7 +164,7 @@ func TestReconcilePagesOrphanUploadsSkipsBusyLeaseAndSourceMismatch(t *testing.T project := createPagesOrphanProject(t, ctx, "busy-orphan") source := createPagesOrphanSource(t, ctx, project.ID) future := realNow.Add(time.Hour) - if err := db.DB(ctx).Create(&model.PagesProjectSourceRuntime{ + if err := repository.DB(ctx).Create(&model.PagesProjectSourceRuntime{ SourceID: source.ID, LeaseToken: "busy-worker", LeaseExpiresAt: &future, @@ -213,7 +213,7 @@ func TestReconcilePagesOrphanUploadsRejectsMalformedMarker(t *testing.T) { metadata := candidate.Metadata metadata.Extra[pagesProjectIDMetadataKey] = "01" candidate.Metadata = metadata - if err := db.DB(ctx).Save(candidate).Error; err != nil { + if err := repository.DB(ctx).Save(candidate).Error; err != nil { t.Fatalf("seed malformed candidate marker error = %v, want nil", err) } @@ -239,7 +239,7 @@ func TestPagesOrphanCleanupAndDeploymentCommitInterleavings(t *testing.T) { if err != nil { t.Fatalf("parsePagesOrphanMarker() error = %v, want nil", err) } - if err := db.DB(ctx).Create(&model.PagesDeployment{ + if err := repository.DB(ctx).Create(&model.PagesDeployment{ ProjectID: project.ID, DeploymentNumber: 1, Checksum: "deployment-first", @@ -276,7 +276,7 @@ func TestPagesOrphanCleanupAndDeploymentCommitInterleavings(t *testing.T) { } target := &model.PagesDeployment{ProjectID: project.ID, UploadID: candidate.ID} - err = db.DB(ctx).Transaction(func(tx *gorm.DB) error { + err = repository.DB(ctx).Transaction(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 @@ -287,7 +287,7 @@ func TestPagesOrphanCleanupAndDeploymentCommitInterleavings(t *testing.T) { t.Errorf("final deployment upload lock after cleanup error = %v, want %v", err, errSourceFinalFence) } var references int64 - if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("upload_id = ?", candidate.ID).Count(&references).Error; err != nil { + if err := repository.DB(ctx).Model(&model.PagesDeployment{}).Where("upload_id = ?", candidate.ID).Count(&references).Error; err != nil { t.Fatalf("count deployment references error = %v, want nil", err) } if references != 0 { @@ -303,7 +303,7 @@ func cleanupOutcomeTotal(summary PagesOrphanCleanupSummary) int { func createPagesOrphanProject(t *testing.T, ctx context.Context, slug string) *model.PagesProject { t.Helper() project := &model.PagesProject{Name: slug, Slug: slug, Enabled: true} - if err := db.DB(ctx).Create(project).Error; err != nil { + if err := repository.DB(ctx).Create(project).Error; err != nil { t.Fatalf("create Pages orphan project %q error = %v, want nil", slug, err) } return project @@ -317,7 +317,7 @@ func createPagesOrphanSource(t *testing.T, ctx context.Context, projectID uint) ConfigVersion: 1, SourceIdentity: "orphan-source-identity", } - if err := db.DB(ctx).Create(source).Error; err != nil { + if err := repository.DB(ctx).Create(source).Error; err != nil { t.Fatalf("create Pages orphan source for project %d error = %v, want nil", projectID, err) } return source @@ -339,21 +339,19 @@ func createPagesOrphanUpload( extra[pagesSourceIDMetadataKey] = strconv.FormatUint(uint64(*sourceID), 10) } candidate := &model.Upload{ - UserID: 999, - FileName: "site.zip", - FilePath: "pages/orphan-site.zip", - FileSize: 64, - MimeType: "application/zip", - Extension: "zip", - Hash: "orphan-checksum", - Type: ofupload.ReservedPagesDeploymentType, - Status: model.UploadStatusUsed, - AccessMode: 0, - Metadata: model.UploadMetadata{Extra: extra}, - CreatedAt: createdAt, - UpdatedAt: createdAt, + UserID: 999, + FileName: "site.zip", + FilePath: "pages/orphan-site.zip", + Size: 64, + MimeType: "application/zip", + Hash: "orphan-checksum", + Type: ofupload.ReservedPagesDeploymentType, + Status: model.UploadStatusUsed, + Metadata: model.UploadMetadata{Extra: extra}, + CreatedAt: createdAt, + UpdatedAt: createdAt, } - if err := db.DB(ctx).Create(candidate).Error; err != nil { + if err := repository.DB(ctx).Table("w_uploads").Create(candidate).Error; err != nil { t.Fatalf("create Pages orphan upload error = %v, want nil", err) } return candidate @@ -362,7 +360,7 @@ func createPagesOrphanUpload( func assertPagesCleanupUploadStatus(t *testing.T, ctx context.Context, uploadID uint64, want model.UploadStatus) { t.Helper() var got model.Upload - if err := db.DB(ctx).First(&got, uploadID).Error; err != nil { + if err := repository.DB(ctx).Table("w_uploads").First(&got, uploadID).Error; err != nil { t.Fatalf("load upload %d error = %v, want nil", uploadID, err) } if got.Status != want { @@ -373,10 +371,10 @@ func assertPagesCleanupUploadStatus(t *testing.T, ctx context.Context, uploadID func assertPagesCleanupTotalStat(t *testing.T, ctx context.Context, want int64) { t.Helper() var stat model.UploadStat - if err := db.DB(ctx).Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").First(&stat).Error; err != nil { + if err := repository.DB(ctx).Where("dimension = ?", model.UploadStatDimensionTotal).First(&stat).Error; err != nil { t.Fatalf("load total upload stat error = %v, want nil", err) } - if stat.FileCount != want { + if int64(stat.FileCount) != want { t.Errorf("total upload stat FileCount = %d, want %d", stat.FileCount, want) } } diff --git a/backend/openflare/plugins/server/domain/pages/source_runtime_test.go b/backend/openflare/plugins/server/domain/pages/source_runtime_test.go index 7500fa83..4228e123 100644 --- a/backend/openflare/plugins/server/domain/pages/source_runtime_test.go +++ b/backend/openflare/plugins/server/domain/pages/source_runtime_test.go @@ -13,7 +13,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) func TestSourceLeaseHeartbeatRenewsAndCancelsOnOwnershipLoss(t *testing.T) { @@ -38,7 +37,7 @@ func TestSourceLeaseHeartbeatRenewsAndCancelsOnOwnershipLoss(t *testing.T) { t.Cleanup(func() { _ = heartbeat.stop() }) var initial model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&initial).Error; err != nil { + if err := repository.DB(ctx).Where("source_id = ?", source.ID).First(&initial).Error; err != nil { t.Fatalf("load initial heartbeat runtime error = %v, want nil", err) } if initial.LeaseExpiresAt == nil { @@ -47,7 +46,7 @@ func TestSourceLeaseHeartbeatRenewsAndCancelsOnOwnershipLoss(t *testing.T) { deadline := time.Now().Add(2 * time.Second) for { var renewedRuntime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&renewedRuntime).Error; err != nil { + if err := repository.DB(ctx).Where("source_id = ?", source.ID).First(&renewedRuntime).Error; err != nil { t.Fatalf("load renewed heartbeat runtime error = %v, want nil", err) } if renewedRuntime.LeaseExpiresAt != nil && renewedRuntime.LeaseExpiresAt.After(*initial.LeaseExpiresAt) { @@ -59,7 +58,7 @@ func TestSourceLeaseHeartbeatRenewsAndCancelsOnOwnershipLoss(t *testing.T) { time.Sleep(10 * time.Millisecond) } - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", source.ID). Update("lease_token", "replacement-owner").Error; err != nil { t.Fatalf("replace heartbeat lease owner error = %v, want nil", err) @@ -166,7 +165,7 @@ func TestAcquireSourceLeaseMutualExclusionExpiryAndTerminalOwnership(t *testing. } past := time.Now().Add(-time.Second) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", source.ID). Update("lease_expires_at", &past).Error; err != nil { t.Fatalf("expire first lease error = %v, want nil", err) @@ -196,7 +195,7 @@ func TestAcquireSourceLeaseMutualExclusionExpiryAndTerminalOwnership(t *testing. t.Fatalf("failSourceLease(expired owner) error = %v, want nil", err) } var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { + if err := repository.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { t.Fatalf("load runtime after takeover error = %v, want nil", err) } if got, want := runtime.LeaseToken, takeover.LeaseToken; got != want { @@ -217,7 +216,7 @@ func TestAcquireSourceLeaseMutualExclusionExpiryAndTerminalOwnership(t *testing. t.Fatalf("failSourceLease(current owner) error = %v, want nil", err) } var failedRuntime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&failedRuntime).Error; err != nil { + if err := repository.DB(ctx).Where("source_id = ?", source.ID).First(&failedRuntime).Error; err != nil { t.Fatalf("load failed runtime error = %v, want nil", err) } if got, want := failedRuntime.SyncStatus, pagesSourceStatusFailed; got != want { @@ -258,7 +257,7 @@ func TestSourceConfigAndProjectContentChangesFenceLease(t *testing.T) { t.Error("renewSourceLease(after source update) = true, want false") } var updatedSource model.PagesProjectSource - if err := db.DB(ctx).Where("id = ?", source.ID).First(&updatedSource).Error; err != nil { + if err := repository.DB(ctx).Where("id = ?", source.ID).First(&updatedSource).Error; err != nil { t.Fatalf("load updated source error = %v, want nil", err) } if got, want := updatedSource.ConfigVersion, source.ConfigVersion+1; got != want { @@ -296,7 +295,7 @@ func TestSourceConfigAndProjectContentChangesFenceLease(t *testing.T) { t.Errorf("ContentConfigVersion after RootDir update = %d, want %d", got, want) } var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { + if err := repository.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { t.Fatalf("load fenced runtime error = %v, want nil", err) } if runtime.LeaseToken != "" || runtime.LeaseExpiresAt != nil { diff --git a/backend/openflare/plugins/server/domain/pages/source_scanner_test.go b/backend/openflare/plugins/server/domain/pages/source_scanner_test.go index 12abc7cc..d718c41c 100644 --- a/backend/openflare/plugins/server/domain/pages/source_scanner_test.go +++ b/backend/openflare/plugins/server/domain/pages/source_scanner_test.go @@ -19,7 +19,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/share/githubrelease" - db "Wavelet/plugins/infra/database" "gorm.io/gorm" ) @@ -42,7 +41,7 @@ func TestGitHubLatestAutoConfigPreservesIdentityAndRuntimeCursor(t *testing.T) { seenRevision := strings.Repeat("a", sourceRevisionHexLength) appliedRevision := strings.Repeat("b", sourceRevisionHexLength) future := time.Now().Add(time.Hour) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", source.ID). Updates(map[string]any{ "etag": `"cursor-etag"`, @@ -67,7 +66,7 @@ func TestGitHubLatestAutoConfigPreservesIdentityAndRuntimeCursor(t *testing.T) { if err := validateGitHubSourceInput(input); err != nil { t.Fatalf("validateGitHubSourceInput(auto latest) error = %v, want nil", err) } - if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error { + if err := repository.DB(ctx).Transaction(func(tx *gorm.DB) error { changed, err := updateGitHubSourceTx(tx, project.ID, input) if err == nil && !changed { return errors.New("auto config update was treated as no-op") @@ -132,7 +131,7 @@ func TestRecoverExpiredPagesSourceLeaseUsesExactCASAndStableJitter(t *testing.T) now := time.Date(2026, 7, 19, 12, 0, 0, 0, time.UTC) usePagesSourceScannerClock(t, now) expiredAt := now.Add(-time.Minute) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", source.ID). Updates(map[string]any{ "sync_status": pagesSourceStatusChecking, @@ -167,7 +166,7 @@ func TestRecoverExpiredPagesSourceLeaseUsesExactCASAndStableJitter(t *testing.T) } renewedExpiry := now.Add(time.Minute) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", source.ID). Updates(map[string]any{ "sync_status": pagesSourceStatusSyncing, @@ -274,24 +273,24 @@ func TestPagesSourceScannerSerialBatchIsolation304AndActualBacklog(t *testing.T) byRepository := make(map[string]int, 22) for index := 1; index <= 22; index++ { project := mustCreatePagesSourceProject(t, ctx, fmt.Sprintf("scanner-batch-%02d", index)) - repository := fmt.Sprintf("scanner/source-%02d", index) + repoPath := fmt.Sprintf("scanner/source-%02d", index) source, runtime := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{ SourceType: PagesSourceTypeGitHubRelease, - RepositoryURL: "https://github.com/" + repository, + RepositoryURL: "https://github.com/" + repoPath, AutoUpdateEnabled: index != 4, CheckIntervalMinutes: 60, }) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", source.ID). Update("next_check_at", &dueAt).Error; err != nil { t.Fatalf("mark source %d due error = %v, want nil", source.ID, err) } - fixtures = append(fixtures, fixture{source: source, runtime: runtime, repository: repository}) - byRepository[repository] = index + fixtures = append(fixtures, fixture{source: source, runtime: runtime, repository: repoPath}) + byRepository[repoPath] = index } busyUntil := now.Add(time.Hour) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", fixtures[0].source.ID). Updates(map[string]any{ "sync_status": pagesSourceStatusChecking, @@ -302,7 +301,7 @@ func TestPagesSourceScannerSerialBatchIsolation304AndActualBacklog(t *testing.T) } stored304Revision := strings.Repeat("3", sourceRevisionHexLength) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", fixtures[2].source.ID). Updates(map[string]any{ "etag": `"stored-etag"`, @@ -313,7 +312,7 @@ func TestPagesSourceScannerSerialBatchIsolation304AndActualBacklog(t *testing.T) t.Fatalf("seed 304 cursor error = %v, want nil", err) } - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", fixtures[4].source.ID). Updates(map[string]any{ "last_applied_revision": strings.Repeat("a", sourceRevisionHexLength), diff --git a/backend/openflare/plugins/server/domain/pages/source_sync_test.go b/backend/openflare/plugins/server/domain/pages/source_sync_test.go index 73d77920..bb4b1e4c 100644 --- a/backend/openflare/plugins/server/domain/pages/source_sync_test.go +++ b/backend/openflare/plugins/server/domain/pages/source_sync_test.go @@ -22,7 +22,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/ofupload" "Wavelet/openflare/share/pagesarchive" - db "Wavelet/plugins/infra/database" "gorm.io/gorm" ) @@ -160,7 +159,7 @@ func TestSyncRemoteSourceAtomicallyActivatesAndReusesChecksum(t *testing.T) { t.Errorf("deployment provenance = label:%q meta:%q, want no query secret", deployment.SourceLabel, deployment.SourceMeta) } var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { + if err := repository.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { t.Fatalf("load source runtime error = %v, want nil", err) } if got, want := runtime.SyncStatus, pagesSourceStatusIdle; got != want { @@ -173,7 +172,7 @@ func TestSyncRemoteSourceAtomicallyActivatesAndReusesChecksum(t *testing.T) { t.Errorf("runtime lease = (%q, %v), want cleared", runtime.LeaseToken, runtime.LeaseExpiresAt) } var uploadRecord model.Upload - if err := db.DB(ctx).First(&uploadRecord, deployment.UploadID).Error; err != nil { + if err := repository.DB(ctx).First(&uploadRecord, deployment.UploadID).Error; err != nil { t.Fatalf("load deployment upload %d error = %v, want nil", deployment.UploadID, err) } if got, want := uploadRecord.Status, model.UploadStatusUsed; got != want { @@ -198,10 +197,10 @@ func TestSyncRemoteSourceAtomicallyActivatesAndReusesChecksum(t *testing.T) { t.Errorf("reused deployment ID = %d, want %d", got, want) } var deploymentCount, uploadCount int64 - if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&deploymentCount).Error; err != nil { + if err := repository.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&deploymentCount).Error; err != nil { t.Fatalf("count source deployments error = %v, want nil", err) } - if err := db.DB(ctx).Model(&model.Upload{}).Count(&uploadCount).Error; err != nil { + if err := repository.DB(ctx).Model(&model.Upload{}).Count(&uploadCount).Error; err != nil { t.Fatalf("count source uploads error = %v, want nil", err) } if got, want := deploymentCount, int64(1); got != want { @@ -267,7 +266,7 @@ func TestSyncRemoteSourceFinalFenceCompensatesIngest(t *testing.T) { packageBytes := testPagesZip(t, map[string]string{"index.html": "never-activate"}) mutationResult := make(chan error, 1) server := newPagesArchiveServer(t, http.StatusOK, packageBytes, func() error { - err := db.DB(context.Background()).Model(&model.PagesProject{}). + err := repository.DB(context.Background()).Model(&model.PagesProject{}). Where("id = ?", project.ID). Update("content_config_version", gorm.Expr("content_config_version + 1")).Error mutationResult <- err @@ -304,14 +303,14 @@ func TestSyncRemoteSourceFinalFenceCompensatesIngest(t *testing.T) { t.Errorf("ActiveDeploymentID after final fence = %v, want %d", storedProject.ActiveDeploymentID, oldActive.ID) } var deployments []model.PagesDeployment - if err := db.DB(ctx).Where("project_id = ?", project.ID).Find(&deployments).Error; err != nil { + if err := repository.DB(ctx).Where("project_id = ?", project.ID).Find(&deployments).Error; err != nil { t.Fatalf("list deployments after final fence error = %v, want nil", err) } if got, want := len(deployments), 1; got != want { t.Errorf("deployment count after final fence = %d, want %d", got, want) } var uploads []model.Upload - if err := db.DB(ctx).Order("id asc").Find(&uploads).Error; err != nil { + if err := repository.DB(ctx).Order("id asc").Find(&uploads).Error; err != nil { t.Fatalf("list uploads after final fence error = %v, want nil", err) } var compensated *model.Upload @@ -328,7 +327,7 @@ func TestSyncRemoteSourceFinalFenceCompensatesIngest(t *testing.T) { t.Errorf("compensated upload Status = %q, want %q", got, want) } var danglingCount int64 - if err := db.DB(ctx).Model(&model.PagesDeployment{}). + if err := repository.DB(ctx).Model(&model.PagesDeployment{}). Where("upload_id = ?", compensated.ID). Count(&danglingCount).Error; err != nil { t.Fatalf("count compensated upload references error = %v, want nil", err) @@ -358,12 +357,12 @@ func TestCommitSourceDeploymentRechecksLeaseAfterUploadLocks(t *testing.T) { if err != nil { t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", first.Deployment.ID, err) } - if err := db.DB(ctx).Model(&model.PagesProject{}). + if err := repository.DB(ctx).Model(&model.PagesProject{}). Where("id = ?", project.ID). Update("active_deployment_id", nil).Error; err != nil { t.Fatalf("clear active deployment error = %v, want nil", err) } - if err := db.DB(ctx).Model(&model.PagesDeployment{}). + if err := repository.DB(ctx).Model(&model.PagesDeployment{}). Where("id = ?", deployment.ID). Update("status", model.PagesDeploymentStatusUploaded).Error; err != nil { t.Fatalf("reset deployment status error = %v, want nil", err) @@ -371,7 +370,7 @@ func TestCommitSourceDeploymentRechecksLeaseAfterUploadLocks(t *testing.T) { snapshot := mustAcquireRemoteSyncLease(t, ctx, source) expiresAt := time.Now().Add(time.Hour) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ? AND lease_token = ?", source.ID, snapshot.LeaseToken). Update("lease_expires_at", &expiresAt).Error; err != nil { t.Fatalf("set deterministic lease expiry error = %v, want nil", err) @@ -465,7 +464,7 @@ func TestCompensateSourceIngestSurvivesCanceledParentContext(t *testing.T) { }) var uploadRecord model.Upload - if err := db.DB(ctx).Where("id = ?", result.Upload.ID).First(&uploadRecord).Error; err != nil { + if err := repository.DB(ctx).Where("id = ?", result.Upload.ID).First(&uploadRecord).Error; err != nil { t.Fatalf("load compensated upload error = %v, want nil", err) } if got, want := uploadRecord.Status, model.UploadStatusDeleted; got != want { @@ -490,7 +489,7 @@ func assertPagesSyncFailureState( t.Errorf("project %d ActiveDeploymentID = %v, want %d", projectID, project.ActiveDeploymentID, oldActiveID) } var deploymentCount int64 - if err := db.DB(ctx).Model(&model.PagesDeployment{}). + if err := repository.DB(ctx).Model(&model.PagesDeployment{}). Where("project_id = ?", projectID). Count(&deploymentCount).Error; err != nil { t.Fatalf("count project %d deployments error = %v, want nil", projectID, err) @@ -499,7 +498,7 @@ func assertPagesSyncFailureState( t.Errorf("project %d deployment count = %d, want %d", projectID, got, want) } var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { + if err := repository.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { t.Fatalf("load source %d runtime error = %v, want nil", sourceID, err) } if got, want := runtime.SyncStatus, pagesSourceStatusFailed; got != want { @@ -529,19 +528,17 @@ func TestCommitSourceDeploymentRejectsDeletedTargetUpload(t *testing.T) { identity := source.SourceIdentity revision := strings.Repeat("d", 64) uploadRecord := &model.Upload{ - ID: 987654321, - UserID: 999, - FileName: "deleted.zip", - FilePath: "deleted.zip", - FileSize: 1, - MimeType: "application/zip", - Extension: "zip", - Hash: revision, - Type: ofupload.ReservedPagesDeploymentType, - Status: model.UploadStatusDeleted, - AccessMode: 0, + ID: 987654321, + UserID: 999, + FileName: "deleted.zip", + FilePath: "deleted.zip", + Size: 1, + MimeType: "application/zip", + Hash: revision, + Type: ofupload.ReservedPagesDeploymentType, + Status: model.UploadStatusDeleted, } - if err := db.DB(ctx).Create(uploadRecord).Error; err != nil { + if err := repository.DB(ctx).Table("w_uploads").Create(uploadRecord).Error; err != nil { t.Fatalf("create deleted upload error = %v, want nil", err) } deployment := &model.PagesDeployment{ @@ -560,10 +557,10 @@ func TestCommitSourceDeploymentRejectsDeletedTargetUpload(t *testing.T) { SourceMeta: `{"provider":"remote_url","display_name":"deleted.zip"}`, TriggerType: pagesSourceTriggerManualSync, } - if err := db.DB(ctx).Create(deployment).Error; err != nil { + if err := repository.DB(ctx).Create(deployment).Error; err != nil { t.Fatalf("create source deployment error = %v, want nil", err) } - if err := db.DB(ctx).Create(&model.PagesDeploymentFile{ + if err := repository.DB(ctx).Create(&model.PagesDeploymentFile{ DeploymentID: deployment.ID, Path: "index.html", Size: 1, @@ -600,7 +597,7 @@ func TestCommitSourceDeploymentRejectsDeletedTargetUpload(t *testing.T) { t.Errorf("ActiveDeploymentID after deleted upload rejection = %v, want nil", storedProject.ActiveDeploymentID) } var activeCount int64 - if err := db.DB(ctx).Model(&model.PagesDeployment{}). + if err := repository.DB(ctx).Model(&model.PagesDeployment{}). Where("project_id = ? AND status = ?", project.ID, model.PagesDeploymentStatusActive). Count(&activeCount).Error; err != nil && !errors.Is(err, gorm.ErrRecordNotFound) { t.Fatalf("count active deployments error = %v, want nil", err) diff --git a/backend/openflare/plugins/server/domain/pages/source_test.go b/backend/openflare/plugins/server/domain/pages/source_test.go index 152aa332..f3f82d6f 100644 --- a/backend/openflare/plugins/server/domain/pages/source_test.go +++ b/backend/openflare/plugins/server/domain/pages/source_test.go @@ -13,16 +13,15 @@ import ( "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) func setupPagesSourceTest(t *testing.T) context.Context { t.Helper() cleanup := setupPagesTestDB(t) t.Cleanup(cleanup) - sqlDB, err := db.DB(t.Context()).DB() + sqlDB, err := repository.DB(t.Context()).DB() if err != nil { - t.Fatalf("db.DB().DB() error = %v, want nil", err) + t.Fatalf("repository.DB().DB() error = %v, want nil", err) } // SQLite :memory: is scoped to one connection. Keeping one connection also // makes lease tests exercise the production CAS without creating empty @@ -93,11 +92,11 @@ func mustConfigureRemoteSource( t.Fatalf("UpdateSource(%d, %q) error = %v, want nil", projectID, remoteURL, err) } var source model.PagesProjectSource - if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { + if err := repository.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { t.Fatalf("load source for project %d error = %v, want nil", projectID, err) } var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { + if err := repository.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil { t.Fatalf("load runtime for source %d error = %v, want nil", source.ID, err) } return &source, &runtime @@ -187,7 +186,7 @@ func TestRemoteSourceCRUDPreservesSecretAndResetsRuntimeByIdentity(t *testing.T) t.Fatalf("UpdateSource(%d, no-op) error = %v, want nil", project.ID, err) } var unchangedSource model.PagesProjectSource - if err := db.DB(ctx).Where("id = ?", source.ID).First(&unchangedSource).Error; err != nil { + if err := repository.DB(ctx).Where("id = ?", source.ID).First(&unchangedSource).Error; err != nil { t.Fatalf("load no-op source error = %v, want nil", err) } if got, want := unchangedSource.ConfigVersion, source.ConfigVersion; got != want { @@ -197,7 +196,7 @@ func TestRemoteSourceCRUDPreservesSecretAndResetsRuntimeByIdentity(t *testing.T) seenRevision := strings.Repeat("a", 64) appliedRevision := strings.Repeat("b", 64) future := time.Now().Add(time.Hour) - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", source.ID). Updates(map[string]any{ "last_seen_revision": seenRevision, @@ -327,10 +326,10 @@ func TestDeleteSourceIsIdempotentAndKeepsDeploymentState(t *testing.T) { SourceType: "manual_upload", TriggerType: "manual_upload", } - if err := db.DB(ctx).Create(deployment).Error; err != nil { + if err := repository.DB(ctx).Create(deployment).Error; err != nil { t.Fatalf("create deployment error = %v, want nil", err) } - if err := db.DB(ctx).Model(&model.PagesProject{}). + if err := repository.DB(ctx).Model(&model.PagesProject{}). Where("id = ?", project.ID). Update("active_deployment_id", deployment.ID).Error; err != nil { t.Fatalf("set active deployment error = %v, want nil", err) @@ -346,13 +345,13 @@ func TestDeleteSourceIsIdempotentAndKeepsDeploymentState(t *testing.T) { } } var sourceCount, runtimeCount, deploymentCount int64 - if err := db.DB(ctx).Model(&model.PagesProjectSource{}).Where("id = ?", source.ID).Count(&sourceCount).Error; err != nil { + if err := repository.DB(ctx).Model(&model.PagesProjectSource{}).Where("id = ?", source.ID).Count(&sourceCount).Error; err != nil { t.Fatalf("count source error = %v, want nil", err) } - if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Count(&runtimeCount).Error; err != nil { + if err := repository.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Count(&runtimeCount).Error; err != nil { t.Fatalf("count runtime error = %v, want nil", err) } - if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("id = ?", deployment.ID).Count(&deploymentCount).Error; err != nil { + if err := repository.DB(ctx).Model(&model.PagesDeployment{}).Where("id = ?", deployment.ID).Count(&deploymentCount).Error; err != nil { t.Fatalf("count deployment error = %v, want nil", err) } if sourceCount != 0 || runtimeCount != 0 || deploymentCount != 1 { diff --git a/backend/openflare/plugins/server/domain/site/apply_log/logics_test.go b/backend/openflare/plugins/server/domain/site/apply_log/logics_test.go index 2e171d5e..9dec97b3 100644 --- a/backend/openflare/plugins/server/domain/site/apply_log/logics_test.go +++ b/backend/openflare/plugins/server/domain/site/apply_log/logics_test.go @@ -11,7 +11,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -29,10 +28,10 @@ func setupApplyLogTestDB(t *testing.T) func() { err = sqliteDB.AutoMigrate(&model.OpenFlareApplyLog{}) require.NoError(t, err) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) } } @@ -49,7 +48,7 @@ func TestListPageAndCleanup(t *testing.T) { {NodeID: "node-logs", Version: "v3", Result: "success", Message: "3", CreatedAt: now}, } for i := range logs { - require.NoError(t, db.DB(ctx).Create(&logs[i]).Error) + require.NoError(t, repository.DB(ctx).Create(&logs[i]).Error) } pageResult, err := ListPage(ctx, ListQuery{ diff --git a/backend/openflare/plugins/server/domain/site/config_version/certificate_snapshot_test.go b/backend/openflare/plugins/server/domain/site/config_version/certificate_snapshot_test.go index fdcf9914..f0986934 100644 --- a/backend/openflare/plugins/server/domain/site/config_version/certificate_snapshot_test.go +++ b/backend/openflare/plugins/server/domain/site/config_version/certificate_snapshot_test.go @@ -20,7 +20,6 @@ import ( oftls "Wavelet/openflare/plugins/server/domain/tls" "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/runtimeconfig" - db "Wavelet/plugins/infra/database" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -29,7 +28,7 @@ import ( func TestBuildCertificateSupportFilesDecryptsSealedPrivateKey(t *testing.T) { cleanup := setupConfigVersionTestDB(t) defer cleanup() - require.NoError(t, db.DB(context.Background()).AutoMigrate(&model.TLSCertificate{})) + require.NoError(t, repository.DB(context.Background()).AutoMigrate(&model.TLSCertificate{})) previous := runtimeconfig.Get() runtimeconfig.SetSessionSecret("test-session-secret-for-tls-seal") @@ -65,7 +64,7 @@ func TestBuildSnapshotReadsZoneDomainCertificates(t *testing.T) { cleanup := setupConfigVersionTestDB(t) defer cleanup() ctx := context.Background() - require.NoError(t, db.DB(ctx).AutoMigrate(&model.TLSCertificate{})) + require.NoError(t, repository.DB(ctx).AutoMigrate(&model.TLSCertificate{})) previous := runtimeconfig.Get() runtimeconfig.SetSessionSecret("test-session-secret-for-zone-domain-snapshots") @@ -81,9 +80,9 @@ func TestBuildSnapshotReadsZoneDomainCertificates(t *testing.T) { route := &model.ProxyRoute{SiteName: "tls-site", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true, EnableHTTPS: true} 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) - require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &route.ID, Domain: "two.example.com", CertID: &second.ID}).Error) + require.NoError(t, repository.DB(ctx).Create(zone).Error) + require.NoError(t, repository.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &route.ID, Domain: "one.example.com", CertID: &first.ID}).Error) + require.NoError(t, repository.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, ProxyRouteID: &route.ID, Domain: "two.example.com", CertID: &second.ID}).Error) bundle, err := buildCurrentConfigBundle(ctx, true) require.NoError(t, err) diff --git a/backend/openflare/plugins/server/domain/site/config_version/logics_test.go b/backend/openflare/plugins/server/domain/site/config_version/logics_test.go index 24032047..def324ae 100644 --- a/backend/openflare/plugins/server/domain/site/config_version/logics_test.go +++ b/backend/openflare/plugins/server/domain/site/config_version/logics_test.go @@ -16,7 +16,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" openrestyrender "Wavelet/openflare/share/render/openresty" "Wavelet/pkg/cache/ram" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -45,9 +44,9 @@ func setupConfigVersionTestDB(t *testing.T) func() { &model.SystemConfig{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) ram.ResetForTest() } } @@ -55,9 +54,9 @@ func setupConfigVersionTestDB(t *testing.T) func() { func createSnapshotZoneDomains(t *testing.T, ctx context.Context, route *model.ProxyRoute, domains ...string) { t.Helper() zone := &model.Zone{Domain: fmt.Sprintf("zone-%d.example", route.ID)} - require.NoError(t, db.DB(ctx).Create(zone).Error) + require.NoError(t, repository.DB(ctx).Create(zone).Error) for _, domain := range domains { - require.NoError(t, db.DB(ctx).Create(&model.ZoneDomain{ + require.NoError(t, repository.DB(ctx).Create(&model.ZoneDomain{ ZoneID: zone.ID, ProxyRouteID: &route.ID, Domain: domain, @@ -69,7 +68,7 @@ func TestListConfigVersionsOrdersByCreatedAtDesc(t *testing.T) { cleanup := setupConfigVersionTestDB(t) defer cleanup() ctx := context.Background() - conn := db.DB(ctx) + conn := repository.DB(ctx) require.NotNil(t, conn) newer := &model.ConfigVersion{ @@ -220,7 +219,7 @@ func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testi graphJSON, err := json.Marshal(snapshotPoWGraph()) require.NoError(t, err) globalGroup.Graph = string(graphJSON) - require.NoError(t, db.DB(ctx).Model(globalGroup).Update("graph", globalGroup.Graph).Error) + require.NoError(t, repository.DB(ctx).Model(globalGroup).Update("graph", globalGroup.Graph).Error) bundle, err := buildCurrentConfigBundle(ctx, true) require.NoError(t, err) diff --git a/backend/openflare/plugins/server/domain/site/config_version/origin_error_page_snapshot_test.go b/backend/openflare/plugins/server/domain/site/config_version/origin_error_page_snapshot_test.go index 58a2cae9..03007971 100644 --- a/backend/openflare/plugins/server/domain/site/config_version/origin_error_page_snapshot_test.go +++ b/backend/openflare/plugins/server/domain/site/config_version/origin_error_page_snapshot_test.go @@ -9,13 +9,12 @@ import ( "testing" "Wavelet/openflare/plugins/server/kernel/model" + "Wavelet/openflare/plugins/server/kernel/repository" + "Wavelet/openflare/plugins/server/kernel/testhelper" "Wavelet/pkg/cache/ram" - db "Wavelet/plugins/infra/database" - "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "gorm.io/gorm" ) func setupOriginErrorPageSnapshotDB(t *testing.T) func() { @@ -25,15 +24,10 @@ func setupOriginErrorPageSnapshotDB(t *testing.T) func() { // 重置,否则 shuffle 下先跑的用例会污染后跑的用例。 ram.ResetForTest() - sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ - DisableForeignKeyConstraintWhenMigrating: true, - }) - require.NoError(t, err) - require.NoError(t, sqliteDB.AutoMigrate(&model.SystemConfig{})) - db.SetDB(sqliteDB) + _, _, cleanup := testhelper.SetupTestEnvironment(t) return func() { - db.SetDB(nil) + cleanup() ram.ResetForTest() } } @@ -58,15 +52,9 @@ func TestBuildOpenRestyConfigSnapshotOriginErrorPageCustom(t *testing.T) { defer cleanup() ctx := context.Background() - require.NoError(t, db.DB(ctx).Create(&model.SystemConfig{ - Key: model.ConfigKeyOriginErrorPageEnabled, Value: "false", Type: "business", - }).Error) - require.NoError(t, db.DB(ctx).Create(&model.SystemConfig{ - Key: model.ConfigKeyOriginErrorPageStatusCodes, Value: `["522","500-502"]`, Type: "business", - }).Error) - require.NoError(t, db.DB(ctx).Create(&model.SystemConfig{ - Key: model.ConfigKeyOriginErrorPageHTML, Value: "

{{status}}

", Type: "business", - }).Error) + require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyOriginErrorPageEnabled, "false")) + require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyOriginErrorPageStatusCodes, `["522","500-502"]`)) + require.NoError(t, repository.SaveOrUpdateSystemConfig(ctx, model.ConfigKeyOriginErrorPageHTML, "

{{status}}

")) snapshot := buildOpenRestyConfigSnapshot(ctx) assert.False(t, snapshot.OriginErrorPageEnabled) diff --git a/backend/openflare/plugins/server/domain/site/config_version/pages_snapshot_test.go b/backend/openflare/plugins/server/domain/site/config_version/pages_snapshot_test.go index ecdbc7db..3ba9202b 100644 --- a/backend/openflare/plugins/server/domain/site/config_version/pages_snapshot_test.go +++ b/backend/openflare/plugins/server/domain/site/config_version/pages_snapshot_test.go @@ -12,7 +12,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" openrestyrender "Wavelet/openflare/share/render/openresty" - db "Wavelet/plugins/infra/database" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -95,7 +94,7 @@ func TestBuildSnapshotPagesDeploymentRejectsUnsafeStoredPaths(t *testing.T) { func requireDB(t *testing.T, ctx context.Context) *gorm.DB { t.Helper() - conn := db.DB(ctx) + conn := repository.DB(ctx) require.NotNil(t, conn) require.NoError(t, conn.AutoMigrate( &model.PagesProject{}, diff --git a/backend/openflare/plugins/server/domain/site/config_version/waf_graph_snapshot_test.go b/backend/openflare/plugins/server/domain/site/config_version/waf_graph_snapshot_test.go index e523d463..76e72eed 100644 --- a/backend/openflare/plugins/server/domain/site/config_version/waf_graph_snapshot_test.go +++ b/backend/openflare/plugins/server/domain/site/config_version/waf_graph_snapshot_test.go @@ -13,7 +13,6 @@ import ( "Wavelet/openflare/plugins/server/domain/waf" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -38,7 +37,7 @@ func TestBuildSnapshotRejectsOversizedAggregateWAFIPGroups(t *testing.T) { Enabled: true, IPList: string(ipList), } - require.NoError(t, db.DB(ctx).Create(group).Error) + require.NoError(t, repository.DB(ctx).Create(group).Error) groupIDs = append(groupIDs, group.ID) } createSnapshotRule(t, ctx, "oversized-aggregate", snapshotIPMatchGraphForGroups(groupIDs)) @@ -59,8 +58,8 @@ func TestWAFGraphSnapshotPreservesOrderAndGraphReferences(t *testing.T) { referenced := &model.OpenFlareWAFIPGroup{Name: "referenced", Type: "manual", Enabled: true, IPList: `["192.0.2.1"]`} unused := &model.OpenFlareWAFIPGroup{Name: "unused", Type: "manual", Enabled: true, IPList: `["198.51.100.1"]`} - require.NoError(t, db.DB(ctx).Create(referenced).Error) - require.NoError(t, db.DB(ctx).Create(unused).Error) + require.NoError(t, repository.DB(ctx).Create(referenced).Error) + require.NoError(t, repository.DB(ctx).Create(unused).Error) customA := createSnapshotRule(t, ctx, "custom-a", waf.DefaultRuleGraph()) customB := createSnapshotRule(t, ctx, "custom-b", snapshotIPMatchGraph(referenced.ID)) @@ -114,7 +113,7 @@ func TestBuildSnapshotRejectsInvalidWAFGraph(t *testing.T) { ctx := context.Background() invalid := &model.OpenFlareWAFRuleGroup{Name: "invalid", Enabled: true, Graph: `{"schema_version":1,"nodes":[],"edges":[]}`, Revision: 1} - require.NoError(t, db.DB(ctx).Create(invalid).Error) + require.NoError(t, repository.DB(ctx).Create(invalid).Error) _, err := buildSnapshotWAFDocument(ctx, nil) require.ErrorContains(t, err, "invalid") } @@ -124,7 +123,7 @@ func createSnapshotRule(t *testing.T, ctx context.Context, name string, graph wa raw, err := json.Marshal(graph) require.NoError(t, err) rule := &model.OpenFlareWAFRuleGroup{Name: name, Enabled: true, Graph: string(raw), Revision: 1} - require.NoError(t, db.DB(ctx).Create(rule).Error) + require.NoError(t, repository.DB(ctx).Create(rule).Error) return rule } diff --git a/backend/openflare/plugins/server/domain/site/origin/logics_test.go b/backend/openflare/plugins/server/domain/site/origin/logics_test.go index 69c80bd6..8eb82507 100644 --- a/backend/openflare/plugins/server/domain/site/origin/logics_test.go +++ b/backend/openflare/plugins/server/domain/site/origin/logics_test.go @@ -8,7 +8,7 @@ import ( "testing" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" + "Wavelet/openflare/plugins/server/kernel/repository" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -25,9 +25,9 @@ func setupOriginTestDB(t *testing.T) func() { require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.Origin{})) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) } } diff --git a/backend/openflare/plugins/server/domain/site/proxy_route/logics_test.go b/backend/openflare/plugins/server/domain/site/proxy_route/logics_test.go index 844e5b44..20b82ece 100644 --- a/backend/openflare/plugins/server/domain/site/proxy_route/logics_test.go +++ b/backend/openflare/plugins/server/domain/site/proxy_route/logics_test.go @@ -9,7 +9,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/repository" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -29,21 +28,21 @@ func setupProxyRouteTestDB(t *testing.T) func() { &model.TLSCertificate{}, &model.PagesProject{}, )) - db.SetDB(sqliteDB) - return func() { db.SetDB(nil) } + repository.SetDBForTest(sqliteDB) + return func() { repository.SetDBForTest(nil) } } func createZoneDomain(t *testing.T, ctx context.Context, domain string, certID *uint) *model.ZoneDomain { t.Helper() zone := &model.Zone{Domain: "example.com"} var existing model.Zone - if err := db.DB(ctx).Where("domain = ?", zone.Domain).First(&existing).Error; err == nil { + if err := repository.DB(ctx).Where("domain = ?", zone.Domain).First(&existing).Error; err == nil { zone = &existing } else { - require.NoError(t, db.DB(ctx).Create(zone).Error) + require.NoError(t, repository.DB(ctx).Create(zone).Error) } item := &model.ZoneDomain{ZoneID: zone.ID, Domain: domain, CertID: certID} - require.NoError(t, db.DB(ctx).Create(item).Error) + require.NoError(t, repository.DB(ctx).Create(item).Error) return item } @@ -102,7 +101,7 @@ func TestPagesRouteLocksAndRevalidatesTargetProject(t *testing.T) { Enabled: true, ActiveDeploymentID: &activeDeploymentID, } - require.NoError(t, db.DB(ctx).Create(project).Error) + require.NoError(t, repository.DB(ctx).Create(project).Error) view, err := CreateProxyRoute(ctx, Input{ SiteName: "pages", @@ -115,7 +114,7 @@ func TestPagesRouteLocksAndRevalidatesTargetProject(t *testing.T) { require.NotNil(t, view.PagesProjectID) assert.Equal(t, project.ID, *view.PagesProjectID) - require.NoError(t, db.DB(ctx).Delete(&model.PagesProject{}, project.ID).Error) + require.NoError(t, repository.DB(ctx).Delete(&model.PagesProject{}, project.ID).Error) route := &model.ProxyRoute{UpstreamType: proxyRouteUpstreamTypePages, PagesProjectID: &project.ID} err = repository.WithProxyRouteTx(ctx, func(tx *gorm.DB) error { return lockPagesProjectsForRouteMutation(tx, 0, route) diff --git a/backend/openflare/plugins/server/domain/site/zone/legacy_import_test.go b/backend/openflare/plugins/server/domain/site/zone/legacy_import_test.go index eb624686..ce12c196 100644 --- a/backend/openflare/plugins/server/domain/site/zone/legacy_import_test.go +++ b/backend/openflare/plugins/server/domain/site/zone/legacy_import_test.go @@ -8,7 +8,7 @@ import ( "database/sql" "testing" - db "Wavelet/plugins/infra/database" + "Wavelet/openflare/plugins/server/kernel/repository" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -65,10 +65,10 @@ func setupLegacyImportDB(t *testing.T) (*sql.DB, func()) { require.NoError(t, err) } - previous := db.DB(context.Background()) - db.SetDB(gormDB) + previous := repository.DB(context.Background()) + repository.SetDBForTest(gormDB) return sqlDB, func() { - db.SetDB(previous) + repository.SetDBForTest(previous) _ = sqlDB.Close() } } diff --git a/backend/openflare/plugins/server/domain/site/zone/logics_test.go b/backend/openflare/plugins/server/domain/site/zone/logics_test.go index 0637a8e3..0e140910 100644 --- a/backend/openflare/plugins/server/domain/site/zone/logics_test.go +++ b/backend/openflare/plugins/server/domain/site/zone/logics_test.go @@ -12,7 +12,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/testhelper" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/require" @@ -24,8 +23,8 @@ func setupZoneDB(t *testing.T) context.Context { conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{DisableForeignKeyConstraintWhenMigrating: true}) require.NoError(t, err) require.NoError(t, conn.AutoMigrate(&model.Zone{}, &model.ZoneDomain{}, &model.TLSCertificate{}, &model.CFPointingGroup{}, &model.CFPointingMember{})) - db.SetDB(conn) - t.Cleanup(func() { db.SetDB(nil) }) + repository.SetDBForTest(conn) + t.Cleanup(func() { repository.SetDBForTest(nil) }) return context.Background() } diff --git a/backend/openflare/plugins/server/domain/tls/logics.go b/backend/openflare/plugins/server/domain/tls/logics.go index cf0af83a..560f3529 100644 --- a/backend/openflare/plugins/server/domain/tls/logics.go +++ b/backend/openflare/plugins/server/domain/tls/logics.go @@ -12,10 +12,10 @@ import ( "mime/multipart" "strings" - "Wavelet/openflare/plugins/server/kernel/repository" - "Wavelet/openflare/plugins/server/kernel/model" + "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/task" + "Wavelet/pkg/util" ) // CertificateInput TLS 证书创建/更新请求。 @@ -202,10 +202,10 @@ func ApplyCertificate(ctx context.Context, input ApplyInput) (*model.TLSCertific returned := sanitizeCertificateForResponse(cert) obtainFn := obtainTLSCertificate // 捕获当前实现,避免 goroutine 内读可变包变量(测试热替换) - go func(c *model.TLSCertificate) { + util.Go(func() { asyncCtx := context.WithoutCancel(ctx) - _ = obtainFn(asyncCtx, c) - }(cert) + _ = obtainFn(asyncCtx, cert) + }) return returned, nil } @@ -233,10 +233,10 @@ func UpdateACMECertificate(ctx context.Context, id uint, input ApplyInput) (*mod returned := sanitizeCertificateForResponse(cert) obtainFn := obtainTLSCertificate // 捕获当前实现,避免 goroutine 内读可变包变量(测试热替换) - go func(c *model.TLSCertificate) { + util.Go(func() { asyncCtx := context.WithoutCancel(ctx) - _ = obtainFn(asyncCtx, c) - }(cert) + _ = obtainFn(asyncCtx, cert) + }) return returned, nil } @@ -266,12 +266,12 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (* } obtainFn := obtainTLSCertificate // 捕获当前实现,避免 goroutine 内读可变包变量(测试热替换) - go func(c *model.TLSCertificate) { + util.Go(func() { asyncCtx := context.WithoutCancel(ctx) - if err := obtainFn(asyncCtx, c); err != nil { + if err := obtainFn(asyncCtx, cert); err != nil { return } - latest, err := repository.GetTLSCertificateByID(asyncCtx, c.ID) + latest, err := repository.GetTLSCertificateByID(asyncCtx, cert.ID) if err != nil { return } @@ -279,7 +279,7 @@ func ConvertCertificateToACME(ctx context.Context, id uint, input ApplyInput) (* latest.ApplyStatus = tlsApplyStatusReady latest.ApplyMessage = "" _ = repository.SaveTLSCertificate(asyncCtx, latest) - }(cert) + }) return sanitizeCertificateForResponse(cert), nil } diff --git a/backend/openflare/plugins/server/domain/tls/logics_test.go b/backend/openflare/plugins/server/domain/tls/logics_test.go index 38d07f92..be4f1dac 100644 --- a/backend/openflare/plugins/server/domain/tls/logics_test.go +++ b/backend/openflare/plugins/server/domain/tls/logics_test.go @@ -24,7 +24,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/runtimeconfig" "Wavelet/pkg/idgen" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -51,7 +50,7 @@ func setupTLSTestDB(t *testing.T) func() { &model.TaskExecution{}, // 异步任务执行记录也需要 migrate )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) require.NoError(t, idgen.Init(1)) previous := runtimeconfig.Get() runtimeconfig.SetSessionSecret("test_session_secret_for_tls_encryption") @@ -59,7 +58,7 @@ func setupTLSTestDB(t *testing.T) func() { oftask.SetService(&testhelper.NoopTaskService{}) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) runtimeconfig.Set(previous) credential.SetSessionSecret(previous.SessionSecret) tlsTestDBMu.Unlock() @@ -75,8 +74,8 @@ func TestDeleteCertificateRejectsZoneDomainReference(t *testing.T) { certificate, err := CreateCertificate(ctx, CertificateInput{Name: "api-cert", CertPEM: certPEM, KeyPEM: keyPEM}) require.NoError(t, err) 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, Domain: "api.example.com", CertID: &certificate.ID}).Error) + require.NoError(t, repository.DB(ctx).Create(zone).Error) + require.NoError(t, repository.DB(ctx).Create(&model.ZoneDomain{ZoneID: zone.ID, Domain: "api.example.com", CertID: &certificate.ID}).Error) err = DeleteCertificate(ctx, certificate.ID) require.EqualError(t, err, errCertificateDeleteReferenced) diff --git a/backend/openflare/plugins/server/domain/tls/ssl_renew_test.go b/backend/openflare/plugins/server/domain/tls/ssl_renew_test.go index fecee873..fe2965a5 100644 --- a/backend/openflare/plugins/server/domain/tls/ssl_renew_test.go +++ b/backend/openflare/plugins/server/domain/tls/ssl_renew_test.go @@ -13,7 +13,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/runtimeconfig" oftask "Wavelet/openflare/plugins/server/kernel/task" "Wavelet/openflare/plugins/server/kernel/testhelper" - db "Wavelet/plugins/infra/database" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -23,7 +22,7 @@ func setupSSLRenewTestDB(t *testing.T) func() { t.Helper() _, _, cleanup := testhelper.SetupTestEnvironment(t) - require.NoError(t, db.DB(context.Background()).AutoMigrate(&model.TLSCertificate{}, &model.TaskExecution{})) + require.NoError(t, repository.DB(context.Background()).AutoMigrate(&model.TLSCertificate{}, &model.TaskExecution{})) previous := runtimeconfig.Get() runtimeconfig.SetSessionSecret("test_session_secret_for_ssl_renew") oftask.SetService(&testhelper.NoopTaskService{}) diff --git a/backend/openflare/plugins/server/domain/waf/ip_group_sync_test.go b/backend/openflare/plugins/server/domain/waf/ip_group_sync_test.go index 83419ee2..1ec9f3f2 100644 --- a/backend/openflare/plugins/server/domain/waf/ip_group_sync_test.go +++ b/backend/openflare/plugins/server/domain/waf/ip_group_sync_test.go @@ -15,7 +15,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/testhelper" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -35,10 +34,10 @@ func setupIPGroupSyncTestDB(t *testing.T) func() { &model.OpenFlareWAFIPGroup{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) testhelper.SetupLogStoresForTest(t) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) } } diff --git a/backend/openflare/plugins/server/domain/waf/logics_test.go b/backend/openflare/plugins/server/domain/waf/logics_test.go index 544fed91..e7d6ac4e 100644 --- a/backend/openflare/plugins/server/domain/waf/logics_test.go +++ b/backend/openflare/plugins/server/domain/waf/logics_test.go @@ -10,7 +10,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" @@ -32,9 +31,9 @@ func setupWAFTestDB(t *testing.T) func() { &model.OriginProxyRoute{}, )) - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) return func() { - db.SetDB(nil) + repository.SetDBForTest(nil) } } diff --git a/backend/openflare/plugins/server/domain/waf/rule_logics_test.go b/backend/openflare/plugins/server/domain/waf/rule_logics_test.go index 88c78edd..2b8d592e 100644 --- a/backend/openflare/plugins/server/domain/waf/rule_logics_test.go +++ b/backend/openflare/plugins/server/domain/waf/rule_logics_test.go @@ -16,7 +16,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/pkg/response" - db "Wavelet/plugins/infra/database" "github.com/gin-gonic/gin" "github.com/stretchr/testify/assert" @@ -84,7 +83,7 @@ func TestRuleHandlersMapFailures(t *testing.T) { require.NoError(t, err) return cleanup }, want: http.StatusConflict}, - {name: "database failure", method: http.MethodGet, path: "/rules", setup: func(t *testing.T) func() { t.Helper(); db.SetDB(nil); return func() {} }, want: http.StatusInternalServerError}, + {name: "database failure", method: http.MethodGet, path: "/rules", setup: func(t *testing.T) func() { t.Helper(); repository.SetDBForTest(nil); return func() {} }, want: http.StatusInternalServerError}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { @@ -157,7 +156,7 @@ func TestReplaceSiteRuleGroupsPreservesOrderAndRejectsGlobal(t *testing.T) { defer cleanup() ctx := context.Background() - require.NoError(t, db.DB(ctx).Create(&model.OriginProxyRoute{ID: 7, Domain: "example.com"}).Error) + require.NoError(t, repository.DB(ctx).Create(&model.OriginProxyRoute{ID: 7, Domain: "example.com"}).Error) first, err := CreateRule(ctx, CreateRuleInput{Name: "first"}) require.NoError(t, err) second, err := CreateRule(ctx, CreateRuleInput{Name: "second"}) diff --git a/backend/openflare/plugins/server/kernel/geoip/runtime_test.go b/backend/openflare/plugins/server/kernel/geoip/runtime_test.go index 3b8e9c40..40554e60 100644 --- a/backend/openflare/plugins/server/kernel/geoip/runtime_test.go +++ b/backend/openflare/plugins/server/kernel/geoip/runtime_test.go @@ -8,8 +8,8 @@ import ( "testing" "Wavelet/openflare/plugins/server/kernel/model" + "Wavelet/openflare/plugins/server/kernel/repository" pkggeoip "Wavelet/openflare/share/geoip" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "gorm.io/gorm" @@ -23,16 +23,16 @@ func TestEnsureRuntimeProviderInitializesConfiguredProvider(t *testing.T) { if err := sqliteDB.AutoMigrate(&model.SystemConfig{}); err != nil { t.Fatalf("migrate: %v", err) } - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) t.Cleanup(func() { - db.SetDB(nil) + repository.SetDBForTest(nil) ResetRuntimeForTest() }) ctx := context.Background() ResetRuntimeForTest() // 通过 SystemConfig 设置 GeoIPProvider 配置 - if err := db.DB(ctx).Create(&model.SystemConfig{ + if err := repository.DB(ctx).Create(&model.SystemConfig{ Key: model.ConfigKeyGeoIPProvider, Value: pkggeoip.ProviderIPInfo, Type: "business", diff --git a/backend/openflare/plugins/server/kernel/model/analytics/user_access_log.go b/backend/openflare/plugins/server/kernel/model/analytics/user_access_log.go index 46ab3759..83a7c8e5 100644 --- a/backend/openflare/plugins/server/kernel/model/analytics/user_access_log.go +++ b/backend/openflare/plugins/server/kernel/model/analytics/user_access_log.go @@ -4,8 +4,34 @@ package analytics import ( - risklogstore "Wavelet/plugins/domain/risk_control/logstore" + "time" ) -// UserAccessLog is Wavelet risk_control's w_user_access_logs entity. -type UserAccessLog = risklogstore.UserAccessLog +const ( + userAccessLogTableName = "w_user_access_logs" + userAccessLogInsertColumns = "id, user_id, path, method, ip, user_agent, headers, status, latency, created_at" +) + +// UserAccessLog represents a user HTTP access log entry. +type UserAccessLog struct { + ID uint64 `gorm:"column:id"` + UserID uint64 `gorm:"column:user_id"` + Path string `gorm:"column:path"` + Method string `gorm:"column:method"` + IP string `gorm:"column:ip"` + UserAgent string `gorm:"column:user_agent"` + Headers string `gorm:"column:headers"` + Status int32 `gorm:"column:status"` + Latency int64 `gorm:"column:latency"` + CreatedAt time.Time `gorm:"column:created_at"` +} + +// TableName returns the table name. +func (UserAccessLog) TableName() string { + return userAccessLogTableName +} + +// InsertColumns returns comma-separated column names for batch insert. +func (UserAccessLog) InsertColumns() string { + return userAccessLogInsertColumns +} diff --git a/backend/openflare/plugins/server/kernel/model/errs.go b/backend/openflare/plugins/server/kernel/model/errs.go index 51da0c48..ed7ac7ab 100644 --- a/backend/openflare/plugins/server/kernel/model/errs.go +++ b/backend/openflare/plugins/server/kernel/model/errs.go @@ -2,16 +2,3 @@ // SPDX-License-Identifier: Apache-2.0 package model - -// Domain validation messages used by model.Validate and other no-IO rules. -// Persistence / data-access messages belong in internal/repository (do not import repository). -const ( - errTemplateKeyRequired = "模板标识符不能为空" - errTemplateNameRequired = "模板名称不能为空" - errTemplateContentRequired = "模板内容不能为空" - errAuthSourceNameRequired = "认证源名称不能为空" - errAuthSourceNameInvalid = "认证源名称只能包含字母、数字、短横线或下划线,且必须以字母或数字开头" - errAuthSourceTypeUnsupported = "认证源类型仅支持 oidc" - errAuthSourceDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL" - errAuthSourceClientCredentialsRequired = "启用认证源前必须配置 Client ID 和 Client Secret" //nolint:gosec // false positive: this is an error message, not hardcoded credentials -) diff --git a/backend/openflare/plugins/server/kernel/model/platform_aliases.go b/backend/openflare/plugins/server/kernel/model/platform_aliases.go index 8a339e52..895030a1 100644 --- a/backend/openflare/plugins/server/kernel/model/platform_aliases.go +++ b/backend/openflare/plugins/server/kernel/model/platform_aliases.go @@ -9,10 +9,9 @@ import ( "encoding/hex" "fmt" - adminmodel "Wavelet/plugins/domain/admin/model" - authmodel "Wavelet/plugins/domain/auth" - uploadmodels "Wavelet/plugins/domain/upload/models" - usermodel "Wavelet/plugins/domain/user" + "time" + + "Wavelet/core/contracts" ) const ( @@ -20,57 +19,161 @@ const ( maskThreshold = 8 ) -// User is the Wavelet w_users entity. -type User = usermodel.User +// User represents a user identity view. +type User struct { + ID uint64 `json:"id,string" gorm:"primaryKey"` + Username string `json:"username"` + Password string `json:"-"` + Nickname string `json:"nickname"` + Email string `json:"email"` + IsAdmin bool `json:"is_admin"` + IsActive bool `json:"is_active"` + LastLoginAt time.Time `json:"last_login_at"` +} -// AccessToken is the Wavelet w_access_tokens entity. -type AccessToken = usermodel.AccessToken +func (User) TableName() string { + return "w_users" +} -// AuthSource is the Wavelet w_auth_sources entity. -type AuthSource = authmodel.AuthSource +func (u *User) SetEncryptedPassword(pwd string) error { + u.Password = pwd + return nil +} -// ExternalAccount is the Wavelet w_external_accounts entity. -type ExternalAccount = authmodel.ExternalAccount +// AccessToken represents an access token view. +type AccessToken struct { + ID uint64 `json:"id" gorm:"primaryKey"` + UserID uint64 `json:"user_id"` + Name string `json:"name"` + Token string `json:"token"` + MaskedToken string `json:"masked_token"` + TokenHash string `json:"token_hash"` + IsAdmin bool `json:"is_admin"` + ExpiredAt time.Time `json:"expired_at"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} -// TaskExecution is the Wavelet w_task_executions entity. -type TaskExecution = adminmodel.TaskExecution +func (AccessToken) TableName() string { + return "w_access_tokens" +} -// Template is the Wavelet w_templates entity. -type Template = adminmodel.Template +// AuthSource represents an authentication source view. +type AuthSource struct { + ID uint64 `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + DisplayName string `json:"display_name"` + IconURL string `json:"icon_url"` + IsActive bool `json:"is_active"` +} -// Schedule is the Wavelet w_schedules entity. -type Schedule = adminmodel.Schedule +// TaskExecution represents task execution entity. +type TaskExecution struct { + ID uint64 `json:"id" gorm:"primaryKey"` + TaskID string `json:"task_id" gorm:"size:64;index"` + TaskType string `json:"task_type" gorm:"size:100;index"` + TaskName string `json:"task_name" gorm:"size:255"` + Status string `json:"status" gorm:"size:20;index"` + Retryable bool `json:"retryable"` + MaxRetry int `json:"max_retry"` + RetryCount int `json:"retry_count"` + Log string `json:"log" gorm:"type:text"` + ErrorMessage string `json:"error_message" gorm:"type:text"` + Result string `json:"result" gorm:"type:text"` + StartedAt *time.Time `json:"started_at"` + FinishedAt *time.Time `json:"finished_at"` + Duration int64 `json:"duration"` + Payload string `json:"payload" gorm:"type:text"` + TriggeredBy string `json:"triggered_by" gorm:"size:100"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} -// Upload is the Wavelet w_uploads entity. -type Upload = uploadmodels.Upload +func (TaskExecution) TableName() string { + return "w_task_executions" +} -// UploadMetadata is the Wavelet upload metadata JSON. -type UploadMetadata = uploadmodels.UploadMetadata - -// UploadStatus is the Wavelet upload status. -type UploadStatus = uploadmodels.UploadStatus - -// UploadStat is the Wavelet w_upload_stats entity. -type UploadStat = uploadmodels.UploadStat +type UploadStatus = string const ( - // UploadStatusPending is a newly stored unused upload. - UploadStatusPending = uploadmodels.UploadStatusPending - // UploadStatusUsed is an in-use upload. - UploadStatusUsed = uploadmodels.UploadStatusUsed - // UploadStatusDeleted is a soft-deleted upload. - UploadStatusDeleted = uploadmodels.UploadStatusDeleted - - // UploadStatDimensionTotal is the total stats dimension. - UploadStatDimensionTotal = uploadmodels.UploadStatDimensionTotal - // UploadStatDimensionType is the type stats dimension. - UploadStatDimensionType = uploadmodels.UploadStatDimensionType - // UploadStatDimensionCategory is the category stats dimension. - UploadStatDimensionCategory = uploadmodels.UploadStatDimensionCategory - // UploadStatDimensionTrend is the trend stats dimension. - UploadStatDimensionTrend = uploadmodels.UploadStatDimensionTrend + UploadStatusPending UploadStatus = "pending" + UploadStatusUsed UploadStatus = "used" + UploadStatusDeleted UploadStatus = "deleted" ) +// UploadMetadata represents upload metadata JSON. +type UploadMetadata = contracts.UploadMetadataDTO + +// Upload represents file upload entity. +type Upload struct { + ID uint64 `json:"id" gorm:"primaryKey"` + UserID uint64 `json:"user_id" gorm:"index"` + FileName string `json:"file_name" gorm:"size:255"` + FilePath string `json:"file_path" gorm:"size:500"` + MimeType string `json:"mime_type" gorm:"size:100"` + Size int64 `json:"size"` + Hash string `json:"hash" gorm:"size:64"` + Status string `json:"status" gorm:"type:varchar(20)"` + Type string `json:"type" gorm:"size:50;index"` + Metadata contracts.UploadMetadataDTO `json:"metadata"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +func (Upload) TableName() string { + return "w_uploads" +} + +func (u *Upload) ToDTO() contracts.UploadDTO { + return contracts.UploadDTO{ + ID: u.ID, + UserID: u.UserID, + FileName: u.FileName, + FilePath: u.FilePath, + MimeType: u.MimeType, + Size: u.Size, + Hash: u.Hash, + Status: u.Status, + Type: u.Type, + Metadata: u.Metadata, + CreatedAt: u.CreatedAt, + UpdatedAt: u.UpdatedAt, + } +} + +func FromUploadDTO(d contracts.UploadDTO) Upload { + return Upload{ + ID: d.ID, + UserID: d.UserID, + FileName: d.FileName, + FilePath: d.FilePath, + MimeType: d.MimeType, + Size: d.Size, + Hash: d.Hash, + Status: d.Status, + Type: d.Type, + Metadata: d.Metadata, + CreatedAt: d.CreatedAt, + UpdatedAt: d.UpdatedAt, + } +} + +const UploadStatDimensionTotal = "total" + +// UploadStat tracks upload statistics by dimension. +type UploadStat struct { + ID uint64 `gorm:"primaryKey"` + Dimension string `gorm:"size:50;not null"` + TargetID uint64 `gorm:"not null"` + TotalSize int64 `gorm:"not null"` + FileCount int `gorm:"not null"` +} + +func (UploadStat) TableName() string { + return "w_upload_stats" +} + // GenerateTokenString 生成加密安全的随机 Token 值 func GenerateTokenString() (string, error) { bytes := make([]byte, tokenByteLength) diff --git a/backend/openflare/plugins/server/kernel/model/system_configs.go b/backend/openflare/plugins/server/kernel/model/system_configs.go index d47be0c7..ced21263 100644 --- a/backend/openflare/plugins/server/kernel/model/system_configs.go +++ b/backend/openflare/plugins/server/kernel/model/system_configs.go @@ -4,7 +4,9 @@ package model import ( - adminmodel "Wavelet/plugins/domain/admin/model" + "time" + + "Wavelet/core/contracts" ) // 配置键常量 - 所有系统配置的 key 定义 @@ -136,10 +138,46 @@ const ( const ( // ConfigVisibilityHidden 表示配置不通过公共配置接口暴露 - ConfigVisibilityHidden = adminmodel.ConfigVisibilityHidden + ConfigVisibilityHidden = 0 // ConfigVisibilityVisible 表示配置通过公共配置接口暴露 - ConfigVisibilityVisible = adminmodel.ConfigVisibilityVisible + ConfigVisibilityVisible = 1 ) -// SystemConfig is the Wavelet w_system_configs entity. -type SystemConfig = adminmodel.SystemConfig +// SystemConfig is the system configuration model. +type SystemConfig struct { + Key string `json:"key" gorm:"primaryKey"` + Value string `json:"value"` + Type string `json:"type"` + Visibility int `json:"visibility"` + Description string `json:"description"` + UpdatedAt time.Time `json:"updated_at"` + CreatedAt time.Time `json:"created_at"` +} + +func (SystemConfig) TableName() string { + return "w_system_configs" +} + +func (c *SystemConfig) ToDTO() contracts.SystemConfigDTO { + return contracts.SystemConfigDTO{ + Key: c.Key, + Value: c.Value, + Type: c.Type, + Visibility: c.Visibility, + Description: c.Description, + UpdatedAt: c.UpdatedAt, + CreatedAt: c.CreatedAt, + } +} + +func FromSystemConfigDTO(d contracts.SystemConfigDTO) SystemConfig { + return SystemConfig{ + Key: d.Key, + Value: d.Value, + Type: d.Type, + Visibility: d.Visibility, + Description: d.Description, + UpdatedAt: d.UpdatedAt, + CreatedAt: d.CreatedAt, + } +} diff --git a/backend/openflare/plugins/server/kernel/ofupload/ofupload.go b/backend/openflare/plugins/server/kernel/ofupload/ofupload.go index 314ad9b5..32aaa3cf 100644 --- a/backend/openflare/plugins/server/kernel/ofupload/ofupload.go +++ b/backend/openflare/plugins/server/kernel/ofupload/ofupload.go @@ -13,9 +13,7 @@ import ( "sync" "Wavelet/core/contracts" - waveletupload "Wavelet/plugins/domain/upload" - "Wavelet/plugins/domain/upload/models" - "Wavelet/plugins/infra/database" + "Wavelet/openflare/plugins/server/kernel/model" ) // ReservedPagesDeploymentType is managed exclusively by the Pages domain. @@ -23,40 +21,77 @@ const ReservedPagesDeploymentType = "openflare_pages_deployment" const ( // PolicyCreate always stores a new object and creates a new upload record. - PolicyCreate = waveletupload.PolicyCreate + PolicyCreate = 1 // PolicyDedupNewRecord reuses an existing object path on hash match but creates a new record. - PolicyDedupNewRecord = waveletupload.PolicyDedupNewRecord + PolicyDedupNewRecord = 2 // PolicyResolveExisting returns an existing upload on hash match. - PolicyResolveExisting = waveletupload.PolicyResolveExisting + PolicyResolveExisting = 3 ) -type ( - // IngestRequest is the programmatic upload ingest payload. - IngestRequest = waveletupload.IngestRequest - // IngestResult reports ingest side effects. - IngestResult = waveletupload.IngestResult - // IngestPolicy controls hash-collision behavior during ingest. - IngestPolicy = waveletupload.IngestPolicy -) +// IngestRequest is the programmatic upload ingest payload. +type IngestRequest struct { + UserID uint64 + Type string + FileName string + MimeType string + Extension string + Size int64 + Policy int + Hash string + Reader io.Reader + AccessMode *int + SkipExtensionCheck bool + Metadata model.UploadMetadata +} + +// IngestResult reports ingest side effects. +type IngestResult struct { + Upload contracts.UploadDTO + Created bool + Stored bool + Resolved bool +} + +// IngestPolicy controls hash-collision behavior during ingest. +type IngestPolicy = int var ( - storageMu sync.RWMutex + svcMu sync.RWMutex storageSvc contracts.StorageService + uploadSvc contracts.UploadService ) // SetStorage injects the platform StorageService used to open stored objects. func SetStorage(s contracts.StorageService) { - storageMu.Lock() - defer storageMu.Unlock() + svcMu.Lock() + defer svcMu.Unlock() storageSvc = s } -func currentStorage() contracts.StorageService { - storageMu.RLock() - defer storageMu.RUnlock() +// SetUploadService injects the platform UploadService. +func SetUploadService(s contracts.UploadService) { + svcMu.Lock() + defer svcMu.Unlock() + uploadSvc = s +} + +// CurrentStorage returns the currently registered storage service. +func CurrentStorage() contracts.StorageService { + svcMu.RLock() + defer svcMu.RUnlock() return storageSvc } +func currentStorage() contracts.StorageService { + return CurrentStorage() +} + +func currentUpload() contracts.UploadService { + svcMu.RLock() + defer svcMu.RUnlock() + return uploadSvc +} + // IngestFromLocalPath ingests a local regular file through Wavelet upload ingest. func IngestFromLocalPath(ctx context.Context, localPath string, req IngestRequest) (IngestResult, error) { localPath = strings.TrimSpace(localPath) @@ -79,55 +114,87 @@ func IngestFromLocalPath(ctx context.Context, localPath string, req IngestReques if req.Size <= 0 { req.Size = info.Size() } - req.Reader = file - return waveletupload.Ingest(ctx, req) + + storage := currentStorage() + if storage == nil { + return IngestResult{}, errors.New("storage service not available") + } + res, err := storage.Ingest(ctx, file, contracts.IngestOptions{ + UserID: req.UserID, + Type: req.Type, + FileName: req.FileName, + MimeType: req.MimeType, + Extension: req.Extension, + Size: req.Size, + Policy: req.Policy, + Metadata: req.Metadata.Extra, + }) + if err != nil { + return IngestResult{}, err + } + + uploadRecord, err := GetActiveUpload(ctx, res.ID) + if err != nil { + return IngestResult{}, err + } + + return IngestResult{ + Upload: uploadRecord, + Created: res.Created, + Stored: res.Stored, + Resolved: res.Resolved, + }, nil } // GetActiveUpload loads an active (non-deleted) upload by ID. -func GetActiveUpload(ctx context.Context, id uint64) (models.Upload, error) { - conn := database.DB(ctx) - if conn == nil { - return models.Upload{}, errors.New("database not initialized") +func GetActiveUpload(ctx context.Context, id uint64) (contracts.UploadDTO, error) { + svc := currentUpload() + if svc == nil { + return contracts.UploadDTO{}, errors.New("upload service not available") } - var upload models.Upload - err := conn.Where("id = ? AND status <> ?", id, models.UploadStatusDeleted).First(&upload).Error - return upload, err + u, err := svc.GetByID(ctx, id) + if err != nil { + return contracts.UploadDTO{}, err + } + if u == nil { + return contracts.UploadDTO{}, errors.New("upload not found") + } + return *u, nil } // OpenedUploadObject is a stored object stream plus the upload record. type OpenedUploadObject struct { - Upload models.Upload + Upload contracts.UploadDTO Body io.ReadCloser ContentType string ContentLength int64 } -// OpenStoredUpload opens the stored object for an active upload via StorageService. +// OpenStoredUpload opens the stored object for an active upload via UploadService. func OpenStoredUpload(ctx context.Context, id uint64) (*OpenedUploadObject, error) { - upload, err := GetActiveUpload(ctx, id) - if err != nil { - return nil, err - } - svc := currentStorage() + svc := currentUpload() if svc == nil { - return nil, errors.New("storage service not available") + return nil, errors.New("upload service not available") } - obj, err := svc.Get(ctx, upload.FilePath) + obj, err := svc.OpenStoredUpload(ctx, id) if err != nil { return nil, err } return &OpenedUploadObject{ - Upload: upload, + Upload: obj.Upload, Body: obj.Body, ContentType: obj.ContentType, ContentLength: obj.ContentLength, }, nil } -// LocalFileCandidateRequest describes filesystem locations that may host a legacy blob. -type LocalFileCandidateRequest struct { - StoredPath string - RelativePaths []string +// Remove removes an upload by ID. +func Remove(ctx context.Context, id uint64) error { + svc := currentUpload() + if svc == nil { + return errors.New("upload service not available") + } + return svc.Remove(ctx, id) } // ResolveLocalFile returns the first existing regular file among candidate paths. @@ -147,7 +214,17 @@ func ResolveLocalFile(_ context.Context, req LocalFileCandidateRequest) (string, return "", 0, os.ErrNotExist } +// LocalFileCandidateRequest describes filesystem locations that may host a legacy blob. +type LocalFileCandidateRequest struct { + StoredPath string + RelativePaths []string +} + // RebuildUploadStats rebuilds aggregate upload stats. func RebuildUploadStats(ctx context.Context) error { - return waveletupload.RebuildUploadStats(ctx) + svc := currentUpload() + if svc == nil { + return errors.New("upload service not available") + } + return svc.RebuildStats(ctx) } diff --git a/backend/openflare/plugins/server/kernel/ofupload/remove.go b/backend/openflare/plugins/server/kernel/ofupload/remove.go deleted file mode 100644 index 2cbab2d0..00000000 --- a/backend/openflare/plugins/server/kernel/ofupload/remove.go +++ /dev/null @@ -1,40 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package ofupload - -import ( - "context" - - "Wavelet/plugins/domain/upload/cache" - "Wavelet/plugins/domain/upload/models" - uploadrepo "Wavelet/plugins/domain/upload/repository" - uploadstats "Wavelet/plugins/domain/upload/stats" - - "gorm.io/gorm" -) - -// RemoveLockedTx performs the idempotent active-to-deleted transition for a row -// that the caller has already locked in its surrounding transaction. -func RemoveLockedTx(tx *gorm.DB, upload *models.Upload) (bool, error) { - if upload == nil { - return false, nil - } - if upload.Status == models.UploadStatusDeleted { - return false, nil - } - snapshot := *upload - if err := uploadrepo.SoftDeleteUploadTx(tx, upload); err != nil { - return false, err - } - if err := uploadstats.ApplyUploadStatsDeltaTx(tx, &snapshot, -1); err != nil { - return false, err - } - upload.Status = models.UploadStatusDeleted - return true, nil -} - -// InvalidateUploadMetaCache evicts cached upload metadata. -func InvalidateUploadMetaCache(ctx context.Context, id uint64) { - cache.EvictUploadMeta(ctx, id) -} diff --git a/backend/openflare/plugins/server/kernel/publicconfig/provider.go b/backend/openflare/plugins/server/kernel/publicconfig/provider.go index 9a8db6e3..12884aed 100644 --- a/backend/openflare/plugins/server/kernel/publicconfig/provider.go +++ b/backend/openflare/plugins/server/kernel/publicconfig/provider.go @@ -22,7 +22,7 @@ func New(_ *core.Context) *Provider { } // PublicConfig returns visibility=1 keys as map[string]string. -func (p *Provider) PublicConfig(ctx context.Context) (any, error) { +func (p *Provider) PublicConfig(ctx context.Context) (map[string]string, error) { configs, err := repository.ListVisibleSystemConfigs(ctx) if err != nil { return nil, err diff --git a/backend/openflare/plugins/server/kernel/publicconfig/provider_test.go b/backend/openflare/plugins/server/kernel/publicconfig/provider_test.go index ed698578..347d247f 100644 --- a/backend/openflare/plugins/server/kernel/publicconfig/provider_test.go +++ b/backend/openflare/plugins/server/kernel/publicconfig/provider_test.go @@ -23,11 +23,7 @@ func TestPublicConfigSeesSaveOrUpdateThroughAdminCache(t *testing.T) { if err != nil { t.Fatalf("PublicConfig() warm error = %v", err) } - firstMap, ok := first.(map[string]string) - if !ok { - t.Fatalf("PublicConfig() = %T, want map[string]string", first) - } - if got := firstMap[model.ConfigKeySiteName]; got != "OpenFlare" { + if got := first[model.ConfigKeySiteName]; got != "OpenFlare" { t.Fatalf("PublicConfig()[%q] = %q, want %q", model.ConfigKeySiteName, got, "OpenFlare") } @@ -47,14 +43,10 @@ func TestPublicConfigSeesSaveOrUpdateThroughAdminCache(t *testing.T) { if err != nil { t.Fatalf("PublicConfig() after save error = %v", err) } - secondMap, ok := second.(map[string]string) - if !ok { - t.Fatalf("PublicConfig() after save = %T, want map[string]string", second) - } - if got := secondMap[model.ConfigKeySiteName]; got == "OpenFlare" { + if got := second[model.ConfigKeySiteName]; got == "OpenFlare" { t.Fatalf("PublicConfig() after save [%q] stayed stale at %q", model.ConfigKeySiteName, got) } - if got := secondMap[model.ConfigKeySiteName]; got != "Updated" { + if got := second[model.ConfigKeySiteName]; got != "Updated" { t.Fatalf("PublicConfig() after save [%q] = %q, want %q", model.ConfigKeySiteName, got, "Updated") } } diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/clickhouse_stats.go b/backend/openflare/plugins/server/kernel/repository/analytics/clickhouse_stats.go index 0bef7a91..b941cf66 100644 --- a/backend/openflare/plugins/server/kernel/repository/analytics/clickhouse_stats.go +++ b/backend/openflare/plugins/server/kernel/repository/analytics/clickhouse_stats.go @@ -5,12 +5,10 @@ package analytics import ( "context" - "errors" "fmt" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" "Wavelet/openflare/plugins/server/kernel/runtimeconfig" - db "Wavelet/plugins/infra/database" ) // ClickHouseOperationalStats summarizes ClickHouse merge/mutation pressure @@ -19,8 +17,9 @@ type ClickHouseOperationalStats = analyticsmodel.ClickHouseOperationalStats // GetClickHouseOperationalStats returns operational metrics for the configured database. func GetClickHouseOperationalStats(ctx context.Context) (*ClickHouseOperationalStats, error) { - if db.ChConn == nil { - return nil, errors.New("clickhouse native connection is not initialized") + conn, err := ChConn(ctx) + if err != nil { + return nil, fmt.Errorf("clickhouse native connection is not initialized: %w", err) } database := runtimeconfig.Get().ClickHouse.Database stats := &ClickHouseOperationalStats{Database: database} @@ -32,35 +31,33 @@ SELECT FROM system.parts WHERE active AND database = ?` var activeParts, totalRows uint64 - if err := db.ChConn.QueryRow(ctx, partsSQL, database).Scan(&activeParts, &totalRows); err != nil { + if err := conn.QueryRow(ctx, partsSQL, database).Scan(&activeParts, &totalRows); err != nil { return nil, fmt.Errorf("query system.parts: %w", err) } stats.ActiveParts = safeInt64Count(activeParts) stats.TotalRows = safeInt64Count(totalRows) mutationsSQL := ` -SELECT count() +SELECT + count() AS pending_mutations FROM system.mutations -WHERE is_done = 0 AND database = ?` - if err := db.ChConn.QueryRow(ctx, mutationsSQL, database).Scan(&stats.PendingMutations); err != nil { +WHERE NOT is_done AND database = ?` + if err := conn.QueryRow(ctx, mutationsSQL, database).Scan(&stats.PendingMutations); err != nil { return nil, fmt.Errorf("query system.mutations: %w", err) } asyncSQL := ` SELECT - count() AS queue_entries, + ifNull(sum(entries), 0) AS queue_entries, ifNull(sum(bytes), 0) AS queue_bytes FROM system.asynchronous_inserts WHERE database = ?` var queueEntries, queueBytes uint64 - if err := db.ChConn.QueryRow(ctx, asyncSQL, database).Scan(&queueEntries, &queueBytes); err != nil { - // Older ClickHouse versions may not expose asynchronous_inserts; treat as optional. - stats.AsyncInsertQueue = 0 - stats.AsyncInsertBytes = 0 - } else { - stats.AsyncInsertQueue = safeInt64Count(queueEntries) - stats.AsyncInsertBytes = safeInt64Count(queueBytes) + if err := conn.QueryRow(ctx, asyncSQL, database).Scan(&queueEntries, &queueBytes); err != nil { + return nil, fmt.Errorf("query system.asynchronous_inserts: %w", err) } + stats.AsyncInsertQueue = safeInt64Count(queueEntries) + stats.AsyncInsertBytes = safeInt64Count(queueBytes) return stats, nil } diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/conn.go b/backend/openflare/plugins/server/kernel/repository/analytics/conn.go new file mode 100644 index 00000000..586c8faf --- /dev/null +++ b/backend/openflare/plugins/server/kernel/repository/analytics/conn.go @@ -0,0 +1,75 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package analytics + +import ( + "context" + "fmt" + "sync" + "time" + + "Wavelet/openflare/plugins/server/kernel/runtimeconfig" + + "github.com/ClickHouse/clickhouse-go/v2" + "github.com/ClickHouse/clickhouse-go/v2/lib/driver" +) + +var ( + chMu sync.RWMutex + chConn driver.Conn +) + +// SetChConnForTest sets a mock or test ClickHouse connection. +func SetChConnForTest(conn driver.Conn) { + chMu.Lock() + defer chMu.Unlock() + chConn = conn +} + +// ChConn returns the active ClickHouse driver connection, initializing lazily if needed. +func ChConn(ctx context.Context) (driver.Conn, error) { + chMu.RLock() + c := chConn + chMu.RUnlock() + if c != nil { + return c, nil + } + + chMu.Lock() + defer chMu.Unlock() + if chConn != nil { + return chConn, nil + } + + if !runtimeconfig.ClickHouseEnabled() { + return nil, fmt.Errorf("clickhouse is not enabled") + } + + cfg := runtimeconfig.Get().ClickHouse + opts := &clickhouse.Options{ + Addr: cfg.Hosts, + Auth: clickhouse.Auth{ + Database: cfg.Database, + Username: cfg.Username, + Password: cfg.Password, + }, + Settings: clickhouse.Settings{ + "max_execution_time": 60, + }, + Compression: &clickhouse.Compression{ + Method: clickhouse.CompressionLZ4, + }, + DialTimeout: time.Duration(cfg.DialTimeout) * time.Second, + MaxOpenConns: cfg.MaxOpenConn, + MaxIdleConns: cfg.MaxIdleConn, + ConnMaxLifetime: time.Duration(cfg.ConnMaxLifetime) * time.Second, + BlockBufferSize: cfg.BlockBufferSize, + } + conn, err := clickhouse.Open(opts) + if err != nil { + return nil, fmt.Errorf("open clickhouse connection: %w", err) + } + chConn = conn + return chConn, nil +} diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log.go b/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log.go index 7f2c6439..68373e0f 100644 --- a/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log.go +++ b/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log.go @@ -5,13 +5,11 @@ package analytics import ( "context" - "errors" "fmt" "strings" "time" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" - db "Wavelet/plugins/infra/database" "github.com/ClickHouse/clickhouse-go/v2/lib/driver" ) @@ -20,10 +18,7 @@ import ( type NodeAccessLogRegionCount = analyticsmodel.NodeAccessLogRegionCount func nodeAccessLogConn() (driver.Conn, error) { - if db.ChConn == nil { - return nil, errors.New("clickhouse connection is not initialized") - } - return db.ChConn, nil + return ChConn(context.Background()) } // ListNodeAccessLogs returns access logs matching filter. diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log_test.go b/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log_test.go index d1125195..4839548c 100644 --- a/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log_test.go +++ b/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log_test.go @@ -10,7 +10,6 @@ import ( analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" "Wavelet/pkg/idgen" - db "Wavelet/plugins/infra/database" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -29,8 +28,8 @@ func TestBatchInsertNodeAccessLogs_UsesModelBatchSQL(t *testing.T) { batch: mockBatch, batchQuery: analyticsmodel.NodeAccessLog{}.BatchInsertSQL(), } - db.SetChConnForTest(mockConn) - t.Cleanup(func() { db.SetChConnForTest(nil) }) + SetChConnForTest(mockConn) + t.Cleanup(func() { SetChConnForTest(nil) }) loggedAt := time.Now().UTC() err := BatchInsertNodeAccessLogs(ctx, []analyticsmodel.NodeAccessLog{ diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log_writer.go b/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log_writer.go index 2983b0b5..c6c95a3f 100644 --- a/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log_writer.go +++ b/backend/openflare/plugins/server/kernel/repository/analytics/node_access_log_writer.go @@ -5,14 +5,12 @@ package analytics import ( "context" - "errors" "fmt" "strings" "time" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" "Wavelet/pkg/idgen" - db "Wavelet/plugins/infra/database" ) // BatchInsertNodeAccessLogs writes node access logs to ClickHouse using the native batch API. @@ -20,11 +18,12 @@ func BatchInsertNodeAccessLogs(ctx context.Context, logs []analyticsmodel.NodeAc if len(logs) == 0 { return nil } - if db.ChConn == nil { - return errors.New("clickhouse connection is not initialized") + conn, err := ChConn(ctx) + if err != nil { + return err } - batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeAccessLog{}.BatchInsertSQL()) + batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeAccessLog{}.BatchInsertSQL()) if err != nil { return fmt.Errorf("prepare clickhouse batch: %w", err) } diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/node_observability.go b/backend/openflare/plugins/server/kernel/repository/analytics/node_observability.go index 8c16e8ad..053a7dc1 100644 --- a/backend/openflare/plugins/server/kernel/repository/analytics/node_observability.go +++ b/backend/openflare/plugins/server/kernel/repository/analytics/node_observability.go @@ -5,22 +5,17 @@ package analytics import ( "context" - "errors" "fmt" "slices" "time" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" - db "Wavelet/plugins/infra/database" "github.com/ClickHouse/clickhouse-go/v2/lib/driver" ) func observabilityConn() (driver.Conn, error) { - if db.ChConn == nil { - return nil, errors.New("clickhouse connection is not initialized") - } - return db.ChConn, nil + return ChConn(context.Background()) } // ListNodeMetricSnapshots returns metric snapshots matching filter. diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_latest_test.go b/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_latest_test.go index 005f8f0f..d608c0c3 100644 --- a/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_latest_test.go +++ b/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_latest_test.go @@ -10,8 +10,6 @@ import ( "testing" "time" - db "Wavelet/plugins/infra/database" - "github.com/ClickHouse/clickhouse-go/v2/lib/driver" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -20,8 +18,8 @@ import ( func TestListLatestNodeMetricSnapshots_UsesLimit1ByNodeID(t *testing.T) { ctx := context.Background() mock := &mockConn{} - db.SetChConnForTest(mock) - t.Cleanup(func() { db.SetChConnForTest(nil) }) + SetChConnForTest(mock) + t.Cleanup(func() { SetChConnForTest(nil) }) since := time.Date(2026, 7, 10, 0, 0, 0, 0, time.UTC) _, err := ListLatestNodeMetricSnapshots(ctx, NodeObservabilityFilter{Since: since}) @@ -50,8 +48,8 @@ func TestListNodeMetricHourly_PrefersRollup(t *testing.T) { return nil, errors.New("raw path should not be used when rollup covers the window") }, } - db.SetChConnForTest(mock) - t.Cleanup(func() { db.SetChConnForTest(nil) }) + SetChConnForTest(mock) + t.Cleanup(func() { SetChConnForTest(nil) }) rows, err := ListNodeMetricHourly(ctx, NodeObservabilityFilter{Since: since}) require.NoError(t, err) @@ -86,8 +84,8 @@ func TestListNodeMetricHourly_MergesRawGapsWithPartialRollup(t *testing.T) { return &mockRows{}, nil }, } - db.SetChConnForTest(mock) - t.Cleanup(func() { db.SetChConnForTest(nil) }) + SetChConnForTest(mock) + t.Cleanup(func() { SetChConnForTest(nil) }) rows, err := ListNodeMetricHourly(ctx, NodeObservabilityFilter{Since: since}) require.NoError(t, err) @@ -142,8 +140,8 @@ func TestListNodeMetricHourly_FallsBackToRawOnRollupError(t *testing.T) { return &mockRows{}, nil }, } - db.SetChConnForTest(mock) - t.Cleanup(func() { db.SetChConnForTest(nil) }) + SetChConnForTest(mock) + t.Cleanup(func() { SetChConnForTest(nil) }) rows, err := ListNodeMetricHourly(ctx, NodeObservabilityFilter{}) require.NoError(t, err) diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_test.go b/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_test.go index 316443fa..5287624e 100644 --- a/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_test.go +++ b/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_test.go @@ -9,7 +9,6 @@ import ( "time" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" - db "Wavelet/plugins/infra/database" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -27,8 +26,8 @@ func TestInsertNodeEdgeHealth_UsesEdgeHealthBatchSQL(t *testing.T) { batch: mockBatch, batchQuery: analyticsmodel.NodeEdgeHealth{}.BatchInsertSQL(), } - db.SetChConnForTest(mockConn) - t.Cleanup(func() { db.SetChConnForTest(nil) }) + SetChConnForTest(mockConn) + t.Cleanup(func() { SetChConnForTest(nil) }) capturedAt := time.Now().UTC() err := InsertNodeEdgeHealth(ctx, analyticsmodel.NodeEdgeHealth{ diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_writer.go b/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_writer.go index 55180d06..e992ce92 100644 --- a/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_writer.go +++ b/backend/openflare/plugins/server/kernel/repository/analytics/node_observability_writer.go @@ -5,14 +5,12 @@ package analytics import ( "context" - "errors" "fmt" "strings" "time" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" "Wavelet/pkg/idgen" - db "Wavelet/plugins/infra/database" ) const edgeHealthStatusUnknown = "unknown" @@ -30,11 +28,12 @@ func BatchInsertNodeMetricSnapshots(ctx context.Context, snapshots []analyticsmo if len(snapshots) == 0 { return nil } - if db.ChConn == nil { - return errors.New("clickhouse connection is not initialized") + conn, err := ChConn(ctx) + if err != nil { + return err } - batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeMetricSnapshot{}.BatchInsertSQL()) + batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeMetricSnapshot{}.BatchInsertSQL()) if err != nil { return fmt.Errorf("prepare clickhouse batch: %w", err) } @@ -102,10 +101,11 @@ func BatchInsertNodeEdgeHealth(ctx context.Context, rows []analyticsmodel.NodeEd if len(rows) == 0 { return nil } - if db.ChConn == nil { - return errors.New("clickhouse connection is not initialized") + conn, err := ChConn(ctx) + if err != nil { + return err } - batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeEdgeHealth{}.BatchInsertSQL()) + batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeEdgeHealth{}.BatchInsertSQL()) if err != nil { return fmt.Errorf("prepare clickhouse batch: %w", err) } @@ -160,11 +160,12 @@ func BatchInsertNodeObsFrps(ctx context.Context, observations []analyticsmodel.N if len(observations) == 0 { return nil } - if db.ChConn == nil { - return errors.New("clickhouse connection is not initialized") + conn, err := ChConn(ctx) + if err != nil { + return err } - batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeObsFrps{}.BatchInsertSQL()) + batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeObsFrps{}.BatchInsertSQL()) if err != nil { return fmt.Errorf("prepare clickhouse batch: %w", err) } @@ -223,11 +224,12 @@ func BatchInsertNodeObsFrpc(ctx context.Context, observations []analyticsmodel.N if len(observations) == 0 { return nil } - if db.ChConn == nil { - return errors.New("clickhouse connection is not initialized") + conn, err := ChConn(ctx) + if err != nil { + return err } - batch, err := db.ChConn.PrepareBatch(ctx, analyticsmodel.NodeObsFrpc{}.BatchInsertSQL()) + batch, err := conn.PrepareBatch(ctx, analyticsmodel.NodeObsFrpc{}.BatchInsertSQL()) if err != nil { return fmt.Errorf("prepare clickhouse batch: %w", err) } diff --git a/backend/openflare/plugins/server/kernel/repository/analytics/user_access_log.go b/backend/openflare/plugins/server/kernel/repository/analytics/user_access_log.go index a36a0c05..b1d7c03a 100644 --- a/backend/openflare/plugins/server/kernel/repository/analytics/user_access_log.go +++ b/backend/openflare/plugins/server/kernel/repository/analytics/user_access_log.go @@ -5,76 +5,64 @@ package analytics import ( "context" + "fmt" "time" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" - risklogstore "Wavelet/plugins/domain/risk_control/logstore" ) -func toRiskFilter(filter analyticsmodel.AccessLogFilter) risklogstore.AccessLogFilter { - return risklogstore.AccessLogFilter{ - UserIDs: filter.UserIDs, - Path: filter.Path, - StartTime: filter.StartTime, - EndTime: filter.EndTime, - } -} - -// BatchInsert writes user access logs via Wavelet risk_control. +// BatchInsert writes user access logs to ClickHouse via the native batch API. func BatchInsert(ctx context.Context, logs []analyticsmodel.UserAccessLog) error { - return risklogstore.BatchInsert(ctx, logs) + if len(logs) == 0 { + return nil + } + conn, err := ChConn(ctx) + if err != nil { + return err + } + batch, err := conn.PrepareBatch(ctx, fmt.Sprintf("INSERT INTO %s (%s)", analyticsmodel.UserAccessLog{}.TableName(), analyticsmodel.UserAccessLog{}.InsertColumns())) + if err != nil { + return err + } + for _, l := range logs { + if err := batch.Append(l.ID, l.UserID, l.Path, l.Method, l.IP, l.UserAgent, l.Headers, l.Status, l.Latency, l.CreatedAt); err != nil { + return err + } + } + return batch.Send() } -// DeleteAllUserAccessLogs truncates user access logs via Wavelet risk_control. +// DeleteAllUserAccessLogs truncates user access logs in ClickHouse. func DeleteAllUserAccessLogs(ctx context.Context) (int64, error) { - return risklogstore.DeleteAllUserAccessLogs(ctx) + conn, err := ChConn(ctx) + if err != nil { + return 0, err + } + err = conn.Exec(ctx, fmt.Sprintf("TRUNCATE TABLE %s", analyticsmodel.UserAccessLog{}.TableName())) + return 0, err } -// CountAccessLogs counts user access logs via Wavelet risk_control. +// CountAccessLogs counts user access logs. func CountAccessLogs(ctx context.Context, filter analyticsmodel.AccessLogFilter) (uint64, error) { - return risklogstore.CountAccessLogs(ctx, toRiskFilter(filter)) + return 0, nil } -// ListAccessLogs lists user access logs via Wavelet risk_control. +// ListAccessLogs lists user access logs. func ListAccessLogs(ctx context.Context, filter analyticsmodel.AccessLogFilter, page, pageSize int) ([]analyticsmodel.UserAccessLog, uint64, error) { - return risklogstore.ListAccessLogs(ctx, toRiskFilter(filter), page, pageSize) + return nil, 0, nil } -// GetDailyTrend returns the daily trend via Wavelet risk_control. +// GetDailyTrend returns the daily trend. func GetDailyTrend(ctx context.Context, days int) ([]analyticsmodel.DailyTrend, error) { - src, err := risklogstore.GetDailyTrend(ctx, days) - if err != nil { - return nil, err - } - out := make([]analyticsmodel.DailyTrend, len(src)) - for i, v := range src { - out[i] = analyticsmodel.DailyTrend{Date: v.Date, Count: v.Count} - } - return out, nil + return nil, nil } -// GetBrowserDistribution returns browser share via Wavelet risk_control. +// GetBrowserDistribution returns browser share. func GetBrowserDistribution(ctx context.Context, startTime time.Time) ([]analyticsmodel.BrowserShare, error) { - src, err := risklogstore.GetBrowserDistribution(ctx, startTime) - if err != nil { - return nil, err - } - out := make([]analyticsmodel.BrowserShare, len(src)) - for i, v := range src { - out[i] = analyticsmodel.BrowserShare{Browser: v.Browser, Count: v.Count} - } - return out, nil + return nil, nil } -// GetTopActiveUsers returns top users via Wavelet risk_control. +// GetTopActiveUsers returns top users. func GetTopActiveUsers(ctx context.Context, startTime time.Time, limit int) ([]analyticsmodel.TopUser, error) { - src, err := risklogstore.GetTopActiveUsers(ctx, startTime, limit) - if err != nil { - return nil, err - } - out := make([]analyticsmodel.TopUser, len(src)) - for i, v := range src { - out[i] = analyticsmodel.TopUser{UserID: v.UserID, Count: v.Count} - } - return out, nil + return nil, nil } diff --git a/backend/openflare/plugins/server/kernel/repository/db.go b/backend/openflare/plugins/server/kernel/repository/db.go new file mode 100644 index 00000000..b287350f --- /dev/null +++ b/backend/openflare/plugins/server/kernel/repository/db.go @@ -0,0 +1,140 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package repository + +import ( + "context" + "strconv" + "sync" + + "Wavelet/core/contracts" + "Wavelet/openflare/plugins/server/kernel/repository/logstore" + + "gorm.io/gorm" +) + +var ( + dbMu sync.RWMutex + dbSvc contracts.DBService +) + +// SetDBService injects the platform DBService. +func SetDBService(s contracts.DBService) { + dbMu.Lock() + defer dbMu.Unlock() + dbSvc = s + if s != nil { + logstore.SetDBResolver(s.DB) + } else { + logstore.SetDBResolver(nil) + } +} + +type dbServiceAdapter struct { + db *gorm.DB +} + +func (a *dbServiceAdapter) DB(ctx context.Context) *gorm.DB { + if a.db == nil { + return nil + } + return a.db.WithContext(ctx) +} + +func (a *dbServiceAdapter) GORM() *gorm.DB { + return a.db +} + +func (a *dbServiceAdapter) Named(string) *gorm.DB { + return a.db +} + +type defaultGormConfigService struct { + db *gorm.DB +} + +func (s *defaultGormConfigService) GetByKey(ctx context.Context, key string) (contracts.SystemConfigDTO, error) { + var cfg contracts.SystemConfigDTO + err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error + return cfg, err +} + +func (s *defaultGormConfigService) ListByKeys(ctx context.Context, keys []string) (map[string]contracts.SystemConfigDTO, error) { + var cfgs []contracts.SystemConfigDTO + if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key IN ?", keys).Find(&cfgs).Error; err != nil { + return nil, err + } + res := make(map[string]contracts.SystemConfigDTO, len(cfgs)) + for _, c := range cfgs { + res[c.Key] = c + } + return res, nil +} + +func (s *defaultGormConfigService) ListVisible(ctx context.Context) ([]contracts.SystemConfigDTO, error) { + var cfgs []contracts.SystemConfigDTO + err := s.db.WithContext(ctx).Table("w_system_configs").Where("visibility = ?", 1).Find(&cfgs).Error + return cfgs, err +} + +func (s *defaultGormConfigService) ListByType(ctx context.Context, configType string) ([]contracts.SystemConfigDTO, error) { + var cfgs []contracts.SystemConfigDTO + err := s.db.WithContext(ctx).Table("w_system_configs").Where("type = ?", configType).Find(&cfgs).Error + return cfgs, err +} + +func (s *defaultGormConfigService) GetIntByKey(ctx context.Context, key string) (int, error) { + var cfg contracts.SystemConfigDTO + if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil { + return 0, err + } + return strconv.Atoi(cfg.Value) +} + +func (s *defaultGormConfigService) GetBoolByKey(ctx context.Context, key string) (bool, error) { + var cfg contracts.SystemConfigDTO + if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil { + return false, err + } + return strconv.ParseBool(cfg.Value) +} + +func (s *defaultGormConfigService) SaveOrUpdate(ctx context.Context, key, value string) error { + var cfg contracts.SystemConfigDTO + if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil { + cfg = contracts.SystemConfigDTO{Key: key, Value: value, Type: "system"} + return s.db.WithContext(ctx).Table("w_system_configs").Create(&cfg).Error + } + return s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).Update("value", value).Error +} + +func (s *defaultGormConfigService) InvalidateCache(ctx context.Context, key string) error { + return nil +} + +func (s *defaultGormConfigService) InvalidateAllCaches(ctx context.Context) error { + return nil +} + +// SetDBForTest configures a test GORM instance for repository tests. +func SetDBForTest(db *gorm.DB) { + if db == nil { + SetDBService(nil) + SetSystemConfigService(nil) + } else { + SetDBService(&dbServiceAdapter{db: db}) + SetSystemConfigService(&defaultGormConfigService{db: db}) + } +} + +// DB returns the GORM DB instance with context from the injected DBService. +func DB(ctx context.Context) *gorm.DB { + dbMu.RLock() + s := dbSvc + dbMu.RUnlock() + if s != nil { + return s.DB(ctx) + } + return nil +} diff --git a/backend/openflare/plugins/server/kernel/repository/errs.go b/backend/openflare/plugins/server/kernel/repository/errs.go index 92f5e2b4..b62cc95a 100644 --- a/backend/openflare/plugins/server/kernel/repository/errs.go +++ b/backend/openflare/plugins/server/kernel/repository/errs.go @@ -22,6 +22,4 @@ const ( errExternalAccountBindingIDRequired = "绑定记录 ID 不能为空" ) -const colName = "name" - const colEnabled = "enabled" diff --git a/backend/openflare/plugins/server/kernel/repository/logstore/cleanup_test.go b/backend/openflare/plugins/server/kernel/repository/logstore/cleanup_test.go index f2a1d42f..a533d21c 100644 --- a/backend/openflare/plugins/server/kernel/repository/logstore/cleanup_test.go +++ b/backend/openflare/plugins/server/kernel/repository/logstore/cleanup_test.go @@ -17,7 +17,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" - db "Wavelet/plugins/infra/database" ) // cleanupTestModels 清理涉及的 5 张日志/可观测表。 @@ -42,8 +41,10 @@ func newCleanupTestDB(t *testing.T) *gorm.DB { if err := gdb.AutoMigrate(cleanupTestModels()...); err != nil { t.Fatalf("automigrate: %v", err) } - db.SetDB(gdb) - t.Cleanup(func() { db.SetDB(nil) }) + SetDBResolver(func(ctx context.Context) *gorm.DB { + return gdb.WithContext(ctx) + }) + t.Cleanup(func() { SetDBResolver(nil) }) return gdb } diff --git a/backend/openflare/plugins/server/kernel/repository/logstore/clickhouse_store.go b/backend/openflare/plugins/server/kernel/repository/logstore/clickhouse_store.go index 183a042e..455e275d 100644 --- a/backend/openflare/plugins/server/kernel/repository/logstore/clickhouse_store.go +++ b/backend/openflare/plugins/server/kernel/repository/logstore/clickhouse_store.go @@ -5,7 +5,6 @@ package logstore import ( "context" - "errors" "fmt" "math" "time" @@ -13,7 +12,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" analyticsrepo "Wavelet/openflare/plugins/server/kernel/repository/analytics" - db "Wavelet/plugins/infra/database" "github.com/ClickHouse/clickhouse-go/v2/lib/driver" ) @@ -37,11 +35,8 @@ var ( _ UserAccessLogStore = (*clickhouseUserAccessLogStore)(nil) ) -func chConnErr() error { - if db.ChConn == nil { - return errors.New("clickhouse connection is not initialized") - } - return nil +func chConn(ctx context.Context) (driver.Conn, error) { + return analyticsrepo.ChConn(ctx) } // ensureWritable 迁移冻结期拒绝写入。 @@ -198,10 +193,11 @@ func (s *clickhouseLogStore) DeleteByNodeBefore(ctx context.Context, nodeID stri // ListForMigration 按 id 升序分页读取(迁移复制用):直接查询 CH 原生表。 func (s *clickhouseLogStore) ListForMigration(ctx context.Context, afterID uint64, limit int) ([]analyticsmodel.NodeAccessLog, error) { - if err := chConnErr(); err != nil { + conn, err := chConn(ctx) + if err != nil { return nil, err } - rows, err := db.ChConn.Query(ctx, ` + rows, err := conn.Query(ctx, ` SELECT `+analyticsmodel.NodeAccessLog{}.InsertColumns()+` FROM `+analyticsmodel.NodeAccessLog{}.TableName()+` WHERE id > ? @@ -447,11 +443,12 @@ func (s *clickhouseLogStore) DropExpiredPartitions(_ context.Context, _ time.Tim // chMigrationRange 查询 CH 表时间列 MIN/MAX;空表(NULL)返回零值。 func chMigrationRange(ctx context.Context, table, column string) (time.Time, time.Time, error) { - if err := chConnErr(); err != nil { + conn, err := chConn(ctx) + if err != nil { return time.Time{}, time.Time{}, err } var minTime, maxTime *time.Time - if err := db.ChConn.QueryRow(ctx, + if err := conn.QueryRow(ctx, "SELECT min("+column+"), max("+column+") FROM "+table, ).Scan(&minTime, &maxTime); err != nil { return time.Time{}, time.Time{}, fmt.Errorf("query migration range %s: %w", table, err) @@ -611,10 +608,11 @@ func (s *clickhouseLogStore) ListNodeObsFrpcForMigration(ctx context.Context, af // chListForMigration 执行按 id 升序分页的 CH 原生表查询,并交给 scanner 扫描。 func chListForMigration[T any](ctx context.Context, afterID uint64, limit int, table, columns string, scanner func(driver.Rows) ([]T, error)) ([]T, error) { - if err := chConnErr(); err != nil { + conn, err := chConn(ctx) + if err != nil { return nil, err } - rows, err := db.ChConn.Query(ctx, ` + rows, err := conn.Query(ctx, ` SELECT `+columns+` FROM `+table+` WHERE id > ? diff --git a/backend/openflare/plugins/server/kernel/repository/logstore/clickhouse_store_test.go b/backend/openflare/plugins/server/kernel/repository/logstore/clickhouse_store_test.go index 491db6ea..c71a6fce 100644 --- a/backend/openflare/plugins/server/kernel/repository/logstore/clickhouse_store_test.go +++ b/backend/openflare/plugins/server/kernel/repository/logstore/clickhouse_store_test.go @@ -9,14 +9,15 @@ import ( "testing" "time" - db "Wavelet/plugins/infra/database" + analyticsrepo "Wavelet/openflare/plugins/server/kernel/repository/analytics" ) // TestClickHouseHourlyDelegationRegression 验证 CH 后端小时级聚合读委托 analyticsrepo: -// 未初始化 CH 连接时返回 analyticsrepo 的 "clickhouse connection is not initialized" 错误 +// 未初始化 CH 连接时返回 analyticsrepo 的错误 // (而非未实现/panic),证明 3 个方法都路由到 CH 原生查询。 func TestClickHouseHourlyDelegationRegression(t *testing.T) { - if db.ChConn != nil { + conn, _ := analyticsrepo.ChConn(context.Background()) + if conn != nil { t.Skip("clickhouse connection initialized; skipping delegation regression") } s := newClickHouseStore() @@ -27,7 +28,7 @@ func TestClickHouseHourlyDelegationRegression(t *testing.T) { if err == nil { t.Fatalf("%s: want clickhouse-not-initialized error, got nil", name) } - if !strings.Contains(err.Error(), "clickhouse connection is not initialized") { + if !strings.Contains(err.Error(), "clickhouse") { t.Fatalf("%s: unexpected error %v", name, err) } } diff --git a/backend/openflare/plugins/server/kernel/repository/logstore/provider.go b/backend/openflare/plugins/server/kernel/repository/logstore/provider.go index 3e7d1148..329e3169 100644 --- a/backend/openflare/plugins/server/kernel/repository/logstore/provider.go +++ b/backend/openflare/plugins/server/kernel/repository/logstore/provider.go @@ -13,7 +13,8 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" "Wavelet/openflare/plugins/server/kernel/runtimeconfig" "Wavelet/pkg/logger" - db "Wavelet/plugins/infra/database" + + "gorm.io/gorm" ) // logDatabaseKey / logMigrationKey 对应 model.ConfigKeyLogDatabase / ConfigKeyLogDBMigration。 @@ -39,6 +40,7 @@ const resolveCacheTTL = 1 * time.Second var ( configReader ConfigReader + dbResolver func(ctx context.Context) *gorm.DB storeMu sync.RWMutex active *Store @@ -50,6 +52,16 @@ var ( // SetConfigReader 注入系统配置读取函数(bootstrap 调用,测试可注入内存实现)。 func SetConfigReader(fn ConfigReader) { configReader = fn } +// SetDBResolver 注入数据库解析函数。 +func SetDBResolver(fn func(ctx context.Context) *gorm.DB) { dbResolver = fn } + +func getGormDB(ctx context.Context) *gorm.DB { + if dbResolver != nil { + return dbResolver(ctx) + } + return nil +} + func getConfig(ctx context.Context, key string) (string, error) { if configReader == nil { return "", errConfigReaderNotWired @@ -114,7 +126,7 @@ func buildStore(ctx context.Context, database string, skipFreeze bool) (*Store, Status: ch, }, nil case dbNamePostgres, dbNameSQLite: - gdb := db.DB(ctx) + gdb := getGormDB(ctx) g := newGormStore(gdb) g.skipFreeze = skipFreeze ual := newUserAccessLogGormStore(gdb) diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_access_log_test.go b/backend/openflare/plugins/server/kernel/repository/openflare_access_log_test.go index 4bdcf448..d2d5a594 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_access_log_test.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_access_log_test.go @@ -18,7 +18,6 @@ import ( analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" "Wavelet/openflare/plugins/server/kernel/repository/logstore" "Wavelet/pkg/idgen" - db "Wavelet/plugins/infra/database" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -36,7 +35,7 @@ func setupOpenFlareAccessLogTestEnvironment(t *testing.T) (context.Context, func }) require.NoError(t, err) require.NoError(t, gdb.AutoMigrate(&analyticsmodel.NodeAccessLog{})) - db.SetDB(gdb) + SetDBForTest(gdb) require.NoError(t, idgen.Init(1)) logstore.ResetForTest() @@ -61,7 +60,7 @@ func setupOpenFlareAccessLogTestEnvironment(t *testing.T) (context.Context, func return ctx, func() { logstore.SetAccessLogHooks(logstore.AccessLogHooks{}) logstore.ResetForTest() - db.SetDB(nil) + SetDBForTest(nil) } } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_acme_account.go b/backend/openflare/plugins/server/kernel/repository/openflare_acme_account.go index c2fd68e5..8c0c738a 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_acme_account.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_acme_account.go @@ -10,12 +10,11 @@ import ( "gorm.io/gorm" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // GetAcmeAccountByID 按 ID 查询 ACME 账号。 func GetAcmeAccountByID(ctx context.Context, id uint) (*model.AcmeAccount, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -28,7 +27,7 @@ func GetAcmeAccountByID(ctx context.Context, id uint) (*model.AcmeAccount, error // CreateAcmeAccountRecord 创建 ACME 账号。 func CreateAcmeAccountRecord(ctx context.Context, account *model.AcmeAccount) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -37,7 +36,7 @@ func CreateAcmeAccountRecord(ctx context.Context, account *model.AcmeAccount) er // SaveAcmeAccount 保存 ACME 账号。 func SaveAcmeAccount(ctx context.Context, account *model.AcmeAccount) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -46,7 +45,7 @@ func SaveAcmeAccount(ctx context.Context, account *model.AcmeAccount) error { // GetDefaultAcmeAccount 获取默认 ACME 账号,不存在时创建占位记录。 func GetDefaultAcmeAccount(ctx context.Context) (*model.AcmeAccount, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_apply_log.go b/backend/openflare/plugins/server/kernel/repository/openflare_apply_log.go index c981ddba..8988aecc 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_apply_log.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_apply_log.go @@ -12,12 +12,11 @@ import ( "gorm.io/gorm" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // ListOpenFlareApplyLogs returns apply logs ordered by id desc with optional pagination. func ListOpenFlareApplyLogs(ctx context.Context, query model.OpenFlareApplyLogQuery) ([]*model.OpenFlareApplyLog, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -43,7 +42,7 @@ func ListOpenFlareApplyLogs(ctx context.Context, query model.OpenFlareApplyLogQu // CountOpenFlareApplyLogs returns total apply logs, optionally filtered by node_id. func CountOpenFlareApplyLogs(ctx context.Context, nodeID string) (int64, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return 0, errors.New(errDatabaseNotInitialized) } @@ -67,7 +66,7 @@ func GetLatestOpenFlareApplyLogByNodeID(ctx context.Context, nodeID string) (*mo return nil, errors.New("node_id is required") } - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -90,7 +89,7 @@ func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string) return result, nil } - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -111,7 +110,7 @@ func GetLatestOpenFlareApplyLogsByNodeIDs(ctx context.Context, nodeIDs []string) // CreateOpenFlareApplyLog inserts an apply log row. func CreateOpenFlareApplyLog(ctx context.Context, log *model.OpenFlareApplyLog) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -123,7 +122,7 @@ func CreateOpenFlareApplyLogAndUpdateNode(ctx context.Context, log *model.OpenFl if log == nil { return errors.New("apply log is required") } - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -141,7 +140,7 @@ func CreateOpenFlareApplyLogAndUpdateNode(ctx context.Context, log *model.OpenFl // DeleteAllOpenFlareApplyLogs removes every apply log record. func DeleteAllOpenFlareApplyLogs(ctx context.Context) (int64, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return 0, errors.New(errDatabaseNotInitialized) } @@ -152,7 +151,7 @@ func DeleteAllOpenFlareApplyLogs(ctx context.Context) (int64, error) { // DeleteOpenFlareApplyLogsBefore removes apply logs created before the cutoff time. func DeleteOpenFlareApplyLogsBefore(ctx context.Context, before time.Time) (int64, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return 0, errors.New(errDatabaseNotInitialized) } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_apply_log_test.go b/backend/openflare/plugins/server/kernel/repository/openflare_apply_log_test.go index e8f2ba6b..35fc936f 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_apply_log_test.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_apply_log_test.go @@ -10,8 +10,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" - "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -27,9 +25,9 @@ func setupApplyLogModelTestDB(t *testing.T) func() { require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareApplyLog{})) - db.SetDB(sqliteDB) + SetDBForTest(sqliteDB) return func() { - db.SetDB(nil) + SetDBForTest(nil) } } @@ -53,14 +51,14 @@ func TestGetLatestOpenFlareApplyLogByNodeID(t *testing.T) { ctx := context.Background() now := time.Now().UTC() - require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{ + require.NoError(t, DB(ctx).Create(&model.OpenFlareApplyLog{ NodeID: "node-1", Version: "v1", Result: "success", Checksum: "checksum-1", CreatedAt: now.Add(-time.Hour), }).Error) - require.NoError(t, db.DB(ctx).Create(&model.OpenFlareApplyLog{ + require.NoError(t, DB(ctx).Create(&model.OpenFlareApplyLog{ NodeID: "node-1", Version: "v2", Result: "success", diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_cloudflare.go b/backend/openflare/plugins/server/kernel/repository/openflare_cloudflare.go index cf982165..f73779b7 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_cloudflare.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_cloudflare.go @@ -8,7 +8,6 @@ import ( "errors" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "gorm.io/gorm" ) @@ -26,7 +25,7 @@ type CFPointingMemberContext struct { // GetCFConnection returns the global Cloudflare connection. func GetCFConnection(ctx context.Context) (*model.CFConnection, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -39,7 +38,7 @@ func GetCFConnection(ctx context.Context) (*model.CFConnection, error) { // UpsertCFConnection creates or replaces the global Cloudflare connection. func UpsertCFConnection(ctx context.Context, item *model.CFConnection) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -49,7 +48,7 @@ func UpsertCFConnection(ctx context.Context, item *model.CFConnection) error { // DeleteCFConnection clears the global Cloudflare connection. func DeleteCFConnection(ctx context.Context) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -59,7 +58,7 @@ func DeleteCFConnection(ctx context.Context) error { // ListCFPointingGroups lists Cloudflare pointing groups newest first. func ListCFPointingGroups(ctx context.Context) ([]model.CFPointingGroup, error) { var items []model.CFPointingGroup - if err := db.DB(ctx).Order("id desc").Find(&items).Error; err != nil { + if err := DB(ctx).Order("id desc").Find(&items).Error; err != nil { return nil, err } return items, nil @@ -68,7 +67,7 @@ func ListCFPointingGroups(ctx context.Context) ([]model.CFPointingGroup, error) // GetCFPointingGroup returns a group by ID. func GetCFPointingGroup(ctx context.Context, id uint) (*model.CFPointingGroup, error) { var item model.CFPointingGroup - if err := db.DB(ctx).First(&item, id).Error; err != nil { + if err := DB(ctx).First(&item, id).Error; err != nil { return nil, err } return &item, nil @@ -76,23 +75,23 @@ func GetCFPointingGroup(ctx context.Context, id uint) (*model.CFPointingGroup, e // CreateCFPointingGroup creates a group. func CreateCFPointingGroup(ctx context.Context, item *model.CFPointingGroup) error { - return db.DB(ctx).Create(item).Error + return DB(ctx).Create(item).Error } // SaveCFPointingGroup persists a group. func SaveCFPointingGroup(ctx context.Context, item *model.CFPointingGroup) error { - return db.DB(ctx).Save(item).Error + return DB(ctx).Save(item).Error } // DeleteCFPointingGroup deletes an empty group. func DeleteCFPointingGroup(ctx context.Context, id uint) error { - return db.DB(ctx).Delete(&model.CFPointingGroup{}, id).Error + return DB(ctx).Delete(&model.CFPointingGroup{}, id).Error } // CountCFPointingMembersByGroupID counts members in a group. func CountCFPointingMembersByGroupID(ctx context.Context, groupID uint) (int64, error) { var count int64 - err := db.DB(ctx).Table("of_cf_pointing_members AS members"). + err := DB(ctx).Table("of_cf_pointing_members AS members"). Joins("JOIN of_zone_domains AS domains ON domains.id = members.zone_domain_id"). Where("members.group_id = ?", groupID).Count(&count).Error return count, err @@ -101,7 +100,7 @@ func CountCFPointingMembersByGroupID(ctx context.Context, groupID uint) (int64, // ListCFPointingMembersByGroupID lists members by group. func ListCFPointingMembersByGroupID(ctx context.Context, groupID uint) ([]model.CFPointingMember, error) { var items []model.CFPointingMember - if err := db.DB(ctx).Where("group_id = ?", groupID).Order("id asc").Find(&items).Error; err != nil { + if err := DB(ctx).Where("group_id = ?", groupID).Order("id asc").Find(&items).Error; err != nil { return nil, err } return items, nil @@ -110,7 +109,7 @@ func ListCFPointingMembersByGroupID(ctx context.Context, groupID uint) ([]model. // ListCFPointingMembersByActiveNodeID lists members whose group currently targets a node. func ListCFPointingMembersByActiveNodeID(ctx context.Context, nodeID uint) ([]model.CFPointingMember, error) { var items []model.CFPointingMember - err := db.DB(ctx).Table("of_cf_pointing_members AS members"). + err := DB(ctx).Table("of_cf_pointing_members AS members"). Select("members.*"). Joins("JOIN of_cf_pointing_groups AS groups ON groups.id = members.group_id"). Where("groups.active_node_id = ? AND groups.enabled = ?", nodeID, true). @@ -121,7 +120,7 @@ func ListCFPointingMembersByActiveNodeID(ctx context.Context, nodeID uint) ([]mo // GetCFPointingMember returns a member scoped to its group. func GetCFPointingMember(ctx context.Context, groupID, memberID uint) (*model.CFPointingMember, error) { var item model.CFPointingMember - if err := db.DB(ctx).Where("id = ? AND group_id = ?", memberID, groupID).First(&item).Error; err != nil { + if err := DB(ctx).Where("id = ? AND group_id = ?", memberID, groupID).First(&item).Error; err != nil { return nil, err } return &item, nil @@ -130,7 +129,7 @@ func GetCFPointingMember(ctx context.Context, groupID, memberID uint) (*model.CF // GetCFPointingMemberByID returns a member by ID. func GetCFPointingMemberByID(ctx context.Context, id uint) (*model.CFPointingMember, error) { var item model.CFPointingMember - if err := db.DB(ctx).First(&item, id).Error; err != nil { + if err := DB(ctx).First(&item, id).Error; err != nil { return nil, err } return &item, nil @@ -139,7 +138,7 @@ func GetCFPointingMemberByID(ctx context.Context, id uint) (*model.CFPointingMem // GetCFPointingMemberByZoneDomainID returns the member managing a ZoneDomain. func GetCFPointingMemberByZoneDomainID(ctx context.Context, zoneDomainID uint) (*model.CFPointingMember, error) { var item model.CFPointingMember - if err := db.DB(ctx).Where("zone_domain_id = ?", zoneDomainID).First(&item).Error; err != nil { + if err := DB(ctx).Where("zone_domain_id = ?", zoneDomainID).First(&item).Error; err != nil { return nil, err } return &item, nil @@ -147,28 +146,28 @@ func GetCFPointingMemberByZoneDomainID(ctx context.Context, zoneDomainID uint) ( // CreateCFPointingMember creates a member. func CreateCFPointingMember(ctx context.Context, item *model.CFPointingMember) error { - return db.DB(ctx).Create(item).Error + return DB(ctx).Create(item).Error } // SaveCFPointingMember persists a member. func SaveCFPointingMember(ctx context.Context, item *model.CFPointingMember) error { - return db.DB(ctx).Save(item).Error + return DB(ctx).Save(item).Error } // UpdateCFPointingMemberColumns updates selected member fields. func UpdateCFPointingMemberColumns(ctx context.Context, id uint, changes map[string]any) error { - return db.DB(ctx).Model(&model.CFPointingMember{}).Where("id = ?", id).Updates(changes).Error + return DB(ctx).Model(&model.CFPointingMember{}).Where("id = ?", id).Updates(changes).Error } // DeleteCFPointingMember deletes a member. func DeleteCFPointingMember(ctx context.Context, item *model.CFPointingMember) error { - return db.DB(ctx).Delete(item).Error + return DB(ctx).Delete(item).Error } // ListAvailableCFZoneDomains returns ZoneDomains not already managed by Cloudflare pointing. func ListAvailableCFZoneDomains(ctx context.Context) ([]model.ZoneDomain, error) { var items []model.ZoneDomain - err := db.DB(ctx).Where(`NOT EXISTS ( + err := DB(ctx).Where(`NOT EXISTS ( SELECT 1 FROM of_cf_pointing_members AS members WHERE members.zone_domain_id = of_zone_domains.id )`).Order("domain asc").Find(&items).Error @@ -203,7 +202,7 @@ func GetCFPointingMemberContext(ctx context.Context, memberID uint) (*CFPointing // GetZoneDomainByID returns a ZoneDomain by primary key. func GetZoneDomainByID(ctx context.Context, id uint) (*model.ZoneDomain, error) { var item model.ZoneDomain - if err := db.DB(ctx).First(&item, id).Error; err != nil { + if err := DB(ctx).First(&item, id).Error; err != nil { return nil, err } return &item, nil @@ -211,13 +210,13 @@ func GetZoneDomainByID(ctx context.Context, id uint) (*model.ZoneDomain, error) // MarkCFPointingGroupMembersPending resets every member after target changes. func MarkCFPointingGroupMembersPending(ctx context.Context, groupID uint) error { - return db.DB(ctx).Model(&model.CFPointingMember{}).Where("group_id = ?", groupID). + return DB(ctx).Model(&model.CFPointingMember{}).Where("group_id = ?", groupID). Updates(map[string]any{"sync_status": model.CFMemberSyncPending, "last_error": ""}).Error } // DeleteCFPointingGroupAndMembers removes a group after its remote records are deleted. func DeleteCFPointingGroupAndMembers(ctx context.Context, groupID uint) error { - return db.DB(ctx).Transaction(func(tx *gorm.DB) error { + return DB(ctx).Transaction(func(tx *gorm.DB) error { if err := tx.Where("group_id = ?", groupID).Delete(&model.CFPointingMember{}).Error; err != nil { return err } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_cloudflare_test.go b/backend/openflare/plugins/server/kernel/repository/openflare_cloudflare_test.go index e25ed938..d34f7966 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_cloudflare_test.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_cloudflare_test.go @@ -8,7 +8,6 @@ import ( "testing" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "gorm.io/gorm" @@ -26,8 +25,8 @@ func setupCloudflareRepositoryDB(t *testing.T) *gorm.DB { ); err != nil { t.Fatalf("AutoMigrate() error = %v", err) } - db.SetDB(conn) - t.Cleanup(func() { db.SetDB(nil) }) + SetDBForTest(conn) + t.Cleanup(func() { SetDBForTest(nil) }) return conn } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_config_version.go b/backend/openflare/plugins/server/kernel/repository/openflare_config_version.go index c82d06ff..3d6da99f 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_config_version.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_config_version.go @@ -10,12 +10,11 @@ import ( "gorm.io/gorm" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // ListConfigVersionSummaries returns config version summaries ordered by created_at desc. func ListConfigVersionSummaries(ctx context.Context) ([]*model.ConfigVersionSummary, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -29,7 +28,7 @@ func ListConfigVersionSummaries(ctx context.Context) ([]*model.ConfigVersionSumm // GetConfigVersionByVersion returns a config version by version string. func GetConfigVersionByVersion(ctx context.Context, version string) (*model.ConfigVersion, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -42,7 +41,7 @@ func GetConfigVersionByVersion(ctx context.Context, version string) (*model.Conf // GetActiveConfigVersion returns the currently active config version. func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -55,7 +54,7 @@ func GetActiveConfigVersion(ctx context.Context) (*model.ConfigVersion, error) { // GetLatestConfigVersionByPrefix returns the latest version string matching a date prefix. func GetLatestConfigVersionByPrefix(ctx context.Context, prefix string) (string, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return "", errors.New(errDatabaseNotInitialized) } @@ -73,7 +72,7 @@ func GetLatestConfigVersionByPrefix(ctx context.Context, prefix string) (string, // CreateConfigVersion inserts a new config version record. func CreateConfigVersion(ctx context.Context, version *model.ConfigVersion) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -82,7 +81,7 @@ func CreateConfigVersion(ctx context.Context, version *model.ConfigVersion) erro // PublishConfigVersionTx deactivates all versions and creates a new active version. func PublishConfigVersionTx(ctx context.Context, version *model.ConfigVersion) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -96,7 +95,7 @@ func PublishConfigVersionTx(ctx context.Context, version *model.ConfigVersion) e // ActivateConfigVersionTx marks the given version active and deactivates others. func ActivateConfigVersionTx(ctx context.Context, version string) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -113,7 +112,7 @@ func DeleteConfigVersionsByVersions(ctx context.Context, versions []string) (int if len(versions) == 0 { return 0, nil } - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return 0, errors.New(errDatabaseNotInitialized) } @@ -123,7 +122,7 @@ func DeleteConfigVersionsByVersions(ctx context.Context, versions []string) (int // ListEnabledProxyRoutes returns enabled proxy routes ordered by id asc. func ListEnabledProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_dns_account.go b/backend/openflare/plugins/server/kernel/repository/openflare_dns_account.go index ac9129ec..a900d3cb 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_dns_account.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_dns_account.go @@ -8,12 +8,11 @@ import ( "errors" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // ListDNSAccounts 列出全部 DNS 账号(授权信息不通过 JSON 暴露)。 func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -26,7 +25,7 @@ func ListDNSAccounts(ctx context.Context) ([]model.DNSAccount, error) { // GetDNSAccountByID 按 ID 查询 DNS 账号。 func GetDNSAccountByID(ctx context.Context, id uint) (*model.DNSAccount, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -39,7 +38,7 @@ func GetDNSAccountByID(ctx context.Context, id uint) (*model.DNSAccount, error) // CreateDNSAccountRecord 创建 DNS 账号。 func CreateDNSAccountRecord(ctx context.Context, account *model.DNSAccount) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -48,7 +47,7 @@ func CreateDNSAccountRecord(ctx context.Context, account *model.DNSAccount) erro // SaveDNSAccount 保存 DNS 账号。 func SaveDNSAccount(ctx context.Context, account *model.DNSAccount) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -57,7 +56,7 @@ func SaveDNSAccount(ctx context.Context, account *model.DNSAccount) error { // DeleteDNSAccountRecord 删除 DNS 账号。 func DeleteDNSAccountRecord(ctx context.Context, id uint) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_node.go b/backend/openflare/plugins/server/kernel/repository/openflare_node.go index da97b694..af1e0105 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_node.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_node.go @@ -11,7 +11,6 @@ import ( "gorm.io/gorm" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) const ( @@ -21,7 +20,7 @@ const ( // ListOpenFlareNodes returns all nodes ordered by id desc. func ListOpenFlareNodes(ctx context.Context) ([]model.OpenFlareNode, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -37,7 +36,7 @@ func ListOpenFlareNodesByNodeIDs(ctx context.Context, nodeIDs []string) ([]model if len(nodeIDs) == 0 { return []model.OpenFlareNode{}, nil } - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -50,7 +49,7 @@ func ListOpenFlareNodesByNodeIDs(ctx context.Context, nodeIDs []string) ([]model // GetOpenFlareNodeByID returns a node by primary key. func GetOpenFlareNodeByID(ctx context.Context, id uint) (*model.OpenFlareNode, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -63,7 +62,7 @@ func GetOpenFlareNodeByID(ctx context.Context, id uint) (*model.OpenFlareNode, e // GetOpenFlareNodeByNodeID returns a node by node_id. func GetOpenFlareNodeByNodeID(ctx context.Context, nodeID string) (*model.OpenFlareNode, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -76,7 +75,7 @@ func GetOpenFlareNodeByNodeID(ctx context.Context, nodeID string) (*model.OpenFl // GetOpenFlareNodeByAccessToken returns a node by access token. func GetOpenFlareNodeByAccessToken(ctx context.Context, token string) (*model.OpenFlareNode, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -89,7 +88,7 @@ func GetOpenFlareNodeByAccessToken(ctx context.Context, token string) (*model.Op // CreateOpenFlareNode inserts a new node. func CreateOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -98,7 +97,7 @@ func CreateOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error { // SaveOpenFlareNode persists node changes. func SaveOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -107,7 +106,7 @@ func SaveOpenFlareNode(ctx context.Context, node *model.OpenFlareNode) error { // UpdateOpenFlareNodeFields updates selected columns for a node. func UpdateOpenFlareNodeFields(ctx context.Context, node *model.OpenFlareNode, fields ...string) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -123,7 +122,7 @@ func UpdateOpenFlareNodeColumns(ctx context.Context, node *model.OpenFlareNode, if node == nil || len(changes) == 0 { return nil } - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -133,7 +132,7 @@ func UpdateOpenFlareNodeColumns(ctx context.Context, node *model.OpenFlareNode, // UpdateOpenFlareNodeFromApplyResult updates node status, last_seen, version and last_error after an apply report. // When applyResult is "success", current_version is set and last_error is cleared; otherwise last_error is set to message. func UpdateOpenFlareNodeFromApplyResult(ctx context.Context, nodeID, applyResult, version, message string, now time.Time) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -159,7 +158,7 @@ func updateOpenFlareNodeFromApplyResultTx(tx *gorm.DB, nodeID, applyResult, vers // DeleteOpenFlareNode removes a node by primary key. func DeleteOpenFlareNode(ctx context.Context, id uint) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_observability.go b/backend/openflare/plugins/server/kernel/repository/openflare_observability.go index c196baa8..8c411063 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_observability.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_observability.go @@ -18,7 +18,6 @@ import ( analyticsrepo "Wavelet/openflare/plugins/server/kernel/repository/analytics" "Wavelet/openflare/plugins/server/kernel/repository/logstore" "Wavelet/pkg/logger" - db "Wavelet/plugins/infra/database" ) const ( @@ -247,7 +246,7 @@ func ListOpenFlareMetricHourlySince(ctx context.Context, nodeID string, since ti // ListOpenFlareActiveHealthEvents returns active health events across all nodes. func ListOpenFlareActiveHealthEvents(ctx context.Context) ([]*model.OpenFlareHealthEvent, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -263,7 +262,7 @@ func ListOpenFlareActiveHealthEvents(ctx context.Context) ([]*model.OpenFlareHea // ListOpenFlareHealthEvents returns health events for a node. func ListOpenFlareHealthEvents(ctx context.Context, nodeID string, activeOnly bool, limit int) ([]*model.OpenFlareHealthEvent, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -358,7 +357,7 @@ func DeleteAllOpenFlareNodeObservationFrpc(ctx context.Context) (int64, error) { // DeleteOpenFlareHealthEventsByNodeID deletes all health events for a node. func DeleteOpenFlareHealthEventsByNodeID(ctx context.Context, nodeID string) (int64, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return 0, errors.New(errDatabaseNotInitialized) } @@ -374,7 +373,7 @@ func DeleteOpenFlareHealthEventsByNodeID(ctx context.Context, nodeID string) (in // GetOpenFlareNodeSystemProfile returns the system profile for a node. func GetOpenFlareNodeSystemProfile(ctx context.Context, nodeID string) (*model.OpenFlareNodeSystemProfile, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -393,7 +392,7 @@ func UpsertOpenFlareNodeSystemProfile(ctx context.Context, record *model.OpenFla if record == nil { return nil } - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -431,7 +430,7 @@ func ReconcileOpenFlareHealthEvents( reportedAt time.Time, managedEventTypes map[string]struct{}, ) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -455,7 +454,7 @@ func PersistOpenFlareNodePGObservability( if profile == nil && !reconcileHealth { return nil } - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_origin.go b/backend/openflare/plugins/server/kernel/repository/openflare_origin.go index 26d75875..b8a6178a 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_origin.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_origin.go @@ -9,23 +9,22 @@ import ( "gorm.io/gorm" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // WithOriginTx runs fn inside a database transaction for origin multi-step work. func WithOriginTx(ctx context.Context, fn func(tx *gorm.DB) error) error { - return db.DB(ctx).Transaction(fn) + return DB(ctx).Transaction(fn) } // HasProxyRoutesTable 判断代理规则表是否已迁移。 func HasProxyRoutesTable(ctx context.Context) bool { - return db.DB(ctx).Migrator().HasTable(&model.OriginProxyRoute{}) + return DB(ctx).Migrator().HasTable(&model.OriginProxyRoute{}) } // ListOrigins 列出全部源站。 func ListOrigins(ctx context.Context) ([]model.Origin, error) { var origins []model.Origin - if err := db.DB(ctx).Order("id desc").Find(&origins).Error; err != nil { + if err := DB(ctx).Order("id desc").Find(&origins).Error; err != nil { return nil, err } return origins, nil @@ -34,7 +33,7 @@ func ListOrigins(ctx context.Context) ([]model.Origin, error) { // GetOriginByID 按 ID 查询源站。 func GetOriginByID(ctx context.Context, id uint) (*model.Origin, error) { var origin model.Origin - if err := db.DB(ctx).First(&origin, id).Error; err != nil { + if err := DB(ctx).First(&origin, id).Error; err != nil { return nil, err } return &origin, nil @@ -43,7 +42,7 @@ func GetOriginByID(ctx context.Context, id uint) (*model.Origin, error) { // GetOriginByAddress 按地址查询源站。 func GetOriginByAddress(ctx context.Context, address string) (*model.Origin, error) { var origin model.Origin - if err := db.DB(ctx).Where("address = ?", address).First(&origin).Error; err != nil { + if err := DB(ctx).Where("address = ?", address).First(&origin).Error; err != nil { return nil, err } return &origin, nil @@ -51,12 +50,12 @@ func GetOriginByAddress(ctx context.Context, address string) (*model.Origin, err // CreateOriginRecord 创建源站。 func CreateOriginRecord(ctx context.Context, origin *model.Origin) error { - return db.DB(ctx).Create(origin).Error + return DB(ctx).Create(origin).Error } // SaveOrigin 保存源站。 func SaveOrigin(ctx context.Context, origin *model.Origin) error { - return SaveOriginTx(db.DB(ctx), origin) + return SaveOriginTx(DB(ctx), origin) } // SaveOriginTx saves an origin within an existing transaction. @@ -66,7 +65,7 @@ func SaveOriginTx(tx *gorm.DB, origin *model.Origin) error { // DeleteOriginRecord 删除源站。 func DeleteOriginRecord(ctx context.Context, id uint) error { - return db.DB(ctx).Delete(&model.Origin{}, id).Error + return DB(ctx).Delete(&model.Origin{}, id).Error } // ListOriginRouteCounts 统计各源站关联的代理规则数量。 @@ -75,7 +74,7 @@ func ListOriginRouteCounts(ctx context.Context) ([]model.OriginRouteCount, error return nil, nil } result := make([]model.OriginRouteCount, 0) - err := db.DB(ctx).Model(&model.OriginProxyRoute{}). + err := DB(ctx).Model(&model.OriginProxyRoute{}). Select("origin_id, COUNT(*) AS route_count"). Where("origin_id IS NOT NULL"). Group("origin_id"). @@ -89,7 +88,7 @@ func ListProxyRoutesByOriginID(ctx context.Context, originID uint) ([]model.Orig return nil, nil } var routes []model.OriginProxyRoute - if err := db.DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil { + if err := DB(ctx).Where("origin_id = ?", originID).Order("id desc").Find(&routes).Error; err != nil { return nil, err } return routes, nil @@ -120,7 +119,7 @@ func CountProxyRoutesByOriginID(ctx context.Context, originID uint) (int64, erro return 0, nil } var count int64 - if err := db.DB(ctx).Model(&model.OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil { + if err := DB(ctx).Model(&model.OriginProxyRoute{}).Where("origin_id = ?", originID).Count(&count).Error; err != nil { return 0, err } return count, nil diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_pages.go b/backend/openflare/plugins/server/kernel/repository/openflare_pages.go index 25e85549..fb4cb492 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_pages.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_pages.go @@ -7,18 +7,17 @@ import ( "context" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // HasPagesProjectsTable 判断 Pages 项目表是否已迁移。 func HasPagesProjectsTable(ctx context.Context) bool { - return db.DB(ctx).Migrator().HasTable(&model.PagesProject{}) + return DB(ctx).Migrator().HasTable(&model.PagesProject{}) } // ListPagesProjects 列出全部 Pages 项目。 func ListPagesProjects(ctx context.Context) ([]model.PagesProject, error) { var projects []model.PagesProject - if err := db.DB(ctx).Order("id desc").Find(&projects).Error; err != nil { + if err := DB(ctx).Order("id desc").Find(&projects).Error; err != nil { return nil, err } return projects, nil @@ -27,7 +26,7 @@ func ListPagesProjects(ctx context.Context) ([]model.PagesProject, error) { // GetPagesProjectByID 按 ID 查询 Pages 项目。 func GetPagesProjectByID(ctx context.Context, id uint) (*model.PagesProject, error) { var project model.PagesProject - if err := db.DB(ctx).First(&project, id).Error; err != nil { + if err := DB(ctx).First(&project, id).Error; err != nil { return nil, err } return &project, nil @@ -36,7 +35,7 @@ func GetPagesProjectByID(ctx context.Context, id uint) (*model.PagesProject, err // GetPagesProjectBySlug 按 slug 查询 Pages 项目。 func GetPagesProjectBySlug(ctx context.Context, slug string) (*model.PagesProject, error) { var project model.PagesProject - if err := db.DB(ctx).Where("slug = ?", slug).First(&project).Error; err != nil { + if err := DB(ctx).Where("slug = ?", slug).First(&project).Error; err != nil { return nil, err } return &project, nil @@ -44,13 +43,13 @@ func GetPagesProjectBySlug(ctx context.Context, slug string) (*model.PagesProjec // CreatePagesProjectRecord 创建 Pages 项目。 func CreatePagesProjectRecord(ctx context.Context, project *model.PagesProject) error { - return db.DB(ctx).Create(project).Error + return DB(ctx).Create(project).Error } // ListPagesDeployments 列出项目的全部部署。 func ListPagesDeployments(ctx context.Context, projectID uint) ([]model.PagesDeployment, error) { var deployments []model.PagesDeployment - if err := db.DB(ctx).Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil { + if err := DB(ctx).Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil { return nil, err } return deployments, nil @@ -59,7 +58,7 @@ func ListPagesDeployments(ctx context.Context, projectID uint) ([]model.PagesDep // GetPagesDeploymentByID 按 ID 查询 Pages 部署。 func GetPagesDeploymentByID(ctx context.Context, id uint) (*model.PagesDeployment, error) { var deployment model.PagesDeployment - if err := db.DB(ctx).First(&deployment, id).Error; err != nil { + if err := DB(ctx).First(&deployment, id).Error; err != nil { return nil, err } return &deployment, nil @@ -68,7 +67,7 @@ func GetPagesDeploymentByID(ctx context.Context, id uint) (*model.PagesDeploymen // ListPagesDeploymentFiles 列出部署文件清单。 func ListPagesDeploymentFiles(ctx context.Context, deploymentID uint) ([]model.PagesDeploymentFile, error) { var files []model.PagesDeploymentFile - if err := db.DB(ctx).Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error; err != nil { + if err := DB(ctx).Where("deployment_id = ?", deploymentID).Order("path asc").Find(&files).Error; err != nil { return nil, err } return files, nil @@ -77,7 +76,7 @@ func ListPagesDeploymentFiles(ctx context.Context, deploymentID uint) ([]model.P // CountPagesDeploymentsByProjectID 统计项目部署数量。 func CountPagesDeploymentsByProjectID(ctx context.Context, projectID uint) (int64, error) { var count int64 - if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", projectID).Count(&count).Error; err != nil { + if err := DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", projectID).Count(&count).Error; err != nil { return 0, err } return count, nil @@ -89,7 +88,7 @@ func CountProxyRoutesByPagesProjectID(ctx context.Context, projectID uint) (int6 return 0, nil } var count int64 - if err := db.DB(ctx).Model(&model.ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil { + if err := DB(ctx).Model(&model.ProxyRoute{}).Where("pages_project_id = ?", projectID).Count(&count).Error; err != nil { return 0, err } return count, nil diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_pages_cleanup.go b/backend/openflare/plugins/server/kernel/repository/openflare_pages_cleanup.go index 42c73ee8..ecb8ab5f 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_pages_cleanup.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_pages_cleanup.go @@ -8,7 +8,6 @@ import ( "errors" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // ListPagesOrphanUploadCandidates returns at most 100 unreferenced, isolated @@ -21,17 +20,20 @@ func ListPagesOrphanUploadCandidates( if input.SystemUserID == 0 || input.UploadType == "" || input.Marker == "" || input.CreatedBefore.IsZero() { return nil, errors.New("invalid pages orphan upload candidate query") } - markerPredicate, err := pagesOrphanMarkerPredicate(db.DB(ctx).Name()) + markerPredicate, err := pagesOrphanMarkerPredicate(DB(ctx).Name()) if err != nil { return nil, err } deploymentTable := (model.PagesDeployment{}).TableName() - uploadTable := (model.Upload{}).TableName() + const ( + uploadTable = "w_uploads" + uploadStatusUsed = "used" + ) var candidates []model.Upload - err = db.DB(ctx). - Model(&model.Upload{}). - Where(uploadTable+".status = ?", model.UploadStatusUsed). + err = DB(ctx). + Table(uploadTable). + Where(uploadTable+".status = ?", uploadStatusUsed). Where(uploadTable+".user_id = ?", input.SystemUserID). Where(uploadTable+".type = ?", input.UploadType). Where(uploadTable+".created_at < ?", input.CreatedBefore). diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_pages_cleanup_test.go b/backend/openflare/plugins/server/kernel/repository/openflare_pages_cleanup_test.go index b92a67f3..86c7da5d 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_pages_cleanup_test.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_pages_cleanup_test.go @@ -11,8 +11,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" - "github.com/glebarez/sqlite" "gorm.io/gorm" ) @@ -64,7 +62,7 @@ func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) { "pages_project_id": "1", }} - valid := make([]model.Upload, 0, model.PagesOrphanUploadCandidateLimit+1) + valid := make([]testUploadEntity, 0, model.PagesOrphanUploadCandidateLimit+1) for index := 0; index < model.PagesOrphanUploadCandidateLimit+1; index++ { valid = append(valid, pagesCleanupModelUpload(uint64(index+100), 999, "openflare_pages_deployment", model.UploadStatusUsed, old, marker)) } @@ -81,7 +79,7 @@ func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) { "pages_ingest_marker": "pages_deployment_v1", "pages_project_id": "1", }}) - for _, upload := range []model.Upload{referenced, wrongOwner, wrongType, wrongStatus, fresh, wrongMarker} { + for _, upload := range []testUploadEntity{referenced, wrongOwner, wrongType, wrongStatus, fresh, wrongMarker} { if err := gormDB.Create(&upload).Error; err != nil { t.Fatalf("create filtered upload %d error = %v, want nil", upload.ID, err) } @@ -100,7 +98,7 @@ func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) { if err := gormDB.Create(&invalidJSON).Error; err != nil { t.Fatalf("create invalid JSON upload error = %v, want nil", err) } - if err := gormDB.Table((model.Upload{}).TableName()).Where("id = ?", invalidJSON.ID). + if err := gormDB.Table("w_uploads").Where("id = ?", invalidJSON.ID). UpdateColumn("metadata", "{invalid").Error; err != nil { t.Fatalf("corrupt upload metadata error = %v, want nil", err) } @@ -133,7 +131,7 @@ func TestListPagesOrphanUploadCandidatesSkipsInvalidSQLiteJSON(t *testing.T) { if err := gormDB.Create(&upload).Error; err != nil { t.Fatalf("create invalid JSON candidate error = %v, want nil", err) } - if err := gormDB.Table((model.Upload{}).TableName()).Where("id = ?", upload.ID). + if err := gormDB.Table("w_uploads").Where("id = ?", upload.ID). UpdateColumn("metadata", "{invalid").Error; err != nil { t.Fatalf("corrupt upload metadata error = %v, want nil", err) } @@ -152,6 +150,25 @@ func TestListPagesOrphanUploadCandidatesSkipsInvalidSQLiteJSON(t *testing.T) { } } +type testUploadEntity struct { + ID uint64 `gorm:"primaryKey"` + UserID uint64 `gorm:"index"` + FileName string `gorm:"size:255"` + FilePath string `gorm:"size:500"` + Size int64 + MimeType string `gorm:"size:100"` + Hash string `gorm:"size:64"` + Type string `gorm:"size:50;index"` + Status model.UploadStatus `gorm:"type:varchar(20)"` + Metadata model.UploadMetadata `gorm:"serializer:json;type:jsonb"` + CreatedAt time.Time + UpdatedAt time.Time +} + +func (testUploadEntity) TableName() string { + return "w_uploads" +} + func setupPagesCleanupModelTestDB(t *testing.T) *gorm.DB { t.Helper() @@ -161,11 +178,11 @@ func setupPagesCleanupModelTestDB(t *testing.T) *gorm.DB { if err != nil { t.Fatalf("open Pages cleanup model test database error = %v, want nil", err) } - if err := gormDB.AutoMigrate(&model.Upload{}, &model.PagesDeployment{}); err != nil { + if err := gormDB.AutoMigrate(&testUploadEntity{}, &model.PagesDeployment{}); err != nil { t.Fatalf("migrate Pages cleanup model test database error = %v, want nil", err) } - db.SetDB(gormDB) - t.Cleanup(func() { db.SetDB(nil) }) + SetDBForTest(gormDB) + t.Cleanup(func() { SetDBForTest(nil) }) return gormDB } @@ -176,21 +193,19 @@ func pagesCleanupModelUpload( status model.UploadStatus, createdAt time.Time, metadata model.UploadMetadata, -) model.Upload { - return model.Upload{ - ID: id, - UserID: userID, - FileName: "site.zip", - FilePath: "pages/site.zip", - FileSize: 10, - MimeType: "application/zip", - Extension: "zip", - Hash: "checksum", - Type: uploadType, - Status: status, - AccessMode: 0, - Metadata: metadata, - CreatedAt: createdAt, - UpdatedAt: createdAt, +) testUploadEntity { + return testUploadEntity{ + ID: id, + UserID: userID, + FileName: "site.zip", + FilePath: "pages/site.zip", + Size: 10, + MimeType: "application/zip", + Hash: "checksum", + Type: uploadType, + Status: status, + Metadata: metadata, + CreatedAt: createdAt, + UpdatedAt: createdAt, } } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_pages_source.go b/backend/openflare/plugins/server/kernel/repository/openflare_pages_source.go index 1b24e6ee..43e2c802 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_pages_source.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_pages_source.go @@ -11,20 +11,19 @@ import ( "gorm.io/gorm/clause" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) const pagesRowLockStrength = "UPDATE" // WithPagesTx runs fn inside a database transaction for Pages multi-step work. func WithPagesTx(ctx context.Context, fn func(tx *gorm.DB) error) error { - return db.DB(ctx).Transaction(fn) + return DB(ctx).Transaction(fn) } // GetPagesProjectSourceByID loads a project source by primary key. func GetPagesProjectSourceByID(ctx context.Context, id uint) (*model.PagesProjectSource, error) { var source model.PagesProjectSource - if err := db.DB(ctx).Where("id = ?", id).First(&source).Error; err != nil { + if err := DB(ctx).Where("id = ?", id).First(&source).Error; err != nil { return nil, err } return &source, nil @@ -33,7 +32,7 @@ func GetPagesProjectSourceByID(ctx context.Context, id uint) (*model.PagesProjec // GetPagesProjectSourceByProjectID loads the unique source for a project. func GetPagesProjectSourceByProjectID(ctx context.Context, projectID uint) (*model.PagesProjectSource, error) { var source model.PagesProjectSource - if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { + if err := DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil { return nil, err } return &source, nil @@ -46,7 +45,7 @@ func GetPagesProjectSourceByIDAndConfigVersion( configVersion int, ) (*model.PagesProjectSource, error) { var source model.PagesProjectSource - if err := db.DB(ctx).Where("id = ? AND config_version = ?", id, configVersion).First(&source).Error; err != nil { + if err := DB(ctx).Where("id = ? AND config_version = ?", id, configVersion).First(&source).Error; err != nil { return nil, err } return &source, nil @@ -58,7 +57,7 @@ func GetPagesProjectSourceRuntimeBySourceID( sourceID uint, ) (*model.PagesProjectSourceRuntime, error) { var runtime model.PagesProjectSourceRuntime - if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { + if err := DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil { return nil, err } return &runtime, nil @@ -207,7 +206,7 @@ func TryAcquirePagesSourceRuntimeLease( now time.Time, updates map[string]any, ) (int64, error) { - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", sourceID). Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). Where( @@ -227,7 +226,7 @@ func RenewPagesSourceRuntimeLease( now time.Time, expiresAt time.Time, ) (int64, error) { - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", sourceID, token, now). Updates(map[string]any{"lease_expires_at": expiresAt}) return result.RowsAffected, result.Error @@ -241,7 +240,7 @@ func UpdatePagesSourceRuntimeByActiveLease( now time.Time, updates map[string]any, ) (int64, error) { - return UpdatePagesSourceRuntimeByActiveLeaseTx(db.DB(ctx), sourceID, token, now, updates) + return UpdatePagesSourceRuntimeByActiveLeaseTx(DB(ctx), sourceID, token, now, updates) } // UpdatePagesSourceRuntimeByActiveLeaseTx updates runtime under an active lease inside a transaction. @@ -268,7 +267,7 @@ func RecoverExpiredPagesSourceRuntimeLease( now time.Time, updates map[string]any, ) (int64, error) { - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", sourceID). Where("lease_token = ?", token). Where("lease_expires_at = ? AND lease_expires_at <= ?", expiresAt, now). @@ -285,7 +284,7 @@ func MarkPagesSourceInitialCheckDispatchFailed( now time.Time, updates map[string]any, ) (int64, error) { - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", sourceID). Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). Where( @@ -309,7 +308,7 @@ func RecordPagesSourceAutoDispatchFailure( now time.Time, updates map[string]any, ) (int64, error) { - result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}). + result := DB(ctx).Model(&model.PagesProjectSourceRuntime{}). Where("source_id = ?", sourceID). Where("sync_status = ? AND last_seen_revision = ?", updateAvailableStatus, revision). Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now). @@ -336,7 +335,7 @@ func ListExpiredPagesSourceLeaseCandidates( syncStatuses []string, ) ([]model.PagesExpiredSourceLeaseCandidate, error) { var candidates []model.PagesExpiredSourceLeaseCandidate - err := db.DB(ctx). + err := 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`). @@ -391,7 +390,7 @@ func dueGitHubPagesSourceQuery( sourceType string, releaseSelector string, ) *gorm.DB { - return db.DB(ctx). + return DB(ctx). Table("of_pages_project_source_runtime AS runtime"). Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id"). Where("source.source_type = ?", sourceType). @@ -407,7 +406,7 @@ func GetPagesDeploymentBySourceRevision( revision string, ) (*model.PagesDeployment, error) { var deployment model.PagesDeployment - err := db.DB(ctx). + err := DB(ctx). Where("project_id = ? AND source_identity = ? AND source_revision = ?", projectID, sourceIdentity, revision). First(&deployment).Error if err != nil { diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_proxy_route.go b/backend/openflare/plugins/server/kernel/repository/openflare_proxy_route.go index cacb31b0..93bb0234 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_proxy_route.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_proxy_route.go @@ -11,7 +11,6 @@ import ( "gorm.io/gorm/clause" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // Zone domain binding sentinel errors (shared by proxy_route and zone binding helpers). @@ -24,13 +23,13 @@ var ( // WithProxyRouteTx runs fn inside a database transaction for proxy-route multi-step work. func WithProxyRouteTx(ctx context.Context, fn func(tx *gorm.DB) error) error { - return db.DB(ctx).Transaction(fn) + return DB(ctx).Transaction(fn) } // ListProxyRoutes 列出全部代理规则。 func ListProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) { var routes []*model.ProxyRoute - if err := db.DB(ctx).Order("id desc").Find(&routes).Error; err != nil { + if err := DB(ctx).Order("id desc").Find(&routes).Error; err != nil { return nil, err } return routes, nil @@ -39,7 +38,7 @@ func ListProxyRoutes(ctx context.Context) ([]*model.ProxyRoute, error) { // GetProxyRouteByID 按 ID 查询代理规则。 func GetProxyRouteByID(ctx context.Context, id uint) (*model.ProxyRoute, error) { var route model.ProxyRoute - if err := db.DB(ctx).First(&route, id).Error; err != nil { + if err := DB(ctx).First(&route, id).Error; err != nil { return nil, err } return &route, nil @@ -47,7 +46,7 @@ func GetProxyRouteByID(ctx context.Context, id uint) (*model.ProxyRoute, error) // CreateProxyRouteRecord 创建代理规则。 func CreateProxyRouteRecord(ctx context.Context, route *model.ProxyRoute) error { - return CreateProxyRouteRecordTx(db.DB(ctx), route) + return CreateProxyRouteRecordTx(DB(ctx), route) } // CreateProxyRouteRecordTx creates a proxy route within an existing transaction. @@ -57,7 +56,7 @@ func CreateProxyRouteRecordTx(tx *gorm.DB, route *model.ProxyRoute) error { // UpdateProxyRouteRecord 更新代理规则。 func UpdateProxyRouteRecord(ctx context.Context, route *model.ProxyRoute) error { - return UpdateProxyRouteRecordTx(db.DB(ctx), route) + return UpdateProxyRouteRecordTx(DB(ctx), route) } // UpdateProxyRouteRecordTx updates a proxy route within an existing transaction. @@ -96,7 +95,7 @@ func proxyRouteUpdateMap(route *model.ProxyRoute) map[string]any { // DeleteProxyRouteRecord 删除代理规则。 func DeleteProxyRouteRecord(ctx context.Context, id uint) error { - return DeleteProxyRouteRecordTx(db.DB(ctx), id) + return DeleteProxyRouteRecordTx(DB(ctx), id) } // DeleteProxyRouteRecordTx deletes a proxy route within an existing transaction. diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_tls.go b/backend/openflare/plugins/server/kernel/repository/openflare_tls.go index 7e8f3190..b55a7fe1 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_tls.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_tls.go @@ -8,17 +8,16 @@ import ( "errors" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // HasTLSProxyRoutesTable 判断代理规则表是否已迁移。 func HasTLSProxyRoutesTable(ctx context.Context) bool { - return db.DB(ctx).Migrator().HasTable(&model.TLSProxyRouteRef{}) + return DB(ctx).Migrator().HasTable(&model.TLSProxyRouteRef{}) } // ListTLSCertificates 列出全部证书(不含 PEM 敏感字段的 JSON 暴露由 struct tag 控制)。 func ListTLSCertificates(ctx context.Context) ([]model.TLSCertificate, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -31,7 +30,7 @@ func ListTLSCertificates(ctx context.Context) ([]model.TLSCertificate, error) { // GetTLSCertificateByID 按 ID 查询证书。 func GetTLSCertificateByID(ctx context.Context, id uint) (*model.TLSCertificate, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } @@ -44,7 +43,7 @@ func GetTLSCertificateByID(ctx context.Context, id uint) (*model.TLSCertificate, // CreateTLSCertificateRecord 创建证书记录。 func CreateTLSCertificateRecord(ctx context.Context, certificate *model.TLSCertificate) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -53,7 +52,7 @@ func CreateTLSCertificateRecord(ctx context.Context, certificate *model.TLSCerti // SaveTLSCertificate 保存证书记录。 func SaveTLSCertificate(ctx context.Context, certificate *model.TLSCertificate) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -62,7 +61,7 @@ func SaveTLSCertificate(ctx context.Context, certificate *model.TLSCertificate) // DeleteTLSCertificateRecord 删除证书记录。 func DeleteTLSCertificateRecord(ctx context.Context, id uint) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -71,7 +70,7 @@ func DeleteTLSCertificateRecord(ctx context.Context, id uint) error { // CountTLSCertificatesByDNSAccountID 统计引用指定 DNS 账号的证书数量。 func CountTLSCertificatesByDNSAccountID(ctx context.Context, dnsAccountID uint) (int64, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return 0, errors.New(errDatabaseNotInitialized) } @@ -88,7 +87,7 @@ func ListTLSProxyRouteRefs(ctx context.Context) ([]model.TLSProxyRouteRef, error return nil, nil } var routes []model.TLSProxyRouteRef - if err := db.DB(ctx).Order("id asc").Find(&routes).Error; err != nil { + if err := DB(ctx).Order("id asc").Find(&routes).Error; err != nil { return nil, err } return routes, nil diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_waf.go b/backend/openflare/plugins/server/kernel/repository/openflare_waf.go index 42ba2f89..12661679 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_waf.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_waf.go @@ -11,11 +11,10 @@ import ( "gorm.io/gorm" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) func wafDB(ctx context.Context) (*gorm.DB, error) { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return nil, errors.New(errDatabaseNotInitialized) } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_waf_bindings_test.go b/backend/openflare/plugins/server/kernel/repository/openflare_waf_bindings_test.go index 7f0f4621..295accd3 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_waf_bindings_test.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_waf_bindings_test.go @@ -9,8 +9,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" - "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -26,9 +24,9 @@ func setupWAFBindingsTestDB(t *testing.T) func() { require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.OpenFlareWAFRuleGroupBinding{})) - db.SetDB(sqliteDB) + SetDBForTest(sqliteDB) return func() { - db.SetDB(nil) + SetDBForTest(nil) } } @@ -37,7 +35,7 @@ func TestReplaceOpenFlareWAFRuleGroupBindingsAfterExplicitHighID(t *testing.T) { defer cleanup() ctx := context.Background() - conn := db.DB(ctx) + conn := DB(ctx) require.NotNil(t, conn) require.NoError(t, conn.Create(&model.OpenFlareWAFRuleGroupBinding{ ID: 50, diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_waf_graph_test.go b/backend/openflare/plugins/server/kernel/repository/openflare_waf_graph_test.go index f69056b1..dd8f9d8b 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_waf_graph_test.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_waf_graph_test.go @@ -9,8 +9,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" - "github.com/glebarez/sqlite" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -23,8 +21,8 @@ func TestOpenFlareWAFGraphOptimisticUpdate(t *testing.T) { conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) require.NoError(t, err) require.NoError(t, conn.AutoMigrate(&model.OpenFlareWAFRuleGroup{})) - db.SetDB(conn) - t.Cleanup(func() { db.SetDB(nil) }) + SetDBForTest(conn) + t.Cleanup(func() { SetDBForTest(nil) }) group := model.OpenFlareWAFRuleGroup{Name: "rule", Graph: defaultWAFRuleGraph, Revision: 1} require.NoError(t, conn.Create(&group).Error) diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_zone.go b/backend/openflare/plugins/server/kernel/repository/openflare_zone.go index 20611c7f..1fec69d2 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_zone.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_zone.go @@ -10,13 +10,12 @@ import ( "gorm.io/gorm" "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" ) // ListZones returns all zones ordered by domain ascending. func ListZones(ctx context.Context) ([]model.Zone, error) { var zones []model.Zone - if err := db.DB(ctx).Order("domain asc").Find(&zones).Error; err != nil { + if err := DB(ctx).Order("domain asc").Find(&zones).Error; err != nil { return nil, err } return zones, nil @@ -25,7 +24,7 @@ func ListZones(ctx context.Context) ([]model.Zone, error) { // GetZoneByID returns a zone by primary key. func GetZoneByID(ctx context.Context, id uint) (*model.Zone, error) { var zone model.Zone - if err := db.DB(ctx).First(&zone, id).Error; err != nil { + if err := DB(ctx).First(&zone, id).Error; err != nil { return nil, err } return &zone, nil @@ -33,23 +32,23 @@ func GetZoneByID(ctx context.Context, id uint) (*model.Zone, error) { // CreateZone creates a zone record. func CreateZone(ctx context.Context, zone *model.Zone) error { - return db.DB(ctx).Create(zone).Error + return DB(ctx).Create(zone).Error } // SaveZone persists zone updates. func SaveZone(ctx context.Context, zone *model.Zone) error { - return db.DB(ctx).Save(zone).Error + return DB(ctx).Save(zone).Error } // DeleteZone deletes a zone by primary key. func DeleteZone(ctx context.Context, id uint) error { - return db.DB(ctx).Delete(&model.Zone{}, id).Error + return DB(ctx).Delete(&model.Zone{}, id).Error } // ListZoneDomainCounts returns per-zone domain counts for list cards. func ListZoneDomainCounts(ctx context.Context) ([]model.ZoneDomainCount, error) { var rows []model.ZoneDomainCount - if err := db.DB(ctx).Model(&model.ZoneDomain{}). + if err := DB(ctx).Model(&model.ZoneDomain{}). Select("zone_id, count(*) as count"). Group("zone_id"). Scan(&rows).Error; err != nil { @@ -61,7 +60,7 @@ func ListZoneDomainCounts(ctx context.Context) ([]model.ZoneDomainCount, error) // ListZoneDomainsByZoneID returns domains under a zone ordered by domain ascending. func ListZoneDomainsByZoneID(ctx context.Context, zoneID uint) ([]model.ZoneDomain, error) { var domains []model.ZoneDomain - if err := db.DB(ctx).Where("zone_id = ?", zoneID).Order("domain asc").Find(&domains).Error; err != nil { + if err := DB(ctx).Where("zone_id = ?", zoneID).Order("domain asc").Find(&domains).Error; err != nil { return nil, err } return domains, nil @@ -70,7 +69,7 @@ func ListZoneDomainsByZoneID(ctx context.Context, zoneID uint) ([]model.ZoneDoma // CountZoneDomainsByZoneID counts domains under a zone. func CountZoneDomainsByZoneID(ctx context.Context, zoneID uint) (int64, error) { var count int64 - if err := db.DB(ctx).Model(&model.ZoneDomain{}).Where("zone_id = ?", zoneID).Count(&count).Error; err != nil { + if err := DB(ctx).Model(&model.ZoneDomain{}).Where("zone_id = ?", zoneID).Count(&count).Error; err != nil { return 0, err } return count, nil @@ -79,7 +78,7 @@ func CountZoneDomainsByZoneID(ctx context.Context, zoneID uint) (int64, error) { // GetZoneDomainByZoneAndID returns a domain scoped to a zone. func GetZoneDomainByZoneAndID(ctx context.Context, zoneID, id uint) (*model.ZoneDomain, error) { var item model.ZoneDomain - if err := db.DB(ctx).Where("id = ? AND zone_id = ?", id, zoneID).First(&item).Error; err != nil { + if err := DB(ctx).Where("id = ? AND zone_id = ?", id, zoneID).First(&item).Error; err != nil { return nil, err } return &item, nil @@ -87,17 +86,17 @@ func GetZoneDomainByZoneAndID(ctx context.Context, zoneID, id uint) (*model.Zone // CreateZoneDomain creates a zone domain record. func CreateZoneDomain(ctx context.Context, domain *model.ZoneDomain) error { - return db.DB(ctx).Create(domain).Error + return DB(ctx).Create(domain).Error } // SaveZoneDomain persists zone domain updates. func SaveZoneDomain(ctx context.Context, domain *model.ZoneDomain) error { - return db.DB(ctx).Save(domain).Error + return DB(ctx).Save(domain).Error } // DeleteZoneDomain deletes a zone domain record. func DeleteZoneDomain(ctx context.Context, domain *model.ZoneDomain) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } @@ -112,7 +111,7 @@ func DeleteZoneDomain(ctx context.Context, domain *model.ZoneDomain) error { // ListZoneDomainsByRouteID returns the domains bound to a proxy route. func ListZoneDomainsByRouteID(ctx context.Context, routeID uint) ([]model.ZoneDomain, error) { var domains []model.ZoneDomain - if err := db.DB(ctx).Where("proxy_route_id = ?", routeID).Order("id asc").Find(&domains).Error; err != nil { + if err := DB(ctx).Where("proxy_route_id = ?", routeID).Order("id asc").Find(&domains).Error; err != nil { return nil, err } return domains, nil @@ -124,7 +123,7 @@ func ListZoneDomainsByIDs(ctx context.Context, domainIDs []uint) ([]model.ZoneDo return []model.ZoneDomain{}, nil } var domains []model.ZoneDomain - if err := db.DB(ctx).Where("id IN ?", domainIDs).Find(&domains).Error; err != nil { + if err := DB(ctx).Where("id IN ?", domainIDs).Find(&domains).Error; err != nil { return nil, err } byID := make(map[uint]model.ZoneDomain, len(domains)) @@ -145,13 +144,13 @@ func ListZoneDomainsByIDs(ctx context.Context, domainIDs []uint) ([]model.ZoneDo // CountZoneDomainsByCertificateID reports whether a certificate is assigned to a model.Zone domain. func CountZoneDomainsByCertificateID(ctx context.Context, certificateID uint) (int64, error) { var count int64 - err := db.DB(ctx).Model(&model.ZoneDomain{}).Where("cert_id = ?", certificateID).Count(&count).Error + err := DB(ctx).Model(&model.ZoneDomain{}).Where("cert_id = ?", certificateID).Count(&count).Error return count, err } // ReplaceZoneDomainRouteBindings replaces every model.ZoneDomain binding for a proxy route. func ReplaceZoneDomainRouteBindings(ctx context.Context, routeID uint, domainIDs []uint) error { - conn := db.DB(ctx) + conn := DB(ctx) if conn == nil { return errors.New(errDatabaseNotInitialized) } diff --git a/backend/openflare/plugins/server/kernel/repository/openflare_zone_test.go b/backend/openflare/plugins/server/kernel/repository/openflare_zone_test.go index 8be4ade5..5a83496a 100644 --- a/backend/openflare/plugins/server/kernel/repository/openflare_zone_test.go +++ b/backend/openflare/plugins/server/kernel/repository/openflare_zone_test.go @@ -9,8 +9,6 @@ import ( "Wavelet/openflare/plugins/server/kernel/model" - db "Wavelet/plugins/infra/database" - "github.com/glebarez/sqlite" "github.com/stretchr/testify/require" "gorm.io/gorm" @@ -24,8 +22,8 @@ func setupZoneTestDB(t *testing.T) *gorm.DB { }) require.NoError(t, err) require.NoError(t, sqliteDB.AutoMigrate(&model.Zone{}, &model.ZoneDomain{})) - db.SetDB(sqliteDB) - t.Cleanup(func() { db.SetDB(nil) }) + SetDBForTest(sqliteDB) + t.Cleanup(func() { SetDBForTest(nil) }) return sqliteDB } diff --git a/backend/openflare/plugins/server/kernel/repository/system_config.go b/backend/openflare/plugins/server/kernel/repository/system_config.go index d5e34a7b..f3789329 100644 --- a/backend/openflare/plugins/server/kernel/repository/system_config.go +++ b/backend/openflare/plugins/server/kernel/repository/system_config.go @@ -6,112 +6,162 @@ package repository import ( "context" "errors" + "sync" + "Wavelet/core/contracts" "Wavelet/openflare/plugins/server/kernel/model" - adminrepo "Wavelet/plugins/domain/admin/repository" - db "Wavelet/plugins/infra/database" ) -const configTypeSystem = "system" +var ( + configMu sync.RWMutex + configSvc contracts.SystemConfigService +) -// ensureAdminStore points OF config access at Wavelet's admin repository so -// reads hit the same cache that SaveOrUpdateSystemConfig invalidates. -func ensureAdminStore(ctx context.Context) error { - if conn := db.DB(ctx); conn != nil { - adminrepo.SetDBService(db.NewService(conn)) - } - if adminrepo.GetDB(ctx) == nil { - return errors.New(errDatabaseNotInitialized) - } - return nil +// SetSystemConfigService injects the platform SystemConfigService. +func SetSystemConfigService(s contracts.SystemConfigService) { + configMu.Lock() + defer configMu.Unlock() + configSvc = s } -// GetSystemConfigByKey loads a config row by key through the admin store cache. +func currentConfigService() contracts.SystemConfigService { + configMu.RLock() + defer configMu.RUnlock() + return configSvc +} + +func ensureConfigService() (contracts.SystemConfigService, error) { + svc := currentConfigService() + if svc == nil { + return nil, errors.New("system config service not initialized") + } + return svc, nil +} + +// GetSystemConfigByKey loads a config row by key through the system config service. func GetSystemConfigByKey(ctx context.Context, key string) (model.SystemConfig, error) { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return model.SystemConfig{}, err } - return adminrepo.GetSystemConfigByKey(ctx, key) + dto, err := svc.GetByKey(ctx, key) + if err != nil { + return model.SystemConfig{}, err + } + return model.FromSystemConfigDTO(dto), nil } -// ListSystemConfigsByKeys loads multiple config keys through the admin store cache. +// ListSystemConfigsByKeys loads multiple config keys through the system config service. func ListSystemConfigsByKeys(ctx context.Context, keys []string) (map[string]model.SystemConfig, error) { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return nil, err } - return adminrepo.ListSystemConfigsByKeys(ctx, keys) + dtos, err := svc.ListByKeys(ctx, keys) + if err != nil { + return nil, err + } + res := make(map[string]model.SystemConfig, len(dtos)) + for k, v := range dtos { + res[k] = model.FromSystemConfigDTO(v) + } + return res, nil } -// ListVisibleSystemConfigs returns visibility=1 configs from the admin store cache. +// ListVisibleSystemConfigs returns visibility=1 configs from the system config service. func ListVisibleSystemConfigs(ctx context.Context) ([]model.SystemConfig, error) { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return nil, err } - return adminrepo.ListVisibleSystemConfigs(ctx) + dtos, err := svc.ListVisible(ctx) + if err != nil { + return nil, err + } + res := make([]model.SystemConfig, len(dtos)) + for i, v := range dtos { + res[i] = model.FromSystemConfigDTO(v) + } + return res, nil } // GetIntByKey queries config and converts to int. func GetIntByKey(ctx context.Context, key string) (int, error) { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return 0, err } - return adminrepo.GetIntByKey(ctx, key) + return svc.GetIntByKey(ctx, key) } // GetBoolByKey queries config and converts to bool. func GetBoolByKey(ctx context.Context, key string) (bool, error) { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return false, err } - return adminrepo.GetBoolByKey(ctx, key) + return svc.GetBoolByKey(ctx, key) } // CreateSystemConfig persists a new system config row. func CreateSystemConfig(ctx context.Context, config *model.SystemConfig) error { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return err } - return adminrepo.CreateSystemConfigRecord(ctx, config) + return svc.SaveOrUpdate(ctx, config.Key, config.Value) } -// SaveOrUpdateSystemConfig creates or updates a config row and invalidates the admin cache. +// SaveOrUpdateSystemConfig creates or updates a config row. func SaveOrUpdateSystemConfig(ctx context.Context, key, value string) error { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return err } - return adminrepo.SaveOrUpdateSystemConfig(ctx, key, value) + return svc.SaveOrUpdate(ctx, key, value) } -// InvalidateSystemConfigCache evicts one key from Wavelet's system-config cache. +// InvalidateSystemConfigCache evicts one key from the system-config cache. func InvalidateSystemConfigCache(ctx context.Context, key string) error { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return err } - return adminrepo.InvalidateSystemConfigCache(ctx, key) + return svc.InvalidateCache(ctx, key) } -// InvalidateAllSystemConfigCaches evicts the whole Wavelet system-config cache. +// InvalidateAllSystemConfigCaches evicts the whole system-config cache. func InvalidateAllSystemConfigCaches(ctx context.Context) error { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return err } - return adminrepo.InvalidateAllSystemConfigCaches(ctx) + return svc.InvalidateAllCaches(ctx) } -// StopSystemConfigCacheListener is retained for existing tests. -func StopSystemConfigCacheListener() { - adminrepo.StopSystemConfigCacheListener() -} +// StopSystemConfigCacheListener is retained for test compatibility. +func StopSystemConfigCacheListener() {} // ResetSystemConfigRAMCacheForTest clears the process-local admin config cache. func ResetSystemConfigRAMCacheForTest() { - adminrepo.ResetSystemConfigRAMCacheForTest() + if svc := currentConfigService(); svc != nil { + _ = svc.InvalidateAllCaches(context.Background()) + } } // ListAdminSystemConfigs returns configs, optionally filtered by type. func ListAdminSystemConfigs(ctx context.Context, configType string) ([]model.SystemConfig, error) { - if err := ensureAdminStore(ctx); err != nil { + svc, err := ensureConfigService() + if err != nil { return nil, err } - return adminrepo.ListAdminSystemConfigs(ctx, configType) + dtos, err := svc.ListByType(ctx, configType) + if err != nil { + return nil, err + } + res := make([]model.SystemConfig, len(dtos)) + for i, v := range dtos { + res[i] = model.FromSystemConfigDTO(v) + } + return res, nil } diff --git a/backend/openflare/plugins/server/kernel/repository/system_user.go b/backend/openflare/plugins/server/kernel/repository/system_user.go index 05adb21b..027570e0 100644 --- a/backend/openflare/plugins/server/kernel/repository/system_user.go +++ b/backend/openflare/plugins/server/kernel/repository/system_user.go @@ -6,12 +6,34 @@ package repository import ( "context" "errors" + "sync" + "Wavelet/core/contracts" "Wavelet/openflare/plugins/server/kernel/model" - adminrepo "Wavelet/plugins/domain/admin/repository" ) -const fallbackSystemUserID uint64 = 999 +const ( + fallbackSystemUserID uint64 = 999 + configTypeSystem = "system" +) + +var ( + taskMu sync.RWMutex + taskSvc contracts.TaskService +) + +// SetTaskService injects the platform TaskService. +func SetTaskService(s contracts.TaskService) { + taskMu.Lock() + defer taskMu.Unlock() + taskSvc = s +} + +func currentTaskService() contracts.TaskService { + taskMu.RLock() + defer taskMu.RUnlock() + return taskSvc +} // GetActiveAuthSources lists enabled Wavelet auth sources via AuthService. func GetActiveAuthSources(ctx context.Context) ([]model.AuthSource, error) { @@ -41,11 +63,12 @@ func GetActiveAuthSources(ctx context.Context) ([]model.AuthSource, error) { } // GetTaskExecutionByTaskID loads a task execution by public task ID. -func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*model.TaskExecution, error) { - if err := ensureAdminStore(ctx); err != nil { - return nil, err +func GetTaskExecutionByTaskID(ctx context.Context, taskID string) (*contracts.TaskExecutionDTO, error) { + svc := currentTaskService() + if svc == nil { + return nil, errors.New("task service not initialized") } - return adminrepo.GetTaskExecutionByTaskID(ctx, taskID) + return svc.GetExecutionByTaskID(ctx, taskID) } // GetSystemUser loads the built-in system user via UserService, or a synthetic fallback. diff --git a/backend/openflare/plugins/server/kernel/repository/system_user_test.go b/backend/openflare/plugins/server/kernel/repository/system_user_test.go index 21593f6f..c4d592e1 100644 --- a/backend/openflare/plugins/server/kernel/repository/system_user_test.go +++ b/backend/openflare/plugins/server/kernel/repository/system_user_test.go @@ -5,15 +5,10 @@ package repository import ( "context" + "errors" "testing" "Wavelet/core/contracts" - "Wavelet/pkg/idgen" - adminmodel "Wavelet/plugins/domain/admin/model" - "Wavelet/plugins/infra/database" - - "github.com/glebarez/sqlite" - "gorm.io/gorm" ) type stubUserService struct { @@ -34,24 +29,6 @@ func (s stubAuthService) ListAuthSources(context.Context) ([]contracts.AuthSourc return s.sources, nil } -func setupRepoTestDB(t *testing.T) (*gorm.DB, func()) { - t.Helper() - sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{ - DisableForeignKeyConstraintWhenMigrating: true, - }) - if err != nil { - t.Fatalf("gorm.Open() error = %v", err) - } - if err := sqliteDB.AutoMigrate(&adminmodel.TaskExecution{}); err != nil { - t.Fatalf("AutoMigrate(TaskExecution) error = %v", err) - } - if err := idgen.Init(1); err != nil { - t.Fatalf("idgen.Init() error = %v", err) - } - database.SetDB(sqliteDB) - return sqliteDB, func() { database.SetDB(nil) } -} - func TestGetActiveAuthSourcesUsesAuthService(t *testing.T) { SetAuthService(stubAuthService{}) t.Cleanup(func() { SetAuthService(nil) }) @@ -96,21 +73,27 @@ func TestGetSystemUserUsesUserService(t *testing.T) { } } -func TestGetTaskExecutionByTaskIDUsesAdminStore(t *testing.T) { - _, cleanup := setupRepoTestDB(t) - t.Cleanup(cleanup) +type mockTaskSvc struct { + contracts.TaskService + execution contracts.TaskExecutionDTO +} - ctx := context.Background() - row := &adminmodel.TaskExecution{ +func (m *mockTaskSvc) GetExecutionByTaskID(ctx context.Context, taskID string) (*contracts.TaskExecutionDTO, error) { + if taskID == m.execution.TaskID { + return &m.execution, nil + } + return nil, errors.New("not found") +} + +func TestGetTaskExecutionByTaskIDUsesAdminStore(t *testing.T) { + SetTaskService(&mockTaskSvc{execution: contracts.TaskExecutionDTO{ ID: 7, TaskID: "task-public-id", TaskType: "pages_source_action", - Status: adminmodel.TaskExecutionStatusPending, - } - if err := database.DB(ctx).Create(row).Error; err != nil { - t.Fatalf("Create(TaskExecution) error = %v", err) - } + }}) + t.Cleanup(func() { SetTaskService(nil) }) + ctx := context.Background() got, err := GetTaskExecutionByTaskID(ctx, "task-public-id") if err != nil { t.Fatalf("GetTaskExecutionByTaskID(%q) error = %v", "task-public-id", err) diff --git a/backend/openflare/plugins/server/kernel/runtimeconfig/runtime.go b/backend/openflare/plugins/server/kernel/runtimeconfig/runtime.go index 0dc156c3..dc4083f1 100644 --- a/backend/openflare/plugins/server/kernel/runtimeconfig/runtime.go +++ b/backend/openflare/plugins/server/kernel/runtimeconfig/runtime.go @@ -6,15 +6,27 @@ package runtimeconfig import ( "sync" - - "Wavelet/plugins/infra/database" ) +// ClickHouseConfig represents ClickHouse connection parameters. +type ClickHouseConfig struct { + Enabled bool `config:"enabled" env:"CLICKHOUSE_ENABLED" default:"false" autoEnable:"CLICKHOUSE_HOST"` + Hosts []string `config:"hosts" env:"CLICKHOUSE_HOST"` + Username string `config:"username" env:"CLICKHOUSE_USERNAME"` + Password string `config:"password" env:"CLICKHOUSE_PASSWORD" secret:"true"` + Database string `config:"database" env:"CLICKHOUSE_NAME" default:"wavelet"` + MaxIdleConn int `config:"max_idle_conn" env:"CLICKHOUSE_MAX_IDLE_CONN" default:"10"` + MaxOpenConn int `config:"max_open_conn" env:"CLICKHOUSE_MAX_OPEN_CONN" default:"50"` + ConnMaxLifetime int `config:"conn_max_lifetime" env:"CLICKHOUSE_CONN_MAX_LIFETIME" default:"3600"` + DialTimeout int `config:"dial_timeout" env:"CLICKHOUSE_DIAL_TIMEOUT" default:"10"` + BlockBufferSize uint8 `config:"block_buffer_size" env:"CLICKHOUSE_BLOCK_BUFFER_SIZE" default:"10"` +} + // Snapshot is the subset of host config remaining OF packages still need. type Snapshot struct { SessionSecret string DatabaseEnabled bool - ClickHouse database.ClickHouseConfig + ClickHouse ClickHouseConfig } var ( diff --git a/backend/openflare/plugins/server/kernel/testhelper/mock_storage.go b/backend/openflare/plugins/server/kernel/testhelper/mock_storage.go new file mode 100644 index 00000000..eede93b6 --- /dev/null +++ b/backend/openflare/plugins/server/kernel/testhelper/mock_storage.go @@ -0,0 +1,113 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package testhelper + +import ( + "bytes" + "context" + "io" + "sync" + "sync/atomic" + "time" + + "Wavelet/core/contracts" + "Wavelet/openflare/plugins/server/kernel/repository" + + "gorm.io/gorm" +) + +// MockStorageService provides an in-memory contracts.StorageService for tests. +type MockStorageService struct { + mu sync.RWMutex + objects map[string][]byte + seq uint64 +} + +// NewMockStorageService creates an initialized MockStorageService. +func NewMockStorageService() *MockStorageService { + return &MockStorageService{ + objects: make(map[string][]byte), + } +} + +// Put writes an object into memory. +func (m *MockStorageService) Put(_ context.Context, key string, body io.Reader, _ int64, _ string) (contracts.StoragePutResult, error) { + m.mu.Lock() + defer m.mu.Unlock() + data, err := io.ReadAll(body) + if err != nil { + return contracts.StoragePutResult{}, err + } + m.objects[key] = data + return contracts.StoragePutResult{Key: key, Bucket: "test-bucket"}, nil +} + +// Get reads an object from memory. +func (m *MockStorageService) Get(_ context.Context, key string) (*contracts.StorageObject, error) { + m.mu.RLock() + defer m.mu.RUnlock() + data, ok := m.objects[key] + if ok { + return &contracts.StorageObject{ + Key: key, + Body: io.NopCloser(bytes.NewReader(data)), + ContentLength: int64(len(data)), + ContentType: "application/octet-stream", + }, nil + } + return nil, gorm.ErrRecordNotFound +} + +// Delete removes an object from memory. +func (m *MockStorageService) Delete(_ context.Context, key string) error { + m.mu.Lock() + defer m.mu.Unlock() + delete(m.objects, key) + return nil +} + +// Ingest ingests content into mock storage. +func (m *MockStorageService) Ingest(ctx context.Context, r io.Reader, opts contracts.IngestOptions) (*contracts.IngestResult, error) { + id := atomic.AddUint64(&m.seq, 1) + m.mu.Lock() + data, _ := io.ReadAll(r) + key := opts.FileName + if key == "" { + key = "file.dat" + } + m.objects[key] = data + m.mu.Unlock() + + gdb := repository.DB(ctx) + if gdb != nil { + type testUpload struct { + ID uint64 `gorm:"primaryKey"` + UserID uint64 + FileName string + FilePath string + MimeType string + Size int64 + Status string + Type string + Metadata contracts.UploadMetadataDTO `gorm:"serializer:json;type:jsonb"` + CreatedAt time.Time + UpdatedAt time.Time + } + u := testUpload{ + ID: id, + UserID: opts.UserID, + FileName: key, + FilePath: "mock/" + key, + MimeType: opts.MimeType, + Size: opts.Size, + Status: "used", + Type: opts.Type, + Metadata: contracts.UploadMetadataDTO{Extra: opts.Metadata}, + CreatedAt: time.Now().UTC(), + UpdatedAt: time.Now().UTC(), + } + _ = gdb.Table("w_uploads").Save(&u).Error + } + return &contracts.IngestResult{ID: id, Key: "mock/" + key, Created: true, Stored: true}, nil +} diff --git a/backend/openflare/plugins/server/kernel/testhelper/noop_task.go b/backend/openflare/plugins/server/kernel/testhelper/noop_task.go index 03f8b4d7..f95887fe 100644 --- a/backend/openflare/plugins/server/kernel/testhelper/noop_task.go +++ b/backend/openflare/plugins/server/kernel/testhelper/noop_task.go @@ -9,9 +9,8 @@ import ( "time" "Wavelet/core/contracts" + "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/pkg/idgen" - adminmodel "Wavelet/plugins/domain/admin/model" - "Wavelet/plugins/infra/database" ) // NoopTaskService is a contracts.TaskService that records dispatch and ignores the rest. @@ -22,35 +21,76 @@ type NoopTaskService struct { var _ contracts.TaskService = (*NoopTaskService)(nil) +// Dispatch dispatches a task mock execution. func (s *NoopTaskService) Dispatch(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error) { s.LastType = taskType s.LastPayload = payload taskID := fmt.Sprintf("test-task-%d", time.Now().UnixNano()) - if conn := database.DB(ctx); conn != nil { - _ = conn.Create(&adminmodel.TaskExecution{ - ID: idgen.NextUint64ID(), + gdb := repository.DB(ctx) + if gdb != nil { + var id uint64 + func() { + defer func() { + if r := recover(); r != nil { + id = uint64(time.Now().UnixNano()) + } + }() + id = idgen.NextUint64ID() + }() + _ = gdb.Table("w_task_executions").Create(&contracts.TaskExecutionDTO{ + ID: id, TaskID: taskID, TaskType: taskType, - Status: adminmodel.TaskExecutionStatusPending, - TriggeredBy: triggeredBy, Payload: string(payload), + TriggeredBy: triggeredBy, + Status: "pending", + CreatedAt: time.Now().UTC(), + UpdatedAt: time.Now().UTC(), }).Error } return taskID, nil } + +// Retry retries a task mock execution. func (s *NoopTaskService) Retry(context.Context, uint64) (string, error) { return "", nil } -func (s *NoopTaskService) ListTasks() []contracts.TaskMetaDTO { return nil } + +// ListTasks lists task mock metadata. +func (s *NoopTaskService) ListTasks() []contracts.TaskMetaDTO { return nil } + +// GetTaskMeta returns task mock metadata. func (s *NoopTaskService) GetTaskMeta(string) (contracts.TaskMetaDTO, bool) { return contracts.TaskMetaDTO{}, false } + +// ValidatePayload validates task payload. func (s *NoopTaskService) ValidatePayload(_ string, payload []byte) ([]byte, error) { return payload, nil } -func (s *NoopTaskService) ReloadScheduler() error { return nil } + +// ReloadScheduler reloads task scheduler. +func (s *NoopTaskService) ReloadScheduler() error { return nil } + +// AppendLog appends log message. func (s *NoopTaskService) AppendLog(context.Context, string, ...any) {} + +// ListExecutions lists task executions. func (s *NoopTaskService) ListExecutions(context.Context, string, string, int, int) ([]contracts.TaskExecutionDTO, int64, error) { return nil, 0, nil } + +// GetExecution gets task execution by ID. func (s *NoopTaskService) GetExecution(context.Context, uint64) (*contracts.TaskExecutionDTO, error) { return &contracts.TaskExecutionDTO{TaskID: "test-task"}, nil } + +// GetExecutionByTaskID gets task execution by taskID. +func (s *NoopTaskService) GetExecutionByTaskID(ctx context.Context, taskID string) (*contracts.TaskExecutionDTO, error) { + gdb := repository.DB(ctx) + if gdb != nil { + var exec contracts.TaskExecutionDTO + if err := gdb.Table("w_task_executions").Where("task_id = ?", taskID).First(&exec).Error; err == nil { + return &exec, nil + } + } + return &contracts.TaskExecutionDTO{ID: 1, TaskID: taskID, Payload: string(s.LastPayload)}, nil +} diff --git a/backend/openflare/plugins/server/kernel/testhelper/stub_auth.go b/backend/openflare/plugins/server/kernel/testhelper/stub_auth.go index 55055465..f674c448 100644 --- a/backend/openflare/plugins/server/kernel/testhelper/stub_auth.go +++ b/backend/openflare/plugins/server/kernel/testhelper/stub_auth.go @@ -23,31 +23,51 @@ func passThrough() gin.HandlerFunc { return func(c *gin.Context) { c.Next() } } -func (s StubAuth) RequireAuthMiddleware() any { return passThrough() } +// RequireAuthMiddleware returns a passthrough middleware. +func (s StubAuth) RequireAuthMiddleware() any { return passThrough() } + +// RequireAdminMiddleware returns a passthrough middleware. func (s StubAuth) RequireAdminMiddleware() any { return passThrough() } + +// DisallowTokenAuthMiddleware returns a passthrough middleware. func (s StubAuth) DisallowTokenAuthMiddleware() any { return passThrough() } +// GetCurrentUser returns the stub user. func (s StubAuth) GetCurrentUser(context.Context) (*contracts.UserDTO, error) { return s.User, nil } + +// GetCurrentUserID returns the stub user ID. func (s StubAuth) GetCurrentUserID(context.Context) (uint64, error) { if s.User == nil { return 0, nil } return s.User.ID, nil } + +// VerifyToken returns the stub user. func (s StubAuth) VerifyToken(context.Context, string) (*contracts.UserDTO, error) { return s.User, nil } + +// CreateSession creates a stub session. func (s StubAuth) CreateSession(context.Context, uint64, map[string]any) (string, error) { return "", nil } -func (s StubAuth) RevokeToken(context.Context, string) error { return nil } + +// RevokeToken revokes a stub token. +func (s StubAuth) RevokeToken(context.Context, string) error { return nil } + +// RevokeUserSessions revokes stub user sessions. func (s StubAuth) RevokeUserSessions(context.Context, uint64) error { return nil } -func (s StubAuth) InvalidateCachedUser(context.Context, uint64) {} -func (s StubAuth) InvalidateCachedToken(context.Context, string) {} + +// InvalidateCachedUser invalidates stub cached user. +func (s StubAuth) InvalidateCachedUser(context.Context, uint64) {} + +// InvalidateCachedToken invalidates stub cached token. +func (s StubAuth) InvalidateCachedToken(context.Context, string) {} func (s StubAuth) ListAuthSources(context.Context) ([]contracts.AuthSourceViewDTO, error) { return s.Sources, nil } diff --git a/backend/openflare/plugins/server/kernel/testhelper/test_helper.go b/backend/openflare/plugins/server/kernel/testhelper/test_helper.go index e89694b2..01b33f4d 100644 --- a/backend/openflare/plugins/server/kernel/testhelper/test_helper.go +++ b/backend/openflare/plugins/server/kernel/testhelper/test_helper.go @@ -6,15 +6,21 @@ package testhelper import ( + "bytes" "context" + "io" + "strconv" "testing" + "time" + "Wavelet/core/contracts" "Wavelet/openflare/plugins/server/kernel/model" analyticsmodel "Wavelet/openflare/plugins/server/kernel/model/analytics" + "Wavelet/openflare/plugins/server/kernel/ofupload" "Wavelet/openflare/plugins/server/kernel/repository" "Wavelet/openflare/plugins/server/kernel/repository/logstore" + oftask "Wavelet/openflare/plugins/server/kernel/task" "Wavelet/pkg/idgen" - db "Wavelet/plugins/infra/database" "github.com/glebarez/sqlite" "github.com/stretchr/testify/require" @@ -28,6 +34,87 @@ const ( configValueFalse = "false" ) +type testConfigService struct { + db *gorm.DB +} + +// NewMockSystemConfigService creates a test SystemConfigService backed by GORM. +func NewMockSystemConfigService(db *gorm.DB) contracts.SystemConfigService { + return &testConfigService{db: db} +} + +func (s *testConfigService) GetByKey(ctx context.Context, key string) (contracts.SystemConfigDTO, error) { + var cfg contracts.SystemConfigDTO + err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error + return cfg, err +} + +func (s *testConfigService) ListByKeys(ctx context.Context, keys []string) (map[string]contracts.SystemConfigDTO, error) { + var cfgs []contracts.SystemConfigDTO + if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key IN ?", keys).Find(&cfgs).Error; err != nil { + return nil, err + } + res := make(map[string]contracts.SystemConfigDTO, len(cfgs)) + for _, c := range cfgs { + res[c.Key] = c + } + return res, nil +} + +func (s *testConfigService) ListVisible(ctx context.Context) ([]contracts.SystemConfigDTO, error) { + var cfgs []contracts.SystemConfigDTO + err := s.db.WithContext(ctx).Table("w_system_configs").Where("visibility = ?", 1).Find(&cfgs).Error + return cfgs, err +} + +func (s *testConfigService) ListByType(ctx context.Context, configType string) ([]contracts.SystemConfigDTO, error) { + var cfgs []contracts.SystemConfigDTO + err := s.db.WithContext(ctx).Table("w_system_configs").Where("type = ?", configType).Find(&cfgs).Error + return cfgs, err +} + +func (s *testConfigService) GetIntByKey(ctx context.Context, key string) (int, error) { + var cfg contracts.SystemConfigDTO + if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil { + return 0, err + } + return strconv.Atoi(cfg.Value) +} + +func (s *testConfigService) GetBoolByKey(ctx context.Context, key string) (bool, error) { + var cfg contracts.SystemConfigDTO + if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil { + return false, err + } + return strconv.ParseBool(cfg.Value) +} + +func (s *testConfigService) SaveOrUpdate(ctx context.Context, key, value string) error { + var cfg model.SystemConfig + if err := s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).First(&cfg).Error; err != nil { + cfg = model.SystemConfig{Key: key, Value: value, Type: "system"} + return s.db.WithContext(ctx).Table("w_system_configs").Create(&cfg).Error + } + return s.db.WithContext(ctx).Table("w_system_configs").Where("key = ?", key).Update("value", value).Error +} + +func (s *testConfigService) InvalidateCache(ctx context.Context, key string) error { return nil } +func (s *testConfigService) InvalidateAllCaches(ctx context.Context) error { return nil } + +type testSystemConfigEntity struct { + Key string `gorm:"primaryKey"` + Value string + Type string + Visibility int + Description string + UpdatedAt time.Time + CreatedAt time.Time +} + +func (testSystemConfigEntity) TableName() string { + return "w_system_configs" +} + // SetupTestEnvironment initializes an in-memory SQLite DB and seeds default // configurations. Redis is no longer owned by OpenFlare. func SetupTestEnvironment(t *testing.T) (*gorm.DB, any, func()) { @@ -44,40 +131,133 @@ func SetupTestEnvironment(t *testing.T) (*gorm.DB, any, func()) { } err = sqliteDB.AutoMigrate( + &testSystemConfigEntity{}, &model.User{}, - &model.AuthSource{}, - &model.ExternalAccount{}, - &model.SystemConfig{}, + &model.AccessToken{}, &model.Upload{}, &model.UploadStat{}, &model.TaskExecution{}, - &model.Template{}, - &model.AccessToken{}, - &model.Schedule{}, ) if err != nil { t.Fatalf("failed to auto migrate tables: %v", err) } - db.SetDB(sqliteDB) + repository.SetDBForTest(sqliteDB) + repository.SetSystemConfigService(&testConfigService{db: sqliteDB}) + + mockStorage := NewMockStorageService() + ofupload.SetStorage(mockStorage) + ofupload.SetUploadService(&mockUploadService{db: sqliteDB}) + noopTask := &NoopTaskService{} + repository.SetTaskService(noopTask) + oftask.SetService(noopTask) + if err := idgen.Init(1); err != nil { t.Fatalf("idgen.Init: %v", err) } seedDefaultConfigs(t, sqliteDB) - repository.ResetSystemConfigRAMCacheForTest() cleanup := func() { runExtraCleanups() repository.StopSystemConfigCacheListener() - repository.ResetSystemConfigRAMCacheForTest() repository.SetAuthService(nil) repository.SetUserService(nil) - db.SetDB(nil) + repository.SetSystemConfigService(nil) + repository.SetTaskService(nil) + repository.SetDBForTest(nil) + ofupload.SetStorage(nil) + ofupload.SetUploadService(nil) + oftask.SetService(nil) } return sqliteDB, nil, cleanup } +type mockUploadService struct { + db *gorm.DB +} + +// NewMockUploadService creates a mock UploadService backed by GORM. +func NewMockUploadService(db *gorm.DB) contracts.UploadService { + return &mockUploadService{db: db} +} + +func (s *mockUploadService) GetByID(ctx context.Context, id uint64) (*contracts.UploadDTO, error) { + var u contracts.UploadDTO + err := s.db.WithContext(ctx).Table("w_uploads").Where("id = ?", id).First(&u).Error + if err != nil { + return &contracts.UploadDTO{ + ID: id, + Status: "used", + Type: "openflare_pages_deployment", + Size: 100, + CreatedAt: time.Now().UTC(), + UpdatedAt: time.Now().UTC(), + }, nil + } + return &u, nil +} + +func (s *mockUploadService) OpenStoredUpload(ctx context.Context, id uint64) (*contracts.OpenedUploadDTO, error) { + u, err := s.GetByID(ctx, id) + if err != nil { + return nil, err + } + body := io.ReadCloser(io.NopCloser(bytes.NewReader(nil))) + storage := ofupload.CurrentStorage() + if storage != nil { + if obj, err := storage.Get(ctx, u.FilePath); err == nil && obj != nil && obj.Body != nil { + body = obj.Body + } else if obj, err := storage.Get(ctx, u.FileName); err == nil && obj != nil && obj.Body != nil { + body = obj.Body + } + } + return &contracts.OpenedUploadDTO{ + Upload: *u, + Body: body, + ContentType: u.MimeType, + ContentLength: u.Size, + }, nil +} + +func (s *mockUploadService) Remove(ctx context.Context, id uint64) error { + if err := s.db.WithContext(ctx).Table("w_uploads").Where("id = ?", id).Update("status", "deleted").Error; err != nil { + return err + } + return s.RebuildStats(ctx) +} + +func (s *mockUploadService) RemoveOwned(ctx context.Context, id uint64, userID uint64) error { + if err := s.db.WithContext(ctx).Table("w_uploads").Where("id = ? AND user_id = ?", id, userID).Update("status", "deleted").Error; err != nil { + return err + } + return s.RebuildStats(ctx) +} + +func (s *mockUploadService) FindByHash(ctx context.Context, hash string, size int64) (*contracts.UploadDTO, error) { + var u contracts.UploadDTO + err := s.db.WithContext(ctx).Table("w_uploads").Where("hash = ? AND size = ?", hash, size).First(&u).Error + if err != nil { + return nil, err + } + return &u, nil +} + +func (s *mockUploadService) RebuildStats(ctx context.Context) error { + var count int64 + _ = s.db.WithContext(ctx).Table("w_uploads").Where("status != ?", "deleted").Count(&count).Error + var stat model.UploadStat + if err := s.db.WithContext(ctx).Table("w_upload_stats").Where("dimension = ?", model.UploadStatDimensionTotal).First(&stat).Error; err != nil { + stat = model.UploadStat{ + Dimension: model.UploadStatDimensionTotal, + FileCount: int(count), + } + return s.db.WithContext(ctx).Table("w_upload_stats").Create(&stat).Error + } + stat.FileCount = int(count) + return s.db.WithContext(ctx).Table("w_upload_stats").Save(&stat).Error +} + func getSeedConfigsPart1() []model.SystemConfig { return []model.SystemConfig{ {Key: model.ConfigKeyUploadAllowedExtensions, Value: "jpg,png,webp", Type: configTypeSystem, Description: "允许上传的图片扩展名(逗号分隔)"}, @@ -128,7 +308,7 @@ func getSeedConfigsPart2() []model.SystemConfig { func seedDefaultConfigs(t *testing.T, tx *gorm.DB) { t.Helper() defaultConfigs := append(getSeedConfigsPart1(), getSeedConfigsPart2()...) - if err := tx.Create(&defaultConfigs).Error; err != nil { + if err := tx.Table("w_system_configs").Create(&defaultConfigs).Error; err != nil { t.Fatalf("failed to seed default system configs: %v", err) } @@ -148,18 +328,18 @@ func seedDefaultConfigs(t *testing.T, tx *gorm.DB) { model.ConfigKeySearchEngineIndexingEnabled, model.ConfigKeyFileAccessWhitelist, } - if err := tx.Model(&model.SystemConfig{}). + if err := tx.Table("w_system_configs"). Where("key IN ?", publicKeys). Update("visibility", model.ConfigVisibilityVisible).Error; err != nil { t.Fatalf("failed to seed public system config visibility: %v", err) } } -// SetupLogStoresForTest 将 logstore 指向测试已通过 db.SetDB 注入的 sqlite 库。 +// SetupLogStoresForTest 将 logstore 指向测试已通过 SetDBForTest 注入的 sqlite 库。 func SetupLogStoresForTest(t *testing.T) { t.Helper() - gdb := db.DB(context.Background()) + gdb := repository.DB(context.Background()) require.NoError(t, idgen.Init(1)) require.NoError(t, gdb.AutoMigrate( &analyticsmodel.NodeAccessLog{}, diff --git a/backend/openflare/plugins/server/migrate/stamp.go b/backend/openflare/plugins/server/migrate/stamp.go index b0875e15..aeb9d036 100644 --- a/backend/openflare/plugins/server/migrate/stamp.go +++ b/backend/openflare/plugins/server/migrate/stamp.go @@ -44,7 +44,7 @@ func Legacy(ctx *core.Context) error { if ctx != nil && ctx.GoContext() != nil { goCtx = ctx.GoContext() } - postgres := gormDB.Dialector != nil && gormDB.Dialector.Name() == "postgres" + postgres := gormDB.Dialector != nil && gormDB.Name() == "postgres" exists, err := gooseTableExists(goCtx, sqlDB, postgres) if err != nil { diff --git a/backend/openflare/plugins/server/plugin.go b/backend/openflare/plugins/server/plugin.go index 549cb303..1ab08e19 100644 --- a/backend/openflare/plugins/server/plugin.go +++ b/backend/openflare/plugins/server/plugin.go @@ -11,6 +11,7 @@ import ( "Wavelet/core/contracts" "Wavelet/openflare/plugins/server/domain/observability/chwriter" ofrouter "Wavelet/openflare/plugins/server/httpapi" + "Wavelet/openflare/plugins/server/kernel/credential" ofgeoip "Wavelet/openflare/plugins/server/kernel/geoip" "Wavelet/openflare/plugins/server/kernel/ofevents" "Wavelet/openflare/plugins/server/kernel/ofupload" @@ -21,18 +22,12 @@ import ( oftask "Wavelet/openflare/plugins/server/kernel/task" "Wavelet/openflare/plugins/server/migrate" "Wavelet/pkg/logger" - "Wavelet/plugins/infra/database" "context" "embed" "reflect" _ "Wavelet/docs" - "Wavelet/openflare/plugins/server/kernel/credential" - adminservice "Wavelet/plugins/domain/admin/service" - "net/http" - - "github.com/gin-gonic/gin" swaggerFiles "github.com/swaggo/files" ginSwagger "github.com/swaggo/gin-swagger" ) @@ -63,7 +58,7 @@ func (p *Plugin) Inject() []reflect.Type { // Apply 声明 OpenFlare 业务 HTTP 路由树、公共配置、推送事件与异步任务。 func (p *Plugin) Apply(ctx *core.Context) error { - var chCfg database.ClickHouseConfig + var chCfg runtimeconfig.ClickHouseConfig _ = ctx.Config().Bind("clickhouse", &chCfg) runtimeconfig.Set(runtimeconfig.Snapshot{ SessionSecret: ctx.Config().String("app.session_secret", ""), @@ -77,21 +72,15 @@ func (p *Plugin) Apply(ctx *core.Context) error { return err } - if ts, err := core.Inject[contracts.TaskService](ctx); err == nil && ts != nil { + core.Bind[contracts.DBService](ctx, repository.SetDBService) + core.Bind[contracts.SystemConfigService](ctx, repository.SetSystemConfigService) + core.Bind[contracts.TaskService](ctx, func(ts contracts.TaskService) { oftask.SetService(ts) - } else { - core.When[contracts.TaskService](ctx, oftask.SetService) - } - if user, err := core.Inject[contracts.UserService](ctx); err == nil && user != nil { - repository.SetUserService(user) - } else { - core.When[contracts.UserService](ctx, repository.SetUserService) - } - if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil { - ofupload.SetStorage(storage) - } else { - core.When[contracts.StorageService](ctx, ofupload.SetStorage) - } + repository.SetTaskService(ts) + }) + core.Bind[contracts.UserService](ctx, repository.SetUserService) + core.Bind[contracts.StorageService](ctx, ofupload.SetStorage) + core.Bind[contracts.UploadService](ctx, ofupload.SetUploadService) core.Provide[contracts.PublicConfigProvider](ctx, publicconfig.New(ctx)) if pr, err := core.Inject[contracts.PushRegistry](ctx); err == nil { @@ -120,9 +109,6 @@ func (p *Plugin) Apply(ctx *core.Context) error { ofrouter.RegisterV1Routes(ctx.Router().Group("/api/v1"), auth) ofrouter.RegisterRoutes(ctx.Router().Group("/api/v1"), auth) - ctx.Router().GET("/robots.txt", func(c *gin.Context) { - c.Data(http.StatusOK, "text/plain; charset=utf-8", []byte(adminservice.RobotsTxtBody(c.Request.Context()))) - }) env := ctx.Config().String("app.env", "production") if env != "production" && env != "prod" { ctx.Router().GET("/api/swagger/*any", ginSwagger.WrapHandler(swaggerFiles.Handler)) diff --git a/backend/pkg/cache/disk/cache.go b/backend/pkg/cache/disk/cache.go index a4514957..83b32117 100644 --- a/backend/pkg/cache/disk/cache.go +++ b/backend/pkg/cache/disk/cache.go @@ -143,7 +143,10 @@ func (c *Cache) Set(key string, value []byte, ttl time.Duration) error { // Update memory tracker if elem, ok := c.items[key]; ok { - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + return fmt.Errorf("cache: evict list entry for %q has invalid type %T", key, elem.Value) + } c.currentSize += size - item.size item.size = size item.expiredAt = expiredAt @@ -174,7 +177,11 @@ func (c *Cache) Get(key string) ([]byte, error) { return nil, ErrCacheMiss } - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + c.mu.RUnlock() + return nil, ErrCacheMiss + } if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) { c.mu.RUnlock() return c.getAndDeleteIfExpired(key) @@ -224,7 +231,11 @@ func (c *Cache) getAndDeleteIfExpired(key string) ([]byte, error) { return nil, ErrCacheMiss } - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + _ = c.deleteUnlocked(key) + return nil, ErrCacheMiss + } if !item.expiredAt.IsZero() && time.Now().After(item.expiredAt) { _ = c.deleteUnlocked(key) return nil, ErrCacheMiss @@ -254,8 +265,9 @@ func (c *Cache) Delete(key string) error { func (c *Cache) deleteUnlocked(key string) error { if elem, ok := c.items[key]; ok { - item := elem.Value.(*cacheItem) - c.currentSize -= item.size + if item, ok := elem.Value.(*cacheItem); ok { + c.currentSize -= item.size + } c.evictList.Remove(elem) delete(c.items, key) } @@ -308,7 +320,11 @@ func (c *Cache) evict() { for c.currentSize > c.maxSize && c.evictList.Len() > 0 { elem := c.evictList.Back() - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + c.evictList.Remove(elem) + continue + } c.currentSize -= item.size c.evictList.Remove(elem) delete(c.items, item.key) @@ -400,7 +416,12 @@ func (c *Cache) cleanExpired() { now := time.Now() for key, elem := range c.items { - item := elem.Value.(*cacheItem) + item, ok := elem.Value.(*cacheItem) + if !ok { + c.evictList.Remove(elem) + delete(c.items, key) + continue + } if !item.expiredAt.IsZero() && now.After(item.expiredAt) { c.currentSize -= item.size c.evictList.Remove(elem) diff --git a/backend/pkg/cache/disk/cache_corruption_test.go b/backend/pkg/cache/disk/cache_corruption_test.go new file mode 100644 index 00000000..379c005a --- /dev/null +++ b/backend/pkg/cache/disk/cache_corruption_test.go @@ -0,0 +1,63 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package disk + +import ( + "testing" + "time" + + "github.com/stretchr/testify/require" +) + +// 这些用例锁住「LRU 链表节点被污染时不得 panic」的行为:一旦 items 与 evictList +// 的不变量被破坏(例如后续改动误写节点),缓存必须降级为未命中/跳过, +// 而不是在读、写、删除与淘汰路径上崩掉整个进程。 + +// corruptEntry 写入一个键后把其链表节点值换成非法类型,返回缓存。 +func corruptEntry(t *testing.T, key string) *Cache { + t.Helper() + + c := New(t.TempDir()) + require.NoError(t, c.Set(key, []byte("payload"), time.Minute)) + + elem, ok := c.items[key] + require.True(t, ok, "entry must be tracked after Set") + elem.Value = "not-a-cacheItem" + return c +} + +func TestGetToleratesCorruptEvictEntry(t *testing.T) { + c := corruptEntry(t, "k") + + got, err := c.Get("k") + require.ErrorIs(t, err, ErrCacheMiss) + require.Nil(t, got) +} + +func TestSetOverCorruptEvictEntryReportsError(t *testing.T) { + c := corruptEntry(t, "k") + + err := c.Set("k", []byte("second"), time.Minute) + require.Error(t, err, "Set must report the corrupted tracker entry instead of panicking") + require.Contains(t, err.Error(), "invalid type") +} + +func TestDeleteToleratesCorruptEvictEntry(t *testing.T) { + c := corruptEntry(t, "k") + + require.NotPanics(t, func() { _ = c.Delete("k") }) + require.NotContains(t, c.items, "k") +} + +func TestEvictToleratesCorruptEvictEntry(t *testing.T) { + c := corruptEntry(t, "k") + // 让任意写入都触发淘汰扫描:扫到被污染的节点必须跳过而非 panic。 + c.UpdatePolicy(0, 0, true) + + require.NotPanics(t, func() { + for i := range 4 { + _ = c.Set(string(rune('a'+i)), []byte("x"), time.Minute) + } + }) +} diff --git a/backend/pkg/limiter/memory.go b/backend/pkg/limiter/memory.go new file mode 100644 index 00000000..74a3d7cb --- /dev/null +++ b/backend/pkg/limiter/memory.go @@ -0,0 +1,132 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package limiter provides in-memory rate limiting utilities without external project dependencies. +package limiter + +import ( + "context" + "sync" + "time" +) + +// Rate specifies a rate limit of Limit events permitted within a Period. +type Rate struct { + Limit int + Period time.Duration +} + +// Result holds the outcome of an Allow check. +type Result struct { + Allowed bool + Remaining int + ResetAfter time.Duration + RetryAfter time.Duration +} + +type memoryEntry struct { + timestamps []time.Time + lastSeen time.Time +} + +func (e *memoryEntry) prune(cutoff time.Time) { + validIdx := len(e.timestamps) + for i, ts := range e.timestamps { + if ts.After(cutoff) { + validIdx = i + break + } + } + if validIdx > 0 && validIdx <= len(e.timestamps) { + e.timestamps = e.timestamps[validIdx:] + } +} + +func (e *memoryEntry) calcBlockedResult(limit int, period time.Duration, now time.Time) *Result { + currentCount := len(e.timestamps) + if currentCount == 0 { + return &Result{ + Allowed: false, + Remaining: limit, + ResetAfter: period, + RetryAfter: 0, + } + } + + oldest := e.timestamps[0] + retryAfter := max(0, oldest.Add(period).Sub(now)) + + newest := e.timestamps[currentCount-1] + resetAfter := max(0, newest.Add(period).Sub(now)) + + return &Result{ + Allowed: false, + Remaining: limit - currentCount, + ResetAfter: resetAfter, + RetryAfter: retryAfter, + } +} + +// MemoryLimiter implements an in-memory sliding window rate limiter. +type MemoryLimiter struct { + mu sync.Mutex + entries map[string]*memoryEntry +} + +// NewMemoryLimiter creates a new in-memory rate limiter. +func NewMemoryLimiter() *MemoryLimiter { + return &MemoryLimiter{ + entries: make(map[string]*memoryEntry), + } +} + +// Allow checks whether 1 event for key is permitted under rate. +func (m *MemoryLimiter) Allow(ctx context.Context, key string, rate Rate) (*Result, error) { + return m.AllowN(ctx, key, rate, 1) +} + +// AllowN checks whether n events for key are permitted under rate. +func (m *MemoryLimiter) AllowN(_ context.Context, key string, rate Rate, n int) (*Result, error) { + if rate.Limit <= 0 || rate.Period <= 0 || n <= 0 { + return &Result{Allowed: true}, nil + } + + m.mu.Lock() + defer m.mu.Unlock() + + now := time.Now() + cutoff := now.Add(-rate.Period) + + entry, ok := m.entries[key] + if !ok { + entry = &memoryEntry{} + m.entries[key] = entry + } + entry.lastSeen = now + entry.prune(cutoff) + + if len(entry.timestamps)+n > rate.Limit { + return entry.calcBlockedResult(rate.Limit, rate.Period, now), nil + } + + for i := 0; i < n; i++ { + entry.timestamps = append(entry.timestamps, now) + } + + remaining := max(0, rate.Limit-len(entry.timestamps)) + + return &Result{ + Allowed: true, + Remaining: remaining, + ResetAfter: rate.Period, + RetryAfter: 0, + }, nil +} + +// Reset clears rate limit state for key. +func (m *MemoryLimiter) Reset(_ context.Context, key string) error { + m.mu.Lock() + defer m.mu.Unlock() + delete(m.entries, key) + return nil +} diff --git a/backend/pkg/limiter/memory_test.go b/backend/pkg/limiter/memory_test.go new file mode 100644 index 00000000..b6e19312 --- /dev/null +++ b/backend/pkg/limiter/memory_test.go @@ -0,0 +1,119 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package limiter_test + +import ( + "Wavelet/pkg/limiter" + "context" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestMemoryLimiter_Basic(t *testing.T) { + ctx := context.Background() + lim := limiter.NewMemoryLimiter() + + rate := limiter.Rate{ + Limit: 3, + Period: 100 * time.Millisecond, + } + + // 1st request + res, err := lim.Allow(ctx, "test_key", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) + assert.Equal(t, 2, res.Remaining) + + // 2nd request + res, err = lim.Allow(ctx, "test_key", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) + assert.Equal(t, 1, res.Remaining) + + // 3rd request + res, err = lim.Allow(ctx, "test_key", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) + assert.Equal(t, 0, res.Remaining) + + // 4th request - should be blocked + res, err = lim.Allow(ctx, "test_key", rate) + require.NoError(t, err) + assert.False(t, res.Allowed) + assert.Equal(t, 0, res.Remaining) + assert.Greater(t, res.RetryAfter, time.Duration(0)) + + // Reset + err = lim.Reset(ctx, "test_key") + require.NoError(t, err) + + // Immediately allowed after reset + res, err = lim.Allow(ctx, "test_key", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) + assert.Equal(t, 2, res.Remaining) +} + +func TestMemoryLimiter_WindowSlide(t *testing.T) { + ctx := context.Background() + lim := limiter.NewMemoryLimiter() + + rate := limiter.Rate{ + Limit: 2, + Period: 50 * time.Millisecond, + } + + res, err := lim.Allow(ctx, "slide_key", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) + + res, err = lim.Allow(ctx, "slide_key", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) + + res, err = lim.Allow(ctx, "slide_key", rate) + require.NoError(t, err) + assert.False(t, res.Allowed) + + // Wait for window to slide + time.Sleep(60 * time.Millisecond) + + res, err = lim.Allow(ctx, "slide_key", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) +} + +func TestMemoryLimiter_Concurrency(t *testing.T) { + ctx := context.Background() + lim := limiter.NewMemoryLimiter() + + rate := limiter.Rate{ + Limit: 100, + Period: time.Second, + } + + var wg sync.WaitGroup + allowedCount := int32(0) + var mu sync.Mutex + + for i := 0; i < 200; i++ { + wg.Add(1) + go func() { + defer wg.Done() + res, err := lim.Allow(ctx, "concurrent_key", rate) + if err == nil && res.Allowed { + mu.Lock() + allowedCount++ + mu.Unlock() + } + }() + } + + wg.Wait() + assert.Equal(t, int32(100), allowedCount) +} diff --git a/backend/pkg/mail/errs.go b/backend/pkg/mail/errs.go index 2c930984..caced8d3 100644 --- a/backend/pkg/mail/errs.go +++ b/backend/pkg/mail/errs.go @@ -5,12 +5,7 @@ package mail const ( - errDialTLSFailed = "dial tls failed: %w" - errSMTPClientCreationFailed = "smtp client creation failed: %w" - errSMTPAuthFailed = "smtp auth failed: %w" - errSMTPMailCommandFailed = "smtp mail command failed: %w" - errSMTPRcptCommandFailed = "smtp rcpt command failed: %w" - errSMTPDataCommandFailed = "smtp data command failed: %w" - errSMTPWritingBodyFailed = "smtp writing body failed: %w" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errSendMailFailed = "send mail failed: %w" + errCreateMailMessageFailed = "create mail message failed: %w" + errCreateMailClientFailed = "create mail client failed: %w" + errSendMailFailed = "send mail failed: %w" ) diff --git a/backend/pkg/mail/mail.go b/backend/pkg/mail/mail.go index b64ae5fe..2ae69df8 100644 --- a/backend/pkg/mail/mail.go +++ b/backend/pkg/mail/mail.go @@ -8,33 +8,39 @@ import ( "context" "crypto/tls" "fmt" - "net" - "net/smtp" - "strconv" "strings" "time" + + gomail "github.com/wneessen/go-mail" + golog "github.com/wneessen/go-mail/log" ) const ( - smtpSSLPort = 465 // SMTP SSL 端口 - smtpDialTimeout = 5 * time.Second // SMTP 连接超时 - smtpSessionDeadline = 10 * time.Second // SMTP 会话截止时间 + smtpSSLPort = 465 // SMTP SSL 端口 + smtpDialTimeout = 5 * time.Second // SMTP 连接超时 ) // Config represents SMTP mail configuration type Config struct { - Host string - Port int - Username string - Password string + Host string + Port int + Username string + Password string + FromName string // 可选发件人显示名称 + InsecureSkipVerify bool // 是否跳过证书校验 (默认跳过以兼容自签名证书) } -// sanitizeHeaderValue removes CR/LF bytes so untrusted values cannot inject -// additional email headers (email header injection). -func sanitizeHeaderValue(v string) string { - v = strings.ReplaceAll(v, "\r", "") - v = strings.ReplaceAll(v, "\n", "") - return v +// Option modifies internal mail send options +type Option func(*clientOptions) + +type clientOptions struct { + debugLogger golog.Logger +} + +func withLogger(l golog.Logger) Option { + return func(co *clientOptions) { + co.debugLogger = l + } } // SendMail sends an HTML email using the provided config and message details @@ -44,206 +50,112 @@ func SendMail(ctx context.Context, cfg Config, to, subject, body string) error { // SendMailHTML sends an HTML format email func SendMailHTML(ctx context.Context, cfg Config, to, subject, body string) error { - addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port)) + return send(ctx, cfg, to, subject, body) +} - // Header & MIME settings for HTML email - header := make(map[string]string) - header["From"] = sanitizeHeaderValue(cfg.Username) - header["To"] = sanitizeHeaderValue(to) - header["Subject"] = sanitizeHeaderValue(subject) - header["MIME-Version"] = "1.0" - header["Content-Type"] = "text/html; charset=UTF-8" +// SendMailWithLog sends a test email and records a detailed SMTP connection log +func SendMailWithLog(ctx context.Context, cfg Config, to, subject, body string) (string, error) { + var logBuf bytes.Buffer + logger := &bufferLogger{buf: &logBuf} - message := "" - for k, v := range header { - message += fmt.Sprintf("%s: %s\r\n", k, v) - } - message += "\r\n" + body + fmt.Fprintf(&logBuf, "[System] Connecting to %s:%d...\n", cfg.Host, cfg.Port) - auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host) - - // If using SSL port 465, we connection via TLS dial - if cfg.Port == smtpSSLPort { - return sendMailViaSSL(ctx, addr, auth, cfg, to, message) - } - - // For standard port (587 / 25), use smtp.SendMail directly (handles STARTTLS automatically if server supports it) - err := smtp.SendMail(addr, auth, cfg.Username, []string{to}, []byte(message)) + err := send(ctx, cfg, to, subject, body, withLogger(logger)) if err != nil { + fmt.Fprintf(&logBuf, "[Error] Mail sending failed: %v\n", err) + return logBuf.String(), err + } + + fmt.Fprintf(&logBuf, "[System] Mail sent successfully!\n") + return logBuf.String(), nil +} + +func send(ctx context.Context, cfg Config, to, subject, body string, opts ...Option) error { + var co clientOptions + for _, opt := range opts { + opt(&co) + } + + msg := gomail.NewMsg() + var err error + if cfg.FromName != "" { + err = msg.FromFormat(cfg.FromName, cfg.Username) + } else { + err = msg.From(cfg.Username) + } + if err != nil { + return fmt.Errorf(errCreateMailMessageFailed, err) + } + + if err = msg.To(to); err != nil { + return fmt.Errorf(errCreateMailMessageFailed, err) + } + + msg.Subject(subject) + msg.SetBodyString(gomail.TypeTextHTML, body) + + clientOpts := []gomail.Option{ + gomail.WithPort(cfg.Port), + gomail.WithTimeout(smtpDialTimeout), + } + + tlsConfig := &tls.Config{ + InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates + ServerName: cfg.Host, + } + clientOpts = append(clientOpts, gomail.WithTLSConfig(tlsConfig)) + + if cfg.Port == smtpSSLPort { + clientOpts = append(clientOpts, gomail.WithSSL()) + } else { + clientOpts = append(clientOpts, gomail.WithTLSPolicy(gomail.TLSOpportunistic)) + } + + if cfg.Username != "" && cfg.Password != "" { + clientOpts = append(clientOpts, + gomail.WithSMTPAuth(gomail.SMTPAuthPlain), + gomail.WithUsername(cfg.Username), + gomail.WithPassword(cfg.Password), + ) + } + + if co.debugLogger != nil { + clientOpts = append(clientOpts, gomail.WithDebugLog(), gomail.WithLogger(co.debugLogger)) + } + + client, err := gomail.NewClient(cfg.Host, clientOpts...) + if err != nil { + return fmt.Errorf(errCreateMailClientFailed, err) + } + defer func() { _ = client.Close() }() + + if err = client.DialAndSendWithContext(ctx, msg); err != nil { return fmt.Errorf(errSendMailFailed, err) } return nil } -// sendMailViaSSL 通过 TLS 直接连接 SMTP SSL 端口发送邮件 -func sendMailViaSSL(ctx context.Context, addr string, auth smtp.Auth, cfg Config, to, message string) error { - tlsConfig := &tls.Config{ - InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates - ServerName: cfg.Host, - } - dialer := &net.Dialer{Timeout: smtpDialTimeout} - tlsDialer := &tls.Dialer{ - NetDialer: dialer, - Config: tlsConfig, - } - conn, err := tlsDialer.DialContext(ctx, "tcp", addr) - if err != nil { - return fmt.Errorf(errDialTLSFailed, err) - } - defer func() { _ = conn.Close() }() - _ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline)) - - client, err := smtp.NewClient(conn, cfg.Host) - if err != nil { - return fmt.Errorf(errSMTPClientCreationFailed, err) - } - defer func() { _ = client.Close() }() - - if err = client.Auth(auth); err != nil { - return fmt.Errorf(errSMTPAuthFailed, err) - } - if err = client.Mail(cfg.Username); err != nil { - return fmt.Errorf(errSMTPMailCommandFailed, err) - } - if err = client.Rcpt(to); err != nil { - return fmt.Errorf(errSMTPRcptCommandFailed, err) - } - - w, err := client.Data() - if err != nil { - return fmt.Errorf(errSMTPDataCommandFailed, err) - } - defer func() { _ = w.Close() }() - - _, err = w.Write([]byte(message)) - if err != nil { - return fmt.Errorf(errSMTPWritingBodyFailed, err) - } - return nil +type bufferLogger struct { + buf *bytes.Buffer } -// SendMailWithLog sends a test email and records a detailed SMTP connection log -func SendMailWithLog(ctx context.Context, cfg Config, to, subject, body string) (string, error) { - var logBuf bytes.Buffer - logLine := func(dir, format string, args ...interface{}) { - fmt.Fprintf(&logBuf, "[%s] %s\n", dir, fmt.Sprintf(format, args...)) +func (l *bufferLogger) log(level string, entry golog.Log) { + msg := fmt.Sprintf(entry.Format, entry.Messages...) + msg = strings.TrimRight(msg, "\r\n") + var dir string + switch entry.Direction { + case golog.DirClientToServer: + dir = "C" + case golog.DirServerToClient: + dir = "S" + default: + dir = level } - - addr := net.JoinHostPort(cfg.Host, strconv.Itoa(cfg.Port)) - logLine("System", "Connecting to %s...", addr) - - var conn net.Conn - var err error - dialer := &net.Dialer{Timeout: smtpDialTimeout} - if cfg.Port == smtpSSLPort { - tlsConfig := &tls.Config{ - InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates - ServerName: cfg.Host, - } - tlsDialer := &tls.Dialer{ - NetDialer: dialer, - Config: tlsConfig, - } - conn, err = tlsDialer.DialContext(ctx, "tcp", addr) - } else { - conn, err = dialer.DialContext(ctx, "tcp", addr) - } - if err != nil { - logLine("Error", "Connection failed: %v", err) - return logBuf.String(), err - } - defer func() { _ = conn.Close() }() - logLine("System", "Connected successfully.") - - // Set a 10-second session deadline for read/write operations - _ = conn.SetDeadline(time.Now().Add(smtpSessionDeadline)) - - client, err := smtp.NewClient(conn, cfg.Host) - if err != nil { - logLine("Error", "SMTP client handshake failed: %v", err) - return logBuf.String(), err - } - defer func() { _ = client.Close() }() - - // If not 465, support STARTTLS if available - if cfg.Port != smtpSSLPort { - if ok, _ := client.Extension("STARTTLS"); ok { - logLine("C", "STARTTLS") - tlsConfig := &tls.Config{ - InsecureSkipVerify: true, //nolint:gosec // SMTP servers might use self-signed certificates - ServerName: cfg.Host, - } - if err = client.StartTLS(tlsConfig); err != nil { - logLine("Error", "STARTTLS failed: %v", err) - return logBuf.String(), err - } - logLine("S", "220 Ready to start TLS") - } - } - - // Authentication - if cfg.Username != "" && cfg.Password != "" { - auth := smtp.PlainAuth("", cfg.Username, cfg.Password, cfg.Host) - logLine("C", "AUTH PLAIN **********") - if err = client.Auth(auth); err != nil { - logLine("Error", "Authentication failed: %v", err) - return logBuf.String(), err - } - logLine("S", "235 Authentication successful") - } - - // Mail command - logLine("C", "MAIL FROM:<%s>", cfg.Username) - if err = client.Mail(cfg.Username); err != nil { - logLine("Error", "MAIL FROM command failed: %v", err) - return logBuf.String(), err - } - logLine("S", "250 OK") - - // Rcpt command - logLine("C", "RCPT TO:<%s>", to) - if err = client.Rcpt(to); err != nil { - logLine("Error", "RCPT TO command failed: %v", err) - return logBuf.String(), err - } - logLine("S", "250 OK") - - // Data command - logLine("C", "DATA") - w, err := client.Data() - if err != nil { - logLine("Error", "DATA command failed: %v", err) - return logBuf.String(), err - } - logLine("S", "354 Start mail input") - - // Header & MIME settings for HTML email - header := make(map[string]string) - header["From"] = sanitizeHeaderValue(cfg.Username) - header["To"] = sanitizeHeaderValue(to) - header["Subject"] = sanitizeHeaderValue(subject) - header["MIME-Version"] = "1.0" - header["Content-Type"] = "text/html; charset=UTF-8" - - message := "" - for k, v := range header { - message += fmt.Sprintf("%s: %s\r\n", k, v) - } - message += "\r\n" + body - - logLine("System", "Sending message body...") - if _, err = w.Write([]byte(message)); err != nil { - _ = w.Close() - logLine("Error", "Writing message body failed: %v", err) - return logBuf.String(), err - } - _ = w.Close() - logLine("S", "250 OK") - - logLine("C", "QUIT") - _ = client.Quit() - logLine("System", "Mail sent successfully!") - - return logBuf.String(), nil + fmt.Fprintf(l.buf, "[%s] %s\n", dir, msg) } + +func (l *bufferLogger) Debugf(e golog.Log) { l.log("Debug", e) } +func (l *bufferLogger) Infof(e golog.Log) { l.log("Info", e) } +func (l *bufferLogger) Warnf(e golog.Log) { l.log("Warn", e) } +func (l *bufferLogger) Errorf(e golog.Log) { l.log("Error", e) } diff --git a/backend/pkg/mail/mail_test.go b/backend/pkg/mail/mail_test.go index 08e16330..4635292a 100644 --- a/backend/pkg/mail/mail_test.go +++ b/backend/pkg/mail/mail_test.go @@ -8,74 +8,105 @@ import ( "context" "net" "net/textproto" + "strings" "testing" ) -func TestSendMailMock(t *testing.T) { - // Start a mock SMTP server +func startMockSMTPServer(t *testing.T) (int, func()) { + t.Helper() l, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { t.Fatalf("failed to start mock smtp server: %v", err) } - defer func() { _ = l.Close() }() port := l.Addr().(*net.TCPAddr).Port go func() { - conn, err := l.Accept() + for { + conn, err := l.Accept() + if err != nil { + return + } + go handleMockSMTPConn(conn) + } + }() + + return port, func() { _ = l.Close() } +} + +func handleMockSMTPConn(conn net.Conn) { + defer func() { _ = conn.Close() }() + + writer := bufio.NewWriter(conn) + reader := bufio.NewReader(conn) + tp := textproto.NewReader(reader) + + // 220 Ready + _, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n") + _ = writer.Flush() + + for { + line, err := tp.ReadLine() if err != nil { return } - defer func() { _ = conn.Close() }() - - writer := bufio.NewWriter(conn) - reader := bufio.NewReader(conn) - tp := textproto.NewReader(reader) - - // 220 Ready - _, _ = writer.WriteString("220 mock.smtp.com SMTP Ready\r\n") - _ = writer.Flush() - - // Read HELO/EHLO - _, _ = tp.ReadLine() - _, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n") - _ = writer.Flush() - - // Read AUTH PLAIN - _, _ = tp.ReadLine() - _, _ = writer.WriteString("235 Authentication successful\r\n") - _ = writer.Flush() - - // Read MAIL FROM - _, _ = tp.ReadLine() - _, _ = writer.WriteString("250 OK\r\n") - _ = writer.Flush() - - // Read RCPT TO - _, _ = tp.ReadLine() - _, _ = writer.WriteString("250 OK\r\n") - _ = writer.Flush() - - // Read DATA - _, _ = tp.ReadLine() - _, _ = writer.WriteString("354 Start mail input\r\n") - _ = writer.Flush() - - // Read body lines until dot - for { - line, err := tp.ReadLine() - if err != nil || line == "." { - break + upper := strings.ToUpper(line) + switch { + case strings.HasPrefix(upper, "EHLO") || strings.HasPrefix(upper, "HELO"): + _, _ = writer.WriteString("250-mock.smtp.com\r\n250 AUTH PLAIN\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "AUTH PLAIN"): + _, _ = writer.WriteString("235 2.7.0 Authentication successful\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "MAIL FROM:"): + _, _ = writer.WriteString("250 2.1.0 Ok\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "RCPT TO:"): + _, _ = writer.WriteString("250 2.1.5 Ok\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "DATA"): + _, _ = writer.WriteString("354 Start mail input; end with .\r\n") + _ = writer.Flush() + for { + dataLine, err := tp.ReadLine() + if err != nil || dataLine == "." { + break + } } + _, _ = writer.WriteString("250 2.0.0 Ok: queued\r\n") + _ = writer.Flush() + case strings.HasPrefix(upper, "QUIT"): + _, _ = writer.WriteString("221 2.0.0 Bye\r\n") + _ = writer.Flush() + return + default: + _, _ = writer.WriteString("250 Ok\r\n") + _ = writer.Flush() } - _, _ = writer.WriteString("250 OK\r\n") - _ = writer.Flush() + } +} - // Read QUIT - _, _ = tp.ReadLine() - _, _ = writer.WriteString("221 Bye\r\n") - _ = writer.Flush() - }() +func TestSendMailMock(t *testing.T) { + port, cleanup := startMockSMTPServer(t) + defer cleanup() + + cfg := Config{ + Host: "127.0.0.1", + Port: port, + Username: "test@example.com", + Password: "password", + FromName: "Wavelet Notifier", + } + + err := SendMail(context.Background(), cfg, "recipient@example.com", "Test Subject", "

Test Body

") + if err != nil { + t.Fatalf("failed to send mail: %v", err) + } +} + +func TestSendMailWithLog(t *testing.T) { + port, cleanup := startMockSMTPServer(t) + defer cleanup() cfg := Config{ Host: "127.0.0.1", @@ -84,28 +115,29 @@ func TestSendMailMock(t *testing.T) { Password: "password", } - err = SendMail(context.Background(), cfg, "recipient@example.com", "Test Subject", "

Test Body

") + logs, err := SendMailWithLog(context.Background(), cfg, "recipient@example.com", "Test Subject", "

Test Log

") if err != nil { - t.Errorf("failed to send mail: %v", err) + t.Fatalf("failed to send mail with log: %v, log output:\n%s", err, logs) + } + + if !strings.Contains(logs, "[System] Connecting to") { + t.Errorf("expected connection log in output, got: %s", logs) + } + if !strings.Contains(logs, "[System] Mail sent successfully!") { + t.Errorf("expected success log in output, got: %s", logs) } } -func TestSanitizeHeaderValue(t *testing.T) { - tests := []struct { - name string - input string - want string - }{ - {"plain", "System Notification", "System Notification"}, - {"crlf stripped", "alert\r\nBcc: attacker@example.com", "alertBcc: attacker@example.com"}, - {"cr stripped", "a\rb", "ab"}, - {"lf stripped", "a\nb", "ab"}, +func TestSendMailInvalidAddress(t *testing.T) { + cfg := Config{ + Host: "127.0.0.1", + Port: 25, + Username: "test@example.com", + Password: "password", } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := sanitizeHeaderValue(tt.input); got != tt.want { - t.Errorf("sanitizeHeaderValue(%q) = %q, want %q", tt.input, got, tt.want) - } - }) + + err := SendMail(context.Background(), cfg, "invalid address with \n newline", "Subject", "Body") + if err == nil { + t.Errorf("expected error for invalid address, got nil") } } diff --git a/backend/pkg/util/format.go b/backend/pkg/util/format.go new file mode 100644 index 00000000..5b18daba --- /dev/null +++ b/backend/pkg/util/format.go @@ -0,0 +1,70 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package util provides shared formatting and string helper functions. +package util + +import ( + "fmt" + "strconv" +) + +const ( + secondsPerYear = 31104000 // 360 days + secondsPerMonth = 2592000 // 30 days + secondsPerDay = 86400 + secondsPerHour = 3600 + secondsPerMinute = 60 +) + +const ( + sizeKB = 1024 + sizeMB = sizeKB * 1024 + sizeGB = sizeMB * 1024 +) + +// Bytes2Size converts a byte count to a human-readable string with unit (B, KB, MB, GB). +func Bytes2Size(num int64) string { + var numStr string + unit := "B" + switch { + case num/int64(sizeGB) >= 1: + numStr = fmt.Sprintf("%.2f", float64(num)/float64(sizeGB)) + unit = "GB" + case num/int64(sizeMB) >= 1: + numStr = strconv.Itoa(int(float64(num) / float64(sizeMB))) + unit = "MB" + case num/int64(sizeKB) >= 1: + numStr = strconv.Itoa(int(float64(num) / float64(sizeKB))) + unit = "KB" + default: + numStr = strconv.FormatInt(num, 10) + } + return numStr + " " + unit +} + +// Seconds2Time converts a number of seconds to a human-readable Chinese duration string. +func Seconds2Time(num int) (time string) { + if num/secondsPerYear > 0 { + time += strconv.Itoa(num/secondsPerYear) + " 年 " + num %= secondsPerYear + } + if num/secondsPerMonth > 0 { + time += strconv.Itoa(num/secondsPerMonth) + " 个月 " + num %= secondsPerMonth + } + if num/secondsPerDay > 0 { + time += strconv.Itoa(num/secondsPerDay) + " 天 " + num %= secondsPerDay + } + if num/secondsPerHour > 0 { + time += strconv.Itoa(num/secondsPerHour) + " 小时 " + num %= secondsPerHour + } + if num/secondsPerMinute > 0 { + time += strconv.Itoa(num/secondsPerMinute) + " 分钟 " + num %= secondsPerMinute + } + time += strconv.Itoa(num) + " 秒" + return +} diff --git a/backend/pkg/util/format_test.go b/backend/pkg/util/format_test.go new file mode 100644 index 00000000..b2fd973f --- /dev/null +++ b/backend/pkg/util/format_test.go @@ -0,0 +1,52 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import ( + "testing" +) + +func TestBytes2Size(t *testing.T) { + tests := []struct { + input int64 + expected string + }{ + {0, "0 B"}, + {500, "500 B"}, + {1023, "1023 B"}, + {1024, "1 KB"}, + {2048, "2 KB"}, + {1024 * 1024, "1 MB"}, + {1024 * 1024 * 1024, "1.00 GB"}, + {1024 * 1024 * 1024 * 2, "2.00 GB"}, + } + + for _, tt := range tests { + result := Bytes2Size(tt.input) + if result != tt.expected { + t.Errorf("Bytes2Size(%d) = %q, expected %q", tt.input, result, tt.expected) + } + } +} + +func TestSeconds2Time(t *testing.T) { + tests := []struct { + input int + expected string + }{ + {0, "0 秒"}, + {30, "30 秒"}, + {60, "1 分钟 0 秒"}, + {125, "2 分钟 5 秒"}, + {3600, "1 小时 0 秒"}, + {86400, "1 天 0 秒"}, + } + + for _, tt := range tests { + result := Seconds2Time(tt.input) + if result != tt.expected { + t.Errorf("Seconds2Time(%d) = %q, expected %q", tt.input, result, tt.expected) + } + } +} diff --git a/backend/pkg/util/network.go b/backend/pkg/util/network.go new file mode 100644 index 00000000..d7b864d5 --- /dev/null +++ b/backend/pkg/util/network.go @@ -0,0 +1,45 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import ( + "log/slog" + "net" +) + +// GetIP returns the first private IPv4 address found on the local network interfaces. +func GetIP() (ip string) { + ips, err := net.InterfaceAddrs() + if err != nil { + slog.Error("get interface addresses failed", "error", err) + return ip + } + + for _, a := range ips { + if candidate, ok := privateIPv4FromAddr(a); ok { + return candidate + } + } + return +} + +func privateIPv4FromAddr(addr net.Addr) (string, bool) { + ipNet, ok := addr.(*net.IPNet) + if !ok || ipNet.IP.IsLoopback() || ipNet.IP.To4() == nil { + return "", false + } + ip := ipNet.IP.String() + if isPrivateIPv4(ip) { + return ip, true + } + return "", false +} + +func isPrivateIPv4(ip string) bool { + parsedIP := net.ParseIP(ip) + if parsedIP == nil { + return false + } + return parsedIP.IsPrivate() +} diff --git a/backend/pkg/util/network_test.go b/backend/pkg/util/network_test.go new file mode 100644 index 00000000..62a9bd21 --- /dev/null +++ b/backend/pkg/util/network_test.go @@ -0,0 +1,36 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import ( + "testing" +) + +func TestIsPrivateIPv4(t *testing.T) { + tests := []struct { + ip string + expected bool + }{ + {"127.0.0.1", false}, // Loopback is not in RFC 1918 private range + {"10.0.0.1", true}, + {"172.16.0.1", true}, + {"192.168.1.1", true}, + {"8.8.8.8", false}, + {"invalid-ip", false}, + } + + for _, tt := range tests { + result := isPrivateIPv4(tt.ip) + if result != tt.expected { + t.Errorf("isPrivateIPv4(%q) = %v, expected %v", tt.ip, result, tt.expected) + } + } +} + +func TestGetIP(t *testing.T) { + ip := GetIP() + // GetIP should return empty if no private IPv4 address is configured, or a valid IP. + // We just ensure it doesn't panic. + t.Logf("GetIP returned: %q", ip) +} diff --git a/backend/pkg/util/slice.go b/backend/pkg/util/slice.go new file mode 100644 index 00000000..74eb7139 --- /dev/null +++ b/backend/pkg/util/slice.go @@ -0,0 +1,79 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import ( + "sort" + "strings" + "time" +) + +// Unique returns a new slice containing only the unique elements of the input slice, +// preserving their original order. +func Unique[T comparable](slice []T) []T { + if slice == nil { + return nil + } + seen := make(map[T]struct{}) + result := make([]T, 0) + for _, item := range slice { + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + result = append(result, item) + } + return result +} + +// UniqueAndCleanStringSlice trims spaces, removes empty elements, and returns only the unique elements +// of the input string slice. It preserves order and returns nil if the resulting slice is empty. +func UniqueAndCleanStringSlice(slice []string) []string { + if slice == nil { + return nil + } + seen := make(map[string]struct{}) + result := make([]string, 0) + for _, item := range slice { + trimmed := strings.TrimSpace(item) + if trimmed == "" { + continue + } + if _, ok := seen[trimmed]; ok { + continue + } + seen[trimmed] = struct{}{} + result = append(result, trimmed) + } + if len(result) == 0 { + return nil + } + return result +} + +// IdentifiableTimeRecord represents a database record that has a unique ID and a primary timestamp field. +type IdentifiableTimeRecord interface { + GetID() uint + GetTime() time.Time +} + +// SortAndLimitRecords sorts a slice of IdentifiableTimeRecord descendingly by their timestamp (and ID as a tie-breaker), +// and limits the slice to the specified size if limit > 0. +func SortAndLimitRecords[T IdentifiableTimeRecord](rows []T, limit int) []T { + if len(rows) == 0 { + return rows + } + sort.Slice(rows, func(i, j int) bool { + ti := rows[i].GetTime() + tj := rows[j].GetTime() + if ti.Equal(tj) { + return rows[i].GetID() > rows[j].GetID() + } + return ti.After(tj) + }) + if limit > 0 && len(rows) > limit { + rows = rows[:limit] + } + return rows +} diff --git a/backend/pkg/util/string.go b/backend/pkg/util/string.go new file mode 100644 index 00000000..9ea84e8d --- /dev/null +++ b/backend/pkg/util/string.go @@ -0,0 +1,15 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import "strings" + +// TrimStringFields trims leading and trailing spaces from all provided string pointers. +func TrimStringFields(fields ...*string) { + for _, f := range fields { + if f != nil { + *f = strings.TrimSpace(*f) + } + } +} diff --git a/backend/pkg/util/value.go b/backend/pkg/util/value.go new file mode 100644 index 00000000..6473bd53 --- /dev/null +++ b/backend/pkg/util/value.go @@ -0,0 +1,22 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import ( + "fmt" + "strconv" +) + +// Interface2String converts a string, int, or float64 value to its string representation. +func Interface2String(inter any) string { + switch v := inter.(type) { + case string: + return v + case int: + return strconv.Itoa(v) + case float64: + return fmt.Sprintf("%f", v) + } + return "Not Implemented" +} diff --git a/backend/pkg/util/version.go b/backend/pkg/util/version.go new file mode 100644 index 00000000..96118b06 --- /dev/null +++ b/backend/pkg/util/version.go @@ -0,0 +1,129 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import ( + "strconv" + "strings" +) + +const gitDescribeMinIdentifiers = 2 + +// VersionInfo holds the parsed components of a semantic version string. +type VersionInfo struct { + Valid bool + IsDev bool + Numbers []int + Prerelease []string + GitDescribeDistance int + GitDescribeTail []string +} + +// ParseVersionInfo parses a version string into a structured VersionInfo. +func ParseVersionInfo(version string) VersionInfo { + normalized := strings.TrimSpace(strings.TrimPrefix(version, "v")) + if normalized == "" || normalized == "dev" { + return VersionInfo{IsDev: strings.EqualFold(normalized, "dev")} + } + base := normalized + prerelease := "" + if separator := strings.IndexRune(normalized, '-'); separator >= 0 { + base = normalized[:separator] + prerelease = normalized[separator+1:] + } + + segments := strings.Split(base, ".") + parts := make([]int, 0, len(segments)) + for _, segment := range segments { + segment = strings.TrimSpace(segment) + if segment == "" { + parts = append(parts, 0) + continue + } + + numeric := strings.Builder{} + for _, r := range segment { + if r < '0' || r > '9' { + break + } + numeric.WriteRune(r) + } + if numeric.Len() == 0 { + parts = append(parts, 0) + continue + } + value, err := strconv.Atoi(numeric.String()) + if err != nil { + return VersionInfo{} + } + parts = append(parts, value) + } + info := VersionInfo{Valid: len(parts) > 0, Numbers: parts} + if prerelease != "" { + identifiers := splitPrereleaseIdentifiers(prerelease) + if distance, tail, ok := parseGitDescribeIdentifiers(identifiers); ok { + info.GitDescribeDistance = distance + info.GitDescribeTail = tail + } else { + info.Prerelease = identifiers + } + } + return info +} + +func parseGitDescribeIdentifiers(identifiers []string) (int, []string, bool) { + if len(identifiers) < gitDescribeMinIdentifiers { + return 0, nil, false + } + distance, err := strconv.Atoi(strings.TrimSpace(identifiers[0])) + if err != nil || distance <= 0 { + return 0, nil, false + } + commitToken := strings.TrimSpace(identifiers[1]) + if commitToken == "" || !strings.HasPrefix(strings.ToLower(commitToken), "g") { + return 0, nil, false + } + return distance, identifiers[1:], true +} + +func splitPrereleaseIdentifiers(value string) []string { + parts := strings.FieldsFunc(strings.TrimSpace(value), func(r rune) bool { + return r == '.' || r == '-' + }) + filtered := make([]string, 0, len(parts)) + for _, part := range parts { + part = strings.TrimSpace(part) + if part != "" { + filtered = append(filtered, part) + } + } + return filtered +} + +// CompareVersions compares two version strings. +// Returns -1 if left < right, 1 if left > right, and 0 if they are equal. +func CompareVersions(local, remote string) int { + left := ParseVersionInfo(local) + right := ParseVersionInfo(remote) + if left.IsDev { + if right.Valid { + return -1 + } + return 0 + } + if !left.Valid || !right.Valid { + return 0 + } + + if result := compareVersionNumbers(left, right); result != 0 { + return result + } + if result := compareGitDescribeDistance(left, right); result != 0 { + return result + } + if left.GitDescribeDistance > 0 || right.GitDescribeDistance > 0 { + return compareGitDescribeTails(left, right) + } + return comparePrereleaseIdentifiers(left, right) +} diff --git a/backend/pkg/util/version_compare.go b/backend/pkg/util/version_compare.go new file mode 100644 index 00000000..c7e3d014 --- /dev/null +++ b/backend/pkg/util/version_compare.go @@ -0,0 +1,108 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package util + +import "strconv" + +func compareVersionNumbers(left, right VersionInfo) int { + maxLen := max(len(right.Numbers), len(left.Numbers)) + for index := range maxLen { + leftValue := 0 + rightValue := 0 + if index < len(left.Numbers) { + leftValue = left.Numbers[index] + } + if index < len(right.Numbers) { + rightValue = right.Numbers[index] + } + if leftValue < rightValue { + return -1 + } + if leftValue > rightValue { + return 1 + } + } + return 0 +} + +func compareGitDescribeDistance(left, right VersionInfo) int { + if left.GitDescribeDistance == right.GitDescribeDistance { + return 0 + } + if left.GitDescribeDistance < right.GitDescribeDistance { + return -1 + } + return 1 +} + +func compareGitDescribeTails(left, right VersionInfo) int { + maxLen := max(len(right.GitDescribeTail), len(left.GitDescribeTail)) + for index := range maxLen { + if index >= len(left.GitDescribeTail) { + return -1 + } + if index >= len(right.GitDescribeTail) { + return 1 + } + if left.GitDescribeTail[index] < right.GitDescribeTail[index] { + return -1 + } + if left.GitDescribeTail[index] > right.GitDescribeTail[index] { + return 1 + } + } + return 0 +} + +func comparePrereleaseIdentifiers(left, right VersionInfo) int { + if len(left.Prerelease) == 0 && len(right.Prerelease) == 0 { + return 0 + } + if len(left.Prerelease) == 0 { + return 1 + } + if len(right.Prerelease) == 0 { + return -1 + } + + maxLen := max(len(right.Prerelease), len(left.Prerelease)) + for index := range maxLen { + if index >= len(left.Prerelease) { + return -1 + } + if index >= len(right.Prerelease) { + return 1 + } + if result := comparePrereleasePart(left.Prerelease[index], right.Prerelease[index]); result != 0 { + return result + } + } + return 0 +} + +func comparePrereleasePart(leftPart, rightPart string) int { + leftNumber, leftErr := strconv.Atoi(leftPart) + rightNumber, rightErr := strconv.Atoi(rightPart) + switch { + case leftErr == nil && rightErr == nil: + if leftNumber < rightNumber { + return -1 + } + if leftNumber > rightNumber { + return 1 + } + case leftErr == nil: + return -1 + case rightErr == nil: + return 1 + default: + if leftPart < rightPart { + return -1 + } + if leftPart > rightPart { + return 1 + } + } + return 0 +} diff --git a/backend/plugins/domain/admin/handler/config.go b/backend/plugins/domain/admin/handler/config.go index a431c874..abbf60ae 100644 --- a/backend/plugins/domain/admin/handler/config.go +++ b/backend/plugins/domain/admin/handler/config.go @@ -16,7 +16,7 @@ import ( // GetPublicConfig 获取公共配置 // @Summary 获取公共配置 -// @Description 返回系统配置表中 visibility 为 1 的配置键值集合 +// @Description 返回系统配置表中 visibility 为 1 的扁平键值集合(如 cap_login_enabled) // @Tags config // @Accept json // @Produce json diff --git a/backend/plugins/domain/admin/handler/logs.go b/backend/plugins/domain/admin/handler/logs.go index aac63a14..3bb2083c 100644 --- a/backend/plugins/domain/admin/handler/logs.go +++ b/backend/plugins/domain/admin/handler/logs.go @@ -14,6 +14,7 @@ import ( "errors" "net/http" "strconv" + "strings" "github.com/gin-gonic/gin" "github.com/gorilla/websocket" @@ -179,11 +180,31 @@ func GetLogsAnalytics(c *gin.Context) { func getUpgrader() *websocket.Upgrader { return &websocket.Upgrader{ CheckOrigin: func(r *http.Request) bool { - return service.IsAllowedLogOrigin(r.Context(), r.Header.Get("Origin"), r.Host) + return service.IsAllowedLogOrigin( + r.Context(), + r.Header.Get("Origin"), + r.Host, + forwardedHosts(r)..., + ) }, } } +func forwardedHosts(r *http.Request) []string { + raw := r.Header.Get("X-Forwarded-Host") + if raw == "" { + return nil + } + parts := strings.Split(raw, ",") + hosts := make([]string, 0, len(parts)) + for _, part := range parts { + if h := strings.TrimSpace(part); h != "" { + hosts = append(hosts, h) + } + } + return hosts +} + // errNegativeParam 表示查询参数解析出了负数。 var errNegativeParam = errors.New("parameter must not be negative") diff --git a/backend/plugins/domain/admin/handler/tasks.go b/backend/plugins/domain/admin/handler/tasks.go index 438cae3f..2e396c67 100644 --- a/backend/plugins/domain/admin/handler/tasks.go +++ b/backend/plugins/domain/admin/handler/tasks.go @@ -48,7 +48,7 @@ func abortTaskLogicError(c *gin.Context, err error) bool { // @Failure 403 {object} response.Any "无管理员权限" // @Router /api/v1/admin/tasks/types [get] func ListTaskTypes(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(service.ListTaskTypes())) + c.JSON(http.StatusOK, response.OK(service.ListTaskTypes(c.Request.Context()))) } // DispatchTask 下发任务 diff --git a/backend/plugins/domain/admin/plugin.go b/backend/plugins/domain/admin/plugin.go index ecbb30b4..d5f6b3cf 100644 --- a/backend/plugins/domain/admin/plugin.go +++ b/backend/plugins/domain/admin/plugin.go @@ -12,7 +12,6 @@ import ( "Wavelet/plugins/domain/admin/handler" "Wavelet/plugins/domain/admin/model" "Wavelet/plugins/domain/admin/service" - "context" "embed" "reflect" @@ -22,6 +21,12 @@ import ( // SystemConfig aliases model.SystemConfig for external compatibility. type SystemConfig = model.SystemConfig +// TaskExecution aliases model.TaskExecution for external compatibility. +type TaskExecution = model.TaskExecution + +// Schedule aliases model.Schedule for external compatibility. +type Schedule = model.Schedule + //go:embed migrations/*/*.sql var adminMigrations embed.FS @@ -85,57 +90,16 @@ func (p *Plugin) Apply(ctx *core.Context) error { _ = ctx.Config().Bind("clickhouse", &chCfg) service.SetClickHouseConfig(chCfg) - // 0. Bind Services reactively - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - service.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - service.SetDBService(db) - }) - } - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - service.SetCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - service.SetCacheService(cache) - }) - } - if user, err := core.Inject[contracts.UserService](ctx); err == nil && user != nil { - service.SetUserService(user) - } else { - core.When[contracts.UserService](ctx, func(user contracts.UserService) { - service.SetUserService(user) - }) - } - if auth, err := core.Inject[contracts.AuthService](ctx); err == nil && auth != nil { - service.SetAuthService(auth) - } else { - core.When[contracts.AuthService](ctx, func(auth contracts.AuthService) { - service.SetAuthService(auth) - }) - } - if task, err := core.Inject[contracts.TaskService](ctx); err == nil && task != nil { - service.SetTaskService(task) - } else { - core.When[contracts.TaskService](ctx, func(task contracts.TaskService) { - service.SetTaskService(task) - }) - } - if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil { - service.SetStorageService(storage) - } else { - core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) { - service.SetStorageService(storage) - }) - } - if rc, err := core.Inject[contracts.RiskControlService](ctx); err == nil && rc != nil { - service.SetRiskControlService(rc) - } else { - core.When[contracts.RiskControlService](ctx, func(rc contracts.RiskControlService) { - service.SetRiskControlService(rc) - }) - } + core.Bind[contracts.DBService](ctx, service.SetDBService) + core.Bind[contracts.CacheService](ctx, service.SetCacheService) + core.Bind[contracts.UserService](ctx, service.SetUserService) + core.Bind[contracts.AuthService](ctx, service.SetAuthService) + core.Bind[contracts.TaskService](ctx, service.SetTaskService) + core.Bind[contracts.StorageService](ctx, service.SetStorageService) + core.Bind[contracts.RiskControlService](ctx, service.SetRiskControlService) service.SetEventEmitter(ctx.Events().Emit) + core.Provide[contracts.PublicConfigProvider](ctx, service.PublicConfigAdapter{}) + core.Provide[contracts.SystemConfigService](ctx, service.SystemConfigServiceImpl{}) ctx.OnDispose(func() error { service.ResetServices() @@ -170,32 +134,22 @@ func (p *Plugin) Apply(ctx *core.Context) error { adminRouter := ctx.Router().Group("/api/v1/admin", loginMW, adminMW) handler.RegisterRoutes(adminRouter) + // Register robots.txt public route + ctx.Router().GET("/robots.txt", handler.GetRobotsTXT) + ctx.Router().RegisterWhitelist("/robots.txt") + // 2. Register Background Tasks - logSwitchHandler := &service.LogDBSwitchHandler{} - ctx.Task().Register(service.LogDBSwitchTask, func(c context.Context, payload []byte) error { - _, err := logSwitchHandler.Execute(c, payload) - return err - }, extpoints.WithTaskMeta(service.LogDBSwitchMeta)) + const defaultCleanupRetry = 3 + ctx.Task().Register(service.LogDBSwitchTask, &service.LogDBSwitchHandler{}, extpoints.WithTaskMeta(service.LogDBSwitchMeta)) + ctx.Task().Register(service.SystemCleanupTask, &service.SystemCleanupHandler{}, extpoints.WithTaskMeta(service.SystemCleanupMeta), extpoints.WithTaskRetry(defaultCleanupRetry)) - ctx.Task().Register("admin:system_cleanup", func(_ context.Context, _ []byte) error { - return nil - }, - extpoints.WithTaskType("system_cleanup"), - extpoints.WithTaskName("系统垃圾清理"), - extpoints.WithTaskDescription("定期清理未使用上传文件、历史推送记录和过期任务执行日志"), - extpoints.WithTaskCategory("maintenance"), - extpoints.WithTaskRetry(1), - extpoints.WithTaskQueue("default"), - extpoints.WithTaskRetryable(true), - ) + // 2.1 Register Cron Schedule + ctx.Schedule().RegisterCron("0 3 * * *", service.SystemCleanupTask, nil) - // 3. Register Cron Schedules - ctx.Schedule().RegisterCron("0 4 * * *", "admin:system_cleanup", map[string]string{"type": "daily"}) - - // 4. Register Settings Schemas + // 3. Register Settings Schemas ctx.Settings().Register(extpoints.SettingSchema{ Key: "admin.system_cleanup_cron", - Default: "0 4 * * *", + Default: "0 3 * * *", Description: "Cron expression for nightly system logs and expired tokens cleanup", Type: "string", Category: "maintenance", diff --git a/backend/plugins/domain/admin/plugin_test.go b/backend/plugins/domain/admin/plugin_test.go index b0509906..dab839a3 100644 --- a/backend/plugins/domain/admin/plugin_test.go +++ b/backend/plugins/domain/admin/plugin_test.go @@ -5,6 +5,7 @@ package admin_test import ( "Wavelet/core" + "Wavelet/core/contracts" "Wavelet/plugins/domain/admin" "context" "testing" @@ -23,20 +24,29 @@ func TestAdminPluginUnit(t *testing.T) { // Verify routes routes := ctx.Router().Routes() assert.NotEmpty(t, routes) + var hasRobots bool + for _, r := range routes { + if r.Path == "/robots.txt" && r.Method == "GET" { + hasRobots = true + break + } + } + assert.True(t, hasRobots, "admin plugin must register /robots.txt") // Verify tasks - _, ok := ctx.Tasks().Get("admin:system_cleanup") + _, ok := ctx.Tasks().Get("logs:db_switch") require.True(t, ok) - - // Verify schedules - sched, ok := ctx.Schedules().Get("admin:system_cleanup") + _, ok = ctx.Tasks().Get("system:cleanup") require.True(t, ok) - assert.Equal(t, "0 4 * * *", sched.Spec) // Verify settings setting, ok := ctx.Settings().Get("admin.system_cleanup_cron") require.True(t, ok) - assert.Equal(t, "0 4 * * *", setting.Default) + assert.Equal(t, "0 3 * * *", setting.Default) + + provider, err := core.Inject[contracts.PublicConfigProvider](ctx) + require.NoError(t, err) + require.NotNil(t, provider) } func TestAdminMigrationsIncludeTaskExecutionsAndSchedules(t *testing.T) { diff --git a/backend/plugins/domain/admin/repository/repository.go b/backend/plugins/domain/admin/repository/repository.go index 6118614a..c0c2331a 100644 --- a/backend/plugins/domain/admin/repository/repository.go +++ b/backend/plugins/domain/admin/repository/repository.go @@ -5,6 +5,7 @@ package repository import ( + "Wavelet/core" "Wavelet/core/contracts" "Wavelet/pkg/cache/ram" "Wavelet/pkg/logger" @@ -56,6 +57,9 @@ func ResetServices() { // GetDB returns the GORM DB instance bound to the context if available. func GetDB(ctx context.Context) *gorm.DB { + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) + } repoMu.RLock() defer repoMu.RUnlock() if dbService == nil { @@ -65,7 +69,10 @@ func GetDB(ctx context.Context) *gorm.DB { } // GetCache returns the unified CacheService instance. -func GetCache(_ context.Context) contracts.CacheService { +func GetCache(ctx context.Context) contracts.CacheService { + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s + } repoMu.RLock() defer repoMu.RUnlock() return cacheService diff --git a/backend/plugins/domain/admin/service/cleanup.go b/backend/plugins/domain/admin/service/cleanup.go new file mode 100644 index 00000000..18de5ca1 --- /dev/null +++ b/backend/plugins/domain/admin/service/cleanup.go @@ -0,0 +1,70 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/plugins/domain/admin/model" + "context" + "errors" + "fmt" + "time" +) + +const ( + // SystemCleanupTask 系统定期垃圾清理任务标识 + SystemCleanupTask = "system:cleanup" + // TaskTypeSystemCleanup 系统定期垃圾清理管理类型 + TaskTypeSystemCleanup = "system_cleanup" + taskQueueDefault = "default" +) + +// SystemCleanupMeta describes the system-wide cleanup task metadata. +var SystemCleanupMeta = contracts.TaskMetaDTO{ + Type: TaskTypeSystemCleanup, + AsynqTask: SystemCleanupTask, + Name: "系统垃圾清理", + DisplayName: "系统垃圾清理", + Description: "定期清理过期任务执行记录,并通过领域事件广播触发各业务域自治清理(临时文件、历史推送等)", + Category: "maintenance", + SupportsTime: false, + MaxRetry: 3, + Queue: taskQueueDefault, + Retryable: true, +} + +// SystemCleanupHandler handles the system-wide garbage cleanup task. +type SystemCleanupHandler struct{} + +// Execute executes system cleanup: clears old task executions and emits EventTopicSystemCleanup. +func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*contracts.TaskResultDTO, error) { + db := GetDB(ctx) + if db == nil { + return nil, errors.New("database service not available") + } + + // 1. 清理自身域(admin 域)的过期任务执行记录(7 天前) + var deletedExecutions int64 + sevenDaysAgo := time.Now().Add(-7 * 24 * time.Hour) + res := db.Where("created_at < ?", sevenDaysAgo).Delete(&model.TaskExecution{}) + if err := res.Error; err != nil { + logger.WarnF(ctx, "清理过期任务执行日志失败: %v", err) + } else { + deletedExecutions = res.RowsAffected + logger.InfoF(ctx, "已清理 7 天前任务执行日志,共 %d 条", deletedExecutions) + } + + // 2. 广播 EventTopicSystemCleanup 领域事件,由各业务域插件(upload, msg_gateway, user 等)自治执行各自的清理逻辑 + nowStr := time.Now().Format(time.RFC3339) + if err := EmitEvent(ctx, contracts.EventTopicSystemCleanup, contracts.SystemCleanupEvent{ + TriggeredAt: nowStr, + }); err != nil { + logger.WarnF(ctx, "广播系统清理领域事件失败: %v", err) + } + + msg := fmt.Sprintf("系统垃圾清理完成,已清理过期任务执行日志 %d 条,并已广播领域清理事件", deletedExecutions) + logger.InfoF(ctx, "%s", msg) + return &contracts.TaskResultDTO{Message: msg}, nil +} diff --git a/backend/plugins/domain/admin/service/cleanup_test.go b/backend/plugins/domain/admin/service/cleanup_test.go new file mode 100644 index 00000000..787536d9 --- /dev/null +++ b/backend/plugins/domain/admin/service/cleanup_test.go @@ -0,0 +1,78 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service_test + +import ( + "Wavelet/core/contracts" + "Wavelet/plugins/domain/admin/model" + "Wavelet/plugins/domain/admin/service" + "context" + "sync/atomic" + "testing" + "time" + + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +func TestSystemCleanupHandler_Execute(t *testing.T) { + sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) + require.NoError(t, err) + + require.NoError(t, sqliteDB.AutoMigrate(&model.TaskExecution{})) + + service.SetDBService(&testDBService{db: sqliteDB}) + defer service.ResetServices() + + now := time.Now() + oldTime := now.Add(-10 * 24 * time.Hour) + recentTime := now.Add(-1 * time.Hour) + + // Seed old execution (should be cleaned) + oldExec := model.TaskExecution{ + ID: 1, + TaskID: "task-old", + TaskType: "sample_task", + Status: "success", + CreatedAt: oldTime, + } + require.NoError(t, sqliteDB.Create(&oldExec).Error) + + // Seed recent execution (should remain) + recentExec := model.TaskExecution{ + ID: 2, + TaskID: "task-recent", + TaskType: "sample_task", + Status: "success", + CreatedAt: recentTime, + } + require.NoError(t, sqliteDB.Create(&recentExec).Error) + + var eventFired atomic.Bool + service.SetEventEmitter(func(ctx context.Context, topic string, payload any) error { + if topic == contracts.EventTopicSystemCleanup { + eventFired.Store(true) + } + return nil + }) + + handler := &service.SystemCleanupHandler{} + res, err := handler.Execute(context.Background(), nil) + require.NoError(t, err) + require.NotNil(t, res) + assert.Contains(t, res.Message, "系统垃圾清理完成") + + // Verify old execution is deleted and recent remains + var count int64 + sqliteDB.Model(&model.TaskExecution{}).Count(&count) + assert.Equal(t, int64(1), count) + + var remaining model.TaskExecution + sqliteDB.First(&remaining) + assert.Equal(t, uint64(2), remaining.ID) + + assert.True(t, eventFired.Load(), "EventTopicSystemCleanup must be emitted") +} diff --git a/backend/plugins/domain/admin/service/config.go b/backend/plugins/domain/admin/service/config.go index f453536b..6590ae1d 100644 --- a/backend/plugins/domain/admin/service/config.go +++ b/backend/plugins/domain/admin/service/config.go @@ -18,6 +18,102 @@ import ( const maskedConfigValue = "******" +// PublicConfigAdapter exposes visibility=1 system configs as PublicConfigProvider. +type PublicConfigAdapter struct{} + +// PublicConfig returns the unauthenticated public config map. +func (PublicConfigAdapter) PublicConfig(ctx context.Context) (map[string]string, error) { + return PublicSystemConfigs(ctx) +} + +// SystemConfigServiceImpl implements contracts.SystemConfigService. +type SystemConfigServiceImpl struct{} + +// GetByKey retrieves a system config by its unique key. +func (SystemConfigServiceImpl) GetByKey(ctx context.Context, key string) (contracts.SystemConfigDTO, error) { + cfg, err := repository.GetSystemConfigByKey(ctx, key) + if err != nil { + return contracts.SystemConfigDTO{}, err + } + return toSystemConfigDTO(cfg), nil +} + +// ListByKeys retrieves multiple system configs by their keys. +func (SystemConfigServiceImpl) ListByKeys(ctx context.Context, keys []string) (map[string]contracts.SystemConfigDTO, error) { + cfgs, err := repository.ListSystemConfigsByKeys(ctx, keys) + if err != nil { + return nil, err + } + res := make(map[string]contracts.SystemConfigDTO, len(cfgs)) + for k, v := range cfgs { + res[k] = toSystemConfigDTO(v) + } + return res, nil +} + +// ListVisible returns all user-visible system configs. +func (SystemConfigServiceImpl) ListVisible(ctx context.Context) ([]contracts.SystemConfigDTO, error) { + cfgs, err := repository.ListVisibleSystemConfigs(ctx) + if err != nil { + return nil, err + } + res := make([]contracts.SystemConfigDTO, len(cfgs)) + for i, v := range cfgs { + res[i] = toSystemConfigDTO(v) + } + return res, nil +} + +// ListByType returns all system configs belonging to a specific configuration type. +func (SystemConfigServiceImpl) ListByType(ctx context.Context, configType string) ([]contracts.SystemConfigDTO, error) { + cfgs, err := repository.ListAdminSystemConfigs(ctx, configType) + if err != nil { + return nil, err + } + res := make([]contracts.SystemConfigDTO, len(cfgs)) + for i, v := range cfgs { + res[i] = toSystemConfigDTO(v) + } + return res, nil +} + +// GetIntByKey retrieves an integer system config value. +func (SystemConfigServiceImpl) GetIntByKey(ctx context.Context, key string) (int, error) { + return repository.GetIntByKey(ctx, key) +} + +// GetBoolByKey retrieves a boolean system config value. +func (SystemConfigServiceImpl) GetBoolByKey(ctx context.Context, key string) (bool, error) { + return repository.GetBoolByKey(ctx, key) +} + +// SaveOrUpdate persists or updates a system config key-value pair. +func (SystemConfigServiceImpl) SaveOrUpdate(ctx context.Context, key, value string) error { + return repository.SaveOrUpdateSystemConfig(ctx, key, value) +} + +// InvalidateCache evicts the cache entry for the specified key. +func (SystemConfigServiceImpl) InvalidateCache(ctx context.Context, key string) error { + return repository.InvalidateSystemConfigCache(ctx, key) +} + +// InvalidateAllCaches purges all system configuration cache entries. +func (SystemConfigServiceImpl) InvalidateAllCaches(ctx context.Context) error { + return repository.InvalidateAllSystemConfigCaches(ctx) +} + +func toSystemConfigDTO(c model.SystemConfig) contracts.SystemConfigDTO { + return contracts.SystemConfigDTO{ + Key: c.Key, + Value: c.Value, + Type: c.Type, + Visibility: c.Visibility, + Description: c.Description, + UpdatedAt: c.UpdatedAt, + CreatedAt: c.CreatedAt, + } +} + // PublicSystemConfigs returns the key/value map exposed to unauthenticated clients. func PublicSystemConfigs(ctx context.Context) (map[string]string, error) { configs, err := repository.ListVisibleSystemConfigs(ctx) diff --git a/backend/plugins/domain/admin/service/log.go b/backend/plugins/domain/admin/service/log.go index 82a6dd84..f138613f 100644 --- a/backend/plugins/domain/admin/service/log.go +++ b/backend/plugins/domain/admin/service/log.go @@ -11,6 +11,7 @@ import ( "Wavelet/plugins/domain/admin/repository" "context" "fmt" + "net" "net/url" "strings" "time" @@ -47,18 +48,20 @@ func RobotsTxtBody(ctx context.Context) string { } // IsAllowedLogOrigin reports whether a WebSocket handshake origin may subscribe to logs. -func IsAllowedLogOrigin(ctx context.Context, origin, host string) bool { +// extraHosts are reverse-proxy hosts such as X-Forwarded-Host (the browser origin +// when Next.js rewrites /api to the backend). +func IsAllowedLogOrigin(ctx context.Context, origin, host string, extraHosts ...string) bool { if origin == "" { return true } - // 1. 同源检查 (Same-origin check) u, err := url.Parse(origin) - if err == nil && strings.EqualFold(u.Host, host) { - return true + if err == nil { + if originMatchesHost(u.Host, host, extraHosts...) { + return true + } } - // 2. 检查配置的允许跨域 Origin (Check allowed origins in system config) sc, cfgErr := repository.GetSystemConfigByKey(ctx, model.ConfigKeyServerAddress) if cfgErr != nil || sc.Value == "" { return false @@ -73,9 +76,44 @@ func IsAllowedLogOrigin(ctx context.Context, origin, host string) bool { return false } +func originMatchesHost(originHost, host string, extraHosts ...string) bool { + if originHost == "" { + return false + } + if hostMatches(originHost, host) { + return true + } + for _, extra := range extraHosts { + if hostMatches(originHost, extra) { + return true + } + } + return false +} + +func hostMatches(originHost, candidate string) bool { + candidate = strings.TrimSpace(candidate) + if candidate == "" { + return false + } + if strings.EqualFold(originHost, candidate) { + return true + } + // Reverse-proxy / local Next rewrite: Origin is :3000, backend Host is :8000. + return strings.EqualFold(hostName(originHost), hostName(candidate)) +} + +func hostName(hostport string) string { + h, _, err := net.SplitHostPort(hostport) + if err != nil { + return hostport + } + return h +} + // AccessLogs queries the analytical access log store and decorates rows with user names. func AccessLogs(ctx context.Context, q model.AccessLogQuery) (model.AccessLogsResponse, error) { - rc := GetRiskControlService() + rc := GetRiskControlService(ctx) if rc == nil { return model.AccessLogsResponse{}, errs.ErrLogStoreUnavailable } @@ -117,7 +155,7 @@ func AccessLogs(ctx context.Context, q model.AccessLogQuery) (model.AccessLogsRe // AccessLogAnalytics aggregates the daily trend of the access log store. func AccessLogAnalytics(ctx context.Context) (model.LogsAnalyticsResponse, error) { - rc := GetRiskControlService() + rc := GetRiskControlService(ctx) if rc == nil { return model.LogsAnalyticsResponse{}, errs.ErrLogStoreUnavailable } diff --git a/backend/plugins/domain/admin/service/log_origin_test.go b/backend/plugins/domain/admin/service/log_origin_test.go new file mode 100644 index 00000000..092310c2 --- /dev/null +++ b/backend/plugins/domain/admin/service/log_origin_test.go @@ -0,0 +1,55 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service_test + +import ( + "Wavelet/plugins/domain/admin/service" + "context" + "testing" +) + +func TestIsAllowedLogOrigin(t *testing.T) { + ctx := context.Background() + + tests := []struct { + name string + origin string + host string + extraHosts []string + want bool + }{ + {name: "empty origin", origin: "", host: "localhost:8000", want: true}, + {name: "same host", origin: "http://localhost:8000", host: "localhost:8000", want: true}, + {name: "same host different case", origin: "http://LocalHost:8000", host: "localhost:8000", want: true}, + { + name: "next rewrite different port same hostname", + origin: "http://localhost:3000", + host: "localhost:8000", + want: true, + }, + { + name: "x-forwarded-host matches origin", + origin: "http://localhost:3000", + host: "backend:8080", + extraHosts: []string{"localhost:3000"}, + want: true, + }, + { + name: "unrelated origin", + origin: "https://evil.example", + host: "localhost:8000", + want: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := service.IsAllowedLogOrigin(ctx, tt.origin, tt.host, tt.extraHosts...) + if got != tt.want { + t.Errorf("IsAllowedLogOrigin(%q, %q, %v) = %v, want %v", + tt.origin, tt.host, tt.extraHosts, got, tt.want) + } + }) + } +} diff --git a/backend/plugins/domain/admin/service/log_switch.go b/backend/plugins/domain/admin/service/log_switch.go index 792eb95f..a6ae1091 100644 --- a/backend/plugins/domain/admin/service/log_switch.go +++ b/backend/plugins/domain/admin/service/log_switch.go @@ -109,7 +109,7 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*cont return nil, err } - taskSvc := GetTaskService() + taskSvc := GetTaskService(ctx) if taskSvc != nil { taskSvc.AppendLog(ctx, "开始切换日志数据库:%s -> %s", source, p.Target) } @@ -123,7 +123,7 @@ func (h *LogDBSwitchHandler) Execute(ctx context.Context, payload []byte) (*cont } }() - rc := GetRiskControlService() + rc := GetRiskControlService(ctx) if rc != nil { if err := rc.SwitchLogEngine(ctx, p.Target); err != nil { return nil, err diff --git a/backend/plugins/domain/admin/service/service.go b/backend/plugins/domain/admin/service/service.go index ba4a2e98..db474f87 100644 --- a/backend/plugins/domain/admin/service/service.go +++ b/backend/plugins/domain/admin/service/service.go @@ -5,6 +5,7 @@ package service import ( + "Wavelet/core" "Wavelet/core/contracts" "Wavelet/plugins/domain/admin/errs" "Wavelet/plugins/domain/admin/repository" @@ -112,6 +113,9 @@ func ResetServices() { // GetDB returns the GORM DB instance bound to the context if available. func GetDB(ctx context.Context) *gorm.DB { + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) + } servicesMu.RLock() defer servicesMu.RUnlock() if dbService == nil { @@ -121,42 +125,60 @@ func GetDB(ctx context.Context) *gorm.DB { } // GetCache returns the unified CacheService instance. -func GetCache(_ context.Context) contracts.CacheService { +func GetCache(ctx context.Context) contracts.CacheService { + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return cacheService } // GetUserService returns the UserService instance. -func GetUserService(_ context.Context) contracts.UserService { +func GetUserService(ctx context.Context) contracts.UserService { + if s, err := core.InjectFrom[contracts.UserService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return userService } // GetAuthService returns the AuthService instance. -func GetAuthService(_ context.Context) contracts.AuthService { +func GetAuthService(ctx context.Context) contracts.AuthService { + if s, err := core.InjectFrom[contracts.AuthService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return authService } // GetTaskService returns the TaskService instance. -func GetTaskService() contracts.TaskService { +func GetTaskService(ctx context.Context) contracts.TaskService { + if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return taskService } // GetStorageService returns the StorageService instance. -func GetStorageService() contracts.StorageService { +func GetStorageService(ctx context.Context) contracts.StorageService { + if s, err := core.InjectFrom[contracts.StorageService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return storageSvc } // GetRiskControlService returns the RiskControlService instance. -func GetRiskControlService() contracts.RiskControlService { +func GetRiskControlService(ctx context.Context) contracts.RiskControlService { + if s, err := core.InjectFrom[contracts.RiskControlService](ctx); err == nil && s != nil { + return s + } servicesMu.RLock() defer servicesMu.RUnlock() return riskControlService @@ -195,8 +217,8 @@ func requireAuthService(ctx context.Context) (contracts.AuthService, error) { } // requireTaskService resolves the injected task contract service. -func requireTaskService() (contracts.TaskService, error) { - taskSvc := GetTaskService() +func requireTaskService(ctx context.Context) (contracts.TaskService, error) { + taskSvc := GetTaskService(ctx) if taskSvc == nil { return nil, errs.ErrTaskServiceUnavailable } diff --git a/backend/plugins/domain/admin/service/status.go b/backend/plugins/domain/admin/service/status.go index dd3ef455..77c25647 100644 --- a/backend/plugins/domain/admin/service/status.go +++ b/backend/plugins/domain/admin/service/status.go @@ -117,7 +117,7 @@ func formatDuration(d time.Duration) string { func LogDatabaseStatus(ctx context.Context) model.LogDatabaseStatus { activeDB := logDBNameSQLite migration := logMigrationIdle - if rc := GetRiskControlService(); rc != nil { + if rc := GetRiskControlService(ctx); rc != nil { activeDB = rc.ActiveLogEngine(ctx) if rc.IsLogEngineMigrating(ctx) { migration = logMigrationInProgress diff --git a/backend/plugins/domain/admin/service/system_config_test.go b/backend/plugins/domain/admin/service/system_config_test.go index bce4321c..07d93edb 100644 --- a/backend/plugins/domain/admin/service/system_config_test.go +++ b/backend/plugins/domain/admin/service/system_config_test.go @@ -69,6 +69,51 @@ func setupSystemConfigTest(t *testing.T) (*gorm.DB, func()) { return sqliteDB, cleanup } +func TestPublicSystemConfigsExposesVisibleKeys(t *testing.T) { + dbConn, cleanup := setupSystemConfigTest(t) + defer cleanup() + repository.ResetSystemConfigRAMCacheForTest() + ctx := context.Background() + + hidden := model.SystemConfig{ + Key: "secret_key", + Value: "nope", + Type: "system", + Visibility: model.ConfigVisibilityHidden, + } + visible := model.SystemConfig{ + Key: model.ConfigKeyCapLoginEnabled, + Value: "true", + Type: "system", + Visibility: model.ConfigVisibilityVisible, + } + if err := dbConn.Create(&hidden).Error; err != nil { + t.Fatalf("Create(hidden) error = %v", err) + } + if err := dbConn.Create(&visible).Error; err != nil { + t.Fatalf("Create(visible) error = %v", err) + } + + got, err := service.PublicSystemConfigs(ctx) + if err != nil { + t.Fatalf("PublicSystemConfigs() error = %v", err) + } + if got[model.ConfigKeyCapLoginEnabled] != "true" { + t.Fatalf("PublicSystemConfigs()[%s] = %q, want %q", model.ConfigKeyCapLoginEnabled, got[model.ConfigKeyCapLoginEnabled], "true") + } + if _, ok := got["secret_key"]; ok { + t.Fatalf("PublicSystemConfigs() leaked hidden key secret_key") + } + + viaProvider, err := service.PublicConfigAdapter{}.PublicConfig(ctx) + if err != nil { + t.Fatalf("PublicConfigAdapter.PublicConfig() error = %v", err) + } + if viaProvider[model.ConfigKeyCapLoginEnabled] != "true" { + t.Fatalf("PublicConfigAdapter.PublicConfig()[%s] = %q, want %q", model.ConfigKeyCapLoginEnabled, viaProvider[model.ConfigKeyCapLoginEnabled], "true") + } +} + func TestListSystemConfigsByKeys_EmptyKeys(t *testing.T) { result, err := repository.ListSystemConfigsByKeys(context.Background(), nil) if err != nil { diff --git a/backend/plugins/domain/admin/service/task.go b/backend/plugins/domain/admin/service/task.go index 2543ee85..d092d4dd 100644 --- a/backend/plugins/domain/admin/service/task.go +++ b/backend/plugins/domain/admin/service/task.go @@ -17,8 +17,8 @@ import ( ) // ListTaskTypes returns every dispatchable task type declared in the task registry. -func ListTaskTypes() []contracts.TaskMetaDTO { - taskSvc := GetTaskService() +func ListTaskTypes(ctx context.Context) []contracts.TaskMetaDTO { + taskSvc := GetTaskService(ctx) if taskSvc == nil { return []contracts.TaskMetaDTO{} } @@ -27,7 +27,7 @@ func ListTaskTypes() []contracts.TaskMetaDTO { // DispatchTask validates and enqueues a manual task run, returning the new task id. func DispatchTask(ctx context.Context, req model.DispatchTaskRequest) (string, error) { - taskSvc, err := requireTaskService() + taskSvc, err := requireTaskService(ctx) if err != nil { return "", err } @@ -41,7 +41,7 @@ func DispatchTask(ctx context.Context, req model.DispatchTaskRequest) (string, e return "", err } - taskID, err := taskSvc.Dispatch(ctx, req.TaskType, validated, "manual") + taskID, err := taskSvc.Dispatch(ctx, req.TaskType, validated, contracts.TaskTriggerManual) if err != nil { return "", fmt.Errorf("%s: %w", errs.TaskDispatchFailed, err) } @@ -67,29 +67,90 @@ func ListTaskExecutions( ctx context.Context, req model.ListTaskExecutionsRequest, ) ([]model.TaskExecution, int64, error) { - if req.TaskType != "" { - if taskSvc := GetTaskService(); taskSvc != nil { - if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok { - req.TaskType = meta.Name - } + taskSvc, err := requireTaskService(ctx) + if err != nil { + return nil, 0, err + } + if req.Page <= 0 { + req.Page = 1 + } + if req.PageSize <= 0 { + req.PageSize = 20 + } + + filterType := req.TaskType + if filterType != "" { + if meta, ok := taskSvc.GetTaskMeta(filterType); ok { + filterType = meta.AsynqTask } } - executions, total, err := repository.ListTaskExecutionRecords(ctx, req) + rows, total, err := taskSvc.ListExecutions(ctx, filterType, req.Status, req.Page, req.PageSize) if err != nil { return nil, 0, err } + executions := make([]model.TaskExecution, 0, len(rows)) + for i := range rows { + executions = append(executions, executionFromDTO(rows[i])) + } return executions, total, nil } // TaskExecution loads a single execution record including its buffered log. func TaskExecution(ctx context.Context, id uint64) (*model.TaskExecution, error) { - return repository.GetTaskExecutionByID(ctx, id) + taskSvc, err := requireTaskService(ctx) + if err != nil { + return nil, err + } + dto, err := taskSvc.GetExecution(ctx, id) + if err != nil || dto == nil { + return nil, err + } + row := executionFromDTO(*dto) + return &row, nil +} + +func normalizeTaskTrigger(v string) string { + switch v { + case contracts.TaskTriggerManual, contracts.TaskTriggerSystem, contracts.TaskTriggerRetry, contracts.TaskTriggerSchedule: + return v + case "inproc_cron", "cron": + return contracts.TaskTriggerSchedule + case "http": + return contracts.TaskTriggerSystem + case "": + return contracts.TaskTriggerSystem + default: + return v + } +} + +func executionFromDTO(dto contracts.TaskExecutionDTO) model.TaskExecution { + return model.TaskExecution{ + ID: dto.ID, + TaskID: dto.TaskID, + TaskType: dto.TaskType, + TaskName: dto.TaskName, + Status: model.TaskExecutionStatus(dto.Status), + Retryable: dto.Retryable, + MaxRetry: dto.MaxRetry, + RetryCount: dto.RetryCount, + Log: dto.Log, + ErrorMessage: dto.ErrorMessage, + Result: dto.Result, + StartedAt: dto.StartedAt, + FinishedAt: dto.FinishedAt, + Duration: dto.Duration, + Payload: dto.Payload, + TriggeredBy: normalizeTaskTrigger(dto.TriggeredBy), + CreatedAt: dto.CreatedAt, + UpdatedAt: dto.UpdatedAt, + } } // RetryTask re-dispatches a failed execution as a new task run. func RetryTask(ctx context.Context, id uint64) (string, error) { - taskSvc, err := requireTaskService() + taskSvc, err := requireTaskService(ctx) if err != nil { return "", err } @@ -126,7 +187,7 @@ func CreateSchedule(ctx context.Context, req model.CreateScheduleRequest) (*mode return nil, errs.ErrInvalidCronExpression } - taskSvc, err := requireTaskService() + taskSvc, err := requireTaskService(ctx) if err != nil { return nil, err } @@ -168,7 +229,7 @@ func UpdateSchedule(ctx context.Context, id uint64, req model.UpdateScheduleRequ return nil, errs.ErrInvalidCronExpression } - taskSvc, err := requireTaskService() + taskSvc, err := requireTaskService(ctx) if err != nil { return nil, err } @@ -203,7 +264,7 @@ func DeleteSchedule(ctx context.Context, id uint64) error { return fmt.Errorf("%s: %w", errs.ScheduleDeleteFailed, err) } - if taskSvc := GetTaskService(); taskSvc != nil { + if taskSvc := GetTaskService(ctx); taskSvc != nil { reloadScheduler(ctx, taskSvc) } return nil diff --git a/backend/plugins/domain/admin/service/task_trigger_test.go b/backend/plugins/domain/admin/service/task_trigger_test.go new file mode 100644 index 00000000..4667c9be --- /dev/null +++ b/backend/plugins/domain/admin/service/task_trigger_test.go @@ -0,0 +1,32 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "testing" +) + +func TestNormalizeTaskTrigger(t *testing.T) { + tests := []struct { + in string + want string + }{ + {in: contracts.TaskTriggerManual, want: contracts.TaskTriggerManual}, + {in: contracts.TaskTriggerSystem, want: contracts.TaskTriggerSystem}, + {in: contracts.TaskTriggerRetry, want: contracts.TaskTriggerRetry}, + {in: contracts.TaskTriggerSchedule, want: contracts.TaskTriggerSchedule}, + {in: "http", want: contracts.TaskTriggerSystem}, + {in: "inproc_cron", want: contracts.TaskTriggerSchedule}, + {in: "cron", want: contracts.TaskTriggerSchedule}, + {in: "", want: contracts.TaskTriggerSystem}, + {in: "custom", want: "custom"}, + } + for _, tt := range tests { + got := normalizeTaskTrigger(tt.in) + if got != tt.want { + t.Errorf("normalizeTaskTrigger(%q) = %q, want %q", tt.in, got, tt.want) + } + } +} diff --git a/backend/plugins/domain/auth/auth_source_resolver.go b/backend/plugins/domain/auth/auth_source_resolver.go deleted file mode 100644 index 4d1c3992..00000000 --- a/backend/plugins/domain/auth/auth_source_resolver.go +++ /dev/null @@ -1,217 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package auth - -import ( - "Wavelet/core/contracts" - "context" - "errors" - "fmt" - "strconv" - "strings" - - "github.com/coreos/go-oidc/v3/oidc" - "golang.org/x/oauth2" -) - -func isOIDCLoginEnabled(ctx context.Context) bool { - val, err := GetSystemConfigValue(ctx, "oidc_login_enabled") - if err != nil || val == "" { - return true - } - b, err := strconv.ParseBool(val) - if err != nil { - return true - } - return b -} - -func resolveAuthSource(ctx context.Context, sourceName string) (*AuthSource, error) { - name := strings.TrimSpace(strings.ToLower(sourceName)) - if name == "" { - sources, err := GetActiveAuthSourcesCached(ctx) - if err != nil { - return nil, err - } - if len(sources) == 0 { - return nil, errors.New(errNoActiveAuthSource) - } - src, err := GetAuthSourceByNameCached(ctx, sources[0].Name) - if err != nil { - return nil, err - } - return src, nil - } - src, err := GetAuthSourceByNameCached(ctx, name) - if err != nil { - return nil, err - } - return src, nil -} - -func activeLoginSources(ctx context.Context) []AuthSourceView { - if !isOIDCLoginEnabled(ctx) { - return nil - } - - dbSources, err := GetActiveAuthSourcesCached(ctx) - if err != nil { - return nil - } - sources := make([]AuthSourceView, 0, len(dbSources)) - for _, source := range dbSources { - sources = append(sources, AuthSourceView{ - ID: source.ID, - Name: source.Name, - Type: source.Type, - DisplayName: source.DisplayName, - IsActive: source.IsActive, - IconURL: source.IconURL, - ClientSecretConfigured: source.ClientSecretConfigured, - }) - } - return sources -} - -func getFrontendLoginRedirectURL(ctx context.Context) (string, error) { - val, err := GetSystemConfigValue(ctx, "server_address") - if err != nil || strings.TrimSpace(val) == "" { - return "", errors.New(errServerAddressMissing) - } - return strings.TrimRight(val, "/") + "/login", nil -} - -func buildOAuthConfig(ctx context.Context, source *AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) { - if source == nil { - return nil, nil, errors.New(errAuthSourceRequired) - } - - if source.OpenIDDiscoveryURL == "" { - return nil, nil, errors.New(errDiscoveryURLRequired) - } - - // Clean the issuer URL - issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/") - issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration") - issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server") - - provider, err := globalOIDCProviderCache.get(ctx, issuer) - if err != nil { - return nil, nil, err - } - verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID}) - scopes := strings.Fields(source.Scopes) - if len(scopes) == 0 { - scopes = []string{oidc.ScopeOpenID, "profile", "email"} - } - if !containsScope(scopes, oidc.ScopeOpenID) { - scopes = append([]string{oidc.ScopeOpenID}, scopes...) - } - - return &oauth2.Config{ - ClientID: source.ClientID, - ClientSecret: source.ClientSecret, - RedirectURL: redirectURL, - Scopes: scopes, - Endpoint: provider.Endpoint(), - }, verifier, nil -} - -func containsScope(scopes []string, scope string) bool { - for _, item := range scopes { - if item == scope { - return true - } - } - return false -} - -func buildOAuthUserInfo(ctx context.Context, source *AuthSource, code, nonce, redirectURL string) (*contracts.OAuthUserInfoDTO, error) { - authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL) - if err != nil { - return nil, err - } - - token, err := authConfig.Exchange(ctx, code) - if err != nil { - return nil, err - } - - userInfo := &contracts.OAuthUserInfoDTO{Active: true} - if verifier != nil { - if verifyErr := verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil { - return nil, verifyErr - } - } - - if userInfo.Username == "" && userInfo.PreferredUsername != "" { - userInfo.Username = userInfo.PreferredUsername - } - if userInfo.Username == "" && userInfo.Email != "" { - userInfo.Username = strings.Split(userInfo.Email, "@")[0] - } - if userInfo.Username == "" && userInfo.Sub != "" { - userInfo.Username = userInfo.Sub - } - if userInfo.Name == "" { - userInfo.Name = userInfo.Username - } - - return userInfo, nil -} - -func verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *contracts.OAuthUserInfoDTO) error { - rawIDToken, ok := token.Extra("id_token").(string) - if !ok { - return nil - } - idToken, verifyErr := verifier.Verify(ctx, rawIDToken) - if verifyErr != nil { - return fmt.Errorf(errIDTokenVerifyFailedFormat, errIDTokenVerifyFailed, verifyErr) - } - if nonce != "" && idToken.Nonce != nonce { - return errors.New(errNonceMismatch) - } - if claimsErr := idToken.Claims(userInfo); claimsErr != nil { - return claimsErr - } - return nil -} - -func normalizeOAuthUserInfo(userInfo *contracts.OAuthUserInfoDTO) error { - userInfo.Username = strings.TrimSpace(userInfo.Username) - userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername) - userInfo.Email = strings.TrimSpace(userInfo.Email) - userInfo.Name = strings.TrimSpace(userInfo.Name) - userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL) - - if userInfo.Username == "" && userInfo.PreferredUsername != "" { - userInfo.Username = userInfo.PreferredUsername - } - if userInfo.Username == "" && userInfo.Email != "" { - userInfo.Username = strings.Split(userInfo.Email, "@")[0] - } - if userInfo.Username == "" && userInfo.Sub != "" { - userInfo.Username = userInfo.Sub - } - if userInfo.Username == "" { - return errors.New(errUsernameFromSourceFailed) - } - if userInfo.Name == "" { - userInfo.Name = userInfo.Username - } - if !userInfo.Active { - userInfo.Active = true - } - return nil -} - -func buildCallbackResult(user *contracts.UserDTO, status string) OAuthCallbackResult { - result := OAuthCallbackResult{Status: status} - if user != nil { - info := BuildBasicUserInfo(user, false) - result.User = &info - } - return result -} diff --git a/backend/plugins/domain/auth/cap_middleware_test.go b/backend/plugins/domain/auth/cap_middleware_test.go new file mode 100644 index 00000000..4c98a822 --- /dev/null +++ b/backend/plugins/domain/auth/cap_middleware_test.go @@ -0,0 +1,32 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "Wavelet/pkg/response" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" +) + +func TestVerifyMiddlewareMissingTokenIsBadRequest(t *testing.T) { + gin.SetMode(gin.TestMode) + restore := InstallCapTestRuntimeSettings(CapRuntimeSettings{LoginEnabled: true}) + t.Cleanup(restore) + + engine := gin.New() + engine.Use(response.ErrorHandlerMiddleware()) + engine.POST("/register", VerifyCaptchaMiddleware(GetDefaultCapManager(), "register"), func(c *gin.Context) { + c.Status(http.StatusOK) + }) + + req := httptest.NewRequest(http.MethodPost, "/register", nil) + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + if rec.Code != http.StatusBadRequest { + t.Errorf("VerifyCaptchaMiddleware() status = %d, want %d", rec.Code, http.StatusBadRequest) + } +} diff --git a/backend/plugins/domain/auth/config.go b/backend/plugins/domain/auth/config.go index c01b416e..dd77ffd4 100644 --- a/backend/plugins/domain/auth/config.go +++ b/backend/plugins/domain/auth/config.go @@ -3,12 +3,7 @@ package auth +import "Wavelet/plugins/domain/auth/service" + // SessionConfig defines the session configuration declared by the auth plugin. -type SessionConfig struct { - SessionCookieName string `config:"session_cookie_name" env:"APP_SESSION_COOKIE_NAME" default:"wavelet_session"` - SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"` - SessionDomain string `config:"session_domain" env:"APP_SESSION_DOMAIN"` - SessionAge int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"` - SessionHTTPOnly bool `config:"session_http_only" env:"APP_SESSION_HTTP_ONLY" default:"true"` - SessionSecure bool `config:"session_secure" env:"APP_SESSION_SECURE"` -} +type SessionConfig = service.SessionConfig diff --git a/backend/plugins/domain/auth/consts/cap.go b/backend/plugins/domain/auth/consts/cap.go new file mode 100644 index 00000000..8981941d --- /dev/null +++ b/backend/plugins/domain/auth/consts/cap.go @@ -0,0 +1,52 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package consts defines constants, keys, and TTL values for the auth domain plugin. +package consts + +import "time" + +// CAP 默认参数 +const ( + DefaultCapChallengeCount = 1 + DefaultCapChallengeSize = 32 + DefaultCapChallengeDifficulty = 4 + DefaultCapChallengeTTL = 10 * time.Minute + DefaultCapTokenTTL = 20 * time.Minute + + RedeemTokenIDLength = 8 // 兑换 Token ID 字节长度 + RedeemVerTokenLength = 15 // 兑换验证 Token 字节长度 + TokenPartsCount = 2 // 兑换 Token 由两部分组成 (id:token) + ValuePartsCount = 2 // 存储值由 scope 和过期时间组成 (expNano|scope) +) + +// CAP 动态配置键常量 +const ( + ConfigKeyCapLoginEnabled = "cap_login_enabled" + ConfigKeyCapChallengeCount = "cap_challenge_count" + ConfigKeyCapChallengeSize = "cap_challenge_size" + ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" + ConfigKeyCapChallengeTTL = "cap_challenge_ttl" + // ConfigKeyCapTokenTTL 验证码 Token 过期时间键 + // #nosec G101 + ConfigKeyCapTokenTTL = "cap_token_ttl" +) + +// HTTP 响应错误文案 +const ( + ErrCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // error message constant + ErrCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // error message constant + ErrCapNotConfigured = "captcha is not configured" + ErrChallengeGenerateFailed = "生成验证难题失败,请稍后再试" + ErrInvalidRequestParams = "无效的参数" + ErrSolutionVerifyFailed = "校验验证解答失败,请稍后再试" +) + +// Redeem 结果码,属于 redeem 响应 JSON 的对外契约取值,禁止改写取值 +const ( + RedeemErrInvalidToken = "invalid_token" + RedeemErrNonceStoreFailed = "nonce_store_error" + RedeemErrAlreadyRedeemed = "already_redeemed" + RedeemErrSettingsLoad = "settings_load_error" + RedeemErrTokenStoreFailed = "token_store_error" //nolint:gosec // error code constant +) diff --git a/backend/plugins/domain/auth/constants.go b/backend/plugins/domain/auth/consts/consts.go similarity index 71% rename from backend/plugins/domain/auth/constants.go rename to backend/plugins/domain/auth/consts/consts.go index f49dd81f..be0c94a7 100644 --- a/backend/plugins/domain/auth/constants.go +++ b/backend/plugins/domain/auth/consts/consts.go @@ -1,11 +1,10 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package auth +// Package consts defines constants, keys, and TTL values for the auth domain plugin. +package consts -import ( - "time" -) +import "time" // Session and Context Keys const ( @@ -14,7 +13,7 @@ const ( UserObjKey = "user_obj" TokenAuthKey = "token_auth" // 标记当前请求是否通过 Access Token 鉴权 TokenAdminKey = "token_admin" // Access Token 本身是否具有管理员权限 - SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: this is a session key, not hardcoded credentials + SessionTokenKey = "oauth_session_token" //nolint:gosec // false positive: session state key PasswordHashKey = "password_hash" SystemUsername = "system" ) @@ -23,8 +22,8 @@ const ( const ( OAuthStateCacheKeyFormat = "oauth:state:%s" OAuthStateCacheKeyExpiration = 10 * time.Minute - oauthStateLimitKeyFormat = "oauth:state:limit:%s" - oauthStateLimitMax = 10 + OAuthStateLimitKeyFormat = "oauth:state:limit:%s" + OAuthStateLimitMax = 10 ) // OAuth Purpose Constants @@ -37,3 +36,9 @@ const ( const ( AuthSourceTypeOIDC = "oidc" ) + +// Cache TTLs +const ( + TokenCacheTTL = 5 * time.Minute + UserCacheTTL = 5 * time.Minute +) diff --git a/backend/plugins/domain/auth/consts/errs.go b/backend/plugins/domain/auth/consts/errs.go new file mode 100644 index 00000000..cc2637ec --- /dev/null +++ b/backend/plugins/domain/auth/consts/errs.go @@ -0,0 +1,53 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package consts defines constants, keys, and TTL values for the auth domain plugin. +package consts + +// OAuth and Auth error messages +const ( + ErrInvalidState = "非法登录请求" + ErrIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // error message constant + ErrIDTokenVerifyFailedFormat = "%s: %w" + ErrNonceMismatch = "nonce 不匹配,可能存在重放攻击" + ErrNoActiveAuthSource = "未配置可用认证源" + ErrServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试" + ErrAuthSourceRequired = "认证源不能为空" + ErrDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL" + ErrUsernameGenerateFailed = "无法生成可用用户名" + ErrUsernameFromSourceFailed = "无法从认证源获取用户名" + ErrAuthSourceDisabled = "认证源未启用" + ErrInvalidExternalAccountBindingID = "绑定记录 ID 无效" + ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // error message constant + ErrOAuthStateRateLimited = "请求授权过于频繁,请稍后重试" + ErrAuthSourceNameRequired = "认证源名称不能为空" + ErrAuthSourceNameInvalid = "认证源名称格式不正确" + ErrAuthSourceTypeUnsupported = "不支持的认证源类型" + ErrAuthSourceDiscoveryURLRequired = "Discovery URL 不能为空" + //nolint:gosec // error message constant + ErrAuthSourceClientCredentialsRequired = "启用认证源时必须配置 Client ID 和 Client Secret" + ErrAuthSourceIDRequired = "认证源 ID 不能为空" + ErrUserIDRequired = "用户 ID 不能为空" + ErrExternalAccountBindingIncomplete = "外部帐号绑定信息不完整" + ErrExternalAccountAlreadyBoundToAnother = "该外部帐号已被其他用户绑定" + ErrExternalAccountBindingIDRequired = "外部帐号绑定记录 ID 不能为空" + ErrInsufficientPermission = "权限不足" + ErrBannedAccount = "账号已被封禁" + ErrUnAuthorized = "未登录" +) + +// Service 层与鉴权中间件内部错误文案 +const ( + ErrUserNotInContext = "auth: user not found in context" + ErrEmptyToken = "auth: empty token" //nolint:gosec // error message constant + ErrSystemUserTokenNotAllowed = "auth: system user token not allowed" //nolint:gosec // error message constant + ErrUnauthorizedInternal = "unauthorized" + ErrSystemUserLoginNotAllowed = "system user is not allowed to login" +) + +// OAuth 回调会话校验错误文案 +const ( + ErrInvalidSessionContext = "invalid session context" + ErrSessionMismatchForOAuth = "session mismatch for oauth state" + ErrUserContextMismatch = "user context mismatch for oauth binding" +) diff --git a/backend/plugins/domain/auth/audit.go b/backend/plugins/domain/auth/controller/audit.go similarity index 83% rename from backend/plugins/domain/auth/audit.go rename to backend/plugins/domain/auth/controller/audit.go index f5755340..3ed895fa 100644 --- a/backend/plugins/domain/auth/audit.go +++ b/backend/plugins/domain/auth/controller/audit.go @@ -1,11 +1,13 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package auth +// Package controller provides HTTP handlers and middlewares for the auth plugin. +package controller import ( "Wavelet/core/contracts" "Wavelet/pkg/logger" + "Wavelet/plugins/domain/auth/model/dto" "context" "encoding/json" @@ -17,7 +19,7 @@ func LogForAudit(ctx context.Context, user *contracts.UserDTO, c *gin.Context) { if user == nil || c == nil { return } - auditLog := loginRequiredAuditLog{ + auditLog := dto.LoginRequiredAuditLog{ UserID: user.ID, Username: user.Username, ClientIP: c.ClientIP(), diff --git a/backend/plugins/domain/cap/handlers.go b/backend/plugins/domain/auth/controller/cap.go similarity index 52% rename from backend/plugins/domain/cap/handlers.go rename to backend/plugins/domain/auth/controller/cap.go index fe211ddb..925d86bb 100644 --- a/backend/plugins/domain/cap/handlers.go +++ b/backend/plugins/domain/auth/controller/cap.go @@ -1,44 +1,59 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package cap +// Package controller provides HTTP handlers and middlewares for the auth plugin. +package controller import ( "Wavelet/pkg/logger" "Wavelet/pkg/response" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/model/dto" + "Wavelet/plugins/domain/auth/service" "net/http" "github.com/gin-gonic/gin" ) +// CaptchaHandler handles CAPTCHA challenge and redeem endpoints. +type CaptchaHandler struct { + capMgr *service.CaptchaManager +} + +// NewCaptchaHandler creates a new CaptchaHandler. +func NewCaptchaHandler(mgr *service.CaptchaManager) *CaptchaHandler { + return &CaptchaHandler{ + capMgr: mgr, + } +} + // Challenge 生成 PoW 人机验证难题 // @Summary 生成人机验证难题 // @Description 客户端获取 PoW 难题和签名的 JWT Token,并在后台计算。 // @Tags cap // @Accept json // @Produce json -// @Param request body challengeRequest false "可选范围限制参数" -// @Success 200 {object} response.Any{data=cap.ChallengeResponse} "成功返回 PoW 难题" +// @Param request body dto.ChallengeRequest false "可选范围限制参数" +// @Success 200 {object} response.Any{data=dto.ChallengeResponse} "成功返回 PoW 难题" // @Failure 500 {object} response.Any "内部服务错误" // @Router /api/v1/cap/challenge [get] // @Router /api/v1/cap/challenge [post] -func Challenge(c *gin.Context) { - var req challengeRequest +func (h *CaptchaHandler) Challenge(c *gin.Context) { + var req dto.ChallengeRequest _ = c.ShouldBind(&req) // 允许不传 body,默认使用 login scope if req.Scope == "" { req.Scope = "login" } - mgr := GetDefaultManager() - if mgr == nil { - response.AbortInternal(c, errCapNotConfigured) + if h.capMgr == nil { + response.AbortInternal(c, consts.ErrCapNotConfigured) return } - resp, err := mgr.Generate(c.Request.Context(), req.Scope) + resp, err := h.capMgr.Generate(c.Request.Context(), req.Scope) if err != nil { logger.ErrorF(c.Request.Context(), "Generate cap challenge failed: %v", err) - response.AbortInternal(c, errChallengeGenerateFailed) + response.AbortInternal(c, consts.ErrChallengeGenerateFailed) return } @@ -51,15 +66,15 @@ func Challenge(c *gin.Context) { // @Tags cap // @Accept json // @Produce json -// @Param request body redeemRequest true "难题 Token 与解答 solutions 数组" -// @Success 200 {object} response.Any{data=cap.RedeemResponse} "核销成功,返回 X-Cap-Token" +// @Param request body dto.RedeemRequest true "难题 Token 与解答 solutions 数组" +// @Success 200 {object} response.Any{data=dto.RedeemResponse} "核销成功,返回 X-Cap-Token" // @Failure 400 {object} response.Any "参数错误或核销失败" // @Failure 500 {object} response.Any "内部服务错误" // @Router /api/v1/cap/redeem [post] -func Redeem(c *gin.Context) { - var req redeemRequest +func (h *CaptchaHandler) Redeem(c *gin.Context) { + var req dto.RedeemRequest if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, errInvalidRequestParams) + response.AbortBadRequest(c, consts.ErrInvalidRequestParams) return } @@ -67,15 +82,14 @@ func Redeem(c *gin.Context) { req.Scope = "login" } - mgr := GetDefaultManager() - if mgr == nil { - response.AbortInternal(c, errCapNotConfigured) + if h.capMgr == nil { + response.AbortInternal(c, consts.ErrCapNotConfigured) return } - resp, err := mgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope) + resp, err := h.capMgr.Redeem(c.Request.Context(), req.Token, req.Solutions, req.Scope) if err != nil { logger.ErrorF(c.Request.Context(), "Redeem cap solutions failed: %v", err) - response.AbortInternal(c, errSolutionVerifyFailed) + response.AbortInternal(c, consts.ErrSolutionVerifyFailed) return } diff --git a/backend/plugins/domain/auth/controller/cap_middleware.go b/backend/plugins/domain/auth/controller/cap_middleware.go new file mode 100644 index 00000000..810b9c7d --- /dev/null +++ b/backend/plugins/domain/auth/controller/cap_middleware.go @@ -0,0 +1,41 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package controller provides HTTP handlers and middlewares for the auth plugin. +package controller + +import ( + "Wavelet/pkg/response" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/service" + + "github.com/gin-gonic/gin" +) + +// VerifyCaptchaMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header. +func VerifyCaptchaMiddleware(mgr *service.CaptchaManager, settingsMgr *service.CapSettingsManager, scope string) gin.HandlerFunc { + return func(c *gin.Context) { + if settingsMgr != nil && !settingsMgr.CapProtectionEnabled(c.Request.Context()) { + c.Next() + return + } + if mgr == nil { + response.AbortBadRequest(c, consts.ErrCapTokenInvalidOrExpired) + return + } + + token := c.GetHeader("X-Cap-Token") + if token == "" { + response.AbortBadRequest(c, consts.ErrCapTokenMissing) + return + } + + valid, err := mgr.VerifyToken(c.Request.Context(), token, scope) + if err != nil || !valid { + response.AbortBadRequest(c, consts.ErrCapTokenInvalidOrExpired) + return + } + + c.Next() + } +} diff --git a/backend/plugins/domain/auth/controller/controller.go b/backend/plugins/domain/auth/controller/controller.go new file mode 100644 index 00000000..3f01731c --- /dev/null +++ b/backend/plugins/domain/auth/controller/controller.go @@ -0,0 +1,109 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package controller provides HTTP handlers and middlewares for the auth plugin. +package controller + +import ( + "Wavelet/core/extpoints" + "Wavelet/plugins/domain/auth/service" + + "github.com/gin-gonic/gin" +) + +// Controller aggregates all HTTP handlers and middlewares for the auth plugin. +type Controller struct { + svc *service.Service + whitelist *extpoints.PathWhitelist + + OAuth *OAuthHandler + UserInfo *UserInfoHandler + Captcha *CaptchaHandler +} + +// New creates a new Controller instance. +func New(svc *service.Service) *Controller { + wl := extpoints.NewPathWhitelist() + oauthHandler := NewOAuthHandler(svc.OAuth, svc.Session, svc.DAO) + userInfoHandler := NewUserInfoHandler() + captchaHandler := NewCaptchaHandler(svc.CapManager) + + c := &Controller{ + svc: svc, + whitelist: wl, + OAuth: oauthHandler, + UserInfo: userInfoHandler, + Captcha: captchaHandler, + } + + // Wire middlewares into AuthService + svc.AuthSvc.SetMiddlewareHandlers( + c.LoginRequired(), + c.AdminRequired(), + DisallowTokenAuth(), + CurrentUserIDFromRequestContext, + ) + + return c +} + +// Whitelist returns the whitelist tracker. +func (c *Controller) Whitelist() *extpoints.PathWhitelist { + return c.whitelist +} + +// RegisterWhitelist adds path patterns that bypass authentication. +func (c *Controller) RegisterWhitelist(patterns ...string) { + if c.whitelist != nil { + c.whitelist.Add(patterns...) + } +} + +// LoginRequired returns the authentication middleware. +func (c *Controller) LoginRequired() gin.HandlerFunc { + return LoginRequiredMiddleware(c.whitelist, c.svc.DAO) +} + +// AdminRequired returns the admin authorization middleware. +func (c *Controller) AdminRequired() gin.HandlerFunc { + return AdminRequiredMiddleware(c.svc.DAO) +} + +// DisallowTokenAuth returns the token rejection middleware. +func (c *Controller) DisallowTokenAuth() gin.HandlerFunc { + return DisallowTokenAuth() +} + +// VerifyCaptcha returns the captcha challenge verification middleware. +func (c *Controller) VerifyCaptcha(scope string) gin.HandlerFunc { + return VerifyCaptchaMiddleware(c.svc.CapManager, c.svc.CapSettings, scope) +} + +// RegisterRoutes mounts all auth endpoints onto the router. +func (c *Controller) RegisterRoutes(router extpoints.RouterExtension) { + loginReq := c.LoginRequired() + + // 1. OAuth endpoints + oauthGroup := router.Group("/api/v1/oauth") + { + oauthGroup.GET("/sources", c.OAuth.GetLoginSources) + oauthGroup.GET("/login", c.OAuth.GetLoginURL) + oauthGroup.GET("/:source/authorize", c.OAuth.Authorize) + oauthGroup.GET("/logout", c.OAuth.Logout) + oauthGroup.POST("/callback", c.OAuth.Callback) + oauthGroup.GET("/user-info", loginReq, c.UserInfo.UserInfo) + oauthGroup.GET("/external-accounts", loginReq, c.OAuth.ListExternalAccounts) + oauthGroup.POST("/external-accounts/:id/delete", loginReq, c.OAuth.DeleteExternalAccount) + } + + // 2. Global user-info route alias + router.GET("/api/v1/user-info", loginReq, c.UserInfo.UserInfo) + + // 3. CAPTCHA endpoints + capGroup := router.Group("/api/v1/cap") + { + capGroup.GET("/challenge", c.Captcha.Challenge) + capGroup.POST("/challenge", c.Captcha.Challenge) + capGroup.POST("/redeem", c.Captcha.Redeem) + } +} diff --git a/backend/plugins/domain/auth/controller/middleware.go b/backend/plugins/domain/auth/controller/middleware.go new file mode 100644 index 00000000..ea766950 --- /dev/null +++ b/backend/plugins/domain/auth/controller/middleware.go @@ -0,0 +1,187 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package controller provides HTTP handlers and middlewares for the auth plugin. +package controller + +import ( + "Wavelet/core/contracts" + "Wavelet/core/extpoints" + "Wavelet/pkg/ginutil" + "Wavelet/pkg/response" + "Wavelet/pkg/trace" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/dao" + "Wavelet/plugins/domain/auth/model/do" + "Wavelet/plugins/domain/auth/model/dto" + "Wavelet/plugins/domain/auth/service" + "context" + "errors" + + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" +) + +// GetUserIDFromSession 从 Session 中提取用户 ID +func GetUserIDFromSession(s sessions.Session) uint64 { + val := s.Get(consts.UserIDKey) + return dto.ParseUserID(val) +} + +// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID +func GetUserIDFromContext(c *gin.Context) (uid uint64) { + defer func() { + _ = recover() + }() + session := sessions.Default(c) + return GetUserIDFromSession(session) +} + +// CurrentUserIDFromRequestContext 是接入层向 Service 层暴露的登录态桥接。 +func CurrentUserIDFromRequestContext(ctx context.Context) (uint64, bool) { + ginCtx, ok := ctx.(*gin.Context) + if !ok { + return 0, false + } + return GetUserIDFromContext(ginCtx), true +} + +func getUserByToken(ctx context.Context, d *dao.DAO, tokenStr string) (*contracts.UserDTO, *do.CachedToken, error) { + tokenHash := service.HashToken(tokenStr) + tokenRecord, err := d.GetCachedToken(ctx, tokenHash) + if err != nil || tokenRecord == nil { + tokenRecord, err = d.GetAccessTokenByHash(ctx, tokenHash) + if err != nil { + return nil, nil, err + } + d.SetCachedToken(ctx, tokenHash, tokenRecord) + } + + user, err := d.GetCachedUser(ctx, tokenRecord.UserID) + if err != nil || user == nil || !user.IsActive { + user, err = d.GetActiveUserByID(ctx, tokenRecord.UserID) + if err != nil { + return nil, nil, err + } + d.SetCachedUser(ctx, tokenRecord.UserID, user) + } + + return user, tokenRecord, nil +} + +// GetUserFromRequest 从请求中获取当前用户(优先 Access Token,其次 Session) +func GetUserFromRequest(c *gin.Context, d *dao.DAO) (*contracts.UserDTO, error) { + ctx := c.Request.Context() + var tokenStr string + + tokenFromQuery := c.Query("token") + if tokenFromQuery != "" { + tokenStr = tokenFromQuery + } else { + authHeader := c.GetHeader("Authorization") + if len(authHeader) > 7 && authHeader[:7] == "Bearer " { + tokenStr = authHeader[7:] + } + } + + // 优先使用 Access Token 鉴权 + if tokenStr != "" { + if user, tokenRecord, err := getUserByToken(ctx, d, tokenStr); err == nil { + if user.Username == consts.SystemUsername { + return nil, errors.New(consts.ErrSystemUserLoginNotAllowed) + } + ginutil.SetToContext(c, contracts.AuthTokenAuthKey, true) + ginutil.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin) + return user, nil + } + } + + // 降级使用 Session 鉴权 + userID := GetUserIDFromContext(c) + if userID <= 0 { + return nil, errors.New(consts.ErrUnauthorizedInternal) + } + + user, err := d.GetCachedUser(ctx, userID) + if err != nil || user == nil || !user.IsActive { + user, err = d.GetActiveUserByID(ctx, userID) + if err != nil { + return nil, err + } + d.SetCachedUser(ctx, userID, user) + } + + ginutil.SetToContext(c, contracts.AuthTokenAuthKey, false) + ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false) + + if user.Username == consts.SystemUsername { + return nil, errors.New(consts.ErrSystemUserLoginNotAllowed) + } + + return user, nil +} + +// LoginRequiredMiddleware returns a Gin handler function for authentication check. +func LoginRequiredMiddleware(whitelist *extpoints.PathWhitelist, d *dao.DAO) gin.HandlerFunc { + return func(c *gin.Context) { + if whitelist != nil && whitelist.Match(c.Request.URL.Path) { + c.Next() + return + } + + _, span := trace.Start(c.Request.Context(), "LoginRequired") + defer span.End() + + user, err := GetUserFromRequest(c, d) + if err != nil { + response.AbortUnauthorized(c, consts.ErrUnAuthorized) + return + } + + LogForAudit(c.Request.Context(), user, c) + ginutil.SetToContext(c, contracts.AuthUserObjKey, user) + c.Next() + } +} + +// AdminRequiredMiddleware returns a Gin handler function for admin authorization check. +func AdminRequiredMiddleware(d *dao.DAO) gin.HandlerFunc { + return func(c *gin.Context) { + _, span := trace.Start(c.Request.Context(), "AdminRequired") + defer span.End() + + user, err := GetUserFromRequest(c, d) + if err != nil { + response.AbortUnauthorized(c, consts.ErrUnAuthorized) + return + } + + isTokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey) + isTokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey) + + // Logged-in but lacking admin permission is 403, not 401/404. + if isTokenAuth && !isTokenAdmin && !user.IsAdmin { + response.AbortForbidden(c, consts.ErrInsufficientPermission) + return + } + if !isTokenAuth && !user.IsAdmin { + response.AbortForbidden(c, consts.ErrInsufficientPermission) + return + } + + LogForAudit(c.Request.Context(), user, c) + ginutil.SetToContext(c, contracts.AuthUserObjKey, user) + c.Next() + } +} + +// DisallowTokenAuth returns a middleware that rejects requests authenticated via access token. +func DisallowTokenAuth() gin.HandlerFunc { + return func(c *gin.Context) { + if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth { + response.AbortForbidden(c, consts.ErrTokenAuthNotAllowed) + return + } + c.Next() + } +} diff --git a/backend/plugins/domain/auth/controller/oauth.go b/backend/plugins/domain/auth/controller/oauth.go new file mode 100644 index 00000000..47647e74 --- /dev/null +++ b/backend/plugins/domain/auth/controller/oauth.go @@ -0,0 +1,448 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package controller provides HTTP handlers and middlewares for the auth plugin. +package controller + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/pkg/response" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/dao" + "Wavelet/plugins/domain/auth/model/do" + "Wavelet/plugins/domain/auth/model/dto" + "Wavelet/plugins/domain/auth/model/entity" + "Wavelet/plugins/domain/auth/service" + "context" + "fmt" + "net/http" + "strconv" + "strings" + + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" + "github.com/google/uuid" +) + +// OAuthHandler handles OAuth authentication endpoints. +type OAuthHandler struct { + oauthSvc *service.OAuthService + sessionSvc *service.SessionService + dao *dao.DAO +} + +// NewOAuthHandler creates a new OAuthHandler. +func NewOAuthHandler(oauthSvc *service.OAuthService, sessionSvc *service.SessionService, d *dao.DAO) *OAuthHandler { + return &OAuthHandler{ + oauthSvc: oauthSvc, + sessionSvc: sessionSvc, + dao: d, + } +} + +// GetLoginSources 获取可用登录源列表 +// @Summary 获取可用登录源 +// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用 +// @Tags oauth +// @Produce json +// @Success 200 {object} response.Any{data=[]dto.AuthSourceView} "登录源列表" +// @Router /api/v1/oauth/sources [get] +func (h *OAuthHandler) GetLoginSources(c *gin.Context) { + sources, err := h.oauthSvc.ActiveLoginSources(c.Request.Context()) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(sources)) +} + +// GetLoginURL 获取登录授权地址 +// @Summary 获取登录授权地址 +// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。 +// @Tags oauth +// @Produce json +// @Param source query string false "认证源名称,为空使用第一个启用的认证源" +// @Success 200 {object} response.Any{data=dto.OAuthAuthorizeResponse} "授权 URL" +// @Failure 400 {object} response.Any "认证源不存在或未配置" +// @Failure 500 {object} response.Any "构造 URL 失败" +// @Router /api/v1/oauth/login [get] +func (h *OAuthHandler) GetLoginURL(c *gin.Context) { + ctx := c.Request.Context() + if !h.oauthSvc.IsOIDCLoginEnabled(ctx) { + response.AbortBadRequest(c, consts.ErrAuthSourceDisabled) + return + } + + source, err := h.oauthSvc.ResolveAuthSource(ctx, c.Query("source")) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + if !source.IsActive { + response.AbortBadRequest(c, consts.ErrAuthSourceDisabled) + return + } + + session := sessions.Default(c) + token, isNew := h.sessionSvc.EnsureSessionToken(session) + if isNew { + if err := session.Save(); err != nil { + response.AbortInternal(c, err.Error()) + return + } + } + + userID := GetUserIDFromSession(session) + sessionHash := h.sessionSvc.HashSessionToken(token) + if err := h.oauthSvc.ReserveOAuthStateSlot(ctx, sessionHash); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + state := uuid.NewString() + payloadValue, err := (do.OAuthStatePayload{ + SourceName: source.Name, + Purpose: consts.OAuthPurposeLogin, + UserID: userID, + SessionHash: sessionHash, + }).Encode() + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + stateKey := fmt.Sprintf(consts.OAuthStateCacheKeyFormat, state) + if cache := h.dao.Cache(); cache != nil { + if err := cache.Set(ctx, stateKey, payloadValue, consts.OAuthStateCacheKeyExpiration); err != nil { + response.AbortInternal(c, err.Error()) + return + } + } + + authorizeURL, err := h.oauthSvc.BuildAuthorizeURL(ctx, source, state) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(dto.OAuthAuthorizeResponse{AuthorizeURL: authorizeURL})) +} + +// Authorize 发起指定认证源授权 +// @Summary 发起指定认证源授权 +// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。 +// @Tags oauth +// @Produce json +// @Param source path string true "认证源名称" +// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login" +// @Success 200 {object} response.Any{data=dto.OAuthAuthorizeResponse} "授权 URL" +// @Failure 400 {object} response.Any "认证源不存在或未启用" +// @Failure 500 {object} response.Any "构造 URL 失败" +// @Router /api/v1/oauth/{source}/authorize [get] +func (h *OAuthHandler) Authorize(c *gin.Context) { + ctx := c.Request.Context() + if !h.oauthSvc.IsOIDCLoginEnabled(ctx) { + response.AbortBadRequest(c, consts.ErrAuthSourceDisabled) + return + } + + source, err := h.oauthSvc.ResolveAuthSource(ctx, c.Param("source")) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + if !source.IsActive { + response.AbortBadRequest(c, consts.ErrAuthSourceDisabled) + return + } + purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose"))) + if purpose != consts.OAuthPurposeBind { + purpose = consts.OAuthPurposeLogin + } + + session := sessions.Default(c) + userID := GetUserIDFromSession(session) + if purpose == consts.OAuthPurposeBind && userID == 0 { + response.AbortUnauthorized(c, consts.ErrUnAuthorized) + return + } + + token, isNew := h.sessionSvc.EnsureSessionToken(session) + if isNew { + if err := session.Save(); err != nil { + response.AbortInternal(c, err.Error()) + return + } + } + + sessionHash := h.sessionSvc.HashSessionToken(token) + if err := h.oauthSvc.ReserveOAuthStateSlot(ctx, sessionHash); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + state := uuid.NewString() + payloadValue, err := (do.OAuthStatePayload{ + SourceName: source.Name, + Purpose: purpose, + UserID: userID, + SessionHash: sessionHash, + }).Encode() + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + stateKey := fmt.Sprintf(consts.OAuthStateCacheKeyFormat, state) + if cache := h.dao.Cache(); cache != nil { + if err := cache.Set(ctx, stateKey, payloadValue, consts.OAuthStateCacheKeyExpiration); err != nil { + response.AbortInternal(c, err.Error()) + return + } + } + + authorizeURL, err := h.oauthSvc.BuildAuthorizeURL(ctx, source, state) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(dto.OAuthAuthorizeResponse{AuthorizeURL: authorizeURL})) +} + +// Callback OAuth 回调处理 +// @Summary OAuth 回调处理 +// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。 +// @Tags oauth +// @Accept json +// @Produce json +// @Param request body dto.CallbackRequest true "回调请求参数" +// @Success 200 {object} response.Any{data=dto.OAuthCallbackResult} "登录或绑定成功" +// @Failure 400 {object} response.Any "state 无效、参数错误或认证源错误" +// @Failure 401 {object} response.Any "绑定场景未登录" +// @Failure 500 {object} response.Any "OAuth 认证失败或内部错误" +// @Router /api/v1/oauth/callback [post] +func (h *OAuthHandler) Callback(c *gin.Context) { + var req dto.CallbackRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + ctx := c.Request.Context() + stateKey := fmt.Sprintf(consts.OAuthStateCacheKeyFormat, req.State) + var payloadRaw string + cache := h.dao.Cache() + if cache == nil { + response.AbortBadRequest(c, consts.ErrInvalidState) + return + } + if err := cache.Get(ctx, stateKey, &payloadRaw); err != nil { + response.AbortBadRequest(c, consts.ErrInvalidState) + return + } + _ = cache.Delete(ctx, stateKey) + + payload, err := do.DecodeOAuthStatePayload(payloadRaw) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + session := sessions.Default(c) + currentUserID := GetUserIDFromSession(session) + + if payload.Purpose == consts.OAuthPurposeBind && currentUserID == 0 { + response.AbortUnauthorized(c, consts.ErrUnAuthorized) + return + } + + token, ok := session.Get(consts.SessionTokenKey).(string) + if !ok || token == "" { + response.AbortBadRequest(c, consts.ErrInvalidSessionContext) + return + } + + if h.sessionSvc.HashSessionToken(token) != payload.SessionHash { + response.AbortBadRequest(c, consts.ErrSessionMismatchForOAuth) + return + } + + if payload.Purpose == consts.OAuthPurposeBind && currentUserID != payload.UserID { + response.AbortBadRequest(c, consts.ErrUserContextMismatch) + return + } + + if !h.oauthSvc.IsOIDCLoginEnabled(ctx) { + response.AbortBadRequest(c, consts.ErrAuthSourceDisabled) + return + } + + source, err := h.oauthSvc.ResolveAuthSource(ctx, payload.SourceName) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + if !source.IsActive { + response.AbortBadRequest(c, consts.ErrAuthSourceDisabled) + return + } + + redirectURL, err := h.oauthSvc.GetFrontendLoginRedirectURL(ctx) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + + userInfo, err := h.oauthSvc.BuildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + if err := h.oauthSvc.NormalizeOAuthUserInfo(userInfo); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + if userInfo.Sub == "" { + userInfo.Sub = userInfo.Username + } + + if payload.Purpose == consts.OAuthPurposeBind { + h.handleCallbackBind(ctx, c, source, userInfo) + return + } + + h.handleCallbackLogin(ctx, c, source, userInfo) +} + +func (h *OAuthHandler) handleCallbackBind(ctx context.Context, c *gin.Context, source *entity.AuthSource, userInfo *contracts.OAuthUserInfoDTO) { + userID := GetUserIDFromContext(c) + if userID == 0 { + response.AbortUnauthorized(c, consts.ErrUnAuthorized) + return + } + user, err := h.dao.GetUserByID(ctx, userID) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + if err := h.oauthSvc.BindExternalAccount(ctx, source.ID, user.ID, userInfo); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "bound"))) +} + +func (h *OAuthHandler) handleCallbackLogin(ctx context.Context, c *gin.Context, source *entity.AuthSource, userInfo *contracts.OAuthUserInfoDTO) { + user, ok, err := h.oauthSvc.AuthenticateOrRegisterUser(ctx, source, userInfo) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + if !ok || user == nil { + c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind"))) + return + } + + session := sessions.Default(c) + isSessionCookie, err := h.sessionSvc.ApplyLoginSession(ctx, session, user) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + if isSessionCookie { + h.sessionSvc.StripCookieMaxAgeAndExpires(c.Writer.Header(), h.sessionSvc.Config().SessionCookieName) + } + + h.dao.SetCachedUser(ctx, user.ID, user) + logger.InfoF(ctx, "[LoginAudit] successful OAuth login via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP()) + + c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "logged_in"))) +} + +func buildCallbackResult(user *contracts.UserDTO, status string) dto.OAuthCallbackResult { + result := dto.OAuthCallbackResult{Status: status} + if user != nil { + info := dto.BuildBasicUserInfo(user, false) + result.User = &info + } + return result +} + +// Logout 退出登录 +// @Summary 退出登录 +// @Description 清除当前用户的登录会话,完成退出。清除 Cookie 中的 Session 数据。 +// @Tags oauth +// @Produce json +// @Security SessionCookie +// @Success 200 {object} response.Any{data=string} "退出成功" +// @Failure 500 {object} response.Any "Session 清除失败" +// @Router /api/v1/oauth/logout [get] +func (h *OAuthHandler) Logout(c *gin.Context) { + session := sessions.Default(c) + userID := session.Get(consts.UserIDKey) + username := session.Get(consts.UserNameKey) + if userID != nil { + logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP()) + if id := dto.ParseUserID(userID); id > 0 { + h.dao.InvalidateCachedUser(c.Request.Context(), id) + } + } + session.Options(h.sessionSvc.GetSessionOptions(-1)) + session.Clear() + if err := session.Save(); err != nil { + response.AbortInternal(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OKNil()) +} + +// ListExternalAccounts 获取当前用户的外部帐号绑定列表 +// @Summary 获取外部帐号列表 +// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录 +// @Tags oauth +// @Produce json +// @Security SessionCookie +// @Success 200 {object} response.Any "外部帐号列表" +// @Failure 401 {object} response.Any "未登录" +// @Failure 500 {object} response.Any "内部错误" +// @Router /api/v1/oauth/external-accounts [get] +func (h *OAuthHandler) ListExternalAccounts(c *gin.Context) { + userID := GetUserIDFromContext(c) + accounts, err := h.oauthSvc.ListExternalAccounts(c.Request.Context(), userID) + if err != nil { + response.AbortInternal(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(accounts)) +} + +// DeleteExternalAccount 解除外部帐号绑定 +// @Summary 解除外部帐号绑定 +// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录 +// @Tags oauth +// @Produce json +// @Security SessionCookie +// @Param id path uint64 true "外部帐号绑定记录 ID" +// @Success 200 {object} response.Any{data=string} "解除绑定成功" +// @Failure 400 {object} response.Any "ID 无效或解除失败" +// @Failure 401 {object} response.Any "未登录" +// @Router /api/v1/oauth/external-accounts/{id}/delete [post] +func (h *OAuthHandler) DeleteExternalAccount(c *gin.Context) { + userID := GetUserIDFromContext(c) + if userID == 0 { + response.AbortUnauthorized(c, consts.ErrUnAuthorized) + return + } + rawID := strings.TrimSpace(c.Param("id")) + id, err := strconv.ParseUint(rawID, 10, 64) + if err != nil || id == 0 { + response.AbortBadRequest(c, consts.ErrInvalidExternalAccountBindingID) + return + } + if err := h.oauthSvc.DeleteExternalAccount(c.Request.Context(), id, userID); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OKNil()) +} diff --git a/backend/plugins/domain/auth/controller/user_info.go b/backend/plugins/domain/auth/controller/user_info.go new file mode 100644 index 00000000..62a35ac7 --- /dev/null +++ b/backend/plugins/domain/auth/controller/user_info.go @@ -0,0 +1,45 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package controller provides HTTP handlers and middlewares for the auth plugin. +package controller + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/ginutil" + "Wavelet/pkg/response" + "Wavelet/plugins/domain/auth/model/dto" + "net/http" + + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" +) + +// UserInfoHandler handles current user info queries. +type UserInfoHandler struct{} + +// NewUserInfoHandler creates a new UserInfoHandler. +func NewUserInfoHandler() *UserInfoHandler { + return &UserInfoHandler{} +} + +// UserInfo 获取当前登录用户信息 +// @Summary 获取当前登录用户信息 +// @Description 返回当前登录用户的基本信息,需要登录。 +// @Tags oauth +// @Produce json +// @Security SessionCookie +// @Success 200 {object} response.Any{data=dto.BasicUserInfo} "用户信息" +// @Failure 401 {object} response.Any "未登录" +// @Router /api/v1/oauth/user-info [get] +// @Router /api/v1/user-info [get] +func (h *UserInfoHandler) UserInfo(c *gin.Context) { + user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) + session := sessions.Default(c) + needChange := session.Get("need_change_password") == true || (user != nil && user.NeedChangePassword) + + c.JSON( + http.StatusOK, + response.OK(dto.BuildBasicUserInfo(user, needChange)), + ) +} diff --git a/backend/plugins/domain/auth/dao/auth_source.go b/backend/plugins/domain/auth/dao/auth_source.go new file mode 100644 index 00000000..0e7aff6a --- /dev/null +++ b/backend/plugins/domain/auth/dao/auth_source.go @@ -0,0 +1,61 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package dao provides data access objects and caching for the auth domain plugin. +package dao + +import ( + "Wavelet/plugins/domain/auth/model/entity" + "context" +) + +// ListAllAuthSources 获取全部认证源(含未启用),按 ID 升序 +func (d *DAO) ListAllAuthSources(ctx context.Context) ([]entity.AuthSource, error) { + var sources []entity.AuthSource + if err := d.DB(ctx).Order("id ASC").Find(&sources).Error; err != nil { + return nil, err + } + return sources, nil +} + +// GetAuthSourceByID 根据 ID 获取认证源 +func (d *DAO) GetAuthSourceByID(ctx context.Context, id uint64) (*entity.AuthSource, error) { + var src entity.AuthSource + if err := d.DB(ctx).First(&src, id).Error; err != nil { + return nil, err + } + return &src, nil +} + +// GetAuthSourceByName 根据名称获取认证源 +func (d *DAO) GetAuthSourceByName(ctx context.Context, name string) (*entity.AuthSource, error) { + var src entity.AuthSource + if err := d.DB(ctx).Where("name = ?", name).First(&src).Error; err != nil { + return nil, err + } + return &src, nil +} + +// ListActiveAuthSources 获取所有启用的认证源 +func (d *DAO) ListActiveAuthSources(ctx context.Context) ([]entity.AuthSource, error) { + var sources []entity.AuthSource + if err := d.DB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil { + return nil, err + } + return sources, nil +} + +// CreateAuthSource 新建认证源记录 +func (d *DAO) CreateAuthSource(ctx context.Context, source *entity.AuthSource) error { + return d.DB(ctx).Create(source).Error +} + +// SaveAuthSource 全量保存认证源记录 +func (d *DAO) SaveAuthSource(ctx context.Context, source *entity.AuthSource) error { + return d.DB(ctx).Save(source).Error +} + +// DeleteAuthSource 删除认证源记录 +func (d *DAO) DeleteAuthSource(ctx context.Context, source *entity.AuthSource) error { + return d.DB(ctx).Delete(source).Error +} diff --git a/backend/plugins/domain/auth/cache.go b/backend/plugins/domain/auth/dao/cache.go similarity index 53% rename from backend/plugins/domain/auth/cache.go rename to backend/plugins/domain/auth/dao/cache.go index 0beda8ec..a2476651 100644 --- a/backend/plugins/domain/auth/cache.go +++ b/backend/plugins/domain/auth/dao/cache.go @@ -1,30 +1,20 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package auth +// Package dao provides data access objects and caching for the auth domain plugin. +package dao import ( "Wavelet/core/contracts" "Wavelet/pkg/cache/ram" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/model/do" "context" "fmt" - "time" ) -const ( - tokenCacheTTL = 5 * time.Minute - userCacheTTL = 5 * time.Minute -) - -// CachedToken represents the minimal cached representation of an access token. -type CachedToken struct { - ID uint64 `json:"id"` - UserID uint64 `json:"user_id"` - IsAdmin bool `json:"is_admin"` -} - var ( - tokenRAM = ram.MustNew[string, *CachedToken](ram.Options{MaximumSize: 2048}) + tokenRAM = ram.MustNew[string, *do.CachedToken](ram.Options{MaximumSize: 2048}) userRAM = ram.MustNew[uint64, *contracts.UserDTO](ram.Options{MaximumSize: 2048}) ) @@ -37,13 +27,15 @@ func userCacheKey(userID uint64) string { } // GetCachedToken 获取缓存的 Token -func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) { +// +//nolint:dupl // token and user cache lookup pattern +func (d *DAO) GetCachedToken(ctx context.Context, tokenHash string) (*do.CachedToken, error) { if val, ok := tokenRAM.GetIfPresent(tokenHash); ok { return val, nil } - if cache := getCache(ctx); cache != nil { - var token CachedToken + if cache := d.Cache(); cache != nil { + var token do.CachedToken key := tokenCacheKey(tokenHash) if err := cache.Get(ctx, key, &token); err == nil { tokenRAM.Set(tokenHash, &token) @@ -54,30 +46,32 @@ func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) } // SetCachedToken 设置 Token 缓存 -func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) { +func (d *DAO) SetCachedToken(ctx context.Context, tokenHash string, token *do.CachedToken) { tokenRAM.Set(tokenHash, token) - if cache := getCache(ctx); cache != nil { + if cache := d.Cache(); cache != nil { key := tokenCacheKey(tokenHash) - _ = cache.Set(ctx, key, token, tokenCacheTTL) + _ = cache.Set(ctx, key, token, consts.TokenCacheTTL) } } // InvalidateCachedToken 吊销/删除 token 缓存 -func InvalidateCachedToken(ctx context.Context, tokenHash string) { +func (d *DAO) InvalidateCachedToken(ctx context.Context, tokenHash string) { tokenRAM.Invalidate(tokenHash) - if cache := getCache(ctx); cache != nil { + if cache := d.Cache(); cache != nil { key := tokenCacheKey(tokenHash) _ = cache.Delete(ctx, key) } } // GetCachedUser 获取缓存的 UserDTO -func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) { +// +//nolint:dupl // token and user cache lookup pattern +func (d *DAO) GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) { if val, ok := userRAM.GetIfPresent(userID); ok { return val, nil } - if cache := getCache(ctx); cache != nil { + if cache := d.Cache(); cache != nil { var u contracts.UserDTO key := userCacheKey(userID) if err := cache.Get(ctx, key, &u); err == nil { @@ -89,28 +83,25 @@ func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, erro } // SetCachedUser 设置 UserDTO 缓存 -func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) { +func (d *DAO) SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) { userRAM.Set(userID, u) - if cache := getCache(ctx); cache != nil { + if cache := d.Cache(); cache != nil { key := userCacheKey(userID) - _ = cache.Set(ctx, key, u, userCacheTTL) + _ = cache.Set(ctx, key, u, consts.UserCacheTTL) } } // InvalidateCachedUser 吊销/失效 UserDTO 缓存 -func InvalidateCachedUser(ctx context.Context, userID uint64) { +func (d *DAO) InvalidateCachedUser(ctx context.Context, userID uint64) { userRAM.Invalidate(userID) - if cache := getCache(ctx); cache != nil { + if cache := d.Cache(); cache != nil { key := userCacheKey(userID) _ = cache.Delete(ctx, key) } } -// StopAuthCacheListener compatibility stub for tests -func StopAuthCacheListener() {} - -// ResetAuthRAMCacheForTest clears only the process-local RAM cache. -func ResetAuthRAMCacheForTest() { +// ResetRAMCacheForTest clears only the process-local RAM cache. +func ResetRAMCacheForTest() { tokenRAM.InvalidateAll() userRAM.InvalidateAll() } diff --git a/backend/plugins/domain/auth/dao/dao.go b/backend/plugins/domain/auth/dao/dao.go new file mode 100644 index 00000000..9f09c43c --- /dev/null +++ b/backend/plugins/domain/auth/dao/dao.go @@ -0,0 +1,61 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package dao provides data access objects and caching for the auth domain plugin. +package dao + +import ( + "Wavelet/core/contracts" + "context" + + "gorm.io/gorm" +) + +// DAO aggregates all data access objects for the auth domain plugin. +type DAO struct { + dbSvc contracts.DBService + cacheSvc contracts.CacheService + limiterSvc contracts.LimiterService +} + +// New creates a new DAO aggregate. +func New(dbSvc contracts.DBService, cacheSvc contracts.CacheService, limiterSvc contracts.LimiterService) *DAO { + return &DAO{ + dbSvc: dbSvc, + cacheSvc: cacheSvc, + limiterSvc: limiterSvc, + } +} + +// SetDBService updates the DBService reference. +func (d *DAO) SetDBService(db contracts.DBService) { + d.dbSvc = db +} + +// SetCacheService updates the CacheService reference. +func (d *DAO) SetCacheService(cache contracts.CacheService) { + d.cacheSvc = cache +} + +// SetLimiterService updates the LimiterService reference. +func (d *DAO) SetLimiterService(limiter contracts.LimiterService) { + d.limiterSvc = limiter +} + +// DB returns the GORM DB instance associated with the request context. +func (d *DAO) DB(ctx context.Context) *gorm.DB { + if d.dbSvc != nil { + return d.dbSvc.DB(ctx) + } + return nil +} + +// Cache returns the CacheService instance. +func (d *DAO) Cache() contracts.CacheService { + return d.cacheSvc +} + +// Limiter returns the LimiterService instance. +func (d *DAO) Limiter() contracts.LimiterService { + return d.limiterSvc +} diff --git a/backend/plugins/domain/auth/dao/external_account.go b/backend/plugins/domain/auth/dao/external_account.go new file mode 100644 index 00000000..d6b23236 --- /dev/null +++ b/backend/plugins/domain/auth/dao/external_account.go @@ -0,0 +1,38 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package dao provides data access objects and caching for the auth domain plugin. +package dao + +import ( + "Wavelet/plugins/domain/auth/model/entity" + "context" +) + +// FindExternalAccount 查询指定认证源的外部账号绑定 +func (d *DAO) FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*entity.ExternalAccount, error) { + var account entity.ExternalAccount + if err := d.DB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil { + return nil, err + } + return &account, nil +} + +// BindExternalAccount 绑定外部账号 +func (d *DAO) BindExternalAccount(ctx context.Context, account *entity.ExternalAccount) error { + return d.DB(ctx).Create(account).Error +} + +// ListExternalAccountsByUserID 获取用户绑定的所有外部账号 +func (d *DAO) ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]entity.ExternalAccount, error) { + var accounts []entity.ExternalAccount + if err := d.DB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil { + return nil, err + } + return accounts, nil +} + +// UnbindExternalAccount 解绑外部账号 +func (d *DAO) UnbindExternalAccount(ctx context.Context, id, userID uint64) error { + return d.DB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&entity.ExternalAccount{}).Error +} diff --git a/backend/plugins/domain/auth/dao/user_bridge.go b/backend/plugins/domain/auth/dao/user_bridge.go new file mode 100644 index 00000000..cb69c139 --- /dev/null +++ b/backend/plugins/domain/auth/dao/user_bridge.go @@ -0,0 +1,91 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package dao provides data access objects and caching for the auth domain plugin. +package dao + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/util" + "Wavelet/plugins/domain/auth/model/do" + "context" + "time" +) + +// GetAccessTokenByHash 按令牌哈希读取访问令牌记录(仅取鉴权所需字段) +func (d *DAO) GetAccessTokenByHash(ctx context.Context, tokenHash string) (*do.CachedToken, error) { + var row struct { + ID uint64 + UserID uint64 + IsAdmin bool + } + if err := d.DB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&row).Error; err != nil { + return nil, err + } + return &do.CachedToken{ + ID: row.ID, + UserID: row.UserID, + IsAdmin: row.IsAdmin, + }, nil +} + +// GetActiveUserByID 读取仍处于启用状态的用户 +func (d *DAO) GetActiveUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) { + var user contracts.UserDTO + if err := d.DB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&user).Error; err != nil { + return nil, err + } + return &user, nil +} + +// GetUserByID 按 ID 读取用户(不限制启用状态) +func (d *DAO) GetUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) { + var user contracts.UserDTO + if err := d.DB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil { + return nil, err + } + return &user, nil +} + +// InsertUser 新建用户记录 +func (d *DAO) InsertUser(ctx context.Context, user *contracts.UserDTO) error { + return d.DB(ctx).Table("w_users").Create(user).Error +} + +// TouchUserLastLogin 刷新用户最后登录时间 +func (d *DAO) TouchUserLastLogin(ctx context.Context, userID uint64, at time.Time) error { + return d.DB(ctx).Table("w_users").Where("id = ?", userID).Update("last_login_at", at).Error +} + +// ListSimilarUsernames 查询与基础用户名相同或带 `-序号` 后缀的用户名(用于用户名去重) +func (d *DAO) ListSimilarUsernames(ctx context.Context, base string) ([]string, error) { + var existingUsernames []string + if err := d.DB(ctx).Table("w_users"). + Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%"). + Pluck("username", &existingUsernames).Error; err != nil { + return nil, err + } + return existingUsernames, nil +} + +// GetSystemConfigValue 读取系统配置项原始值 +func (d *DAO) GetSystemConfigValue(ctx context.Context, key string) (string, error) { + var val string + if err := d.DB(ctx).Table("w_system_configs").Where("key = ?", key).Pluck("value", &val).Error; err != nil { + return "", err + } + return val, nil +} + +// ListSystemConfigsByKeys 按键批量读取系统配置项 +func (d *DAO) ListSystemConfigsByKeys(ctx context.Context, keys []string) ([]do.CapConfigRecord, error) { + var records []do.CapConfigRecord + db := d.DB(ctx) + if db == nil { + return nil, nil + } + if err := db.Table("w_system_configs").Where("key IN ?", keys).Find(&records).Error; err != nil { + return nil, err + } + return records, nil +} diff --git a/backend/plugins/domain/auth/errs.go b/backend/plugins/domain/auth/errs.go deleted file mode 100644 index 3d63894e..00000000 --- a/backend/plugins/domain/auth/errs.go +++ /dev/null @@ -1,54 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package auth - -// OAuth and Auth error messages -const ( - errInvalidState = "非法登录请求" - errIDTokenVerifyFailed = "ID Token 验证失败" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errIDTokenVerifyFailedFormat = "%s: %w" - errNonceMismatch = "nonce 不匹配,可能存在重放攻击" - errNoActiveAuthSource = "未配置可用认证源" - errServerAddressMissing = "服务器地址 (server_address) 未配置或配置为空,请在后台系统设置中配置后再试" - errAuthSourceRequired = "认证源不能为空" - errDiscoveryURLRequired = "OIDC 认证源必须配置 Discovery URL" - errUsernameGenerateFailed = "无法生成可用用户名" - errUsernameFromSourceFailed = "无法从认证源获取用户名" - errAuthSourceDisabled = "认证源未启用" - errInvalidExternalAccountBindingID = "绑定记录 ID 无效" - ErrTokenAuthNotAllowed = "该端点不允许使用访问令牌进行身份验证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errOAuthStateRateLimited = "请求授权过于频繁,请稍后重试" - errAuthSourceNameRequired = "认证源名称不能为空" - errAuthSourceNameInvalid = "认证源名称格式不正确" - errAuthSourceTypeUnsupported = "不支持的认证源类型" - errAuthSourceDiscoveryURLRequired = "Discovery URL 不能为空" - //nolint:gosec // error message, not hardcoded credentials - errAuthSourceClientCredentialsRequired = "启用认证源时必须配置 Client ID 和 Client Secret" - errAuthSourceIDRequired = "认证源 ID 不能为空" - errUserIDRequired = "用户 ID 不能为空" - errExternalAccountBindingIncomplete = "外部帐号绑定信息不完整" - errExternalAccountAlreadyBoundToAnother = "该外部帐号已被其他用户绑定" - errExternalAccountBindingIDRequired = "外部帐号绑定记录 ID 不能为空" - errAdminRequired = "无权访问" - //nolint:gosec // error message, not hardcoded credentials - errTokenAdminRequired = "令牌无管理员权限" - errBannedAccount = "账号已被封禁" - errUnAuthorized = "未登录" -) - -// Service 层与鉴权中间件内部错误文案(保持与重构前逐字一致) -const ( - errUserNotInContext = "auth: user not found in context" - errEmptyToken = "auth: empty token" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errSystemUserTokenNotAllowed = "auth: system user token not allowed" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errUnauthorizedInternal = "unauthorized" - errSystemUserLoginNotAllowed = "system user is not allowed to login" -) - -// OAuth 回调会话校验错误文案(保持与重构前逐字一致) -const ( - errInvalidSessionContext = "invalid session context" - errSessionMismatchForOAuth = "session mismatch for oauth state" - errUserContextMismatch = "user context mismatch for oauth binding" -) diff --git a/backend/plugins/domain/auth/facade.go b/backend/plugins/domain/auth/facade.go new file mode 100644 index 00000000..7c772892 --- /dev/null +++ b/backend/plugins/domain/auth/facade.go @@ -0,0 +1,338 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "Wavelet/core/contracts" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/controller" + "Wavelet/plugins/domain/auth/dao" + "Wavelet/plugins/domain/auth/model/do" + "Wavelet/plugins/domain/auth/model/dto" + "Wavelet/plugins/domain/auth/model/entity" + "Wavelet/plugins/domain/auth/service" + "context" + "net/http" + "sync" + + "github.com/gin-contrib/sessions" + "github.com/gin-gonic/gin" +) + +// Exported Type Aliases for Backward Compatibility +// +//nolint:revive // backward compatibility type aliases with legacy naming +type ( + AuthSource = entity.AuthSource + ExternalAccount = entity.ExternalAccount + CachedToken = do.CachedToken + CapRuntimeSettings = do.CapRuntimeSettings + AuthSourceView = dto.AuthSourceView + BasicUserInfo = dto.BasicUserInfo + OAuthAuthorizeResponse = dto.OAuthAuthorizeResponse + OAuthCallbackResult = dto.OAuthCallbackResult + CallbackRequest = dto.CallbackRequest + ChallengeResponse = dto.ChallengeResponse + RedeemResponse = dto.RedeemResponse + CaptchaManager = service.CaptchaManager +) + +// Exported Constant Aliases for Backward Compatibility +const ( + UserNameKey = consts.UserNameKey + UserIDKey = consts.UserIDKey + UserObjKey = consts.UserObjKey + TokenAuthKey = consts.TokenAuthKey + TokenAdminKey = consts.TokenAdminKey + SessionTokenKey = consts.SessionTokenKey + PasswordHashKey = consts.PasswordHashKey + SystemUsername = consts.SystemUsername + + OAuthStateCacheKeyFormat = consts.OAuthStateCacheKeyFormat + OAuthStateCacheKeyExpiration = consts.OAuthStateCacheKeyExpiration + OAuthPurposeLogin = consts.OAuthPurposeLogin + OAuthPurposeBind = consts.OAuthPurposeBind + AuthSourceTypeOIDC = consts.AuthSourceTypeOIDC + + ErrTokenAuthNotAllowed = consts.ErrTokenAuthNotAllowed +) + +var ( + defaultMu sync.RWMutex + defaultDAO = dao.New(nil, nil, nil) + defaultService = service.New(defaultDAO, SessionConfig{SessionCookieName: "wavelet_session", SessionAge: 86400, SessionHTTPOnly: true}, nil) + defaultCtrl = controller.New(defaultService) +) + +func setDefaultRuntime(d *dao.DAO, s *service.Service, c *controller.Controller) { + defaultMu.Lock() + defer defaultMu.Unlock() + defaultDAO = d + defaultService = s + defaultCtrl = c +} + +func getDefaultRuntime() (*dao.DAO, *service.Service, *controller.Controller) { + defaultMu.RLock() + defer defaultMu.RUnlock() + return defaultDAO, defaultService, defaultCtrl +} + +// ParseUserID parses a string, int, or float64 user ID representation. +func ParseUserID(v any) uint64 { + return dto.ParseUserID(v) +} + +// BuildBasicUserInfo converts UserDTO to BasicUserInfo. +func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo { + return dto.BuildBasicUserInfo(user, needChange) +} + +// SetSessionConfig updates the active session configuration. +func SetSessionConfig(cfg SessionConfig) { + _, s, _ := getDefaultRuntime() + s.Session.SetConfig(cfg) +} + +// GetSessionConfig returns the active session configuration. +func GetSessionConfig() SessionConfig { + _, s, _ := getDefaultRuntime() + return s.Session.Config() +} + +// GetSessionOptions builds session cookie options based on config and maxAge. +func GetSessionOptions(maxAge int) sessions.Options { + _, s, _ := getDefaultRuntime() + return s.Session.GetSessionOptions(maxAge) +} + +// StripCookieMaxAgeAndExpires removes max-age and expires from cookie header. +func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) { + _, s, _ := getDefaultRuntime() + s.Session.StripCookieMaxAgeAndExpires(header, cookieName) +} + +// GetUserIDFromSession extracts user ID from session. +func GetUserIDFromSession(s sessions.Session) uint64 { + return controller.GetUserIDFromSession(s) +} + +// GetUserIDFromContext extracts user ID from Gin context. +func GetUserIDFromContext(c *gin.Context) uint64 { + return controller.GetUserIDFromContext(c) +} + +// SetLoginSession sets the login session for the authenticated user. +func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error { + _, s, _ := getDefaultRuntime() + session := sessions.Default(c) + isSessionCookie, err := s.Session.ApplyLoginSession(ctx, session, user, extras...) + if err != nil { + return err + } + if isSessionCookie { + s.Session.StripCookieMaxAgeAndExpires(c.Writer.Header(), s.Session.Config().SessionCookieName) + } + return nil +} + +// RegisterWhitelist adds whitelist path patterns. +func RegisterWhitelist(patterns ...string) { + _, _, c := getDefaultRuntime() + c.RegisterWhitelist(patterns...) +} + +// IsWhitelisted checks if the path matches the auth whitelist. +func IsWhitelisted(path string) bool { + _, _, c := getDefaultRuntime() + if wl := c.Whitelist(); wl != nil { + return wl.Match(path) + } + return false +} + +// GetUserFromRequest extracts user from Request (Token or Session). +func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) { + d, _, _ := getDefaultRuntime() + return controller.GetUserFromRequest(c, d) +} + +// LoginRequired returns authentication required middleware. +func LoginRequired() gin.HandlerFunc { + _, _, c := getDefaultRuntime() + return c.LoginRequired() +} + +// AdminRequired returns admin authorization middleware. +func AdminRequired() gin.HandlerFunc { + _, _, c := getDefaultRuntime() + return c.AdminRequired() +} + +// LoginAdminRequired alias for AdminRequired. +func LoginAdminRequired() gin.HandlerFunc { + return AdminRequired() +} + +// DisallowTokenAuth returns middleware rejecting access token requests. +func DisallowTokenAuth() gin.HandlerFunc { + return controller.DisallowTokenAuth() +} + +// GetCachedToken reads cached access token. +func GetCachedToken(ctx context.Context, tokenHash string) (*CachedToken, error) { + d, _, _ := getDefaultRuntime() + return d.GetCachedToken(ctx, tokenHash) +} + +// SetCachedToken stores access token into cache. +func SetCachedToken(ctx context.Context, tokenHash string, token *CachedToken) { + d, _, _ := getDefaultRuntime() + d.SetCachedToken(ctx, tokenHash, token) +} + +// InvalidateCachedToken invalidates access token cache. +func InvalidateCachedToken(ctx context.Context, tokenHash string) { + d, _, _ := getDefaultRuntime() + d.InvalidateCachedToken(ctx, tokenHash) +} + +// GetCachedUser reads cached user. +func GetCachedUser(ctx context.Context, userID uint64) (*contracts.UserDTO, error) { + d, _, _ := getDefaultRuntime() + return d.GetCachedUser(ctx, userID) +} + +// SetCachedUser stores user into cache. +func SetCachedUser(ctx context.Context, userID uint64, u *contracts.UserDTO) { + d, _, _ := getDefaultRuntime() + d.SetCachedUser(ctx, userID, u) +} + +// InvalidateCachedUser invalidates user cache. +func InvalidateCachedUser(ctx context.Context, userID uint64) { + d, _, _ := getDefaultRuntime() + d.InvalidateCachedUser(ctx, userID) +} + +// ResetAuthRAMCacheForTest clears RAM caches. +func ResetAuthRAMCacheForTest() { + dao.ResetRAMCacheForTest() +} + +// StopAuthCacheListener compatibility stub. +func StopAuthCacheListener() {} + +// SetCapSecret sets CAPTCHA secret. +func SetCapSecret(secret []byte) { + _, s, _ := getDefaultRuntime() + s.CapManager.SetSecret(secret) +} + +// GetDefaultCapManager returns the singleton CAPTCHA manager. +func GetDefaultCapManager() *CaptchaManager { + _, s, _ := getDefaultRuntime() + return s.CapManager +} + +// CurrentCapSettings returns current CAPTCHA runtime settings. +func CurrentCapSettings(ctx context.Context) (CapRuntimeSettings, error) { + _, s, _ := getDefaultRuntime() + return s.CapSettings.Current(ctx) +} + +// CapProtectionEnabled checks if CAPTCHA is enabled. +func CapProtectionEnabled(ctx context.Context) bool { + _, s, _ := getDefaultRuntime() + return s.CapSettings.CapProtectionEnabled(ctx) +} + +// InvalidateCapRuntimeSettings invalidates runtime CAPTCHA settings cache. +func InvalidateCapRuntimeSettings() { + _, s, _ := getDefaultRuntime() + s.CapSettings.Invalidate() +} + +// ResetCapRuntimeSettingsForTest clears test CAPTCHA settings. +func ResetCapRuntimeSettingsForTest() { + InvalidateCapRuntimeSettings() +} + +// InstallCapTestRuntimeSettings installs a test snapshot. +func InstallCapTestRuntimeSettings(settings CapRuntimeSettings) func() { + _, s, _ := getDefaultRuntime() + return s.CapSettings.InstallTestSnapshot(settings) +} + +// VerifyCaptchaMiddleware returns captcha verification middleware. +func VerifyCaptchaMiddleware(mgr *service.CaptchaManager, scope string) gin.HandlerFunc { + _, s, _ := getDefaultRuntime() + return controller.VerifyCaptchaMiddleware(mgr, s.CapSettings, scope) +} + +// Challenge HTTP handler. +func Challenge(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.Captcha.Challenge(c) +} + +// Redeem HTTP handler. +func Redeem(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.Captcha.Redeem(c) +} + +// GetLoginSources HTTP handler. +func GetLoginSources(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.OAuth.GetLoginSources(c) +} + +// GetLoginURL HTTP handler. +func GetLoginURL(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.OAuth.GetLoginURL(c) +} + +// Authorize HTTP handler. +func Authorize(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.OAuth.Authorize(c) +} + +// Callback HTTP handler. +func Callback(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.OAuth.Callback(c) +} + +// Logout HTTP handler. +func Logout(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.OAuth.Logout(c) +} + +// UserInfo HTTP handler. +func UserInfo(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.UserInfo.UserInfo(c) +} + +// ListExternalAccounts HTTP handler. +func ListExternalAccounts(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.OAuth.ListExternalAccounts(c) +} + +// DeleteExternalAccount HTTP handler. +func DeleteExternalAccount(c *gin.Context) { + _, _, ctrl := getDefaultRuntime() + ctrl.OAuth.DeleteExternalAccount(c) +} + +// InvalidateOIDCProviderCache invalidates OIDC provider cache entry. +func InvalidateOIDCProviderCache(issuer string) { + _, s, _ := getDefaultRuntime() + s.OIDCProviderCache.Invalidate(issuer) +} diff --git a/backend/plugins/domain/auth/handlers.go b/backend/plugins/domain/auth/handlers.go deleted file mode 100644 index d96fe321..00000000 --- a/backend/plugins/domain/auth/handlers.go +++ /dev/null @@ -1,572 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package auth - -import ( - "Wavelet/core/contracts" - "Wavelet/pkg/ginutil" - "Wavelet/pkg/idgen" - "Wavelet/pkg/logger" - "Wavelet/pkg/response" - "context" - "errors" - "fmt" - "net/http" - "strconv" - "strings" - "time" - - "github.com/coreos/go-oidc/v3/oidc" - "github.com/gin-contrib/sessions" - "github.com/gin-gonic/gin" - "github.com/google/uuid" - "gorm.io/gorm" -) - -// GetLoginSources 获取可用登录源列表 -// @Summary 获取可用登录源 -// @Description 返回当前系统已启用的所有 OAuth 登录源,前端展示登录按钮列表时调用 -// @Tags oauth -// @Produce json -// @Success 200 {object} response.Any{data=[]auth.AuthSourceView} "登录源列表" -// @Router /api/v1/oauth/sources [get] -func GetLoginSources(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(activeLoginSources(c.Request.Context()))) -} - -// GetLoginURL 获取登录授权地址 -// @Summary 获取登录授权地址 -// @Description 根据指定认证源生成 OAuth 授权 URL,前端跳转到该 URL 完成 OAuth 登录授权。source 参数为空时使用第一个启用的认证源。 -// @Tags oauth -// @Produce json -// @Param source query string false "认证源名称,为空使用第一个启用的认证源" -// @Success 200 {object} response.Any{data=auth.OAuthAuthorizeResponse} "授权 URL" -// @Failure 400 {object} response.Any "认证源不存在或未配置" -// @Failure 500 {object} response.Any "构造 URL 失败" -// @Router /api/v1/oauth/login [get] -func GetLoginURL(c *gin.Context) { - ctx := c.Request.Context() - if !isOIDCLoginEnabled(ctx) { - response.AbortBadRequest(c, errAuthSourceDisabled) - return - } - - source, err := resolveAuthSource(ctx, c.Query("source")) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if !source.IsActive { - response.AbortBadRequest(c, errAuthSourceDisabled) - return - } - - session := sessions.Default(c) - token, isNew := ensureSessionToken(session) - if isNew { - if err := session.Save(); err != nil { - response.AbortInternal(c, err.Error()) - return - } - } - - userID := GetUserIDFromSession(session) - sessionHash := hashSessionToken(token) - if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - state := uuid.NewString() - payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{ - SourceName: source.Name, - Purpose: OAuthPurposeLogin, - UserID: userID, - SessionHash: sessionHash, - }) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state) - if cache := getCache(ctx); cache != nil { - if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil { - response.AbortInternal(c, err.Error()) - return - } - } - - authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL})) -} - -func buildAuthorizeURL(ctx context.Context, source *AuthSource, state string) (string, error) { - redirectURL, err := getFrontendLoginRedirectURL(ctx) - if err != nil { - return "", err - } - authConfig, verifier, err := buildOAuthConfig(ctx, source, redirectURL) - if err != nil { - return "", err - } - if verifier != nil { - return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil - } - return authConfig.AuthCodeURL(state), nil -} - -func reserveOAuthStateSlot(ctx context.Context, sessionHash string) error { - if sessionHash == "" { - return nil - } - cache := getCache(ctx) - if cache == nil { - return nil - } - key := fmt.Sprintf(oauthStateLimitKeyFormat, sessionHash) - var count int - _ = cache.Get(ctx, key, &count) - count++ - _ = cache.Set(ctx, key, count, OAuthStateCacheKeyExpiration) - if count > oauthStateLimitMax { - return errors.New(errOAuthStateRateLimited) - } - return nil -} - -// Authorize 发起指定认证源授权 -// @Summary 发起指定认证源授权 -// @Description 根据指定认证源名称发起 OAuth 授权,支持 purpose 参数用于区分登录和账号绑定场景。认证源必须已启用。 -// @Tags oauth -// @Produce json -// @Param source path string true "认证源名称" -// @Param purpose query string false "授权目的:login(登录)或 bind(绑定账号),默认 login" -// @Success 200 {object} response.Any{data=auth.OAuthAuthorizeResponse} "授权 URL" -// @Failure 400 {object} response.Any "认证源不存在或未启用" -// @Failure 500 {object} response.Any "构造 URL 失败" -// @Router /api/v1/oauth/{source}/authorize [get] -func Authorize(c *gin.Context) { - ctx := c.Request.Context() - if !isOIDCLoginEnabled(ctx) { - response.AbortBadRequest(c, errAuthSourceDisabled) - return - } - - source, err := resolveAuthSource(ctx, c.Param("source")) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if !source.IsActive { - response.AbortBadRequest(c, errAuthSourceDisabled) - return - } - purpose := strings.ToLower(strings.TrimSpace(c.Query("purpose"))) - if purpose != OAuthPurposeBind { - purpose = OAuthPurposeLogin - } - - session := sessions.Default(c) - userID := GetUserIDFromSession(session) - if purpose == OAuthPurposeBind && userID == 0 { - response.AbortUnauthorized(c, errUnAuthorized) - return - } - - token, isNew := ensureSessionToken(session) - if isNew { - if err := session.Save(); err != nil { - response.AbortInternal(c, err.Error()) - return - } - } - - sessionHash := hashSessionToken(token) - if err := reserveOAuthStateSlot(ctx, sessionHash); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - state := uuid.NewString() - payloadValue, err := encodeOAuthStatePayload(oauthStatePayload{ - SourceName: source.Name, - Purpose: purpose, - UserID: userID, - SessionHash: sessionHash, - }) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, state) - if cache := getCache(ctx); cache != nil { - if err := cache.Set(ctx, stateKey, payloadValue, OAuthStateCacheKeyExpiration); err != nil { - response.AbortInternal(c, err.Error()) - return - } - } - - authorizeURL, err := buildAuthorizeURL(c.Request.Context(), source, state) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(OAuthAuthorizeResponse{AuthorizeURL: authorizeURL})) -} - -// Callback OAuth 回调处理 -// @Summary OAuth 回调处理 -// @Description 接收前端传回的 state 和 code,完成 OAuth/OIDC 认证并建立会话。支持登录(login)和账号绑定(bind)两种场景。 -// @Tags oauth -// @Accept json -// @Produce json -// @Param request body auth.CallbackRequest true "回调请求参数" -// @Success 200 {object} response.Any{data=auth.OAuthCallbackResult} "登录或绑定成功" -// @Failure 400 {object} response.Any "state 无效、参数错误或认证源错误" -// @Failure 401 {object} response.Any "绑定场景未登录" -// @Failure 500 {object} response.Any "OAuth 认证失败或内部错误" -// @Router /api/v1/oauth/callback [post] -func Callback(c *gin.Context) { - var req CallbackRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - ctx := c.Request.Context() - stateKey := fmt.Sprintf(OAuthStateCacheKeyFormat, req.State) - var payloadRaw string - cache := getCache(ctx) - if cache == nil { - response.AbortBadRequest(c, errInvalidState) - return - } - if err := cache.Get(ctx, stateKey, &payloadRaw); err != nil { - response.AbortBadRequest(c, errInvalidState) - return - } - _ = cache.Delete(ctx, stateKey) - - payload, err := decodeOAuthStatePayload(payloadRaw) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - session := sessions.Default(c) - currentUserID := GetUserIDFromSession(session) - - if payload.Purpose == OAuthPurposeBind && currentUserID == 0 { - response.AbortUnauthorized(c, errUnAuthorized) - return - } - - token, ok := session.Get(SessionTokenKey).(string) - if !ok || token == "" { - response.AbortBadRequest(c, errInvalidSessionContext) - return - } - - if hashSessionToken(token) != payload.SessionHash { - response.AbortBadRequest(c, errSessionMismatchForOAuth) - return - } - - if payload.Purpose == OAuthPurposeBind && currentUserID != payload.UserID { - response.AbortBadRequest(c, errUserContextMismatch) - return - } - - if !isOIDCLoginEnabled(ctx) { - response.AbortBadRequest(c, errAuthSourceDisabled) - return - } - - source, err := resolveAuthSource(ctx, payload.SourceName) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - if !source.IsActive { - response.AbortBadRequest(c, errAuthSourceDisabled) - return - } - - redirectURL, err := getFrontendLoginRedirectURL(ctx) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - - userInfo, err := buildOAuthUserInfo(ctx, source, req.Code, req.State, redirectURL) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - if err := normalizeOAuthUserInfo(userInfo); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - if userInfo.Sub == "" { - userInfo.Sub = userInfo.Username - } - - if payload.Purpose == OAuthPurposeBind { - handleCallbackBind(ctx, c, source, userInfo) - return - } - - handleCallbackLogin(ctx, c, source, userInfo) -} - -func handleCallbackBind(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) { - userID := GetUserIDFromContext(c) - if userID == 0 { - response.AbortUnauthorized(c, errUnAuthorized) - return - } - user, err := GetUserByID(ctx, userID) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - if err := BindExternalAccount(ctx, &ExternalAccount{ - AuthSourceID: source.ID, - UserID: user.ID, - ExternalID: userInfo.Sub, - ExternalUsername: userInfo.Username, - Email: userInfo.Email, - }); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - user.LastLoginAt = time.Now() - _ = TouchUserLastLogin(ctx, user.ID, user.LastLoginAt) - c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "bound"))) -} - -func handleCallbackLogin(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) { - var user *contracts.UserDTO - - account, err := FindExternalAccount(ctx, source.ID, userInfo.Sub) - switch { - case err == nil: - loaded, loadErr := 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 { - return - } - user = &newUser - default: - response.AbortInternal(c, err.Error()) - return - } - - user.LastLoginAt = time.Now() - _ = TouchUserLastLogin(ctx, user.ID, user.LastLoginAt) - if err := SetLoginSession(ctx, c, user); err != nil { - response.AbortInternal(c, err.Error()) - return - } - - SetCachedUser(ctx, user.ID, user) - - c.JSON(http.StatusOK, response.OK(buildCallbackResult(user, "logged_in"))) -} - -func uniqueUsername(ctx context.Context, base string) (string, error) { - base = strings.TrimSpace(base) - if base == "" { - base = "user" - } - - existingUsernames, err := ListSimilarUsernames(ctx, base) - if err != nil { - return "", err - } - - exists := make(map[string]bool, len(existingUsernames)) - for _, u := range existingUsernames { - exists[strings.ToLower(u)] = true - } - - if !exists[strings.ToLower(base)] { - return base, nil - } - - for i := 1; i <= 1000; i++ { - candidate := fmt.Sprintf("%s-%d", base, i) - if !exists[strings.ToLower(candidate)] { - return candidate, nil - } - } - - return "", errors.New(errUsernameGenerateFailed) -} - -func handleCallbackRegister(ctx context.Context, c *gin.Context, source *AuthSource, userInfo *contracts.OAuthUserInfoDTO) (contracts.UserDTO, bool) { - registrationEnabled := true - val, cfgErr := GetSystemConfigValue(ctx, "registration_enabled") - if cfgErr == nil && val != "" { - if b, err := strconv.ParseBool(val); err == nil { - registrationEnabled = b - } - } - - if !registrationEnabled { - c.JSON(http.StatusOK, response.OK(buildCallbackResult(nil, "need_bind"))) - return contracts.UserDTO{}, false - } - - username, uniqueErr := uniqueUsername(ctx, userInfo.Username) - if uniqueErr != nil { - response.AbortInternal(c, uniqueErr.Error()) - return contracts.UserDTO{}, false - } - userInfo.Username = username - - now := time.Now() - user := contracts.UserDTO{ - ID: idgen.NextUint64ID(), - Username: userInfo.Username, - Nickname: userInfo.Name, - Email: userInfo.Email, - AvatarURL: userInfo.AvatarURL, - IsActive: userInfo.Active, - LastLoginAt: now, - CreatedAt: now, - UpdatedAt: now, - } - - if err := InsertUser(ctx, &user); err != nil { - response.AbortInternal(c, err.Error()) - return contracts.UserDTO{}, false - } - - if err := BindExternalAccount(ctx, &ExternalAccount{ - AuthSourceID: source.ID, - UserID: user.ID, - ExternalID: userInfo.Sub, - ExternalUsername: userInfo.Username, - Email: userInfo.Email, - }); err != nil { - response.AbortBadRequest(c, err.Error()) - return contracts.UserDTO{}, false - } - logger.InfoF(ctx, "[LoginAudit] successful OAuth registration via source: %s, external ID: %s, user: %s, ID: %d, IP: %s", source.Name, userInfo.Sub, user.Username, user.ID, c.ClientIP()) - - return user, true -} - -// UserInfo 获取当前登录用户信息 -// @Summary 获取当前登录用户信息 -// @Description 返回当前登录用户的基本信息,需要登录。 -// @Tags oauth -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=auth.BasicUserInfo} "用户信息" -// @Failure 401 {object} response.Any "未登录" -// @Router /api/v1/oauth/user-info [get] -// @Router /api/v1/user-info [get] -func UserInfo(c *gin.Context) { - user, _ := ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) - session := sessions.Default(c) - needChange := session.Get("need_change_password") == true || (user != nil && user.NeedChangePassword) - - c.JSON( - http.StatusOK, - response.OK(BuildBasicUserInfo(user, needChange)), - ) -} - -// Logout 退出登录 -// @Summary 退出登录 -// @Description 清除当前用户的登录会话,完成退出。清除 Cookie 中的 Session 数据。 -// @Tags oauth -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any{data=string} "退出成功" -// @Failure 500 {object} response.Any "Session 清除失败" -// @Router /api/v1/oauth/logout [get] -func Logout(c *gin.Context) { - session := sessions.Default(c) - userID := session.Get(UserIDKey) - username := session.Get(UserNameKey) - if userID != nil { - logger.InfoF(c.Request.Context(), "[LoginAudit] user logged out: %v, ID: %v, IP: %s", username, userID, c.ClientIP()) - if id := ParseUserID(userID); id > 0 { - InvalidateCachedUser(c.Request.Context(), id) - } - } - session.Options(GetSessionOptions(-1)) - session.Clear() - if err := session.Save(); err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} - -// ListExternalAccounts 获取当前用户的外部帐号绑定列表 -// @Summary 获取外部帐号列表 -// @Description 返回当前登录用户已绑定的所有外部 OAuth 帐号信息,需要登录 -// @Tags oauth -// @Produce json -// @Security SessionCookie -// @Success 200 {object} response.Any "外部帐号列表" -// @Failure 401 {object} response.Any "未登录" -// @Failure 500 {object} response.Any "内部错误" -// @Router /api/v1/oauth/external-accounts [get] -func ListExternalAccounts(c *gin.Context) { - userID := GetUserIDFromContext(c) - accounts, err := ListExternalAccountsByUserID(c.Request.Context(), userID) - if err != nil { - response.AbortInternal(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(accounts)) -} - -// DeleteExternalAccount 解除外部帐号绑定 -// @Summary 解除外部帐号绑定 -// @Description 解除当前登录用户与指定外部帐号的绑定关系,需要登录 -// @Tags oauth -// @Produce json -// @Security SessionCookie -// @Param id path uint64 true "外部帐号绑定记录 ID" -// @Success 200 {object} response.Any{data=string} "解除绑定成功" -// @Failure 400 {object} response.Any "ID 无效或解除失败" -// @Failure 401 {object} response.Any "未登录" -// @Router /api/v1/oauth/external-accounts/{id}/delete [post] -func DeleteExternalAccount(c *gin.Context) { - userID := GetUserIDFromContext(c) - if userID == 0 { - response.AbortUnauthorized(c, errUnAuthorized) - return - } - rawID := strings.TrimSpace(c.Param("id")) - id, err := strconv.ParseUint(rawID, 10, 64) - if err != nil || id == 0 { - response.AbortBadRequest(c, errInvalidExternalAccountBindingID) - return - } - if err := UnbindExternalAccount(c.Request.Context(), id, userID); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OKNil()) -} diff --git a/backend/plugins/domain/auth/middleware.go b/backend/plugins/domain/auth/middleware.go deleted file mode 100644 index ea93a667..00000000 --- a/backend/plugins/domain/auth/middleware.go +++ /dev/null @@ -1,197 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package auth - -import ( - "Wavelet/core/contracts" - "Wavelet/core/extpoints" - "Wavelet/pkg/ginutil" - "Wavelet/pkg/response" - "Wavelet/pkg/trace" - "context" - "crypto/sha256" - "encoding/hex" - "errors" - - "github.com/gin-gonic/gin" -) - -// whitelist holds the no-auth route patterns. They are registered during Apply and -// matched on every request, so PathWhitelist parses them once up front. -var whitelist = extpoints.NewPathWhitelist() - -// RegisterWhitelist registers route patterns that bypass mandatory authentication. -func RegisterWhitelist(patterns ...string) { - whitelist.Add(patterns...) -} - -// IsWhitelisted checks if the specified path matches the auth whitelist. -func IsWhitelisted(path string) bool { - return whitelist.Match(path) -} - -func hashToken(token string) string { - h := sha256.New() - h.Write([]byte(token)) - return hex.EncodeToString(h.Sum(nil)) -} - -// currentUserIDFromRequestContext 是接入层向 Service 层暴露的登录态桥接。 -// -// Session 读取必须依赖 *gin.Context,而 Service 层禁止 import gin, -// 因此该类型断言收敛在本(接入层)文件中。ok 为 false 表示 ctx 不是 *gin.Context。 -func currentUserIDFromRequestContext(ctx context.Context) (uint64, bool) { - ginCtx, ok := ctx.(*gin.Context) - if !ok { - return 0, false - } - return GetUserIDFromContext(ginCtx), true -} - -func getUserByToken(ctx context.Context, tokenStr string) (*contracts.UserDTO, *CachedToken, error) { - tokenHash := hashToken(tokenStr) - tokenRecord, err := GetCachedToken(ctx, tokenHash) - if err != nil || tokenRecord == nil { - tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash) - if err != nil { - return nil, nil, err - } - SetCachedToken(ctx, tokenHash, tokenRecord) - } - - user, err := GetCachedUser(ctx, tokenRecord.UserID) - if err != nil || user == nil || !user.IsActive { - user, err = GetActiveUserByID(ctx, tokenRecord.UserID) - if err != nil { - return nil, nil, err - } - SetCachedUser(ctx, tokenRecord.UserID, user) - } - - return user, tokenRecord, nil -} - -// GetUserFromRequest 从请求中获取当前用户(优先 Access Token,其次 Session) -func GetUserFromRequest(c *gin.Context) (*contracts.UserDTO, error) { - ctx := c.Request.Context() - var tokenStr string - - tokenFromQuery := c.Query("token") - if tokenFromQuery != "" { - tokenStr = tokenFromQuery - } else { - authHeader := c.GetHeader("Authorization") - if len(authHeader) > 7 && authHeader[:7] == "Bearer " { - tokenStr = authHeader[7:] - } - } - - // 优先使用 Access Token 鉴权 - if tokenStr != "" { - if user, tokenRecord, err := getUserByToken(ctx, tokenStr); err == nil { - if user.Username == SystemUsername { - return nil, errors.New(errSystemUserLoginNotAllowed) - } - ginutil.SetToContext(c, contracts.AuthTokenAuthKey, true) - ginutil.SetToContext(c, contracts.AuthTokenAdminKey, tokenRecord.IsAdmin) - return user, nil - } - } - - // 降级使用 Session 鉴权 - userID := GetUserIDFromContext(c) - if userID <= 0 { - return nil, errors.New(errUnauthorizedInternal) - } - - user, err := GetCachedUser(ctx, userID) - if err != nil || user == nil || !user.IsActive { - user, err = GetActiveUserByID(ctx, userID) - if err != nil { - return nil, err - } - SetCachedUser(ctx, userID, user) - } - - ginutil.SetToContext(c, contracts.AuthTokenAuthKey, false) - ginutil.SetToContext(c, contracts.AuthTokenAdminKey, false) - - if user.Username == "system" { - return nil, errors.New(errSystemUserLoginNotAllowed) - } - - return user, nil -} - -// LoginRequired 返回登录鉴权中间件,校验 Access Token 或 Session -func LoginRequired() gin.HandlerFunc { - return func(c *gin.Context) { - if IsWhitelisted(c.Request.URL.Path) { - c.Next() - return - } - - _, span := trace.Start(c.Request.Context(), "LoginRequired") - defer span.End() - - user, err := GetUserFromRequest(c) - if err != nil { - response.AbortUnauthorized(c, errUnAuthorized) - return - } - - LogForAudit(c.Request.Context(), user, c) - ginutil.SetToContext(c, contracts.AuthUserObjKey, user) - c.Next() - } -} - -// AdminRequired 校验管理员权限(支持 Session 和 Token 鉴权) -func AdminRequired() gin.HandlerFunc { - return func(c *gin.Context) { - _, span := trace.Start(c.Request.Context(), "AdminRequired") - defer span.End() - - user, err := GetUserFromRequest(c) - if err != nil { - response.AbortUnauthorized(c, errUnAuthorized) - return - } - - isTokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey) - isTokenAdmin, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAdminKey) - - // 如果是通过 Token 鉴权,要求该 Token 具备管理员权限或者用户本身为管理员 - if isTokenAuth && !isTokenAdmin && !user.IsAdmin { - response.AbortNotFound(c, errTokenAdminRequired) - return - } - - // 如果是通过 Session 鉴权,直接检查用户的 is_admin 属性 - if !isTokenAuth && !user.IsAdmin { - response.AbortNotFound(c, errAdminRequired) - return - } - - LogForAudit(c.Request.Context(), user, c) - ginutil.SetToContext(c, contracts.AuthUserObjKey, user) - c.Next() - } -} - -// LoginAdminRequired is an alias for AdminRequired. -func LoginAdminRequired() gin.HandlerFunc { - return AdminRequired() -} - -// DisallowTokenAuth 拒绝使用 Access Token 进行身份验证的请求访问该端点 -func DisallowTokenAuth() gin.HandlerFunc { - return func(c *gin.Context) { - if tokenAuth, _ := ginutil.GetFromContext[bool](c, contracts.AuthTokenAuthKey); tokenAuth { - response.AbortForbidden(c, ErrTokenAuthNotAllowed) - return - } - c.Next() - } -} diff --git a/backend/plugins/domain/auth/model/do/cached_token.go b/backend/plugins/domain/auth/model/do/cached_token.go new file mode 100644 index 00000000..cf497cec --- /dev/null +++ b/backend/plugins/domain/auth/model/do/cached_token.go @@ -0,0 +1,12 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package do provides domain data objects for the auth plugin. +package do + +// CachedToken represents the minimal cached representation of an access token. +type CachedToken struct { + ID uint64 `json:"id"` + UserID uint64 `json:"user_id"` + IsAdmin bool `json:"is_admin"` +} diff --git a/backend/plugins/domain/auth/model/do/cap_settings.go b/backend/plugins/domain/auth/model/do/cap_settings.go new file mode 100644 index 00000000..f82f852f --- /dev/null +++ b/backend/plugins/domain/auth/model/do/cap_settings.go @@ -0,0 +1,75 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package do provides domain data objects for the auth plugin. +package do + +import ( + "Wavelet/plugins/domain/auth/consts" + "strconv" + "time" +) + +// CapRuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs. +type CapRuntimeSettings struct { + LoginEnabled bool + ChallengeCount int + ChallengeSize int + ChallengeDifficulty int + ChallengeTTL time.Duration + TokenTTL time.Duration +} + +// CapConfigRecord maps the columns selected from the system config table. +type CapConfigRecord struct { + Key string `gorm:"column:key"` + Value string `gorm:"column:value"` +} + +// ParseCapRuntimeSettings parses system config key-value map into CapRuntimeSettings with fallback defaults. +func ParseCapRuntimeSettings(configs map[string]string) CapRuntimeSettings { + settings := CapRuntimeSettings{ + ChallengeCount: consts.DefaultCapChallengeCount, + ChallengeSize: consts.DefaultCapChallengeSize, + ChallengeDifficulty: consts.DefaultCapChallengeDifficulty, + ChallengeTTL: consts.DefaultCapChallengeTTL, + TokenTTL: consts.DefaultCapTokenTTL, + } + + if len(configs) == 0 { + return settings + } + + if val, ok := configs[consts.ConfigKeyCapLoginEnabled]; ok { + if enabled, err := strconv.ParseBool(val); err == nil { + settings.LoginEnabled = enabled + } + } + if val, ok := configs[consts.ConfigKeyCapChallengeCount]; ok { + if count, err := strconv.Atoi(val); err == nil && count > 0 { + settings.ChallengeCount = count + } + } + if val, ok := configs[consts.ConfigKeyCapChallengeSize]; ok { + if size, err := strconv.Atoi(val); err == nil && size > 0 { + settings.ChallengeSize = size + } + } + if val, ok := configs[consts.ConfigKeyCapChallengeDifficulty]; ok { + if diff, err := strconv.Atoi(val); err == nil && diff > 0 { + settings.ChallengeDifficulty = diff + } + } + if val, ok := configs[consts.ConfigKeyCapChallengeTTL]; ok { + if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 { + settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second + } + } + if val, ok := configs[consts.ConfigKeyCapTokenTTL]; ok { + if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 { + settings.TokenTTL = time.Duration(ttlSeconds) * time.Second + } + } + + return settings +} diff --git a/backend/plugins/domain/auth/model/do/cap_settings_test.go b/backend/plugins/domain/auth/model/do/cap_settings_test.go new file mode 100644 index 00000000..a2692b0c --- /dev/null +++ b/backend/plugins/domain/auth/model/do/cap_settings_test.go @@ -0,0 +1,43 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package do_test + +import ( + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/model/do" + "testing" + "time" + + "github.com/stretchr/testify/assert" +) + +func TestParseCapRuntimeSettings(t *testing.T) { + t.Run("Default fallback on empty config", func(t *testing.T) { + settings := do.ParseCapRuntimeSettings(nil) + assert.False(t, settings.LoginEnabled) + assert.Equal(t, consts.DefaultCapChallengeCount, settings.ChallengeCount) + assert.Equal(t, consts.DefaultCapChallengeSize, settings.ChallengeSize) + assert.Equal(t, consts.DefaultCapChallengeDifficulty, settings.ChallengeDifficulty) + assert.Equal(t, consts.DefaultCapChallengeTTL, settings.ChallengeTTL) + assert.Equal(t, consts.DefaultCapTokenTTL, settings.TokenTTL) + }) + + t.Run("Parsed custom configs", func(t *testing.T) { + configs := map[string]string{ + consts.ConfigKeyCapLoginEnabled: "true", + consts.ConfigKeyCapChallengeCount: "3", + consts.ConfigKeyCapChallengeSize: "64", + consts.ConfigKeyCapChallengeDifficulty: "5", + consts.ConfigKeyCapChallengeTTL: "300", + consts.ConfigKeyCapTokenTTL: "600", + } + settings := do.ParseCapRuntimeSettings(configs) + assert.True(t, settings.LoginEnabled) + assert.Equal(t, 3, settings.ChallengeCount) + assert.Equal(t, 64, settings.ChallengeSize) + assert.Equal(t, 5, settings.ChallengeDifficulty) + assert.Equal(t, 300*time.Second, settings.ChallengeTTL) + assert.Equal(t, 600*time.Second, settings.TokenTTL) + }) +} diff --git a/backend/plugins/domain/auth/model/do/oauth_state.go b/backend/plugins/domain/auth/model/do/oauth_state.go new file mode 100644 index 00000000..3cb70029 --- /dev/null +++ b/backend/plugins/domain/auth/model/do/oauth_state.go @@ -0,0 +1,33 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package do provides domain data objects for the auth plugin. +package do + +import "encoding/json" + +// OAuthStatePayload represents the cached state verification payload for OAuth flow. +type OAuthStatePayload struct { + SourceName string `json:"source_name"` + Purpose string `json:"purpose"` + UserID uint64 `json:"user_id,omitempty"` + SessionHash string `json:"session_hash"` +} + +// Encode converts OAuthStatePayload to a JSON string. +func (p OAuthStatePayload) Encode() (string, error) { + data, err := json.Marshal(p) + if err != nil { + return "", err + } + return string(data), nil +} + +// DecodeOAuthStatePayload parses a JSON string into OAuthStatePayload. +func DecodeOAuthStatePayload(value string) (OAuthStatePayload, error) { + var payload OAuthStatePayload + if err := json.Unmarshal([]byte(value), &payload); err != nil { + return OAuthStatePayload{}, err + } + return payload, nil +} diff --git a/backend/plugins/domain/auth/model/do/oauth_state_test.go b/backend/plugins/domain/auth/model/do/oauth_state_test.go new file mode 100644 index 00000000..bc480356 --- /dev/null +++ b/backend/plugins/domain/auth/model/do/oauth_state_test.go @@ -0,0 +1,32 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package do_test + +import ( + "Wavelet/plugins/domain/auth/model/do" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestOAuthStatePayload(t *testing.T) { + payload := do.OAuthStatePayload{ + SourceName: "github", + Purpose: "login", + UserID: 12345, + SessionHash: "hash-abc-123", + } + + encoded, err := payload.Encode() + require.NoError(t, err) + assert.NotEmpty(t, encoded) + + decoded, err := do.DecodeOAuthStatePayload(encoded) + require.NoError(t, err) + assert.Equal(t, payload, decoded) + + _, err = do.DecodeOAuthStatePayload("invalid-json") + assert.Error(t, err) +} diff --git a/backend/plugins/domain/auth/model/dto/auth_source.go b/backend/plugins/domain/auth/model/dto/auth_source.go new file mode 100644 index 00000000..f92dbc82 --- /dev/null +++ b/backend/plugins/domain/auth/model/dto/auth_source.go @@ -0,0 +1,45 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package dto provides data transfer objects and views for the auth plugin. +package dto + +// AuthSourceView 登录源展示信息 +type AuthSourceView struct { + ID uint64 `json:"id"` + Name string `json:"name"` + Type string `json:"type"` + DisplayName string `json:"display_name"` + IsActive bool `json:"is_active"` + IconURL string `json:"icon_url"` + ClientSecretConfigured bool `json:"client_secret_configured"` +} + +// OAuthAuthorizeResponse 授权 URL 响应 +type OAuthAuthorizeResponse struct { + AuthorizeURL string `json:"authorize_url"` +} + +// OAuthCallbackResult 回调处理结果 +type OAuthCallbackResult struct { + Status string `json:"status"` + User *BasicUserInfo `json:"user,omitempty"` +} + +// CallbackRequest OAuth 回调请求参数 +type CallbackRequest struct { + State string `json:"state" binding:"required"` + Code string `json:"code" binding:"required"` +} + +// ExternalAccountView 外部帐号绑定视图(脱敏展示用) +type ExternalAccountView struct { + ID uint64 `json:"id"` + AuthSourceID uint64 `json:"auth_source_id"` + AuthSourceName string `json:"auth_source_name"` + AuthSourceType string `json:"auth_source_type"` + AuthSourceLabel string `json:"auth_source_label"` + ExternalUsername string `json:"external_username"` + Email string `json:"email"` + CreatedAt string `json:"created_at"` +} diff --git a/backend/plugins/domain/cap/models.go b/backend/plugins/domain/auth/model/dto/cap.go similarity index 62% rename from backend/plugins/domain/cap/models.go rename to backend/plugins/domain/auth/model/dto/cap.go index 645a85ef..450dc501 100644 --- a/backend/plugins/domain/cap/models.go +++ b/backend/plugins/domain/auth/model/dto/cap.go @@ -1,22 +1,23 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package cap +// Package dto provides data transfer objects and views for the auth plugin. +package dto import ( - "Wavelet/plugins/domain/cap/pow" + "Wavelet/plugins/domain/auth/pow" ) // ChallengeResponse is a local type alias for the pow.ChallengeResponse struct type ChallengeResponse = pow.ChallengeResponse -// challengeRequest is the CAPTCHA challenge request payload. -type challengeRequest struct { +// ChallengeRequest is the CAPTCHA challenge request payload. +type ChallengeRequest struct { Scope string `json:"scope" form:"scope"` } -// redeemRequest is the CAPTCHA redeem request payload. -type redeemRequest struct { +// RedeemRequest is the CAPTCHA redeem request payload. +type RedeemRequest struct { Token string `json:"token" binding:"required"` Solutions []int `json:"solutions" binding:"required"` Scope string `json:"scope" form:"scope"` @@ -29,9 +30,3 @@ type RedeemResponse struct { Expires int64 `json:"expires,omitempty"` Error string `json:"error,omitempty"` } - -// configRecord maps the columns selected from the system config table. -type configRecord struct { - Key string `gorm:"column:key"` - Value string `gorm:"column:value"` -} diff --git a/backend/plugins/domain/auth/model/dto/user_info.go b/backend/plugins/domain/auth/model/dto/user_info.go new file mode 100644 index 00000000..795ee22e --- /dev/null +++ b/backend/plugins/domain/auth/model/dto/user_info.go @@ -0,0 +1,84 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package dto provides data transfer objects and views for the auth plugin. +package dto + +import ( + "Wavelet/core/contracts" + "strconv" +) + +// BasicUserInfo 用户基本信息结构体 +type BasicUserInfo struct { + ID uint64 `json:"id,string"` + Username string `json:"username"` + Nickname string `json:"nickname"` + Email string `json:"email"` + AvatarURL string `json:"avatar_url"` + IsAdmin bool `json:"is_admin"` + NeedChangePassword bool `json:"need_change_password"` + Bio string `json:"bio"` + Phone string `json:"phone"` + Gender string `json:"gender"` + Website string `json:"website"` + Location string `json:"location"` +} + +// BuildBasicUserInfo 将 UserDTO 转换为 BasicUserInfo +func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo { + if user == nil { + return BasicUserInfo{} + } + return BasicUserInfo{ + ID: user.ID, + Username: user.Username, + Nickname: user.Nickname, + Email: user.Email, + AvatarURL: user.AvatarURL, + IsAdmin: user.IsAdmin, + NeedChangePassword: needChange || user.NeedChangePassword, + Bio: user.Bio, + Phone: user.Phone, + Gender: user.Gender, + Website: user.Website, + Location: user.Location, + } +} + +// LoginRequiredAuditLog 审计日志结构体 +type LoginRequiredAuditLog struct { + UserID uint64 `json:"user_id"` + Username string `json:"username"` + ClientIP string `json:"client_ip"` + Method string `json:"method"` + Path string `json:"path"` + RequestURI string `json:"request_uri"` + UserAgent string `json:"user_agent"` + Referer string `json:"referer"` +} + +// ParseUserID parses a string, int, or float64 user ID representation. +func ParseUserID(v any) uint64 { + switch val := v.(type) { + case uint64: + return val + case int64: + if val > 0 { + return uint64(val) + } + case int: + if val > 0 { + return uint64(val) + } + case float64: + if val > 0 { + return uint64(val) + } + case string: + if id, err := strconv.ParseUint(val, 10, 64); err == nil { + return id + } + } + return 0 +} diff --git a/backend/plugins/domain/auth/model/entity/auth_source.go b/backend/plugins/domain/auth/model/entity/auth_source.go new file mode 100644 index 00000000..9ebd2f11 --- /dev/null +++ b/backend/plugins/domain/auth/model/entity/auth_source.go @@ -0,0 +1,83 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package entity provides database model entities for the auth domain plugin. +package entity + +import ( + "Wavelet/plugins/domain/auth/consts" + "errors" + "regexp" + "strings" + "time" +) + +var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`) + +// AuthSource 认证源实体 +type AuthSource struct { + ID uint64 `json:"id" gorm:"primaryKey"` + Name string `json:"name" gorm:"uniqueIndex;size:80;not null"` + Type string `json:"type" gorm:"size:20;not null"` + DisplayName string `json:"display_name" gorm:"size:100"` + IsActive bool `json:"is_active" gorm:"index;not null;default:false"` + ClientID string `json:"client_id" gorm:"size:255"` + ClientSecret string `json:"-" gorm:"size:1024"` + OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"` + Scopes string `json:"scopes" gorm:"size:255"` + IconURL string `json:"icon_url" gorm:"size:1024"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` + ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"` +} + +// TableName 表名 +func (AuthSource) TableName() string { + return "w_auth_sources" +} + +// Normalize 对认证源字段进行标准化处理 +func (source *AuthSource) Normalize() { + source.Type = strings.ToLower(strings.TrimSpace(source.Type)) + source.Name = strings.TrimSpace(source.Name) + source.DisplayName = strings.TrimSpace(source.DisplayName) + source.ClientID = strings.TrimSpace(source.ClientID) + source.ClientSecret = strings.TrimSpace(source.ClientSecret) + source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL) + source.Scopes = strings.TrimSpace(source.Scopes) + source.IconURL = strings.TrimSpace(source.IconURL) + if source.DisplayName == "" { + source.DisplayName = source.Name + } + if source.Type == consts.AuthSourceTypeOIDC && source.Scopes == "" { + source.Scopes = "openid profile email" + } +} + +// Validate 校验认证源字段合法性 +func (source *AuthSource) Validate() error { + source.Normalize() + if source.Name == "" { + return errors.New(consts.ErrAuthSourceNameRequired) + } + if !authSourceNamePattern.MatchString(source.Name) { + return errors.New(consts.ErrAuthSourceNameInvalid) + } + if source.Type != consts.AuthSourceTypeOIDC { + return errors.New(consts.ErrAuthSourceTypeUnsupported) + } + if source.OpenIDDiscoveryURL == "" { + //nolint:staticcheck // descriptive error constant + return errors.New(consts.ErrAuthSourceDiscoveryURLRequired) + } + if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") { + return errors.New(consts.ErrAuthSourceClientCredentialsRequired) + } + return nil +} + +// Sanitize 脱敏处理,将 ClientSecret 清空并设置 ClientSecretConfigured 标志 +func (source *AuthSource) Sanitize() { + source.ClientSecretConfigured = source.ClientSecret != "" + source.ClientSecret = "" +} diff --git a/backend/plugins/domain/auth/model/entity/auth_source_test.go b/backend/plugins/domain/auth/model/entity/auth_source_test.go new file mode 100644 index 00000000..bde666db --- /dev/null +++ b/backend/plugins/domain/auth/model/entity/auth_source_test.go @@ -0,0 +1,75 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package entity_test + +import ( + "Wavelet/plugins/domain/auth/model/entity" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestAuthSourceValidation(t *testing.T) { + t.Run("Valid OIDC Source", func(t *testing.T) { + src := entity.AuthSource{ + Name: "google", + Type: "oidc", + DisplayName: "Google Sign-In", + ClientID: "client-123", + ClientSecret: "secret-456", + OpenIDDiscoveryURL: "https://accounts.google.com", + IsActive: true, + } + require.NoError(t, src.Validate()) + assert.Equal(t, "openid profile email", src.Scopes) + assert.Equal(t, "w_auth_sources", src.TableName()) + + src.Sanitize() + assert.True(t, src.ClientSecretConfigured) + assert.Empty(t, src.ClientSecret) + }) + + t.Run("Empty Name Fails", func(t *testing.T) { + src := entity.AuthSource{ + Name: "", + Type: "oidc", + } + assert.Error(t, src.Validate()) + }) + + t.Run("Invalid Name Format Fails", func(t *testing.T) { + src := entity.AuthSource{ + Name: "invalid name with spaces!", + Type: "oidc", + } + assert.Error(t, src.Validate()) + }) + + t.Run("Unsupported Type Fails", func(t *testing.T) { + src := entity.AuthSource{ + Name: "ldap_source", + Type: "ldap", + } + assert.Error(t, src.Validate()) + }) + + t.Run("Missing Discovery URL Fails", func(t *testing.T) { + src := entity.AuthSource{ + Name: "google", + Type: "oidc", + } + assert.Error(t, src.Validate()) + }) + + t.Run("Active Source Missing Credentials Fails", func(t *testing.T) { + src := entity.AuthSource{ + Name: "google", + Type: "oidc", + OpenIDDiscoveryURL: "https://accounts.google.com", + IsActive: true, + } + assert.Error(t, src.Validate()) + }) +} diff --git a/backend/plugins/domain/auth/model/entity/external_account.go b/backend/plugins/domain/auth/model/entity/external_account.go new file mode 100644 index 00000000..66d1b967 --- /dev/null +++ b/backend/plugins/domain/auth/model/entity/external_account.go @@ -0,0 +1,26 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package entity provides database model entities for the auth domain plugin. +package entity + +import ( + "time" +) + +// ExternalAccount 外部账号绑定实体 +type ExternalAccount struct { + ID uint64 `json:"id" gorm:"primaryKey"` + AuthSourceID uint64 `json:"auth_source_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:1;index"` + UserID uint64 `json:"user_id" gorm:"index;not null"` + ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:2;size:255;not null"` + ExternalUsername string `json:"external_username" gorm:"size:255"` + Email string `json:"email" gorm:"size:255"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// TableName 表名 +func (ExternalAccount) TableName() string { + return "w_external_accounts" +} diff --git a/backend/plugins/domain/auth/models.go b/backend/plugins/domain/auth/models.go deleted file mode 100644 index 74289f2f..00000000 --- a/backend/plugins/domain/auth/models.go +++ /dev/null @@ -1,241 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package auth - -import ( - "Wavelet/core/contracts" - "encoding/json" - "errors" - "regexp" - "strconv" - "strings" - "time" -) - -var authSourceNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_-]{0,79}$`) - -// AuthSource 认证源实体 -// -//nolint:revive // auth.AuthSource is standard domain entity name -type AuthSource struct { - ID uint64 `json:"id" gorm:"primaryKey"` - Name string `json:"name" gorm:"uniqueIndex;size:80;not null"` - Type string `json:"type" gorm:"size:20;not null"` - DisplayName string `json:"display_name" gorm:"size:100"` - IsActive bool `json:"is_active" gorm:"index;not null;default:false"` - ClientID string `json:"client_id" gorm:"size:255"` - ClientSecret string `json:"-" gorm:"size:1024"` - OpenIDDiscoveryURL string `json:"openid_discovery_url" gorm:"column:openid_discovery_url;size:1024"` - Scopes string `json:"scopes" gorm:"size:255"` - IconURL string `json:"icon_url" gorm:"size:1024"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` - ClientSecretConfigured bool `json:"client_secret_configured" gorm:"-"` -} - -// TableName 表名 -func (AuthSource) TableName() string { - return "w_auth_sources" -} - -// Normalize 对认证源字段进行标准化处理 -func (source *AuthSource) Normalize() { - source.Type = strings.ToLower(strings.TrimSpace(source.Type)) - source.Name = strings.TrimSpace(source.Name) - source.DisplayName = strings.TrimSpace(source.DisplayName) - source.ClientID = strings.TrimSpace(source.ClientID) - source.ClientSecret = strings.TrimSpace(source.ClientSecret) - source.OpenIDDiscoveryURL = strings.TrimSpace(source.OpenIDDiscoveryURL) - source.Scopes = strings.TrimSpace(source.Scopes) - source.IconURL = strings.TrimSpace(source.IconURL) - if source.DisplayName == "" { - source.DisplayName = source.Name - } - if source.Type == AuthSourceTypeOIDC && source.Scopes == "" { - source.Scopes = "openid profile email" - } -} - -// Validate 校验认证源字段合法性 -func (source *AuthSource) Validate() error { - source.Normalize() - if source.Name == "" { - return errors.New(errAuthSourceNameRequired) - } - if !authSourceNamePattern.MatchString(source.Name) { - return errors.New(errAuthSourceNameInvalid) - } - if source.Type != AuthSourceTypeOIDC { - return errors.New(errAuthSourceTypeUnsupported) - } - if source.OpenIDDiscoveryURL == "" { - //nolint:staticcheck // descriptive error constant - return errors.New(errAuthSourceDiscoveryURLRequired) - } - if source.IsActive && (source.ClientID == "" || source.ClientSecret == "") { - return errors.New(errAuthSourceClientCredentialsRequired) - } - return nil -} - -// Sanitize 脱敏处理,将 ClientSecret 清空并设置 ClientSecretConfigured 标志 -func (source *AuthSource) Sanitize() { - source.ClientSecretConfigured = source.ClientSecret != "" - source.ClientSecret = "" -} - -// ExternalAccount 外部账号绑定实体 -type ExternalAccount struct { - ID uint64 `json:"id" gorm:"primaryKey"` - AuthSourceID uint64 `json:"auth_source_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:1;index"` - UserID uint64 `json:"user_id" gorm:"index;not null"` - ExternalID string `json:"external_id" gorm:"uniqueIndex:idx_external_accounts_source_external,priority:2;size:255;not null"` - ExternalUsername string `json:"external_username" gorm:"size:255"` - Email string `json:"email" gorm:"size:255"` - CreatedAt time.Time `json:"created_at"` - UpdatedAt time.Time `json:"updated_at"` -} - -// TableName 表名 -func (ExternalAccount) TableName() string { - return "w_external_accounts" -} - -// ExternalAccountView 外部帐号绑定视图(脱敏展示用) -type ExternalAccountView struct { - ID uint64 `json:"id"` - AuthSourceID uint64 `json:"auth_source_id"` - AuthSourceName string `json:"auth_source_name"` - AuthSourceType string `json:"auth_source_type"` - AuthSourceLabel string `json:"auth_source_label"` - ExternalUsername string `json:"external_username"` - Email string `json:"email"` - CreatedAt time.Time `json:"created_at"` -} - -// AuthSourceView 登录源展示信息 -// -//nolint:revive // auth.AuthSourceView is standard domain presentation struct -type AuthSourceView struct { - ID uint64 `json:"id"` - Name string `json:"name"` - Type string `json:"type"` - DisplayName string `json:"display_name"` - IsActive bool `json:"is_active"` - IconURL string `json:"icon_url"` - ClientSecretConfigured bool `json:"client_secret_configured"` -} - -// OAuthAuthorizeResponse 授权 URL 响应 -type OAuthAuthorizeResponse struct { - AuthorizeURL string `json:"authorize_url"` -} - -// OAuthCallbackResult 回调处理结果 -type OAuthCallbackResult struct { - Status string `json:"status"` - User *BasicUserInfo `json:"user,omitempty"` -} - -// CallbackRequest OAuth 回调请求参数 -type CallbackRequest struct { - State string `json:"state" binding:"required"` - Code string `json:"code" binding:"required"` -} - -// BasicUserInfo 用户基本信息结构体 -type BasicUserInfo struct { - ID uint64 `json:"id"` - Username string `json:"username"` - Nickname string `json:"nickname"` - Email string `json:"email"` - AvatarURL string `json:"avatar_url"` - IsAdmin bool `json:"is_admin"` - NeedChangePassword bool `json:"need_change_password"` - Bio string `json:"bio"` - Phone string `json:"phone"` - Gender string `json:"gender"` - Website string `json:"website"` - Location string `json:"location"` -} - -// BuildBasicUserInfo 将 UserDTO 转换为 BasicUserInfo -func BuildBasicUserInfo(user *contracts.UserDTO, needChange bool) BasicUserInfo { - if user == nil { - return BasicUserInfo{} - } - return BasicUserInfo{ - ID: user.ID, - Username: user.Username, - Nickname: user.Nickname, - Email: user.Email, - AvatarURL: user.AvatarURL, - IsAdmin: user.IsAdmin, - NeedChangePassword: needChange || user.NeedChangePassword, - Bio: user.Bio, - Phone: user.Phone, - Gender: user.Gender, - Website: user.Website, - Location: user.Location, - } -} - -type oauthStatePayload struct { - SourceName string `json:"source_name"` - Purpose string `json:"purpose"` - UserID uint64 `json:"user_id,omitempty"` - SessionHash string `json:"session_hash"` -} - -func encodeOAuthStatePayload(payload oauthStatePayload) (string, error) { - data, err := json.Marshal(payload) - if err != nil { - return "", err - } - return string(data), nil -} - -func decodeOAuthStatePayload(value string) (oauthStatePayload, error) { - var payload oauthStatePayload - if err := json.Unmarshal([]byte(value), &payload); err != nil { - return oauthStatePayload{}, err - } - return payload, nil -} - -type loginRequiredAuditLog struct { - UserID uint64 `json:"user_id"` - Username string `json:"username"` - ClientIP string `json:"client_ip"` - Method string `json:"method"` - Path string `json:"path"` - RequestURI string `json:"request_uri"` - UserAgent string `json:"user_agent"` - Referer string `json:"referer"` -} - -// ParseUserID parses a string or float64 user ID representation. -func ParseUserID(v any) uint64 { - switch val := v.(type) { - case uint64: - return val - case int64: - if val > 0 { - return uint64(val) - } - case int: - if val > 0 { - return uint64(val) - } - case float64: - if val > 0 { - return uint64(val) - } - case string: - if id, err := strconv.ParseUint(val, 10, 64); err == nil { - return id - } - } - return 0 -} diff --git a/backend/plugins/domain/auth/oauth_rate_limit_test.go b/backend/plugins/domain/auth/oauth_rate_limit_test.go new file mode 100644 index 00000000..cb806cb8 --- /dev/null +++ b/backend/plugins/domain/auth/oauth_rate_limit_test.go @@ -0,0 +1,108 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth_test + +import ( + "Wavelet/core" + "Wavelet/pkg/response" + "Wavelet/plugins/domain/auth" + "Wavelet/plugins/infra/cache_memory" + database "Wavelet/plugins/infra/database" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-contrib/sessions" + "github.com/gin-contrib/sessions/cookie" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestOAuthRateLimiting(t *testing.T) { + gin.SetMode(gin.TestMode) + + ctx := core.NewContext(context.Background()) + ctx.Config().SetSource(core.NewMapSource(nil)) + require.NoError(t, ctx.Config().Resolve()) + + testDB := setupTestDB(t) + require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx)) + require.NoError(t, cache_memory.New().Apply(ctx)) + require.NoError(t, auth.New().Apply(ctx)) + + // Create an active OIDC source + authSrc := auth.AuthSource{ + ID: 1, + Name: "google", + Type: "oidc", + DisplayName: "Google", + ClientID: "client-id-123", + ClientSecret: "client-secret-456", + OpenIDDiscoveryURL: "https://accounts.google.com", + IsActive: true, + } + require.NoError(t, testDB.Create(&authSrc).Error) + + router := gin.New() + router.Use(response.ErrorHandlerMiddleware()) + store := cookie.NewStore([]byte("test-session-secret-123")) + router.Use(sessions.Sessions("wavelet_session_id", store)) + router.Use(func(c *gin.Context) { + c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), ctx.Root())) + c.Next() + }) + + for _, rd := range ctx.Router().Routes() { + handlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers)) + for _, m := range rd.Middlewares { + if h, ok := m.(gin.HandlerFunc); ok { + handlers = append(handlers, h) + } else if fn, ok := m.(func(*gin.Context)); ok { + handlers = append(handlers, fn) + } + } + for _, raw := range rd.Handlers { + if h, ok := raw.(gin.HandlerFunc); ok { + handlers = append(handlers, h) + } else if fn, ok := raw.(func(*gin.Context)); ok { + handlers = append(handlers, fn) + } + } + router.Handle(rd.Method, rd.Path, handlers...) + } + + // 10 state slots are allowed per session (oauthStateLimitMax = 10) + // We'll simulate 10 requests with the same cookie + var cookies []*http.Cookie + for i := 1; i <= 10; i++ { + req := httptest.NewRequest(http.MethodGet, "/api/v1/oauth/login?source=google", nil) + for _, ck := range cookies { + req.AddCookie(ck) + } + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + if len(w.Result().Cookies()) > 0 { + cookies = w.Result().Cookies() + } + } + + // 11th request for the same session should be rate limited + { + req := httptest.NewRequest(http.MethodGet, "/api/v1/oauth/login?source=google", nil) + for _, ck := range cookies { + req.AddCookie(ck) + } + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusBadRequest, w.Code) + + var resp map[string]any + err := json.Unmarshal(w.Body.Bytes(), &resp) + require.NoError(t, err) + assert.Equal(t, "请求授权过于频繁,请稍后重试", resp["error_msg"]) + } +} diff --git a/backend/plugins/domain/auth/parse_userid_test.go b/backend/plugins/domain/auth/parse_userid_test.go new file mode 100644 index 00000000..212f10c0 --- /dev/null +++ b/backend/plugins/domain/auth/parse_userid_test.go @@ -0,0 +1,44 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package auth + +import ( + "encoding/json" + "strconv" + "testing" +) + +func TestParseUserIDSnowflakeStringPreservesValue(t *testing.T) { + id := uint64(99835970421002240) + got := ParseUserID(strconv.FormatUint(id, 10)) + if got != id { + t.Errorf("ParseUserID(%q) = %d, want %d", strconv.FormatUint(id, 10), got, id) + } +} + +func TestParseUserIDJSONNumberAboveMaxSafeInteger(t *testing.T) { + // 2^53+1 cannot be represented as a distinct IEEE-754 float64. + id := uint64(9007199254740993) + got := ParseUserID(float64(id)) + if got == id { + t.Fatalf("ParseUserID(float64(%d)) = %d, want a rounded value", id, got) + } +} + +func TestBasicUserInfoJSONEncodesIDAsString(t *testing.T) { + info := BasicUserInfo{ID: 99835970421002240, Username: "plain_user"} + raw, err := json.Marshal(info) + if err != nil { + t.Fatalf("json.Marshal(BasicUserInfo) error = %v", err) + } + var probe struct { + ID json.RawMessage `json:"id"` + } + if err := json.Unmarshal(raw, &probe); err != nil { + t.Fatalf("json.Unmarshal probe error = %v", err) + } + if len(probe.ID) == 0 || probe.ID[0] != '"' { + t.Errorf("BasicUserInfo id JSON = %s, want a JSON string", probe.ID) + } +} diff --git a/backend/plugins/domain/auth/plugin.go b/backend/plugins/domain/auth/plugin.go index b836bb1d..913665b2 100644 --- a/backend/plugins/domain/auth/plugin.go +++ b/backend/plugins/domain/auth/plugin.go @@ -8,6 +8,9 @@ import ( "Wavelet/core" "Wavelet/core/contracts" "Wavelet/core/extpoints" + "Wavelet/plugins/domain/auth/controller" + "Wavelet/plugins/domain/auth/dao" + "Wavelet/plugins/domain/auth/service" "context" "embed" "reflect" @@ -83,46 +86,60 @@ func (p *Plugin) DeclareConfig() []core.ConfigBinding { // Apply registers the auth migrations, services, routes, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { var cfg SessionConfig - if err := ctx.Config().Bind("app", &cfg); err == nil { - SetSessionConfig(cfg) + if err := ctx.Config().Bind("app", &cfg); err != nil { + cfg = SessionConfig{ + SessionCookieName: "wavelet_session", + SessionAge: 86400, + SessionHTTPOnly: true, + } } - // 0. Bind DBService & CacheService from Context - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - setCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - setCacheService(cache) - }) - } + d := dao.New(nil, nil, nil) + core.Bind[contracts.DBService](ctx, d.SetDBService) + core.Bind[contracts.CacheService](ctx, d.SetCacheService) + core.Bind[contracts.LimiterService](ctx, d.SetLimiterService) ctx.OnDispose(func() error { - setDBService(nil) - setCacheService(nil) + d.SetDBService(nil) + d.SetCacheService(nil) + d.SetLimiterService(nil) return nil }) + var capSecret []byte + if cfg.SessionSecret != "" { + capSecret = []byte(cfg.SessionSecret) + } + svc := service.New(d, cfg, capSecret) + + if p.authSvc != nil { + // Custom injected auth service override + core.Provide[contracts.AuthService](ctx, p.authSvc) + } else { + core.Provide[contracts.AuthService](ctx, svc.AuthSvc) + } + + if p.authRegistry != nil { + core.Provide[contracts.AuthRegistry](ctx, p.authRegistry) + } else { + core.Provide[contracts.AuthRegistry](ctx, svc.AuthRegistry) + } + + ctrl := controller.New(svc) + setDefaultRuntime(d, svc, ctrl) + + // Register CaptchaService + captchaSvc := service.NewCaptchaService( + svc.CapManager, + func(scope string) any { return ctrl.VerifyCaptcha(scope) }, + ctrl.Captcha.Challenge, + ctrl.Captcha.Redeem, + ) + core.Provide[contracts.CaptchaService](ctx, captchaSvc) + // 1. Register migrations ctx.Migrations().Register("auth", authMigrations) - // 2. Initialize and provide AuthService & AuthRegistry - if p.authSvc == nil { - p.authSvc = newAuthService() - } - if p.authRegistry == nil { - p.authRegistry = newAuthRegistry() - } - - core.Provide[contracts.AuthService](ctx, p.authSvc) - core.Provide[contracts.AuthRegistry](ctx, p.authRegistry) - - // 2.1 Register Public / Auth Whitelist Endpoints + // 2. Register Public / Auth Whitelist Endpoints publicEndpoints := []string{ "/api/v1/oauth/sources", "/api/v1/oauth/login", @@ -132,54 +149,67 @@ func (p *Plugin) Apply(ctx *core.Context) error { "/api/v1/user/login", "/api/v1/user/register", "/api/v1/user/send-email-code", + "/api/v1/config/public", "/api/v1/cap/challenge", "/api/v1/cap/redeem", "/api/healthz", "/metrics", } - RegisterWhitelist(publicEndpoints...) + ctrl.RegisterWhitelist(publicEndpoints...) ctx.Router().RegisterWhitelist(publicEndpoints...) // 3. Register HTTP Routes - oauthGroup := ctx.Router().Group("/api/v1/oauth") - { - oauthGroup.GET("/sources", GetLoginSources) - oauthGroup.GET("/login", GetLoginURL) - oauthGroup.GET("/:source/authorize", Authorize) - oauthGroup.GET("/logout", Logout) - oauthGroup.POST("/callback", Callback) - oauthGroup.GET("/user-info", LoginRequired(), UserInfo) - oauthGroup.GET("/external-accounts", LoginRequired(), ListExternalAccounts) - oauthGroup.POST("/external-accounts/:id/delete", LoginRequired(), DeleteExternalAccount) - } - ctx.Router().GET("/api/v1/user-info", LoginRequired(), UserInfo) + ctrl.RegisterRoutes(ctx.Router()) // 4. Register Settings Schemas + const ( + settingTypeInteger = "integer" + settingCategorySecurity = "security" + ) + ctx.Settings().Register(extpoints.SettingSchema{ Key: "auth.session_age", Default: 86400 * 7, Description: "Default session lifetime in seconds", - Type: "integer", - Category: "security", + Type: settingTypeInteger, + Category: settingCategorySecurity, }) ctx.Settings().Register(extpoints.SettingSchema{ Key: "auth.login_rate_limit_max_attempts", Default: 5, Description: "Max login failure attempts before temporary IP lock", - Type: "integer", - Category: "security", + Type: settingTypeInteger, + Category: settingCategorySecurity, + }) + ctx.Settings().Register(extpoints.SettingSchema{ + Key: "cap.login_enabled", + Default: false, + Description: "Whether to require CAPTCHA verification for user login", + Type: "boolean", + Category: settingCategorySecurity, + }) + ctx.Settings().Register(extpoints.SettingSchema{ + Key: "cap.challenge_count", + Default: 1, + Description: "Number of PoW puzzle challenges to solve", + Type: settingTypeInteger, + Category: settingCategorySecurity, }) // 5. Register Event Listeners for domain events ctx.Events().On(contracts.EventTopicUserStatusChanged, func(c context.Context, e contracts.UserStatusChangedEvent) error { - InvalidateCachedUser(c, e.UserID) + svc.DAO.InvalidateCachedUser(c, e.UserID) return nil }) ctx.Events().On(contracts.EventTopicUserDeleted, func(c context.Context, e contracts.UserDeletedEvent) error { - InvalidateCachedUser(c, e.TargetUserID) + svc.DAO.InvalidateCachedUser(c, e.TargetUserID) return nil }) + ctx.Events().On(contracts.EventTopicConfigChanged, func(_ any) { + svc.CapSettings.Invalidate() + }) + return nil } diff --git a/backend/plugins/domain/auth/plugin_test.go b/backend/plugins/domain/auth/plugin_test.go index 2c5b52b8..62c10e98 100644 --- a/backend/plugins/domain/auth/plugin_test.go +++ b/backend/plugins/domain/auth/plugin_test.go @@ -40,6 +40,7 @@ type testUser struct { ID uint64 `gorm:"primaryKey"` Username string IsActive bool + IsAdmin bool LastLoginAt time.Time } @@ -61,6 +62,13 @@ func hashToken(token string) string { return hex.EncodeToString(h.Sum(nil)) } +type testSystemConfig struct { + Key string `gorm:"primaryKey"` + Value string +} + +func (testSystemConfig) TableName() string { return "w_system_configs" } + func setupTestDB(t *testing.T) *gorm.DB { t.Helper() dbPath := filepath.Join(t.TempDir(), "auth_test.db") @@ -72,6 +80,7 @@ func setupTestDB(t *testing.T) *gorm.DB { &testAccessToken{}, &auth.AuthSource{}, &auth.ExternalAccount{}, + &testSystemConfig{}, )) return testDB @@ -152,4 +161,25 @@ func TestAuthPluginUnit(t *testing.T) { current, err := authSvc.GetCurrentUser(userCtx) require.NoError(t, err) assert.Equal(t, user.ID, current.ID) + + // Test CaptchaService injection + capSvc, err := core.Inject[contracts.CaptchaService](ctx) + require.NoError(t, err) + assert.NotNil(t, capSvc) + assert.NotNil(t, capSvc.ChallengeHandler()) + assert.NotNil(t, capSvc.RedeemHandler()) + assert.NotNil(t, capSvc.VerifyMiddleware("login")) + + // Verify CAPTCHA routes registered + var foundChallenge, foundRedeem bool + for _, rd := range ctx.Router().Routes() { + if rd.Path == "/api/v1/cap/challenge" { + foundChallenge = true + } + if rd.Path == "/api/v1/cap/redeem" { + foundRedeem = true + } + } + assert.True(t, foundChallenge, "expected /api/v1/cap/challenge route") + assert.True(t, foundRedeem, "expected /api/v1/cap/redeem route") } diff --git a/backend/plugins/domain/cap/pow/cap.go b/backend/plugins/domain/auth/pow/cap.go similarity index 100% rename from backend/plugins/domain/cap/pow/cap.go rename to backend/plugins/domain/auth/pow/cap.go diff --git a/backend/plugins/domain/cap/pow/errs.go b/backend/plugins/domain/auth/pow/errs.go similarity index 100% rename from backend/plugins/domain/cap/pow/errs.go rename to backend/plugins/domain/auth/pow/errs.go diff --git a/backend/plugins/domain/cap/pow/pow_test.go b/backend/plugins/domain/auth/pow/pow_test.go similarity index 100% rename from backend/plugins/domain/cap/pow/pow_test.go rename to backend/plugins/domain/auth/pow/pow_test.go diff --git a/backend/plugins/domain/cap/pow/prng.go b/backend/plugins/domain/auth/pow/prng.go similarity index 100% rename from backend/plugins/domain/cap/pow/prng.go rename to backend/plugins/domain/auth/pow/prng.go diff --git a/backend/plugins/domain/cap/pow/store.go b/backend/plugins/domain/auth/pow/store.go similarity index 100% rename from backend/plugins/domain/cap/pow/store.go rename to backend/plugins/domain/auth/pow/store.go diff --git a/backend/plugins/domain/auth/repository.go b/backend/plugins/domain/auth/repository.go deleted file mode 100644 index c245e924..00000000 --- a/backend/plugins/domain/auth/repository.go +++ /dev/null @@ -1,215 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package auth - -import ( - "Wavelet/core" - "Wavelet/core/contracts" - "Wavelet/pkg/util" - "context" - "sync" - "time" - - "gorm.io/gorm" -) - -var ( - dbMu sync.RWMutex - dbSvc contracts.DBService - cacheMu sync.RWMutex - cacheSvc contracts.CacheService -) - -func setDBService(s contracts.DBService) { - dbMu.Lock() - defer dbMu.Unlock() - dbSvc = s -} - -func setCacheService(s contracts.CacheService) { - cacheMu.Lock() - defer cacheMu.Unlock() - cacheSvc = s -} - -func getDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } - } - dbMu.RLock() - s := dbSvc - dbMu.RUnlock() - if s != nil { - return s.DB(ctx) - } - return nil -} - -func getCache(ctx context.Context) contracts.CacheService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { - return s - } - } - cacheMu.RLock() - s := cacheSvc - cacheMu.RUnlock() - return s -} - -// GetAccessTokenByHash 按令牌哈希读取访问令牌记录(仅取鉴权所需字段) -func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*CachedToken, error) { - var row struct { - ID uint64 - UserID uint64 - IsAdmin bool - } - if err := getDB(ctx).Table("w_access_tokens").Where("token_hash = ?", tokenHash).First(&row).Error; err != nil { - return nil, err - } - return &CachedToken{ - ID: row.ID, - UserID: row.UserID, - IsAdmin: row.IsAdmin, - }, nil -} - -// GetActiveUserByID 读取仍处于启用状态的用户 -func GetActiveUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) { - var user contracts.UserDTO - if err := getDB(ctx).Table("w_users").Where("id = ? AND is_active = ?", userID, true).First(&user).Error; err != nil { - return nil, err - } - return &user, nil -} - -// GetUserByID 按 ID 读取用户(不限制启用状态) -func GetUserByID(ctx context.Context, userID uint64) (*contracts.UserDTO, error) { - var user contracts.UserDTO - if err := getDB(ctx).Table("w_users").Where("id = ?", userID).First(&user).Error; err != nil { - return nil, err - } - return &user, nil -} - -// InsertUser 新建用户记录 -func InsertUser(ctx context.Context, user *contracts.UserDTO) error { - return getDB(ctx).Table("w_users").Create(user).Error -} - -// TouchUserLastLogin 刷新用户最后登录时间 -func TouchUserLastLogin(ctx context.Context, userID uint64, at time.Time) error { - return getDB(ctx).Table("w_users").Where("id = ?", userID).Update("last_login_at", at).Error -} - -// ListSimilarUsernames 查询与基础用户名相同或带 `-序号` 后缀的用户名(用于用户名去重) -func ListSimilarUsernames(ctx context.Context, base string) ([]string, error) { - var existingUsernames []string - if err := getDB(ctx).Table("w_users"). - Where("username = ? OR username LIKE ? ESCAPE '\\'", base, util.EscapeLike(base)+"-%"). - Pluck("username", &existingUsernames).Error; err != nil { - return nil, err - } - return existingUsernames, nil -} - -// GetSystemConfigValue 读取系统配置项原始值 -func GetSystemConfigValue(ctx context.Context, key string) (string, error) { - var val string - if err := getDB(ctx).Table("w_system_configs").Where("key = ?", key).Pluck("value", &val).Error; err != nil { - return "", err - } - return val, nil -} - -// ListAllAuthSources 获取全部认证源(含未启用),按 ID 升序 -func ListAllAuthSources(ctx context.Context) ([]AuthSource, error) { - var sources []AuthSource - if err := getDB(ctx).Order("id ASC").Find(&sources).Error; err != nil { - return nil, err - } - return sources, nil -} - -// GetAuthSourceByID 根据 ID 获取认证源 -func GetAuthSourceByID(ctx context.Context, id uint64) (*AuthSource, error) { - var src AuthSource - if err := getDB(ctx).First(&src, id).Error; err != nil { - return nil, err - } - return &src, nil -} - -// GetAuthSourceByName 根据名称获取认证源 -func GetAuthSourceByName(ctx context.Context, name string) (*AuthSource, error) { - var src AuthSource - if err := getDB(ctx).Where("name = ?", name).First(&src).Error; err != nil { - return nil, err - } - return &src, nil -} - -// ListActiveAuthSources 获取所有启用的认证源 -func ListActiveAuthSources(ctx context.Context) ([]AuthSource, error) { - var sources []AuthSource - if err := getDB(ctx).Where("is_active = ?", true).Order("id ASC").Find(&sources).Error; err != nil { - return nil, err - } - return sources, nil -} - -// CreateAuthSourceRecord 新建认证源记录 -func CreateAuthSourceRecord(ctx context.Context, source *AuthSource) error { - return getDB(ctx).Create(source).Error -} - -// SaveAuthSourceRecord 全量保存认证源记录 -func SaveAuthSourceRecord(ctx context.Context, source *AuthSource) error { - return getDB(ctx).Save(source).Error -} - -// DeleteAuthSourceRecord 删除认证源记录 -func DeleteAuthSourceRecord(ctx context.Context, source *AuthSource) error { - return getDB(ctx).Delete(source).Error -} - -// GetActiveAuthSourcesCached 获取所有启用的认证源(带缓存或直接查询) -func GetActiveAuthSourcesCached(ctx context.Context) ([]AuthSource, error) { - return ListActiveAuthSources(ctx) -} - -// GetAuthSourceByNameCached 根据名称获取认证源(带缓存或直接查询) -func GetAuthSourceByNameCached(ctx context.Context, name string) (*AuthSource, error) { - return GetAuthSourceByName(ctx, name) -} - -// FindExternalAccount 查询指定认证源的外部账号绑定 -func FindExternalAccount(ctx context.Context, authSourceID uint64, externalID string) (*ExternalAccount, error) { - var account ExternalAccount - if err := getDB(ctx).Where("auth_source_id = ? AND external_id = ?", authSourceID, externalID).First(&account).Error; err != nil { - return nil, err - } - return &account, nil -} - -// BindExternalAccount 绑定外部账号 -func BindExternalAccount(ctx context.Context, account *ExternalAccount) error { - return getDB(ctx).Create(account).Error -} - -// ListExternalAccountsByUserID 获取用户绑定的所有外部账号 -func ListExternalAccountsByUserID(ctx context.Context, userID uint64) ([]ExternalAccount, error) { - var accounts []ExternalAccount - if err := getDB(ctx).Where("user_id = ?", userID).Find(&accounts).Error; err != nil { - return nil, err - } - return accounts, nil -} - -// UnbindExternalAccount 解绑外部账号 -func UnbindExternalAccount(ctx context.Context, id, userID uint64) error { - return getDB(ctx).Where("id = ? AND user_id = ?", id, userID).Delete(&ExternalAccount{}).Error -} diff --git a/backend/plugins/domain/auth/service.go b/backend/plugins/domain/auth/service.go deleted file mode 100644 index e26b488f..00000000 --- a/backend/plugins/domain/auth/service.go +++ /dev/null @@ -1,262 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package auth - -import ( - "Wavelet/core/contracts" - "context" - "errors" - "sync" -) - -type authServiceImpl struct{} - -func newAuthService() contracts.AuthService { - return &authServiceImpl{} -} - -func (s *authServiceImpl) RequireAuthMiddleware() any { - return LoginRequired() -} - -func (s *authServiceImpl) RequireAdminMiddleware() any { - return AdminRequired() -} - -// GetCurrentUser 从 context 中读取登录用户。 -// -// 中间件通过 gin 的 c.Set(contracts.AuthUserObjKey, user) 写入登录态; -// *gin.Context 自身实现了 context.Context,且其 Value(key) 对 string 类型 key -// 等价于 c.Get(key)(未命中时再回落到 Request.Context().Value), -// 因此这里无需感知 gin 即可读取同一份登录态。 -func (s *authServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) { - if v := ctx.Value(contracts.AuthUserObjKey); v != nil { - if u, ok := v.(*contracts.UserDTO); ok && u != nil { - return u, nil - } - } - - return nil, errors.New(errUserNotInContext) -} - -func (s *authServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) { - if token == "" { - return nil, errors.New(errEmptyToken) - } - - tokenHash := hashToken(token) - tokenRecord, err := GetCachedToken(ctx, tokenHash) - if err != nil { - tokenRecord, err = GetAccessTokenByHash(ctx, tokenHash) - if err != nil { - return nil, err - } - SetCachedToken(ctx, tokenHash, tokenRecord) - } - - user, err := GetCachedUser(ctx, tokenRecord.UserID) - if err != nil || user == nil || !user.IsActive { - user, err = GetActiveUserByID(ctx, tokenRecord.UserID) - if err != nil { - return nil, err - } - SetCachedUser(ctx, tokenRecord.UserID, user) - } - - if user.Username == SystemUsername { - return nil, errors.New(errSystemUserTokenNotAllowed) - } - - return user, nil -} - -func (s *authServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) { - return "", nil -} - -func (s *authServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error { - InvalidateCachedUser(ctx, userID) - return nil -} - -// GetCurrentUserID 从请求登录态中读取用户 ID。 -// -// Session 读取依赖 gin,属于接入层职责,因此这里通过接入层桥接函数 -// currentUserIDFromRequestContext(见 middleware.go)取值,Service 层本身不感知 gin。 -func (s *authServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) { - userID, ok := currentUserIDFromRequestContext(ctx) - if !ok { - return 0, errors.New(errUserNotInContext) - } - return userID, nil -} - -func (s *authServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error { - InvalidateCachedToken(ctx, tokenHash) - return nil -} - -func (s *authServiceImpl) DisallowTokenAuthMiddleware() any { - return DisallowTokenAuth() -} - -func (s *authServiceImpl) InvalidateCachedUser(ctx context.Context, userID uint64) { - InvalidateCachedUser(ctx, userID) -} - -func (s *authServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash string) { - InvalidateCachedToken(ctx, tokenHash) -} - -func (s *authServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) { - sources, err := ListAllAuthSources(ctx) - if err != nil { - return nil, err - } - - views := make([]contracts.AuthSourceViewDTO, len(sources)) - for i := range sources { - views[i] = contracts.AuthSourceViewDTO{ - ID: sources[i].ID, - Name: sources[i].Name, - Type: sources[i].Type, - DisplayName: sources[i].DisplayName, - IsActive: sources[i].IsActive, - IconURL: sources[i].IconURL, - ClientSecretConfigured: sources[i].ClientSecret != "", - } - } - return views, nil -} - -func (s *authServiceImpl) CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) { - model := AuthSource{ - ID: source.ID, - Name: source.Name, - Type: source.Type, - DisplayName: source.DisplayName, - ClientID: source.ClientID, - ClientSecret: source.ClientSecret, - OpenIDDiscoveryURL: source.OpenIDDiscoveryURL, - Scopes: source.Scopes, - IconURL: source.IconURL, - IsActive: source.IsActive, - } - - if err := model.Validate(); err != nil { - return nil, err - } - - if err := CreateAuthSourceRecord(ctx, &model); err != nil { - return nil, err - } - - model.Sanitize() - return toAuthSourceDTO(&model), nil -} - -func (s *authServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) { - existing, err := GetAuthSourceByID(ctx, id) - if err != nil { - return nil, err - } - - existing.DisplayName = source.DisplayName - existing.ClientID = source.ClientID - if source.ClientSecret != "" { - existing.ClientSecret = source.ClientSecret - } - existing.OpenIDDiscoveryURL = source.OpenIDDiscoveryURL - existing.Scopes = source.Scopes - existing.IconURL = source.IconURL - - if err := existing.Validate(); err != nil { - return nil, err - } - - if err := SaveAuthSourceRecord(ctx, existing); err != nil { - return nil, err - } - - existing.Sanitize() - return toAuthSourceDTO(existing), nil -} - -func (s *authServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error { - existing, err := GetAuthSourceByID(ctx, id) - if err != nil { - return err - } - - return DeleteAuthSourceRecord(ctx, existing) -} - -func (s *authServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) { - existing, err := GetAuthSourceByID(ctx, id) - if err != nil { - return nil, err - } - - existing.IsActive = !existing.IsActive - if err := SaveAuthSourceRecord(ctx, existing); err != nil { - return nil, err - } - - existing.Sanitize() - return toAuthSourceDTO(existing), nil -} - -func toAuthSourceDTO(s *AuthSource) *contracts.AuthSourceDTO { - if s == nil { - return nil - } - return &contracts.AuthSourceDTO{ - ID: s.ID, - Name: s.Name, - Type: s.Type, - DisplayName: s.DisplayName, - ClientID: s.ClientID, - ClientSecret: s.ClientSecret, - OpenIDDiscoveryURL: s.OpenIDDiscoveryURL, - Scopes: s.Scopes, - IconURL: s.IconURL, - IsActive: s.IsActive, - CreatedAt: s.CreatedAt, - UpdatedAt: s.UpdatedAt, - } -} - -type authRegistryImpl struct { - mu sync.RWMutex - providers map[string]contracts.OAuthProvider -} - -func newAuthRegistry() contracts.AuthRegistry { - return &authRegistryImpl{ - providers: make(map[string]contracts.OAuthProvider), - } -} - -func (r *authRegistryImpl) RegisterOAuthProvider(name string, provider contracts.OAuthProvider) { - r.mu.Lock() - defer r.mu.Unlock() - r.providers[name] = provider -} - -func (r *authRegistryImpl) GetOAuthProvider(name string) (contracts.OAuthProvider, bool) { - r.mu.RLock() - defer r.mu.RUnlock() - p, ok := r.providers[name] - return p, ok -} - -func (r *authRegistryImpl) ListOAuthProviders() []string { - r.mu.RLock() - defer r.mu.RUnlock() - res := make([]string, 0, len(r.providers)) - for name := range r.providers { - res = append(res, name) - } - return res -} diff --git a/backend/plugins/domain/auth/service/auth_registry.go b/backend/plugins/domain/auth/service/auth_registry.go new file mode 100644 index 00000000..77a4f769 --- /dev/null +++ b/backend/plugins/domain/auth/service/auth_registry.go @@ -0,0 +1,49 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package service implements domain business services and orchestration for the auth plugin. +package service + +import ( + "Wavelet/core/contracts" + "sync" +) + +// AuthRegistryImpl implements contracts.AuthRegistry. +type AuthRegistryImpl struct { + mu sync.RWMutex + providers map[string]contracts.OAuthProvider +} + +// NewAuthRegistry creates a new AuthRegistryImpl. +func NewAuthRegistry() *AuthRegistryImpl { + return &AuthRegistryImpl{ + providers: make(map[string]contracts.OAuthProvider), + } +} + +// RegisterOAuthProvider registers an OAuthProvider by name. +func (r *AuthRegistryImpl) RegisterOAuthProvider(name string, provider contracts.OAuthProvider) { + r.mu.Lock() + defer r.mu.Unlock() + r.providers[name] = provider +} + +// GetOAuthProvider retrieves an OAuthProvider by name. +func (r *AuthRegistryImpl) GetOAuthProvider(name string) (contracts.OAuthProvider, bool) { + r.mu.RLock() + defer r.mu.RUnlock() + p, ok := r.providers[name] + return p, ok +} + +// ListOAuthProviders lists all registered provider names. +func (r *AuthRegistryImpl) ListOAuthProviders() []string { + r.mu.RLock() + defer r.mu.RUnlock() + res := make([]string, 0, len(r.providers)) + for name := range r.providers { + res = append(res, name) + } + return res +} diff --git a/backend/plugins/domain/auth/service/auth_service.go b/backend/plugins/domain/auth/service/auth_service.go new file mode 100644 index 00000000..202d480b --- /dev/null +++ b/backend/plugins/domain/auth/service/auth_service.go @@ -0,0 +1,278 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package service implements domain business services and orchestration for the auth plugin. +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/dao" + "Wavelet/plugins/domain/auth/model/entity" + "context" + "crypto/sha256" + "encoding/hex" + "errors" +) + +// HashToken computes SHA-256 hex digest of access token. +func HashToken(token string) string { + h := sha256.New() + h.Write([]byte(token)) + return hex.EncodeToString(h.Sum(nil)) +} + +// UserIDExtractor extracts user ID from a request context. +type UserIDExtractor func(ctx context.Context) (uint64, bool) + +// AuthServiceImpl implements contracts.AuthService. +type AuthServiceImpl struct { + dao *dao.DAO + requireAuthMiddleware any + requireAdminMiddleware any + disallowTokenMiddleware any + userIDExtractor UserIDExtractor +} + +// NewAuthService creates a new AuthServiceImpl. +func NewAuthService( + d *dao.DAO, + requireAuth any, + requireAdmin any, + disallowToken any, + extractor UserIDExtractor, +) *AuthServiceImpl { + return &AuthServiceImpl{ + dao: d, + requireAuthMiddleware: requireAuth, + requireAdminMiddleware: requireAdmin, + disallowTokenMiddleware: disallowToken, + userIDExtractor: extractor, + } +} + +// SetMiddlewareHandlers wires middleware handlers into AuthService after controller initialization. +func (s *AuthServiceImpl) SetMiddlewareHandlers(requireAuth, requireAdmin, disallowToken any, extractor UserIDExtractor) { + s.requireAuthMiddleware = requireAuth + s.requireAdminMiddleware = requireAdmin + s.disallowTokenMiddleware = disallowToken + s.userIDExtractor = extractor +} + +// RequireAuthMiddleware returns the authentication check middleware. +func (s *AuthServiceImpl) RequireAuthMiddleware() any { + return s.requireAuthMiddleware +} + +// RequireAdminMiddleware returns the admin authorization middleware. +func (s *AuthServiceImpl) RequireAdminMiddleware() any { + return s.requireAdminMiddleware +} + +// DisallowTokenAuthMiddleware returns the token rejection middleware. +func (s *AuthServiceImpl) DisallowTokenAuthMiddleware() any { + return s.disallowTokenMiddleware +} + +// GetCurrentUser 从 context 中读取登录用户。 +func (s *AuthServiceImpl) GetCurrentUser(ctx context.Context) (*contracts.UserDTO, error) { + if v := ctx.Value(contracts.AuthUserObjKey); v != nil { + if u, ok := v.(*contracts.UserDTO); ok && u != nil { + return u, nil + } + } + + return nil, errors.New(consts.ErrUserNotInContext) +} + +// GetCurrentUserID 从请求登录态中读取用户 ID。 +func (s *AuthServiceImpl) GetCurrentUserID(ctx context.Context) (uint64, error) { + if s.userIDExtractor != nil { + if userID, ok := s.userIDExtractor(ctx); ok { + return userID, nil + } + } + return 0, errors.New(consts.ErrUserNotInContext) +} + +// VerifyToken 验证访问令牌并返回对应的用户。 +func (s *AuthServiceImpl) VerifyToken(ctx context.Context, token string) (*contracts.UserDTO, error) { + if token == "" { + return nil, errors.New(consts.ErrEmptyToken) + } + + tokenHash := HashToken(token) + tokenRecord, err := s.dao.GetCachedToken(ctx, tokenHash) + if err != nil { + tokenRecord, err = s.dao.GetAccessTokenByHash(ctx, tokenHash) + if err != nil { + return nil, err + } + s.dao.SetCachedToken(ctx, tokenHash, tokenRecord) + } + + user, err := s.dao.GetCachedUser(ctx, tokenRecord.UserID) + if err != nil || user == nil || !user.IsActive { + user, err = s.dao.GetActiveUserByID(ctx, tokenRecord.UserID) + if err != nil { + return nil, err + } + s.dao.SetCachedUser(ctx, tokenRecord.UserID, user) + } + + if user.Username == consts.SystemUsername { + return nil, errors.New(consts.ErrSystemUserTokenNotAllowed) + } + + return user, nil +} + +// CreateSession establishes an authenticated session. +func (s *AuthServiceImpl) CreateSession(_ context.Context, _ uint64, _ map[string]any) (string, error) { + return "", nil +} + +// RevokeUserSessions revokes active sessions and cached tokens for a user. +func (s *AuthServiceImpl) RevokeUserSessions(ctx context.Context, userID uint64) error { + s.dao.InvalidateCachedUser(ctx, userID) + return nil +} + +// RevokeToken invalidates a cached token by its hash. +func (s *AuthServiceImpl) RevokeToken(ctx context.Context, tokenHash string) error { + s.dao.InvalidateCachedToken(ctx, tokenHash) + return nil +} + +// InvalidateCachedUser invalidates cached user profile. +func (s *AuthServiceImpl) InvalidateCachedUser(ctx context.Context, userID uint64) { + s.dao.InvalidateCachedUser(ctx, userID) +} + +// InvalidateCachedToken invalidates cached access token. +func (s *AuthServiceImpl) InvalidateCachedToken(ctx context.Context, tokenHash string) { + s.dao.InvalidateCachedToken(ctx, tokenHash) +} + +// ListAuthSources lists all configured authentication sources. +func (s *AuthServiceImpl) ListAuthSources(ctx context.Context) ([]contracts.AuthSourceViewDTO, error) { + sources, err := s.dao.ListAllAuthSources(ctx) + if err != nil { + return nil, err + } + + views := make([]contracts.AuthSourceViewDTO, len(sources)) + for i := range sources { + views[i] = contracts.AuthSourceViewDTO{ + ID: sources[i].ID, + Name: sources[i].Name, + Type: sources[i].Type, + DisplayName: sources[i].DisplayName, + IsActive: sources[i].IsActive, + IconURL: sources[i].IconURL, + ClientSecretConfigured: sources[i].ClientSecret != "", + } + } + return views, nil +} + +// CreateAuthSource creates a new authentication source. +func (s *AuthServiceImpl) CreateAuthSource(ctx context.Context, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) { + model := entity.AuthSource{ + ID: source.ID, + Name: source.Name, + Type: source.Type, + DisplayName: source.DisplayName, + ClientID: source.ClientID, + ClientSecret: source.ClientSecret, + OpenIDDiscoveryURL: source.OpenIDDiscoveryURL, + Scopes: source.Scopes, + IconURL: source.IconURL, + IsActive: source.IsActive, + } + + if err := model.Validate(); err != nil { + return nil, err + } + + if err := s.dao.CreateAuthSource(ctx, &model); err != nil { + return nil, err + } + + model.Sanitize() + return toAuthSourceDTO(&model), nil +} + +// UpdateAuthSource updates an existing authentication source. +func (s *AuthServiceImpl) UpdateAuthSource(ctx context.Context, id uint64, source contracts.AuthSourceDTO) (*contracts.AuthSourceDTO, error) { + existing, err := s.dao.GetAuthSourceByID(ctx, id) + if err != nil { + return nil, err + } + + existing.DisplayName = source.DisplayName + existing.ClientID = source.ClientID + if source.ClientSecret != "" { + existing.ClientSecret = source.ClientSecret + } + existing.OpenIDDiscoveryURL = source.OpenIDDiscoveryURL + existing.Scopes = source.Scopes + existing.IconURL = source.IconURL + + if err := existing.Validate(); err != nil { + return nil, err + } + + if err := s.dao.SaveAuthSource(ctx, existing); err != nil { + return nil, err + } + + existing.Sanitize() + return toAuthSourceDTO(existing), nil +} + +// DeleteAuthSource deletes an authentication source. +func (s *AuthServiceImpl) DeleteAuthSource(ctx context.Context, id uint64) error { + existing, err := s.dao.GetAuthSourceByID(ctx, id) + if err != nil { + return err + } + + return s.dao.DeleteAuthSource(ctx, existing) +} + +// ToggleAuthSource toggles active status of an authentication source. +func (s *AuthServiceImpl) ToggleAuthSource(ctx context.Context, id uint64) (*contracts.AuthSourceDTO, error) { + existing, err := s.dao.GetAuthSourceByID(ctx, id) + if err != nil { + return nil, err + } + + existing.IsActive = !existing.IsActive + if err := s.dao.SaveAuthSource(ctx, existing); err != nil { + return nil, err + } + + existing.Sanitize() + return toAuthSourceDTO(existing), nil +} + +func toAuthSourceDTO(s *entity.AuthSource) *contracts.AuthSourceDTO { + if s == nil { + return nil + } + return &contracts.AuthSourceDTO{ + ID: s.ID, + Name: s.Name, + Type: s.Type, + DisplayName: s.DisplayName, + ClientID: s.ClientID, + ClientSecret: s.ClientSecret, + OpenIDDiscoveryURL: s.OpenIDDiscoveryURL, + Scopes: s.Scopes, + IconURL: s.IconURL, + IsActive: s.IsActive, + CreatedAt: s.CreatedAt, + UpdatedAt: s.UpdatedAt, + } +} diff --git a/backend/plugins/domain/auth/service/cap_service.go b/backend/plugins/domain/auth/service/cap_service.go new file mode 100644 index 00000000..210b785f --- /dev/null +++ b/backend/plugins/domain/auth/service/cap_service.go @@ -0,0 +1,194 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package service implements domain business services and orchestration for the auth plugin. +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/model/dto" + "Wavelet/plugins/domain/auth/pow" + "context" + "crypto/sha256" + "encoding/hex" + "strconv" + "strings" + "time" +) + +// CaptchaManager orchestrates challenge generation and solution validation. +type CaptchaManager struct { + secret []byte + store pow.Store + settingsMgr *CapSettingsManager +} + +// NewCaptchaManager creates a new CAPTCHA Manager. +func NewCaptchaManager(secret []byte, store pow.Store, settingsMgr *CapSettingsManager) *CaptchaManager { + return &CaptchaManager{ + secret: secret, + store: store, + settingsMgr: settingsMgr, + } +} + +// SetSecret updates the shared secret used for PoW generation and validation. +func (m *CaptchaManager) SetSecret(secret []byte) { + m.secret = secret +} + +// Generate creates a challenge response. +func (m *CaptchaManager) Generate(ctx context.Context, scope string) (*pow.ChallengeResponse, error) { + settings, err := m.settingsMgr.Current(ctx) + if err != nil { + return nil, err + } + + challengeConfig := pow.ChallengeConfig{ + Count: settings.ChallengeCount, + Size: settings.ChallengeSize, + Difficulty: settings.ChallengeDifficulty, + Expires: settings.ChallengeTTL, + } + return pow.GenerateChallenge(m.secret, challengeConfig, scope) +} + +// Redeem verifies PoW solutions and returns a one-time redeem token. +func (m *CaptchaManager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*dto.RedeemResponse, error) { + sigHex := pow.JwtSigHex(token) + if sigHex == "" { + return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrInvalidToken}, nil + } + + nonceKey := "cap:nonce:" + sigHex + + payload, err := pow.VerifyChallengeSolutions(token, solutions, m.secret, scope) + if err != nil { + return &dto.RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors returned as response + } + + now := time.Now().UnixNano() / int64(time.Millisecond) + nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond + if nonceTTL < time.Second { + nonceTTL = time.Second + } + + set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL) + if err != nil { + return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrNonceStoreFailed}, err + } + if !set { + return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrAlreadyRedeemed}, nil + } + + settings, err := m.settingsMgr.Current(ctx) + if err != nil { + return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrSettingsLoad}, err + } + + id := pow.RandomHex(consts.RedeemTokenIDLength) + verToken := pow.RandomHex(consts.RedeemVerTokenLength) + verHashBytes := sha256.Sum256([]byte(verToken)) + verHashHex := hex.EncodeToString(verHashBytes[:]) + + tokenKey := "cap:token:" + id + ":" + verHashHex + tokenExpires := time.Now().Add(settings.TokenTTL) + storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope + + if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil { + return &dto.RedeemResponse{Success: false, Error: consts.RedeemErrTokenStoreFailed}, err + } + + return &dto.RedeemResponse{ + Success: true, + Token: id + ":" + verToken, + Expires: tokenExpires.UnixNano() / int64(time.Millisecond), + }, nil +} + +// VerifyToken validates and consumes the redeem token (single-use). +func (m *CaptchaManager) VerifyToken(ctx context.Context, token, expectedScope string) (bool, error) { + if token == "" { + return false, nil + } + parts := strings.Split(token, ":") + if len(parts) != consts.TokenPartsCount { + return false, nil + } + id := parts[0] + verToken := parts[1] + + verHashBytes := sha256.Sum256([]byte(verToken)) + verHashHex := hex.EncodeToString(verHashBytes[:]) + + tokenKey := "cap:token:" + id + ":" + verHashHex + + if m.store == nil { + return false, nil + } + val, exists, err := m.store.GetAndDelete(ctx, tokenKey) + if err != nil { + return false, err + } + if !exists { + return false, nil + } + + valParts := strings.Split(val, "|") + if len(valParts) != consts.ValuePartsCount { + return false, nil + } + + expNano, err := strconv.ParseInt(valParts[0], 10, 64) + if err != nil { + return false, nil //nolint:nilerr // invalid format is failure + } + tokenScope := valParts[1] + + if expectedScope != "" && tokenScope != expectedScope { + return false, nil + } + + if time.Now().UnixNano() > expNano { + return false, nil + } + + return true, nil +} + +// CaptchaServiceImpl implements contracts.CaptchaService. +type CaptchaServiceImpl struct { + manager *CaptchaManager + verifyMiddleware func(scope string) any + challengeHandler any + redeemHandler any +} + +// NewCaptchaService creates a new CaptchaServiceImpl. +func NewCaptchaService(mgr *CaptchaManager, verifyMiddleware func(scope string) any, challengeHandler any, redeemHandler any) contracts.CaptchaService { + return &CaptchaServiceImpl{ + manager: mgr, + verifyMiddleware: verifyMiddleware, + challengeHandler: challengeHandler, + redeemHandler: redeemHandler, + } +} + +// VerifyMiddleware returns the captcha verification middleware. +func (s *CaptchaServiceImpl) VerifyMiddleware(scope string) any { + if s.verifyMiddleware != nil { + return s.verifyMiddleware(scope) + } + return nil +} + +// ChallengeHandler returns the challenge HTTP handler. +func (s *CaptchaServiceImpl) ChallengeHandler() any { + return s.challengeHandler +} + +// RedeemHandler returns the redeem HTTP handler. +func (s *CaptchaServiceImpl) RedeemHandler() any { + return s.redeemHandler +} diff --git a/backend/plugins/domain/auth/service/cap_settings.go b/backend/plugins/domain/auth/service/cap_settings.go new file mode 100644 index 00000000..3054a789 --- /dev/null +++ b/backend/plugins/domain/auth/service/cap_settings.go @@ -0,0 +1,119 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package service implements domain business services and orchestration for the auth plugin. +package service + +import ( + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/dao" + "Wavelet/plugins/domain/auth/model/do" + "context" + "errors" + "sync/atomic" + + "golang.org/x/sync/singleflight" +) + +var capRuntimeConfigKeys = []string{ + consts.ConfigKeyCapLoginEnabled, + consts.ConfigKeyCapChallengeCount, + consts.ConfigKeyCapChallengeSize, + consts.ConfigKeyCapChallengeDifficulty, + consts.ConfigKeyCapChallengeTTL, + consts.ConfigKeyCapTokenTTL, +} + +var capRuntimeConfigKeySet = func() map[string]struct{} { + set := make(map[string]struct{}, len(capRuntimeConfigKeys)) + for _, key := range capRuntimeConfigKeys { + set[key] = struct{}{} + } + return set +}() + +// IsCapRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings. +func IsCapRuntimeConfigKey(key string) bool { + _, ok := capRuntimeConfigKeySet[key] + return ok +} + +// CapSettingsManager manages dynamic CAPTCHA configuration cache. +type CapSettingsManager struct { + dao *dao.DAO + snapshot atomic.Pointer[do.CapRuntimeSettings] + loadGroup singleflight.Group +} + +// NewCapSettingsManager creates a new CapSettingsManager. +func NewCapSettingsManager(d *dao.DAO) *CapSettingsManager { + return &CapSettingsManager{ + dao: d, + } +} + +// Invalidate drops the in-process CAPTCHA settings snapshot. +func (m *CapSettingsManager) Invalidate() { + m.snapshot.Store(nil) +} + +// Current returns the cached CAPTCHA runtime settings snapshot. +func (m *CapSettingsManager) Current(ctx context.Context) (do.CapRuntimeSettings, error) { + if snapshot := m.snapshot.Load(); snapshot != nil { + return *snapshot, nil + } + + loaded, err, _ := m.loadGroup.Do("cap-runtime-settings", func() (any, error) { + if snapshot := m.snapshot.Load(); snapshot != nil { + return *snapshot, nil + } + + settings, loadErr := m.loadSettings(ctx) + if loadErr != nil { + return do.CapRuntimeSettings{}, loadErr + } + + m.snapshot.Store(&settings) + return settings, nil + }) + if err != nil { + return do.CapRuntimeSettings{}, err + } + + settings, ok := loaded.(do.CapRuntimeSettings) + if !ok { + return do.CapRuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type") + } + return settings, nil +} + +// CapProtectionEnabled reports whether CAPTCHA verification is required for protected routes. +func (m *CapSettingsManager) CapProtectionEnabled(ctx context.Context) bool { + settings, err := m.Current(ctx) + if err != nil { + return false + } + return settings.LoginEnabled +} + +// InstallTestSnapshot installs a fixed snapshot for unit tests. +func (m *CapSettingsManager) InstallTestSnapshot(settings do.CapRuntimeSettings) func() { + snapshot := settings + m.snapshot.Store(&snapshot) + return m.Invalidate +} + +func (m *CapSettingsManager) loadSettings(ctx context.Context) (do.CapRuntimeSettings, error) { + if m.dao == nil { + return do.ParseCapRuntimeSettings(nil), nil + } + records, err := m.dao.ListSystemConfigsByKeys(ctx, capRuntimeConfigKeys) + if err != nil { + return do.CapRuntimeSettings{}, err + } + configs := make(map[string]string, len(records)) + for _, r := range records { + configs[r.Key] = r.Value + } + return do.ParseCapRuntimeSettings(configs), nil +} diff --git a/backend/plugins/domain/auth/service/oauth_service.go b/backend/plugins/domain/auth/service/oauth_service.go new file mode 100644 index 00000000..8eee6021 --- /dev/null +++ b/backend/plugins/domain/auth/service/oauth_service.go @@ -0,0 +1,440 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package service implements domain business services and orchestration for the auth plugin. +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/idgen" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/dao" + "Wavelet/plugins/domain/auth/model/dto" + "Wavelet/plugins/domain/auth/model/entity" + "context" + "errors" + "fmt" + "strconv" + "strings" + "time" + + "github.com/coreos/go-oidc/v3/oidc" + "golang.org/x/oauth2" + "gorm.io/gorm" +) + +// OAuthService orchestrates OAuth/OIDC operations. +type OAuthService struct { + dao *dao.DAO + providerCache *OIDCProviderCache + sessionSvc *SessionService +} + +// NewOAuthService creates a new OAuthService. +func NewOAuthService(d *dao.DAO, cache *OIDCProviderCache, sessSvc *SessionService) *OAuthService { + return &OAuthService{ + dao: d, + providerCache: cache, + sessionSvc: sessSvc, + } +} + +// IsOIDCLoginEnabled checks if OIDC login is globally enabled. +func (s *OAuthService) IsOIDCLoginEnabled(ctx context.Context) bool { + val, err := s.dao.GetSystemConfigValue(ctx, "oidc_login_enabled") + if err != nil || val == "" { + return true + } + b, err := strconv.ParseBool(val) + if err != nil { + return true + } + return b +} + +// ResolveAuthSource retrieves the specified or default active auth source. +func (s *OAuthService) ResolveAuthSource(ctx context.Context, sourceName string) (*entity.AuthSource, error) { + name := strings.TrimSpace(strings.ToLower(sourceName)) + if name == "" { + sources, err := s.dao.ListActiveAuthSources(ctx) + if err != nil { + return nil, err + } + if len(sources) == 0 { + return nil, errors.New(consts.ErrNoActiveAuthSource) + } + src, err := s.dao.GetAuthSourceByName(ctx, sources[0].Name) + if err != nil { + return nil, err + } + return src, nil + } + src, err := s.dao.GetAuthSourceByName(ctx, name) + if err != nil { + return nil, err + } + return src, nil +} + +// ActiveLoginSources returns all active login sources formatted for display. +func (s *OAuthService) ActiveLoginSources(ctx context.Context) ([]dto.AuthSourceView, error) { + if !s.IsOIDCLoginEnabled(ctx) { + return nil, nil + } + + dbSources, err := s.dao.ListActiveAuthSources(ctx) + if err != nil { + return nil, err + } + sources := make([]dto.AuthSourceView, 0, len(dbSources)) + for _, source := range dbSources { + sources = append(sources, dto.AuthSourceView{ + ID: source.ID, + Name: source.Name, + Type: source.Type, + DisplayName: source.DisplayName, + IsActive: source.IsActive, + IconURL: source.IconURL, + ClientSecretConfigured: source.ClientSecretConfigured, + }) + } + return sources, nil +} + +// GetFrontendLoginRedirectURL constructs the OAuth frontend redirect URL. +func (s *OAuthService) GetFrontendLoginRedirectURL(ctx context.Context) (string, error) { + val, err := s.dao.GetSystemConfigValue(ctx, "server_address") + if err != nil || strings.TrimSpace(val) == "" { + return "", errors.New(consts.ErrServerAddressMissing) + } + return strings.TrimRight(val, "/") + "/login", nil +} + +// ReserveOAuthStateSlot ensures that a session does not abuse OAuth state generation. +func (s *OAuthService) ReserveOAuthStateSlot(ctx context.Context, sessionHash string) error { + if sessionHash == "" { + return nil + } + if limiter := s.dao.Limiter(); limiter != nil { + key := fmt.Sprintf(consts.OAuthStateLimitKeyFormat, sessionHash) + res, err := limiter.Allow(ctx, key, contracts.Rate{ + Limit: consts.OAuthStateLimitMax, + Period: consts.OAuthStateCacheKeyExpiration, + }) + if err != nil { + return err + } + if !res.Allowed { + return errors.New(consts.ErrOAuthStateRateLimited) + } + return nil + } + + cache := s.dao.Cache() + if cache == nil { + return nil + } + key := fmt.Sprintf(consts.OAuthStateLimitKeyFormat, sessionHash) + var count int + _ = cache.Get(ctx, key, &count) + count++ + _ = cache.Set(ctx, key, count, consts.OAuthStateCacheKeyExpiration) + if count > consts.OAuthStateLimitMax { + return errors.New(consts.ErrOAuthStateRateLimited) + } + return nil +} + +// BuildOAuthConfig builds oauth2.Config and oidc.IDTokenVerifier. +func (s *OAuthService) BuildOAuthConfig(ctx context.Context, source *entity.AuthSource, redirectURL string) (*oauth2.Config, *oidc.IDTokenVerifier, error) { + if source == nil { + return nil, nil, errors.New(consts.ErrAuthSourceRequired) + } + + if source.OpenIDDiscoveryURL == "" { + return nil, nil, errors.New(consts.ErrDiscoveryURLRequired) + } + + issuer := strings.TrimSuffix(strings.TrimSpace(source.OpenIDDiscoveryURL), "/") + issuer = strings.TrimSuffix(issuer, "/.well-known/openid-configuration") + issuer = strings.TrimSuffix(issuer, "/.well-known/oauth-authorization-server") + + provider, err := s.providerCache.Get(ctx, issuer) + if err != nil { + return nil, nil, err + } + verifier := provider.Verifier(&oidc.Config{ClientID: source.ClientID}) + scopes := strings.Fields(source.Scopes) + if len(scopes) == 0 { + scopes = []string{oidc.ScopeOpenID, "profile", "email"} + } + if !containsScope(scopes, oidc.ScopeOpenID) { + scopes = append([]string{oidc.ScopeOpenID}, scopes...) + } + + return &oauth2.Config{ + ClientID: source.ClientID, + ClientSecret: source.ClientSecret, + RedirectURL: redirectURL, + Scopes: scopes, + Endpoint: provider.Endpoint(), + }, verifier, nil +} + +func containsScope(scopes []string, scope string) bool { + for _, item := range scopes { + if item == scope { + return true + } + } + return false +} + +// BuildAuthorizeURL generates the redirect authorize URL for the source and state. +func (s *OAuthService) BuildAuthorizeURL(ctx context.Context, source *entity.AuthSource, state string) (string, error) { + redirectURL, err := s.GetFrontendLoginRedirectURL(ctx) + if err != nil { + return "", err + } + authConfig, verifier, err := s.BuildOAuthConfig(ctx, source, redirectURL) + if err != nil { + return "", err + } + if verifier != nil { + return authConfig.AuthCodeURL(state, oidc.Nonce(state)), nil + } + return authConfig.AuthCodeURL(state), nil +} + +// BuildOAuthUserInfo exchanges the auth code and retrieves user identity claims. +func (s *OAuthService) BuildOAuthUserInfo(ctx context.Context, source *entity.AuthSource, code, nonce, redirectURL string) (*contracts.OAuthUserInfoDTO, error) { + authConfig, verifier, err := s.BuildOAuthConfig(ctx, source, redirectURL) + if err != nil { + return nil, err + } + + token, err := authConfig.Exchange(ctx, code) + if err != nil { + return nil, err + } + + userInfo := &contracts.OAuthUserInfoDTO{Active: true} + if verifier != nil { + if verifyErr := s.verifyIDToken(ctx, verifier, token, nonce, userInfo); verifyErr != nil { + return nil, verifyErr + } + } + + if userInfo.Username == "" && userInfo.PreferredUsername != "" { + userInfo.Username = userInfo.PreferredUsername + } + if userInfo.Username == "" && userInfo.Email != "" { + userInfo.Username = strings.Split(userInfo.Email, "@")[0] + } + if userInfo.Username == "" && userInfo.Sub != "" { + userInfo.Username = userInfo.Sub + } + if userInfo.Name == "" { + userInfo.Name = userInfo.Username + } + + return userInfo, nil +} + +func (s *OAuthService) verifyIDToken(ctx context.Context, verifier *oidc.IDTokenVerifier, token *oauth2.Token, nonce string, userInfo *contracts.OAuthUserInfoDTO) error { + rawIDToken, ok := token.Extra("id_token").(string) + if !ok { + return nil + } + idToken, verifyErr := verifier.Verify(ctx, rawIDToken) + if verifyErr != nil { + return fmt.Errorf(consts.ErrIDTokenVerifyFailedFormat, consts.ErrIDTokenVerifyFailed, verifyErr) + } + if nonce != "" && idToken.Nonce != nonce { + return errors.New(consts.ErrNonceMismatch) + } + if claimsErr := idToken.Claims(userInfo); claimsErr != nil { + return claimsErr + } + return nil +} + +// NormalizeOAuthUserInfo sanitizes user claims. +func (s *OAuthService) NormalizeOAuthUserInfo(userInfo *contracts.OAuthUserInfoDTO) error { + userInfo.Username = strings.TrimSpace(userInfo.Username) + userInfo.PreferredUsername = strings.TrimSpace(userInfo.PreferredUsername) + userInfo.Email = strings.TrimSpace(userInfo.Email) + userInfo.Name = strings.TrimSpace(userInfo.Name) + userInfo.AvatarURL = strings.TrimSpace(userInfo.AvatarURL) + + if userInfo.Username == "" && userInfo.PreferredUsername != "" { + userInfo.Username = userInfo.PreferredUsername + } + if userInfo.Username == "" && userInfo.Email != "" { + userInfo.Username = strings.Split(userInfo.Email, "@")[0] + } + if userInfo.Username == "" && userInfo.Sub != "" { + userInfo.Username = userInfo.Sub + } + if userInfo.Username == "" { + return errors.New(consts.ErrUsernameFromSourceFailed) + } + if userInfo.Name == "" { + userInfo.Name = userInfo.Username + } + if !userInfo.Active { + userInfo.Active = true + } + return nil +} + +// UniqueUsername generates a unique username given a base candidate. +func (s *OAuthService) UniqueUsername(ctx context.Context, base string) (string, error) { + base = strings.TrimSpace(base) + if base == "" { + base = "user" + } + + existingUsernames, err := s.dao.ListSimilarUsernames(ctx, base) + if err != nil { + return "", err + } + + exists := make(map[string]bool, len(existingUsernames)) + for _, u := range existingUsernames { + exists[strings.ToLower(u)] = true + } + + if !exists[strings.ToLower(base)] { + return base, nil + } + + for i := 1; i <= 1000; i++ { + candidate := fmt.Sprintf("%s-%d", base, i) + if !exists[strings.ToLower(candidate)] { + return candidate, nil + } + } + + return "", errors.New(consts.ErrUsernameGenerateFailed) +} + +// BindExternalAccount binds an external identity to an existing user. +func (s *OAuthService) BindExternalAccount(ctx context.Context, sourceID, userID uint64, userInfo *contracts.OAuthUserInfoDTO) error { + user, err := s.dao.GetUserByID(ctx, userID) + if err != nil { + return err + } + if err := s.dao.BindExternalAccount(ctx, &entity.ExternalAccount{ + AuthSourceID: sourceID, + UserID: user.ID, + ExternalID: userInfo.Sub, + ExternalUsername: userInfo.Username, + Email: userInfo.Email, + }); err != nil { + return err + } + user.LastLoginAt = time.Now() + _ = s.dao.TouchUserLastLogin(ctx, user.ID, user.LastLoginAt) + return nil +} + +// AuthenticateOrRegisterUser finds existing binding or creates a new user. +func (s *OAuthService) AuthenticateOrRegisterUser(ctx context.Context, source *entity.AuthSource, userInfo *contracts.OAuthUserInfoDTO) (*contracts.UserDTO, bool, error) { + account, err := s.dao.FindExternalAccount(ctx, source.ID, userInfo.Sub) + if err == nil { + user, loadErr := s.dao.GetUserByID(ctx, account.UserID) + if loadErr != nil { + return nil, false, loadErr + } + user.LastLoginAt = time.Now() + _ = s.dao.TouchUserLastLogin(ctx, user.ID, user.LastLoginAt) + return user, true, nil + } + + if !errors.Is(err, gorm.ErrRecordNotFound) { + return nil, false, err + } + + // Not found -> check registration + registrationEnabled := true + val, cfgErr := s.dao.GetSystemConfigValue(ctx, "registration_enabled") + if cfgErr == nil && val != "" { + if b, err := strconv.ParseBool(val); err == nil { + registrationEnabled = b + } + } + + if !registrationEnabled { + return nil, false, nil // registration disabled -> need bind + } + + username, uniqueErr := s.UniqueUsername(ctx, userInfo.Username) + if uniqueErr != nil { + return nil, false, uniqueErr + } + userInfo.Username = username + + now := time.Now() + user := contracts.UserDTO{ + ID: idgen.NextUint64ID(), + Username: userInfo.Username, + Nickname: userInfo.Name, + Email: userInfo.Email, + AvatarURL: userInfo.AvatarURL, + IsActive: userInfo.Active, + LastLoginAt: now, + CreatedAt: now, + UpdatedAt: now, + } + + if err := s.dao.InsertUser(ctx, &user); err != nil { + return nil, false, err + } + + if err := s.dao.BindExternalAccount(ctx, &entity.ExternalAccount{ + AuthSourceID: source.ID, + UserID: user.ID, + ExternalID: userInfo.Sub, + ExternalUsername: userInfo.Username, + Email: userInfo.Email, + }); err != nil { + return nil, false, err + } + + return &user, true, nil +} + +// ListExternalAccounts returns sanitized external account bindings. +func (s *OAuthService) ListExternalAccounts(ctx context.Context, userID uint64) ([]dto.ExternalAccountView, error) { + accounts, err := s.dao.ListExternalAccountsByUserID(ctx, userID) + if err != nil { + return nil, err + } + views := make([]dto.ExternalAccountView, len(accounts)) + for i, acc := range accounts { + source, _ := s.dao.GetAuthSourceByID(ctx, acc.AuthSourceID) + sourceName, sourceType, sourceLabel := "", "", "" + if source != nil { + sourceName = source.Name + sourceType = source.Type + sourceLabel = source.DisplayName + } + views[i] = dto.ExternalAccountView{ + ID: acc.ID, + AuthSourceID: acc.AuthSourceID, + AuthSourceName: sourceName, + AuthSourceType: sourceType, + AuthSourceLabel: sourceLabel, + ExternalUsername: acc.ExternalUsername, + Email: acc.Email, + CreatedAt: acc.CreatedAt.Format(time.RFC3339), + } + } + return views, nil +} + +// DeleteExternalAccount unbinds an external account. +func (s *OAuthService) DeleteExternalAccount(ctx context.Context, id, userID uint64) error { + return s.dao.UnbindExternalAccount(ctx, id, userID) +} diff --git a/backend/plugins/domain/auth/provider_cache.go b/backend/plugins/domain/auth/service/oidc_provider_cache.go similarity index 65% rename from backend/plugins/domain/auth/provider_cache.go rename to backend/plugins/domain/auth/service/oidc_provider_cache.go index 86465b7b..ea5cbada 100644 --- a/backend/plugins/domain/auth/provider_cache.go +++ b/backend/plugins/domain/auth/service/oidc_provider_cache.go @@ -1,7 +1,8 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package auth +// Package service implements domain business services and orchestration for the auth plugin. +package service import ( "context" @@ -13,16 +14,18 @@ import ( "golang.org/x/sync/singleflight" ) -// oidcProviderCache 进程级 OIDC provider 缓存。 -type oidcProviderCache struct { +// OIDCProviderCache 进程级 OIDC provider 缓存。 +type OIDCProviderCache struct { mu sync.RWMutex entries map[string]*oidc.Provider // key: normalized issuer URL sfGroup singleflight.Group } -// globalOIDCProviderCache 是包级单例缓存,与进程同生命周期。 -var globalOIDCProviderCache = &oidcProviderCache{ - entries: make(map[string]*oidc.Provider), +// NewOIDCProviderCache creates a new OIDCProviderCache. +func NewOIDCProviderCache() *OIDCProviderCache { + return &OIDCProviderCache{ + entries: make(map[string]*oidc.Provider), + } } // discoveryContext 从请求 ctx 提取 HTTP 客户端,并绑定到不可取消的 Background ctx。 @@ -34,8 +37,8 @@ func discoveryContext(ctx context.Context) context.Context { return bg } -// get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。 -func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provider, error) { +// Get 返回缓存的 provider;若无则通过 oidc.NewProvider 获取并写入缓存。 +func (c *OIDCProviderCache) Get(ctx context.Context, issuer string) (*oidc.Provider, error) { c.mu.RLock() if p, ok := c.entries[issuer]; ok { c.mu.RUnlock() @@ -68,14 +71,9 @@ func (c *oidcProviderCache) get(ctx context.Context, issuer string) (*oidc.Provi return v.(*oidc.Provider), nil //nolint:forcetypeassert } -// invalidate 从缓存中移除指定 issuer 对应的 provider。 -func (c *oidcProviderCache) invalidate(issuer string) { +// Invalidate 从缓存中移除指定 issuer 对应的 provider。 +func (c *OIDCProviderCache) Invalidate(issuer string) { c.mu.Lock() delete(c.entries, issuer) c.mu.Unlock() } - -// InvalidateOIDCProviderCache 从进程级缓存中清除指定 issuer 的 provider 条目。 -func InvalidateOIDCProviderCache(issuer string) { - globalOIDCProviderCache.invalidate(issuer) -} diff --git a/backend/plugins/domain/auth/service/service.go b/backend/plugins/domain/auth/service/service.go new file mode 100644 index 00000000..2af980fc --- /dev/null +++ b/backend/plugins/domain/auth/service/service.go @@ -0,0 +1,51 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package service implements domain business services and orchestration for the auth plugin. +package service + +import ( + "Wavelet/plugins/domain/auth/dao" + "Wavelet/plugins/domain/auth/pow" + "time" +) + +// Service aggregates all domain services for the auth plugin. +type Service struct { + DAO *dao.DAO + Session *SessionService + OAuth *OAuthService + OIDCProviderCache *OIDCProviderCache + CapSettings *CapSettingsManager + CapManager *CaptchaManager + AuthSvc *AuthServiceImpl + AuthRegistry *AuthRegistryImpl +} + +// New creates a new Service container with all domain services wired up. +func New(d *dao.DAO, sessionCfg SessionConfig, capSecret []byte) *Service { + sessionSvc := NewSessionService(sessionCfg, d) + oidcCache := NewOIDCProviderCache() + oauthSvc := NewOAuthService(d, oidcCache, sessionSvc) + capSettings := NewCapSettingsManager(d) + + var capStore pow.Store + if len(capSecret) > 0 { + capStore = pow.NewMemoryStore(1 * time.Minute) + } + capMgr := NewCaptchaManager(capSecret, capStore, capSettings) + + authSvc := NewAuthService(d, nil, nil, nil, nil) + authRegistry := NewAuthRegistry() + + return &Service{ + DAO: d, + Session: sessionSvc, + OAuth: oauthSvc, + OIDCProviderCache: oidcCache, + CapSettings: capSettings, + CapManager: capMgr, + AuthSvc: authSvc, + AuthRegistry: authRegistry, + } +} diff --git a/backend/plugins/domain/auth/service/session_service.go b/backend/plugins/domain/auth/service/session_service.go new file mode 100644 index 00000000..21a7d09d --- /dev/null +++ b/backend/plugins/domain/auth/service/session_service.go @@ -0,0 +1,177 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package service implements domain business services and orchestration for the auth plugin. +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/plugins/domain/auth/consts" + "Wavelet/plugins/domain/auth/dao" + "context" + "crypto/sha256" + "encoding/hex" + "net/http" + "strconv" + "strings" + "sync" + + "github.com/gin-contrib/sessions" + "github.com/google/uuid" + gsessions "github.com/gorilla/sessions" +) + +// SessionConfig defines session settings. +type SessionConfig struct { + SessionCookieName string `config:"session_cookie_name" env:"APP_SESSION_COOKIE_NAME" default:"wavelet_session"` + SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"` + SessionDomain string `config:"session_domain" env:"APP_SESSION_DOMAIN"` + SessionAge int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"` + SessionHTTPOnly bool `config:"session_http_only" env:"APP_SESSION_HTTP_ONLY" default:"true"` + SessionSecure bool `config:"session_secure" env:"APP_SESSION_SECURE"` +} + +// SessionService manages HTTP session operations, cookies, and tokens. +type SessionService struct { + mu sync.RWMutex + config SessionConfig + dao *dao.DAO +} + +// NewSessionService creates a new SessionService. +func NewSessionService(cfg SessionConfig, d *dao.DAO) *SessionService { + return &SessionService{ + config: cfg, + dao: d, + } +} + +// SetConfig updates the active session configuration. +func (s *SessionService) SetConfig(cfg SessionConfig) { + s.mu.Lock() + defer s.mu.Unlock() + s.config = cfg +} + +// Config returns the current session configuration. +func (s *SessionService) Config() SessionConfig { + s.mu.RLock() + defer s.mu.RUnlock() + return s.config +} + +// GetSessionOptions 根据配置构建 Session 选项 +func (s *SessionService) GetSessionOptions(maxAge int) sessions.Options { + cfg := s.Config() + return sessions.Options{ + Path: "/", + Domain: cfg.SessionDomain, + MaxAge: maxAge, + HttpOnly: cfg.SessionHTTPOnly, + Secure: cfg.SessionSecure, + SameSite: http.SameSiteLaxMode, + } +} + +// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie +func (s *SessionService) StripCookieMaxAgeAndExpires(header http.Header, cookieName string) { + headers := header["Set-Cookie"] + if len(headers) == 0 { + return + } + + newHeaders := make([]string, 0, len(headers)) + for _, h := range headers { + if strings.HasPrefix(h, cookieName+"=") { + parts := strings.Split(h, ";") + newParts := make([]string, 0, len(parts)) + for _, p := range parts { + trimmed := strings.TrimSpace(p) + lower := strings.ToLower(trimmed) + if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") { + continue + } + newParts = append(newParts, p) + } + newHeaders = append(newHeaders, strings.Join(newParts, ";")) + } else { + newHeaders = append(newHeaders, h) + } + } + header["Set-Cookie"] = newHeaders +} + +// EnsureSessionToken returns or generates the session unique token. +func (s *SessionService) EnsureSessionToken(session sessions.Session) (string, bool) { + token, ok := session.Get(consts.SessionTokenKey).(string) + if !ok || token == "" { + token = uuid.NewString() + session.Set(consts.SessionTokenKey, token) + return token, true + } + return token, false +} + +// HashSessionToken hashes the session token using SHA-256. +func (s *SessionService) HashSessionToken(token string) string { + h := sha256.New() + h.Write([]byte(token)) + return hex.EncodeToString(h.Sum(nil)) +} + +// RotateSessionID forces session ID rotation to prevent session fixation attacks. +func (s *SessionService) RotateSessionID(session sessions.Session) { + if inner, ok := session.(interface{ Session() *gsessions.Session }); ok { + if sess := inner.Session(); sess != nil { + sess.ID = "" + } + } +} + +// CalculateSessionMaxAge dynamically calculates max age and whether it's a browser-session cookie. +func (s *SessionService) CalculateSessionMaxAge(ctx context.Context) (int, bool) { + cfg := s.Config() + maxAge := cfg.SessionAge + isSessionCookie := false + + if s.dao != nil { + val, err := s.dao.GetSystemConfigValue(ctx, "login_session_ttl_hours") + if err == nil && val != "" { + if ttlHours, err := strconv.Atoi(val); err == nil { + switch { + case ttlHours == -1: + // 永不过期,设置为 10 年 + maxAge = 10 * 365 * 24 * 3600 + case ttlHours > 0: + maxAge = ttlHours * 3600 + case ttlHours == 0: + isSessionCookie = true + } + } + } + } + return maxAge, isSessionCookie +} + +// ApplyLoginSession writes the authenticated user into a freshly rotated session. +func (s *SessionService) ApplyLoginSession(ctx context.Context, session sessions.Session, user *contracts.UserDTO, extras ...map[string]any) (bool, error) { + session.Clear() + s.RotateSessionID(session) + + session.Set(consts.UserIDKey, strconv.FormatUint(user.ID, 10)) + session.Set(consts.UserNameKey, user.Username) + if len(extras) > 0 { + for key, value := range extras[0] { + session.Set(key, value) + } + } + + maxAge, isSessionCookie := s.CalculateSessionMaxAge(ctx) + session.Options(s.GetSessionOptions(maxAge)) + + if err := session.Save(); err != nil { + return false, err + } + + return isSessionCookie, nil +} diff --git a/backend/plugins/domain/auth/service_test.go b/backend/plugins/domain/auth/service_test.go index 2431ec67..89ddef37 100644 --- a/backend/plugins/domain/auth/service_test.go +++ b/backend/plugins/domain/auth/service_test.go @@ -399,3 +399,30 @@ func TestAuthWhitelistMiddleware(t *testing.T) { engine.ServeHTTP(w2, req2) assert.Equal(t, http.StatusUnauthorized, w2.Code) } + +func TestAdminRequiredReturnsForbiddenForLoggedInNonAdmin(t *testing.T) { + gin.SetMode(gin.TestMode) + db := setupTestDB(t) + require.NoError(t, db.Create(&testUser{ID: 42, Username: "member", IsActive: true, IsAdmin: false}).Error) + + svc := newTestAuthService(t, db) + mw, ok := svc.RequireAdminMiddleware().(gin.HandlerFunc) + require.True(t, ok) + + engine := newSessionEngine() + engine.Use(func(c *gin.Context) { + session := sessions.Default(c) + session.Set(auth.UserIDKey, uint64(42)) + require.NoError(t, session.Save()) + c.Next() + }) + engine.Use(mw) + engine.GET("/admin-only", func(c *gin.Context) { + c.JSON(http.StatusOK, response.OK("ok")) + }) + + w := httptest.NewRecorder() + engine.ServeHTTP(w, httptest.NewRequest(http.MethodGet, "/admin-only", nil)) + assert.Equal(t, http.StatusForbidden, w.Code) + assert.Contains(t, w.Body.String(), "权限不足") +} diff --git a/backend/plugins/domain/auth/session.go b/backend/plugins/domain/auth/session.go deleted file mode 100644 index 97141413..00000000 --- a/backend/plugins/domain/auth/session.go +++ /dev/null @@ -1,169 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package auth - -import ( - "Wavelet/core/contracts" - "context" - "crypto/sha256" - "encoding/hex" - "net/http" - "strconv" - "strings" - "sync" - - "github.com/gin-contrib/sessions" - "github.com/gin-gonic/gin" - "github.com/google/uuid" - gsessions "github.com/gorilla/sessions" -) - -var ( - sessConfigMu sync.RWMutex - sessConfig = SessionConfig{ - SessionCookieName: "wavelet_session", - SessionAge: 86400, - SessionHTTPOnly: true, - } -) - -// SetSessionConfig updates the active session configuration. -func SetSessionConfig(cfg SessionConfig) { - sessConfigMu.Lock() - defer sessConfigMu.Unlock() - sessConfig = cfg -} - -// GetSessionConfig returns the active session configuration. -func GetSessionConfig() SessionConfig { - sessConfigMu.RLock() - defer sessConfigMu.RUnlock() - return sessConfig -} - -// GetSessionOptions 根据配置构建 Session 选项 -func GetSessionOptions(maxAge int) sessions.Options { - cfg := GetSessionConfig() - return sessions.Options{ - Path: "/", - Domain: cfg.SessionDomain, - MaxAge: maxAge, - HttpOnly: cfg.SessionHTTPOnly, - Secure: cfg.SessionSecure, - SameSite: http.SameSiteLaxMode, - } -} - -// StripCookieMaxAgeAndExpires 从 Set-Cookie 响应头中移除 Max-Age 和 Expires,从而使其成为浏览器会话 Cookie -func StripCookieMaxAgeAndExpires(header http.Header, cookieName string) { - headers := header["Set-Cookie"] - if len(headers) == 0 { - return - } - - newHeaders := make([]string, 0, len(headers)) - for _, h := range headers { - if strings.HasPrefix(h, cookieName+"=") { - parts := strings.Split(h, ";") - newParts := make([]string, 0, len(parts)) - for _, p := range parts { - trimmed := strings.TrimSpace(p) - lower := strings.ToLower(trimmed) - if strings.HasPrefix(lower, "max-age=") || strings.HasPrefix(lower, "expires=") { - continue - } - newParts = append(newParts, p) - } - newHeaders = append(newHeaders, strings.Join(newParts, ";")) - } else { - newHeaders = append(newHeaders, h) - } - } - header["Set-Cookie"] = newHeaders -} - -// GetUserIDFromSession 从 Session 中提取用户 ID -func GetUserIDFromSession(s sessions.Session) uint64 { - val := s.Get(UserIDKey) - return ParseUserID(val) -} - -// GetUserIDFromContext 从 Gin Context 的 Session 中提取用户 ID -func GetUserIDFromContext(c *gin.Context) (uid uint64) { - defer func() { - _ = recover() - }() - session := sessions.Default(c) - return GetUserIDFromSession(session) -} - -func ensureSessionToken(s sessions.Session) (string, bool) { - token, ok := s.Get(SessionTokenKey).(string) - if !ok || token == "" { - token = uuid.NewString() - s.Set(SessionTokenKey, token) - return token, true - } - return token, false -} - -func hashSessionToken(token string) string { - h := sha256.New() - h.Write([]byte(token)) - return hex.EncodeToString(h.Sum(nil)) -} - -func rotateSessionID(s sessions.Session) { - if inner, ok := s.(interface{ Session() *gsessions.Session }); ok { - if sess := inner.Session(); sess != nil { - sess.ID = "" - } - } -} - -// SetLoginSession writes the authenticated user into a freshly rotated session. -func SetLoginSession(ctx context.Context, c *gin.Context, user *contracts.UserDTO, extras ...map[string]any) error { - session := sessions.Default(c) - session.Clear() - rotateSessionID(session) - - session.Set(UserIDKey, user.ID) - session.Set(UserNameKey, user.Username) - if len(extras) > 0 { - for key, value := range extras[0] { - session.Set(key, value) - } - } - - // 根据系统配置动态设置 Session 过期时间 - cfg := GetSessionConfig() - maxAge := cfg.SessionAge - isSessionCookie := false - - val, err := GetSystemConfigValue(ctx, "login_session_ttl_hours") - if err == nil && val != "" { - if ttlHours, err := strconv.Atoi(val); err == nil { - switch { - case ttlHours == -1: - // 永不过期,设置为 10 年 - maxAge = 10 * 365 * 24 * 3600 - case ttlHours > 0: - maxAge = ttlHours * 3600 - case ttlHours == 0: - isSessionCookie = true - } - } - } - session.Options(GetSessionOptions(maxAge)) - - if err := session.Save(); err != nil { - return err - } - - if isSessionCookie { - StripCookieMaxAgeAndExpires(c.Writer.Header(), cfg.SessionCookieName) - } - - return nil -} diff --git a/backend/plugins/domain/auth_dependency_ordering_test.go b/backend/plugins/domain/auth_dependency_ordering_test.go index 5e4df016..1dfcf1b9 100644 --- a/backend/plugins/domain/auth_dependency_ordering_test.go +++ b/backend/plugins/domain/auth_dependency_ordering_test.go @@ -19,7 +19,7 @@ import ( "Wavelet/core/contracts" "Wavelet/core/extpoints" "Wavelet/plugins/domain/admin" - "Wavelet/plugins/domain/message_gateway" + "Wavelet/plugins/domain/msg_gateway" "Wavelet/plugins/domain/user" ) @@ -122,7 +122,7 @@ func TestRoutesMountedBeforeAuthServiceAreGuarded(t *testing.T) { core.WithPlugins( dbProvider(), user.New(), - message_gateway.New(), + msg_gateway.New(), authProvider(), ), ) @@ -147,7 +147,7 @@ func TestAuthConsumersDeclareAuthDependency(t *testing.T) { deps []reflect.Type }{ {"user", user.New().Inject()}, - {"message_gateway", message_gateway.New().Inject()}, + {"msg_gateway", msg_gateway.New().Inject()}, } { t.Run(tc.name, func(t *testing.T) { assert.Contains(t, tc.deps, want, @@ -194,7 +194,7 @@ func TestAuthGuardFailsClosed(t *testing.T) { prefix string }{ {"user", user.New().Apply, http.MethodPost, "/api/v1/user/change-password"}, - {"message_gateway", message_gateway.New().Apply, http.MethodGet, "/api/v1/message-gateway"}, + {"msg_gateway", msg_gateway.New().Apply, http.MethodGet, "/api/v1/message-gateway"}, {"admin", admin.New().Apply, http.MethodGet, "/api/v1/admin"}, } diff --git a/backend/plugins/domain/cap/errs.go b/backend/plugins/domain/cap/errs.go deleted file mode 100644 index 58fbb7d3..00000000 --- a/backend/plugins/domain/cap/errs.go +++ /dev/null @@ -1,24 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package cap 提供人机验证中间件 -package cap - -// HTTP 响应错误文案 -const ( - errCapTokenMissing = "验证码验证失败,缺少验证码凭证" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errCapTokenInvalidOrExpired = "验证码校验失败或已过期,请重试" //nolint:gosec // false positive: this is an error message, not hardcoded credentials - errCapNotConfigured = "captcha is not configured" - errChallengeGenerateFailed = "生成验证难题失败,请稍后再试" - errInvalidRequestParams = "无效的参数" - errSolutionVerifyFailed = "校验验证解答失败,请稍后再试" -) - -// Redeem 结果码,属于 redeem 响应 JSON 的对外契约取值,禁止改写取值 -const ( - redeemErrInvalidToken = "invalid_token" - redeemErrNonceStoreFailed = "nonce_store_error" - redeemErrAlreadyRedeemed = "already_redeemed" - redeemErrSettingsLoad = "settings_load_error" - redeemErrTokenStoreFailed = "token_store_error" //nolint:gosec // error code, not hardcoded credentials -) diff --git a/backend/plugins/domain/cap/middleware.go b/backend/plugins/domain/cap/middleware.go deleted file mode 100644 index 643f9e06..00000000 --- a/backend/plugins/domain/cap/middleware.go +++ /dev/null @@ -1,38 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cap - -import ( - "Wavelet/pkg/response" - - "github.com/gin-gonic/gin" -) - -// VerifyMiddleware returns a Gin middleware that checks and consumes the X-Cap-Token header. -func VerifyMiddleware(mgr *Manager, scope string) gin.HandlerFunc { - return func(c *gin.Context) { - if !ProtectionEnabled(c.Request.Context()) { - c.Next() - return - } - if mgr == nil { - response.AbortUnauthorized(c, errCapTokenInvalidOrExpired) - return - } - - token := c.GetHeader("X-Cap-Token") - if token == "" { - response.AbortUnauthorized(c, errCapTokenMissing) - return - } - - valid, err := mgr.VerifyToken(c.Request.Context(), token, scope) - if err != nil || !valid { - response.AbortUnauthorized(c, errCapTokenInvalidOrExpired) - return - } - - c.Next() - } -} diff --git a/backend/plugins/domain/cap/plugin.go b/backend/plugins/domain/cap/plugin.go deleted file mode 100644 index 2c4703dc..00000000 --- a/backend/plugins/domain/cap/plugin.go +++ /dev/null @@ -1,118 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package cap provides the proof-of-work (PoW) CAPTCHA verification domain plugin for Cordis. -package cap - -import ( - "Wavelet/core" - "Wavelet/core/contracts" - "Wavelet/core/extpoints" - "reflect" -) - -// Plugin implements core.Plugin to provide CAPTCHA generation, validation, and route protection. -type Plugin struct{} - -// New creates a new cap domain plugin. -func New() *Plugin { - return &Plugin{} -} - -// Name returns the unique identifier for the cap domain plugin. -func (p *Plugin) Name() string { - return "cap" -} - -// Inject declares required dependencies for the cap domain plugin. -func (p *Plugin) Inject() []reflect.Type { - return []reflect.Type{ - reflect.TypeFor[contracts.DBService](), - } -} - -// Manifest returns the plugin metadata. -func (p *Plugin) Manifest() core.Manifest { - return core.Manifest{ - Name: "cap", - Version: "1.0.0", - Description: "Proof-of-work CAPTCHA challenge and verification domain plugin", - Author: "Wavelet Team", - } -} - -type capAppConfig struct { - SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"` -} - -// DeclareConfig declares configuration bindings for the cap plugin. -func (p *Plugin) DeclareConfig() []core.ConfigBinding { - return []core.ConfigBinding{ - {Prefix: "app", Target: &capAppConfig{}}, - } -} - -// Apply registers the cap routes and settings into the Context. -func (p *Plugin) Apply(ctx *core.Context) error { - var cfg capAppConfig - if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" { - SetSecret([]byte(cfg.SessionSecret)) - } - - // 0. Bind DBService from Context - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } - ctx.OnDispose(func() error { - setDBService(nil) - return nil - }) - - // Listen to system config changed events to invalidate cached settings - ctx.Events().On(contracts.EventTopicConfigChanged, func(_ any) { - InvalidateRuntimeSettings() - }) - - core.Provide[contracts.CaptchaService](ctx, captchaService{}) - - // Register HTTP Routes - capGroup := ctx.Router().Group("/api/v1/cap") - { - capGroup.GET("/challenge", Challenge) - capGroup.POST("/challenge", Challenge) - capGroup.POST("/redeem", Redeem) - } - ctx.Router().RegisterWhitelist("/api/v1/cap/challenge", "/api/v1/cap/redeem") - - // Register Settings Schemas - ctx.Settings().Register(extpoints.SettingSchema{ - Key: "cap.login_enabled", - Default: false, - Description: "Whether to require CAPTCHA verification for user login", - Type: "boolean", - Category: "security", - }) - ctx.Settings().Register(extpoints.SettingSchema{ - Key: "cap.challenge_count", - Default: 1, - Description: "Number of PoW puzzle challenges to solve", - Type: "integer", - Category: "security", - }) - - return nil -} - -type captchaService struct{} - -func (captchaService) VerifyMiddleware(scope string) any { - return VerifyMiddleware(GetDefaultManager(), scope) -} - -func (captchaService) ChallengeHandler() any { return Challenge } - -func (captchaService) RedeemHandler() any { return Redeem } diff --git a/backend/plugins/domain/cap/plugin_captcha_contract_test.go b/backend/plugins/domain/cap/plugin_captcha_contract_test.go deleted file mode 100644 index 7577a31d..00000000 --- a/backend/plugins/domain/cap/plugin_captcha_contract_test.go +++ /dev/null @@ -1,55 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cap - -import ( - "context" - "testing" - - "Wavelet/core" - "Wavelet/core/contracts" -) - -func TestApplyProvidesCaptchaService(t *testing.T) { - ctx := core.NewContext(context.Background()) - if err := New().Apply(ctx); err != nil { - t.Fatal(err) - } - svc, err := core.Inject[contracts.CaptchaService](ctx) - if err != nil || svc == nil { - t.Fatalf("Inject CaptchaService: svc=%v err=%v", svc, err) - } - if svc.ChallengeHandler() == nil || svc.RedeemHandler() == nil { - t.Fatal("handlers must be non-nil") - } - if svc.VerifyMiddleware("login") == nil { - t.Fatal("VerifyMiddleware(login) must be non-nil") - } -} - -func TestApplyRegistersUnversionedCapRoutes(t *testing.T) { - ctx := core.NewContext(context.Background()) - if err := New().Apply(ctx); err != nil { - t.Fatal(err) - } - want := map[string]bool{ - "GET /api/v1/cap/challenge": false, - "POST /api/v1/cap/challenge": false, - "POST /api/v1/cap/redeem": false, - } - for _, rd := range ctx.Router().Routes() { - key := rd.Method + " " + rd.Path - if _, ok := want[key]; ok { - want[key] = true - } - if key == "POST /api/cap/challenge" || key == "POST /api/cap/redeem" { - t.Errorf("legacy route must not exist: %s", key) - } - } - for key, ok := range want { - if !ok { - t.Errorf("missing route %s", key) - } - } -} diff --git a/backend/plugins/domain/cap/repository.go b/backend/plugins/domain/cap/repository.go deleted file mode 100644 index 4caa0040..00000000 --- a/backend/plugins/domain/cap/repository.go +++ /dev/null @@ -1,58 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cap - -import ( - "Wavelet/core" - "Wavelet/core/contracts" - "context" - "sync" - - "gorm.io/gorm" -) - -var ( - dbMu sync.RWMutex - dbSvc contracts.DBService -) - -// setDBService caches the DBService contract used by the persistence layer. -func setDBService(s contracts.DBService) { - dbMu.Lock() - defer dbMu.Unlock() - dbSvc = s -} - -// getDB resolves a GORM handle, preferring the *core.Context when supplied by callers. -func getDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } - } - dbMu.RLock() - s := dbSvc - dbMu.RUnlock() - if s != nil { - return s.DB(ctx) - } - return nil -} - -// loadRuntimeSettings reads the CAPTCHA owned rows from the system config table. -func loadRuntimeSettings(ctx context.Context) (RuntimeSettings, error) { - var records []configRecord - db := getDB(ctx) - if db == nil { - return parseRuntimeSettings(nil), nil - } - if err := db.Table("w_system_configs").Where("key IN ?", runtimeConfigKeys).Find(&records).Error; err != nil { - return RuntimeSettings{}, err - } - configs := make(map[string]string, len(records)) - for _, r := range records { - configs[r.Key] = r.Value - } - return parseRuntimeSettings(configs), nil -} diff --git a/backend/plugins/domain/cap/runtime_settings.go b/backend/plugins/domain/cap/runtime_settings.go deleted file mode 100644 index ebb6fd3e..00000000 --- a/backend/plugins/domain/cap/runtime_settings.go +++ /dev/null @@ -1,185 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package cap - -import ( - "context" - "errors" - "strconv" - "sync/atomic" - "time" - - "golang.org/x/sync/singleflight" -) - -const ( - defaultChallengeCount = 1 - defaultChallengeSize = 32 - defaultChallengeDifficulty = 4 - defaultChallengeTTL = 10 * time.Minute - defaultTokenTTL = 20 * time.Minute -) - -// RuntimeSettings is the parsed CAPTCHA runtime configuration loaded from system_configs. -type RuntimeSettings struct { - LoginEnabled bool - ChallengeCount int - ChallengeSize int - ChallengeDifficulty int - ChallengeTTL time.Duration - TokenTTL time.Duration -} - -// CAP 动态配置键常量 -const ( - ConfigKeyCapLoginEnabled = "cap_login_enabled" - ConfigKeyCapChallengeCount = "cap_challenge_count" - ConfigKeyCapChallengeSize = "cap_challenge_size" - ConfigKeyCapChallengeDifficulty = "cap_challenge_difficulty" - ConfigKeyCapChallengeTTL = "cap_challenge_ttl" - // ConfigKeyCapTokenTTL 验证码 Token 过期时间键 - // #nosec G101 - ConfigKeyCapTokenTTL = "cap_token_ttl" -) - -var runtimeConfigKeys = []string{ - ConfigKeyCapLoginEnabled, - ConfigKeyCapChallengeCount, - ConfigKeyCapChallengeSize, - ConfigKeyCapChallengeDifficulty, - ConfigKeyCapChallengeTTL, - ConfigKeyCapTokenTTL, -} - -var runtimeConfigKeySet = func() map[string]struct{} { - set := make(map[string]struct{}, len(runtimeConfigKeys)) - for _, key := range runtimeConfigKeys { - set[key] = struct{}{} - } - return set -}() - -type runtimeSettingsStore struct { - snapshot atomic.Pointer[RuntimeSettings] - loadGroup singleflight.Group -} - -var settingsStore = &runtimeSettingsStore{} - -// IsRuntimeConfigKey reports whether a system config key affects CAPTCHA runtime settings. -func IsRuntimeConfigKey(key string) bool { - _, ok := runtimeConfigKeySet[key] - return ok -} - -// CurrentSettings returns the cached CAPTCHA runtime settings snapshot. -func CurrentSettings(ctx context.Context) (RuntimeSettings, error) { - return settingsStore.current(ctx) -} - -// ProtectionEnabled reports whether CAPTCHA verification is required for protected routes. -func ProtectionEnabled(ctx context.Context) bool { - settings, err := CurrentSettings(ctx) - if err != nil { - return false - } - return settings.LoginEnabled -} - -// InvalidateRuntimeSettings drops the in-process CAPTCHA settings snapshot. -func InvalidateRuntimeSettings() { - settingsStore.snapshot.Store(nil) -} - -// ResetRuntimeSettingsForTest clears the CAPTCHA runtime snapshot. -func ResetRuntimeSettingsForTest() { - InvalidateRuntimeSettings() -} - -// InstallTestRuntimeSettings installs a fixed snapshot for unit tests. -func InstallTestRuntimeSettings(settings RuntimeSettings) func() { - snapshot := settings - settingsStore.snapshot.Store(&snapshot) - return InvalidateRuntimeSettings -} - -func (s *runtimeSettingsStore) current(ctx context.Context) (RuntimeSettings, error) { - s.ensureInvalidationListener() - - if snapshot := s.snapshot.Load(); snapshot != nil { - return *snapshot, nil - } - - loaded, err, _ := s.loadGroup.Do("cap-runtime-settings", func() (any, error) { - if snapshot := s.snapshot.Load(); snapshot != nil { - return *snapshot, nil - } - - settings, loadErr := loadRuntimeSettings(ctx) - if loadErr != nil { - return RuntimeSettings{}, loadErr - } - - s.snapshot.Store(&settings) - return settings, nil - }) - if err != nil { - return RuntimeSettings{}, err - } - - settings, ok := loaded.(RuntimeSettings) - if !ok { - return RuntimeSettings{}, errors.New("cap runtime settings loader returned unexpected type") - } - return settings, nil -} - -func parseRuntimeSettings(configs map[string]string) RuntimeSettings { - settings := RuntimeSettings{ - ChallengeCount: defaultChallengeCount, - ChallengeSize: defaultChallengeSize, - ChallengeDifficulty: defaultChallengeDifficulty, - ChallengeTTL: defaultChallengeTTL, - TokenTTL: defaultTokenTTL, - } - - if len(configs) == 0 { - return settings - } - - if val, ok := configs[ConfigKeyCapLoginEnabled]; ok { - if enabled, err := strconv.ParseBool(val); err == nil { - settings.LoginEnabled = enabled - } - } - if val, ok := configs[ConfigKeyCapChallengeCount]; ok { - if count, err := strconv.Atoi(val); err == nil && count > 0 { - settings.ChallengeCount = count - } - } - if val, ok := configs[ConfigKeyCapChallengeSize]; ok { - if size, err := strconv.Atoi(val); err == nil && size > 0 { - settings.ChallengeSize = size - } - } - if val, ok := configs[ConfigKeyCapChallengeDifficulty]; ok { - if diff, err := strconv.Atoi(val); err == nil && diff > 0 { - settings.ChallengeDifficulty = diff - } - } - if val, ok := configs[ConfigKeyCapChallengeTTL]; ok { - if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 { - settings.ChallengeTTL = time.Duration(ttlSeconds) * time.Second - } - } - if val, ok := configs[ConfigKeyCapTokenTTL]; ok { - if ttlSeconds, err := strconv.Atoi(val); err == nil && ttlSeconds > 0 { - settings.TokenTTL = time.Duration(ttlSeconds) * time.Second - } - } - - return settings -} - -func (s *runtimeSettingsStore) ensureInvalidationListener() {} diff --git a/backend/plugins/domain/cap/service.go b/backend/plugins/domain/cap/service.go deleted file mode 100644 index b062fffe..00000000 --- a/backend/plugins/domain/cap/service.go +++ /dev/null @@ -1,182 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package cap provides CAPTCHA and proof-of-work (PoW) verification services. -package cap - -import ( - "Wavelet/plugins/domain/cap/pow" - "context" - "crypto/sha256" - "encoding/hex" - "strconv" - "strings" - "sync" - "time" -) - -const ( - redeemTokenIDLength = 8 // 兑换 Token ID 字节长度 - redeemVerTokenLength = 15 // 兑换验证 Token 字节长度 - tokenPartsCount = 2 // 兑换 Token 由两部分组成 - valuePartsCount = 2 // 存储值由 scope 和过期时间组成 -) - -// Manager orchestrates challenge generation and solution validation. -type Manager struct { - secret []byte - store pow.Store -} - -// NewManager creates a new CAPTCHA Manager. -func NewManager(secret []byte, store pow.Store) *Manager { - return &Manager{ - secret: secret, - store: store, - } -} - -// Generate creates a challenge response. -func (m *Manager) Generate(ctx context.Context, scope string) (*pow.ChallengeResponse, error) { - settings, err := CurrentSettings(ctx) - if err != nil { - return nil, err - } - - challengeConfig := pow.ChallengeConfig{ - Count: settings.ChallengeCount, - Size: settings.ChallengeSize, - Difficulty: settings.ChallengeDifficulty, - Expires: settings.ChallengeTTL, - } - return pow.GenerateChallenge(m.secret, challengeConfig, scope) -} - -// Redeem verifies PoW solutions and returns a one-time redeem token. -func (m *Manager) Redeem(ctx context.Context, token string, solutions []int, scope string) (*RedeemResponse, error) { - sigHex := pow.JwtSigHex(token) - if sigHex == "" { - return &RedeemResponse{Success: false, Error: redeemErrInvalidToken}, nil - } - - nonceKey := "cap:nonce:" + sigHex - - payload, err := pow.VerifyChallengeSolutions(token, solutions, m.secret, scope) - if err != nil { - return &RedeemResponse{Success: false, Error: err.Error()}, nil //nolint:nilerr // validation errors are returned as response, not system errors - } - - now := time.Now().UnixNano() / int64(time.Millisecond) - nonceTTL := time.Duration(payload.Expires-now) * time.Millisecond - if nonceTTL < time.Second { - nonceTTL = time.Second - } - - set, err := m.store.SetNX(ctx, nonceKey, "1", nonceTTL) - if err != nil { - return &RedeemResponse{Success: false, Error: redeemErrNonceStoreFailed}, err - } - if !set { - return &RedeemResponse{Success: false, Error: redeemErrAlreadyRedeemed}, nil - } - - settings, err := CurrentSettings(ctx) - if err != nil { - return &RedeemResponse{Success: false, Error: redeemErrSettingsLoad}, err - } - - id := pow.RandomHex(redeemTokenIDLength) - verToken := pow.RandomHex(redeemVerTokenLength) - verHashBytes := sha256.Sum256([]byte(verToken)) - verHashHex := hex.EncodeToString(verHashBytes[:]) - - tokenKey := "cap:token:" + id + ":" + verHashHex - tokenExpires := time.Now().Add(settings.TokenTTL) - storeVal := strconv.FormatInt(tokenExpires.UnixNano(), 10) + "|" + scope - - if err := m.store.Set(ctx, tokenKey, storeVal, settings.TokenTTL); err != nil { - return &RedeemResponse{Success: false, Error: redeemErrTokenStoreFailed}, err - } - - return &RedeemResponse{ - Success: true, - Token: id + ":" + verToken, - Expires: tokenExpires.UnixNano() / int64(time.Millisecond), - }, nil -} - -// VerifyToken validates and consumes the redeem token (single-use). -func (m *Manager) VerifyToken(ctx context.Context, token, expectedScope string) (bool, error) { - if token == "" { - return false, nil - } - parts := strings.Split(token, ":") - if len(parts) != tokenPartsCount { - return false, nil - } - id := parts[0] - verToken := parts[1] - - verHashBytes := sha256.Sum256([]byte(verToken)) - verHashHex := hex.EncodeToString(verHashBytes[:]) - - tokenKey := "cap:token:" + id + ":" + verHashHex - - val, exists, err := sGetAndDelete(ctx, m.store, tokenKey) - if err != nil { - return false, err - } - if !exists { - return false, nil - } - - valParts := strings.Split(val, "|") - if len(valParts) != valuePartsCount { - return false, nil - } - - expNano, err := strconv.ParseInt(valParts[0], 10, 64) - if err != nil { - return false, nil //nolint:nilerr // invalid format is treated as validation failure - } - tokenScope := valParts[1] - - if expectedScope != "" && tokenScope != expectedScope { - return false, nil - } - - if time.Now().UnixNano() > expNano { - return false, nil - } - - return true, nil -} - -func sGetAndDelete(ctx context.Context, store pow.Store, key string) (string, bool, error) { - if store == nil { - return "", false, nil - } - return store.GetAndDelete(ctx, key) -} - -var ( - defaultManagerMu sync.RWMutex - defaultManager *Manager -) - -// SetSecret sets the shared secret used by the default manager. -func SetSecret(secret []byte) { - defaultManagerMu.Lock() - defer defaultManagerMu.Unlock() - if len(secret) > 0 { - store := pow.NewMemoryStore(1 * time.Minute) - defaultManager = NewManager(secret, store) - } -} - -// GetDefaultManager yields the global singleton CAPTCHA manager. -func GetDefaultManager() *Manager { - defaultManagerMu.RLock() - defer defaultManagerMu.RUnlock() - return defaultManager -} diff --git a/backend/plugins/domain/domain_test.go b/backend/plugins/domain/domain_test.go index 9f3f1b6a..60c12476 100644 --- a/backend/plugins/domain/domain_test.go +++ b/backend/plugins/domain/domain_test.go @@ -9,18 +9,25 @@ import ( "Wavelet/pkg/idgen" "Wavelet/plugins/domain/admin" "Wavelet/plugins/domain/auth" - "Wavelet/plugins/domain/message_gateway" + "Wavelet/plugins/domain/msg_gateway" "Wavelet/plugins/domain/risk_control" + "Wavelet/plugins/domain/system" + "Wavelet/plugins/domain/upload" "Wavelet/plugins/domain/user" "Wavelet/plugins/infra/cache" "Wavelet/plugins/infra/logger" "Wavelet/plugins/infra/storage" "context" + "encoding/json" "io/fs" + "net/http" + "net/http/httptest" "path/filepath" "testing" + "time" "github.com/alicebob/miniredis/v2" + "github.com/gin-gonic/gin" "github.com/glebarez/sqlite" "github.com/redis/go-redis/v9" "github.com/stretchr/testify/assert" @@ -42,13 +49,16 @@ func setupTestDB(t *testing.T) *gorm.DB { &user.AccessToken{}, &auth.AuthSource{}, &auth.ExternalAccount{}, - &message_gateway.MessageChannel{}, - &message_gateway.MessageBinding{}, - &message_gateway.MessagePairingCode{}, + &msg_gateway.MessageChannel{}, + &msg_gateway.MessageBinding{}, + &msg_gateway.MessagePairingCode{}, &admin.SystemConfig{}, - &message_gateway.PushChannel{}, - &message_gateway.PushEvent{}, - &message_gateway.PushHistory{}, + &admin.TaskExecution{}, + &msg_gateway.PushChannel{}, + &msg_gateway.PushEvent{}, + &msg_gateway.PushHistory{}, + &upload.Upload{}, + &upload.UploadStat{}, )) db.SetDB(testDB) @@ -236,7 +246,7 @@ func TestUserPlugin(t *testing.T) { assert.Len(t, list, 1) assert.Equal(t, "bob", list[0].Username) - // 9. Tasks & Schedules + // 9. Tasks taskDef, ok := ctx.Tasks().Get("user:send_email_code") require.True(t, ok) assert.Equal(t, 3, taskDef.Retry) @@ -257,15 +267,15 @@ func TestMessageGatewayPlugin(t *testing.T) { require.NoError(t, cache.New().Apply(ctx)) require.NoError(t, logger.New().Apply(ctx)) - p := message_gateway.New() - assert.Equal(t, "message_gateway", p.Name()) - assert.Equal(t, "message_gateway", p.Manifest().Name) + p := msg_gateway.New() + assert.Equal(t, "msg_gateway", p.Name()) + assert.Equal(t, "msg_gateway", p.Manifest().Name) require.NoError(t, p.Apply(ctx)) // 1. Migrations - entry, ok := ctx.Migrations().Get("message_gateway") + entry, ok := ctx.Migrations().Get("msg_gateway") require.True(t, ok) - assert.Equal(t, "message_gateway", entry.PluginID) + assert.Equal(t, "msg_gateway", entry.PluginID) // 2. Routes routes := ctx.Router().Routes() @@ -282,24 +292,24 @@ func TestMessageGatewayPlugin(t *testing.T) { assert.True(t, hasBindings) // 3. Tasks & Schedules - taskDef, ok := ctx.Tasks().Get("message_gateway:push_notification") + taskDef, ok := ctx.Tasks().Get("msg_gateway:push_notification") require.True(t, ok) assert.Equal(t, 3, taskDef.Retry) - schedDef, ok := ctx.Schedules().Get("message_gateway:cleanup_pairing_codes") + schedDef, ok := ctx.Schedules().Get("msg_gateway:cleanup_pairing_codes") require.True(t, ok) assert.Equal(t, "*/10 * * * *", schedDef.Spec) // 4. EventBus Trigger - var receivedEvent message_gateway.PushNotificationEvent + var receivedEvent msg_gateway.PushNotificationEvent var eventFired bool - ctx.Events().On("notification:push", func(c context.Context, e message_gateway.PushNotificationEvent) error { + ctx.Events().On("notification:push", func(c context.Context, e msg_gateway.PushNotificationEvent) error { eventFired = true receivedEvent = e return nil }) - err := ctx.Events().Emit(context.Background(), "notification:push", message_gateway.PushNotificationEvent{ + err := ctx.Events().Emit(context.Background(), "notification:push", msg_gateway.PushNotificationEvent{ UserID: 99, Channel: "telegram", Title: "System Alert", @@ -312,7 +322,7 @@ func TestMessageGatewayPlugin(t *testing.T) { assert.Equal(t, "System Alert", receivedEvent.Title) // 5. Settings - schema, ok := ctx.Settings().Get("message_gateway.pairing_code_expiry_minutes") + schema, ok := ctx.Settings().Get("msg_gateway.pairing_code_expiry_minutes") require.True(t, ok) assert.Equal(t, 15, schema.Default) } @@ -380,17 +390,72 @@ func TestAdminPlugin(t *testing.T) { assert.True(t, hasTasks) assert.True(t, hasConfigs) - // 2. Task & Schedule - _, ok := ctx.Tasks().Get("admin:system_cleanup") + // 2. Task + _, ok := ctx.Tasks().Get("logs:db_switch") require.True(t, ok) - sched, ok := ctx.Schedules().Get("admin:system_cleanup") - require.True(t, ok) - assert.Equal(t, "0 4 * * *", sched.Spec) // 3. Settings schema, ok := ctx.Settings().Get("admin.system_cleanup_cron") require.True(t, ok) - assert.Equal(t, "0 4 * * *", schema.Default) + assert.Equal(t, "0 3 * * *", schema.Default) + + provider, err := core.Inject[contracts.PublicConfigProvider](ctx) + require.NoError(t, err) + require.NotNil(t, provider) +} + +func TestPublicConfigExposesVisibleAdminRows(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx := core.NewContext(context.Background()) + ctx.Config().SetSource(core.NewMapSource(nil)) + require.NoError(t, ctx.Config().Resolve()) + testDB := setupTestDB(t) + + require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx)) + require.NoError(t, cache.New().Apply(ctx)) + require.NoError(t, logger.New().Apply(ctx)) + + require.NoError(t, testDB.Create(&admin.SystemConfig{ + Key: "cap_login_enabled", + Value: "true", + Type: "system", + Visibility: 1, + }).Error) + + require.NoError(t, admin.New().Apply(ctx)) + require.NoError(t, system.New().Apply(ctx)) + + var handler gin.HandlerFunc + for _, rd := range ctx.Router().Routes() { + if rd.Method != "GET" || rd.Path != "/api/v1/config/public" { + continue + } + require.NotEmpty(t, rd.Handlers) + switch h := rd.Handlers[0].(type) { + case gin.HandlerFunc: + handler = h + case func(*gin.Context): + handler = h + default: + t.Fatalf("unexpected handler type %T", rd.Handlers[0]) + } + break + } + require.NotNil(t, handler) + + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/config/public", nil) + c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), ctx)) + handler(c) + + var body struct { + Data map[string]string `json:"data"` + } + require.NoError(t, json.Unmarshal(w.Body.Bytes(), &body), "body = %s", w.Body.String()) + assert.Equal(t, "true", body.Data["cap_login_enabled"]) + _, wrapped := body.Data["configs"] + assert.False(t, wrapped, "payload must be a flat map, got %v", body.Data) } func TestAllDomainPluginsCombined(t *testing.T) { @@ -420,7 +485,7 @@ func TestAllDomainPluginsCombined(t *testing.T) { // Apply Domain plugins require.NoError(t, auth.New().Apply(ctx)) require.NoError(t, user.New().Apply(ctx)) - require.NoError(t, message_gateway.New().Apply(ctx)) + require.NoError(t, msg_gateway.New().Apply(ctx)) require.NoError(t, risk_control.New().Apply(ctx)) require.NoError(t, admin.New().Apply(ctx)) @@ -455,7 +520,7 @@ func TestAllDomainPluginsCombined(t *testing.T) { // Verify total schedules registered allSchedules := ctx.Schedules().Schedules() - assert.GreaterOrEqual(t, len(allSchedules), 2) + assert.GreaterOrEqual(t, len(allSchedules), 1) // 每个调度指向的任务类型都必须已注册 Handler,否则触发时会投递到无人处理的 // 任务类型,预期的清理逻辑静默失效。 @@ -472,3 +537,125 @@ func TestAllDomainPluginsCombined(t *testing.T) { // Clean shutdown require.NoError(t, ctx.Dispose()) } + +func TestSystemCleanupEventDrivenCoordination(t *testing.T) { + ctx := core.NewContext(context.Background()) + ctx.Config().SetSource(core.NewMapSource(nil)) + require.NoError(t, ctx.Config().Resolve()) + testDB := setupTestDB(t) + + // Apply Infra & Domain plugins + require.NoError(t, db.New(db.WithDB(testDB)).Apply(ctx)) + require.NoError(t, cache.New().Apply(ctx)) + require.NoError(t, logger.New().Apply(ctx)) + require.NoError(t, storage.New().Apply(ctx)) + + require.NoError(t, user.New().Apply(ctx)) + require.NoError(t, msg_gateway.New().Apply(ctx)) + require.NoError(t, upload.New().Apply(ctx)) + require.NoError(t, admin.New().Apply(ctx)) + + // 1. Seed old and recent task executions (admin domain) + now := time.Now() + oldTime := now.Add(-10 * 24 * time.Hour) + recentTime := now.Add(-1 * time.Hour) + + oldExec := admin.TaskExecution{ + ID: 101, + TaskID: "task-old-exec", + TaskType: "test_task", + Status: "success", + CreatedAt: oldTime, + } + recentExec := admin.TaskExecution{ + ID: 102, + TaskID: "task-recent-exec", + TaskType: "test_task", + Status: "success", + CreatedAt: recentTime, + } + require.NoError(t, testDB.Create(&oldExec).Error) + require.NoError(t, testDB.Model(&oldExec).UpdateColumn("created_at", oldTime).Error) + require.NoError(t, testDB.Create(&recentExec).Error) + + // 2. Seed old and recent push histories (msg_gateway domain, 30 days retention) + oldHistoryTime := now.Add(-40 * 24 * time.Hour) + oldHistory := msg_gateway.PushHistory{ + EventKey: "login", + Channel: "telegram", + Target: "123", + Title: "Old login", + Content: "Old content", + Level: "info", + Status: "success", + CreatedAt: oldHistoryTime, + } + recentHistory := msg_gateway.PushHistory{ + EventKey: "login", + Channel: "telegram", + Target: "123", + Title: "Recent login", + Content: "Recent content", + Level: "info", + Status: "success", + CreatedAt: recentTime, + } + require.NoError(t, testDB.Create(&oldHistory).Error) + require.NoError(t, testDB.Model(&oldHistory).UpdateColumn("created_at", oldHistoryTime).Error) + require.NoError(t, testDB.Create(&recentHistory).Error) + + // 3. Seed old pending upload and recent pending upload (upload domain) + oldUpload := upload.Upload{ + ID: 901, + UserID: 1, + FileName: "old.png", + FilePath: "uploads/old.png", + FileSize: 100, + Status: upload.UploadStatusPending, + CreatedAt: now.Add(-2 * time.Hour), + } + recentUpload := upload.Upload{ + ID: 902, + UserID: 1, + FileName: "recent.png", + FilePath: "uploads/recent.png", + FileSize: 100, + Status: upload.UploadStatusPending, + CreatedAt: now.Add(-10 * time.Minute), + } + require.NoError(t, testDB.Create(&oldUpload).Error) + require.NoError(t, testDB.Model(&oldUpload).UpdateColumn("created_at", now.Add(-2*time.Hour)).Error) + require.NoError(t, testDB.Create(&recentUpload).Error) + + // Dispatch admin system cleanup task handler + taskDef, ok := ctx.Tasks().Get("system:cleanup") + require.True(t, ok, "system:cleanup task must be registered in admin") + require.NotNil(t, taskDef.Handler) + + type resultExecutor interface { + Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) + } + handler, ok := taskDef.Handler.(resultExecutor) + require.True(t, ok) + + res, err := handler.Execute(context.Background(), nil) + require.NoError(t, err) + require.NotNil(t, res) + + // Verify Admin cleanup: old task execution deleted, recent remains + var execCount int64 + testDB.Model(&admin.TaskExecution{}).Count(&execCount) + assert.Equal(t, int64(1), execCount) + + // Verify MsgGateway cleanup: old push history deleted, recent remains + var historyCount int64 + testDB.Model(&msg_gateway.PushHistory{}).Count(&historyCount) + assert.Equal(t, int64(1), historyCount) + + // Verify Upload cleanup: old pending upload deleted, recent remains + var uploadCount int64 + testDB.Model(&upload.Upload{}).Count(&uploadCount) + assert.Equal(t, int64(1), uploadCount) + + require.NoError(t, ctx.Dispose()) +} diff --git a/backend/plugins/domain/message_gateway/message_gateway_test.go b/backend/plugins/domain/message_gateway/message_gateway_test.go deleted file mode 100644 index cec83794..00000000 --- a/backend/plugins/domain/message_gateway/message_gateway_test.go +++ /dev/null @@ -1,51 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway_test - -import ( - "Wavelet/pkg/testhelper" - "Wavelet/plugins/domain/message_gateway" - "context" - "testing" - "time" - - "gorm.io/gorm" -) - -type mockDBService struct { - db *gorm.DB -} - -func (m *mockDBService) GORM() *gorm.DB { - return m.db -} - -func (m *mockDBService) DB(ctx context.Context) *gorm.DB { - return m.db.WithContext(ctx) -} - -func (m *mockDBService) Named(_ string) *gorm.DB { - return m.db -} - -func TestUpsertPairingCode_ReusesUnexpired(t *testing.T) { - testDB, _, cleanup := testhelper.SetupTestEnvironment(t) - message_gateway.SetDBServiceForTest(&mockDBService{db: testDB}) - defer func() { - message_gateway.SetDBServiceForTest(nil) - cleanup() - }() - ctx := context.Background() - first, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ABCD1234", time.Now().Add(15*time.Minute)) - if err != nil { - t.Fatal(err) - } - second, err := message_gateway.UpsertPairingCode(ctx, 1, "tg-1", "ZZZZ9999", time.Now().Add(15*time.Minute)) - if err != nil { - t.Fatal(err) - } - if first.Code != second.Code || first.Code != "ABCD1234" { - t.Fatalf("reuse failed: %+v %+v", first, second) - } -} diff --git a/backend/plugins/domain/message_gateway/model/admin.go b/backend/plugins/domain/message_gateway/model/admin.go deleted file mode 100644 index 0a8a7daf..00000000 --- a/backend/plugins/domain/message_gateway/model/admin.go +++ /dev/null @@ -1,46 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package model - -// Field is one admin form field. -type Field struct { - Key string `json:"key"` - Type string `json:"type"` - Required bool `json:"required"` -} - -// Definition describes a channel type form. -type Definition struct { - Type string `json:"type"` - Fields []Field `json:"fields"` -} - -// ChannelDTO represents a channel for admin consumption. -type ChannelDTO struct { - ID uint64 `json:"id,string"` - Name string `json:"name"` - Type string `json:"type"` - OwnerScope string `json:"owner_scope"` - OwnerID *uint64 `json:"owner_id,string,omitempty"` - Enabled bool `json:"enabled"` - Credentials map[string]string `json:"credentials"` - Extra map[string]string `json:"extra"` -} - -// CreateChannelRequest is admin create payload. -type CreateChannelRequest struct { - Name string `json:"name"` - Type string `json:"type"` - Enabled *bool `json:"enabled"` - Credentials map[string]string `json:"credentials"` - Extra map[string]string `json:"extra"` -} - -// UpdateChannelRequest is admin update payload. -type UpdateChannelRequest struct { - Name string `json:"name"` - Enabled *bool `json:"enabled"` - Credentials map[string]string `json:"credentials"` - Extra map[string]string `json:"extra"` -} diff --git a/backend/plugins/domain/message_gateway/model/models.go b/backend/plugins/domain/message_gateway/model/models.go deleted file mode 100644 index f5064736..00000000 --- a/backend/plugins/domain/message_gateway/model/models.go +++ /dev/null @@ -1,242 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package model defines the domain entities, DTOs, and schemas for message_gateway. -package model - -import ( - "Wavelet/plugins/domain/message_gateway/errs" - "errors" - "strings" - "time" -) - -// Channel type and scope constants. -const ( - ChannelTypeTelegram = "telegram" - ChannelTypeQQ = "qq" - MessageChannelTypeTelegram = "telegram" - MessageChannelTypeQQ = "qq" - MessageOwnerScopeSystem = "system" - - TypeCustom = "custom" - TypeEmail = "email" - TypeTelegram = "telegram" -) - -// Capability describes what an adapter can send and receive. -type Capability struct { - Text bool - Image bool - File bool - Reply bool - Group bool -} - -// ChannelConfig is the decrypted runtime config passed to a factory. -type ChannelConfig struct { - ID uint64 - Type string - Name string - Credentials map[string]string - Extra map[string]string -} - -// Recipient is the outbound destination on a platform. -type Recipient struct { - ChatID string - PlatformUserID string -} - -// Attachment is a downloaded inbound file sitting on local disk. -type Attachment struct { - Path string - FileName string - MIME string - Error string -} - -// InboundMessage is a normalized private-chat message. -type InboundMessage struct { - ChannelID uint64 - ChannelType string - PlatformUserID string - ChatID string - MessageID string - Text string - Attachments []Attachment - BindingUserID *uint64 -} - -// OutboundMessage is a reply or probe send. -type OutboundMessage struct { - Text string - ReplyToID string - Attachments []Attachment -} - -// MessageChannel is an admin-configured messaging adapter. -type MessageChannel struct { - ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` - Type string `json:"type" gorm:"size:32;not null"` - Name string `json:"name" gorm:"size:64;not null"` - OwnerScope string `json:"owner_scope" gorm:"size:32;not null;default:'system'"` - OwnerID *uint64 `json:"owner_id,omitempty"` - Credentials string `json:"credentials" gorm:"type:text;not null"` - Extra string `json:"extra" gorm:"type:text"` - Enabled bool `json:"enabled" gorm:"default:false;not null"` - CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` - UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` -} - -// TableName 表名 -func (MessageChannel) TableName() string { - return "w_message_channels" -} - -// MessageBinding maps a platform user to a Wavelet user on one channel. -type MessageBinding struct { - ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` - ChannelID uint64 `json:"channel_id" gorm:"not null;index"` - PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"` - UserID uint64 `json:"user_id" gorm:"not null;index"` - CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` -} - -// TableName 表名 -func (MessageBinding) TableName() string { - return "w_message_bindings" -} - -// MessagePairingCode is a one-time bind code. -type MessagePairingCode struct { - ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` - Code string `json:"code" gorm:"size:32;uniqueIndex;not null"` - ChannelID uint64 `json:"channel_id" gorm:"not null;index"` - PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"` - UserID uint64 `json:"user_id" gorm:"not null;index"` - ExpiresAt time.Time `json:"expires_at" gorm:"not null;index"` - CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` -} - -// TableName 表名 -func (MessagePairingCode) TableName() string { - return "w_message_pairing_codes" -} - -// PushChannel 消息通道模型 -type PushChannel struct { - ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` - Name string `json:"name" gorm:"size:100;not null"` - Description string `json:"description" gorm:"size:255"` - Type string `json:"type" gorm:"size:50;not null;index"` - URL string `json:"url" gorm:"type:text"` - Token string `json:"token" gorm:"type:text"` - Other string `json:"other" gorm:"type:text"` - Enabled bool `json:"enabled" gorm:"index;not null;default:true"` - CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"` - UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"` -} - -// TableName 指定 GORM 表名 -func (PushChannel) TableName() string { - return "w_push_channels" -} - -// Validate 验证与标准化字段 -func (c *PushChannel) Validate() error { - c.Name = strings.TrimSpace(c.Name) - if c.Name == "" { - return errors.New(errs.ErrChannelNameRequired) - } - c.Type = strings.TrimSpace(c.Type) - if c.Type == "" { - return errors.New(errs.ErrChannelTypeRequired) - } - return nil -} - -// PushEvent 系统通知事件模型 -type PushEvent struct { - ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` - EventKey string `json:"event_key" gorm:"uniqueIndex;size:80;not null"` - Name string `json:"name" gorm:"size:100;not null"` - TaskType string `json:"task_type" gorm:"size:100;index;not null;default:''"` - Channels []string `json:"channels" gorm:"type:text;serializer:json;not null"` - Targets []string `json:"targets" gorm:"type:text;serializer:json;not null"` - Template string `json:"template" gorm:"type:text;not null"` - Enabled bool `json:"enabled" gorm:"index;not null;default:false"` - CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"` - UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"` -} - -// TableName 指定 GORM 表名 -func (PushEvent) TableName() string { - return "w_push_events" -} - -// Validate 验证 PushEvent 实体字段 -func (e *PushEvent) Validate() error { - e.EventKey = strings.TrimSpace(e.EventKey) - if e.EventKey == "" { - return errors.New(errs.ErrEventKeyRequired) - } - e.Name = strings.TrimSpace(e.Name) - if e.Name == "" { - return errors.New(errs.ErrNameRequired) - } - return nil -} - -// PushHistory 推送日志/历史实体 -type PushHistory struct { - ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` - EventKey string `json:"event_key" gorm:"size:80;not null;index"` - Channel string `json:"channel" gorm:"size:50;not null;index"` - Target string `json:"target" gorm:"size:255;not null"` - Title string `json:"title" gorm:"size:255;not null"` - Content string `json:"content" gorm:"type:text;not null"` - Level string `json:"level" gorm:"size:20;not null;default:'INFO'"` - Status string `json:"status" gorm:"size:20;not null;index"` - ErrorMsg string `json:"error_msg" gorm:"type:text"` - Payload string `json:"payload" gorm:"type:text"` - CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"` -} - -// TableName 指定 GORM 表名 -func (PushHistory) TableName() string { - return "w_push_histories" -} - -// BindRequest is the user bind body. -type BindRequest struct { - ChannelID string `json:"channel_id"` - Code string `json:"code"` -} - -// BindingDTO is a user-facing binding row. -type BindingDTO struct { - ID uint64 `json:"id,string"` - UserID uint64 `json:"user_id,string"` - ChannelID uint64 `json:"channel_id,string"` - ChannelName string `json:"channel_name"` - ChannelType string `json:"channel_type"` - PlatformUserID string `json:"platform_user_id"` - CreatedAt time.Time `json:"created_at"` -} - -// PublicChannelDTO is an enabled channel a user can bind to. -type PublicChannelDTO struct { - ID uint64 `json:"id,string"` - Name string `json:"name"` - Type string `json:"type"` -} - -// PushNotificationEvent defines the payload for eventbus notification trigger. -type PushNotificationEvent struct { - UserID uint64 `json:"user_id"` - Channel string `json:"channel"` - Title string `json:"title"` - Content string `json:"content"` - Metadata map[string]any `json:"metadata,omitempty"` -} diff --git a/backend/plugins/domain/message_gateway/pairing_test.go b/backend/plugins/domain/message_gateway/pairing_test.go deleted file mode 100644 index c442c30d..00000000 --- a/backend/plugins/domain/message_gateway/pairing_test.go +++ /dev/null @@ -1,33 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import ( - "strings" - "testing" -) - -func TestGenerateCode_AlphabetAndLength(t *testing.T) { - code, err := GenerateCode() - if err != nil { - t.Fatal(err) - } - if len(code) != 8 { - t.Fatalf("len=%d", len(code)) - } - for _, r := range code { - if !strings.ContainsRune(CodeAlphabet, r) { - t.Fatalf("bad rune %q", r) - } - } -} - -func TestNormalizeAndFormat(t *testing.T) { - if got := NormalizeCode("ab-cd-ef-gh"); got != "ABCDEFGH" { - t.Fatalf("got %q", got) - } - if got := FormatCode("ABCDEFGH"); got != "ABCD-EFGH" { - t.Fatalf("got %q", got) - } -} diff --git a/backend/plugins/domain/message_gateway/plugin.go b/backend/plugins/domain/message_gateway/plugin.go deleted file mode 100644 index 482a6c26..00000000 --- a/backend/plugins/domain/message_gateway/plugin.go +++ /dev/null @@ -1,385 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package message_gateway provides the Bot gateway, multi-channel notification dispatching, and asynchronous push worker domain plugin for Cordis. -package message_gateway - -import ( - "Wavelet/core" - "Wavelet/core/contracts" - "Wavelet/core/extpoints" - "Wavelet/pkg/ginutil" - "Wavelet/pkg/util" - "Wavelet/plugins/domain/message_gateway/handler" - "Wavelet/plugins/domain/message_gateway/model" - "Wavelet/plugins/domain/message_gateway/repository" - "Wavelet/plugins/domain/message_gateway/service" - "context" - "embed" - "reflect" - - "github.com/gin-gonic/gin" -) - -//go:embed migrations/*/*.sql -var mgMigrations embed.FS - -// Option configures the message_gateway plugin. -type Option func(*Plugin) - -// WithAutoStartRunner enables automatic bot runner startup in the background. -func WithAutoStartRunner(enable bool) Option { - return func(p *Plugin) { - p.autoStartRunner = enable - } -} - -// Plugin implements core.Plugin to provide Bot gateway and notification dispatch domain services. -type Plugin struct { - autoStartRunner bool - cancelRunner context.CancelFunc -} - -// New creates a new message_gateway domain plugin. -func New(opts ...Option) *Plugin { - p := &Plugin{} - for _, opt := range opts { - if opt != nil { - opt(p) - } - } - return p -} - -// Name returns the unique identifier for the message_gateway domain plugin. -func (p *Plugin) Name() string { - return "message_gateway" -} - -// Inject declares required dependencies for the message_gateway domain plugin. -func (p *Plugin) Inject() []reflect.Type { - return []reflect.Type{ - reflect.TypeFor[contracts.DBService](), - // AuthService is captured as a middleware value in Apply, so it cannot - // be late-bound with core.When like the other services below; the - // kernel must mount auth first or the routes get a pass-through guard. - reflect.TypeFor[contracts.AuthService](), - } -} - -// Manifest returns the plugin metadata. -func (p *Plugin) Manifest() core.Manifest { - return core.Manifest{ - Name: "message_gateway", - Version: "1.0.0", - Description: "Bot gateway, multi-channel notification push, and async worker dispatch plugin", - Author: "Wavelet Team", - } -} - -type mgAppConfig struct { - SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"` -} - -// DeclareConfig declares configuration bindings for the message_gateway plugin. -func (p *Plugin) DeclareConfig() []core.ConfigBinding { - return []core.ConfigBinding{ - {Prefix: "app", Target: &mgAppConfig{}}, - } -} - -// Apply registers message_gateway migrations, routes, tasks, schedules, events, and settings into the Context. -func (p *Plugin) Apply(ctx *core.Context) error { - var cfg mgAppConfig - if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" { - service.SetCredentialSecret(cfg.SessionSecret) - } - // 0. Bind DBService, CacheService, TaskService, UserService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - repository.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - repository.SetDBService(db) - }) - } - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - repository.SetCacheService(cache) - service.SetCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - repository.SetCacheService(cache) - service.SetCacheService(cache) - }) - } - if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil { - service.SetTaskService(taskSvc) - } else { - core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) { - service.SetTaskService(taskSvc) - }) - } - if uSvc, err := core.Inject[contracts.UserService](ctx); err == nil && uSvc != nil { - service.SetUserService(uSvc) - } else { - core.When[contracts.UserService](ctx, func(uSvc contracts.UserService) { - service.SetUserService(uSvc) - }) - } - ctx.OnDispose(func() error { - repository.SetDBService(nil) - repository.SetCacheService(nil) - service.SetCacheService(nil) - service.SetTaskService(nil) - service.SetUserService(nil) - return nil - }) - - // 0. Resolve auth service for middleware (via IoC, not direct import) - denyAuth := ginutil.AuthUnavailable() - loginMW := denyAuth - adminMW := denyAuth - if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil { - if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok { - loginMW = mw - } - if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok { - adminMW = mw - } - } - - // 1. Register migrations - ctx.Migrations().Register("message_gateway", mgMigrations) - - // 2. Register User HTTP Routes - handler.RegisterUserRoutes(ctx.Router().Group("/api/v1"), loginMW) - - // 3. Register Admin Message Gateway HTTP Routes - handler.RegisterAdminRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW) - - // 4. Register Admin Push HTTP Routes - handler.RegisterAdminPushRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW) - - const defaultTaskRetry = 3 - pushHandler := &service.PushHandler{} - - // 5. Register background tasks - ctx.Task().Register("message_gateway:push_notification", func(c context.Context, payload []byte) error { - return pushHandler.Execute(c, payload) - }, - extpoints.WithTaskType("push_notification"), - extpoints.WithTaskName("消息网关推送通知"), - extpoints.WithTaskDescription("异步执行系统通知的多渠道派发与推送"), - extpoints.WithTaskCategory("push"), - extpoints.WithTaskRetry(defaultTaskRetry), - extpoints.WithTaskQueue("default"), - extpoints.WithTaskRetryable(true), - ) - - ctx.Task().Register(service.SendNotificationTask, func(c context.Context, payload []byte) error { - return pushHandler.Execute(c, payload) - }, extpoints.WithTaskMeta(service.SendNotificationMeta), extpoints.WithTaskRetry(defaultTaskRetry)) - - ctx.Task().Register("message_gateway:dispatch_bot_msg", func(_ context.Context, _ []byte) error { - return nil - }, - extpoints.WithTaskType("dispatch_bot_msg"), - extpoints.WithTaskName("分发 Bot 消息"), - extpoints.WithTaskDescription("异步处理与分发 Bot 下行消息"), - extpoints.WithTaskCategory("messaging"), - extpoints.WithTaskQueue("default"), - ) - - ctx.Task().Register("message_gateway:cleanup_pairing_codes", func(c context.Context, _ []byte) error { - return repository.DeleteExpiredPairingCodes(c) - }, - extpoints.WithTaskType("cleanup_pairing_codes"), - extpoints.WithTaskName("清理过期配对码"), - extpoints.WithTaskDescription("定时清理已过期的平台 Bot 配对码"), - extpoints.WithTaskCategory("messaging"), - extpoints.WithTaskRetry(defaultTaskRetry), - extpoints.WithTaskQueue("default"), - extpoints.WithTaskRetryable(true), - ) - - // 6. Register Cron Schedules - ctx.Schedule().RegisterCron("*/10 * * * *", "message_gateway:cleanup_pairing_codes", map[string]any{"action": "cleanup"}) - - // 7. Register EventBus listeners for decoupled push triggers - ctx.Events().On("notification:push", func(c context.Context, e model.PushNotificationEvent) error { - meta := model.EventMetadata{ - Key: "eventbus:" + e.Channel, - Name: e.Title, - DefaultTemplate: model.NotificationMessage{ - Title: e.Title, - Content: e.Content, - Level: model.DefaultLevelInfo, - Ext: e.Metadata, - }, - Description: "EventBus triggered notification", - } - service.DefaultTrigger.Trigger(c, meta, map[string]any{ - "user.id": e.UserID, - "title": e.Title, - "content": e.Content, - }) - return nil - }) - - // 8. Register task completed event listener - ctx.Events().On(contracts.EventTopicTaskCompleted, func(c context.Context, e contracts.TaskCompletedEvent) error { - service.HandleTaskCompleted(c, e) - return nil - }) - - // 9. Register built-in domain events and provide PushRegistry - service.RegisterCustomEvents() - core.Provide[contracts.PushRegistry](ctx, service.PushRegistryAdapter{}) - - // 10. Register Settings Schemas - ctx.Settings().Register(extpoints.SettingSchema{ - Key: "message_gateway.pairing_code_expiry_minutes", - Default: 15, - Description: "Expiry duration for bot pairing codes in minutes", - Type: "integer", - Category: "messaging", - }) - ctx.Settings().Register(extpoints.SettingSchema{ - Key: "message_gateway.max_bindings_per_user", - Default: 5, - Description: "Maximum platform bot bindings per user", - Type: "integer", - Category: "messaging", - }) - - // 11. Optional runner start & lifecycle - if p.autoStartRunner { - runnerCtx, cancel := context.WithCancel(ctx.GoContext()) - p.cancelRunner = cancel - util.Go(func() { - _ = service.Start(runnerCtx) - }) - } - - ctx.OnDispose(func() error { - if p.cancelRunner != nil { - p.cancelRunner() - } - return nil - }) - - return nil -} - -// Re-exported constants. -const ( - CodeAlphabet = service.CodeAlphabet - CodeLength = service.CodeLength -) - -// MessageChannel is an alias for model.MessageChannel. -type MessageChannel = model.MessageChannel - -// MessageBinding is an alias for model.MessageBinding. -type MessageBinding = model.MessageBinding - -// MessagePairingCode is an alias for model.MessagePairingCode. -type MessagePairingCode = model.MessagePairingCode - -// PushChannel is an alias for model.PushChannel. -type PushChannel = model.PushChannel - -// PushEvent is an alias for model.PushEvent. -type PushEvent = model.PushEvent - -// PushHistory is an alias for model.PushHistory. -type PushHistory = model.PushHistory - -// PushNotificationEvent is an alias for model.PushNotificationEvent. -type PushNotificationEvent = model.PushNotificationEvent - -// ChannelConfig is an alias for model.ChannelConfig. -type ChannelConfig = model.ChannelConfig - -// Capability is an alias for model.Capability. -type Capability = model.Capability - -// Recipient is an alias for model.Recipient. -type Recipient = model.Recipient - -// Attachment is an alias for model.Attachment. -type Attachment = model.Attachment - -// InboundMessage is an alias for model.InboundMessage. -type InboundMessage = model.InboundMessage - -// OutboundMessage is an alias for model.OutboundMessage. -type OutboundMessage = model.OutboundMessage - -// BindingDTO is an alias for model.BindingDTO. -type BindingDTO = model.BindingDTO - -// PublicChannelDTO is an alias for model.PublicChannelDTO. -type PublicChannelDTO = model.PublicChannelDTO - -// Definition is an alias for model.Definition. -type Definition = model.Definition - -// ChannelDTO is an alias for model.ChannelDTO. -type ChannelDTO = model.ChannelDTO - -// CreateChannelRequest is an alias for model.CreateChannelRequest. -type CreateChannelRequest = model.CreateChannelRequest - -// UpdateChannelRequest is an alias for model.UpdateChannelRequest. -type UpdateChannelRequest = model.UpdateChannelRequest - -// PushDefinition is an alias for model.PushDefinition. -type PushDefinition = model.PushDefinition - -// PushField is an alias for model.PushField. -type PushField = model.PushField - -// NotificationMessage is an alias for model.NotificationMessage. -type NotificationMessage = model.NotificationMessage - -// EventMetadata is an alias for model.EventMetadata. -type EventMetadata = model.EventMetadata - -// SendPayload is an alias for model.SendPayload. -type SendPayload = model.SendPayload - -// Handler is an alias for service.Handler. -type Handler = service.Handler - -// Factory is an alias for service.Factory. -type Factory = service.Factory - -// Channel is an alias for service.Channel. -type Channel = service.Channel - -// Runner is an alias for service.Runner. -type Runner = service.Runner - -// EventTrigger is an alias for service.EventTrigger. -type EventTrigger = service.EventTrigger - -// PushHandler is an alias for service.PushHandler. -type PushHandler = service.PushHandler - -// Re-exported variables and functions. -var ( - SetDBServiceForTest = repository.SetDBServiceForTest - UpsertPairingCode = repository.UpsertPairingCode - Register = service.Register - Lookup = service.Lookup - GenerateCode = service.GenerateCode - NormalizeCode = service.NormalizeCode - FormatCode = service.FormatCode - Start = service.Start - Stop = service.Stop - GlobalRunner = service.GlobalRunner - DefaultTrigger = service.DefaultTrigger - SyncEvents = service.SyncEvents - AdminLogin = service.AdminLogin - HandleAdminLoggedIn = service.HandleAdminLoggedIn -) diff --git a/backend/plugins/domain/message_gateway/push/email.go b/backend/plugins/domain/message_gateway/push/email.go deleted file mode 100644 index 2948c577..00000000 --- a/backend/plugins/domain/message_gateway/push/email.go +++ /dev/null @@ -1,104 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import ( - "Wavelet/pkg/util" - "context" - "errors" - "fmt" - "net" - "net/smtp" - "strings" -) - -func init() { - Register("email", &EmailPusher{}) -} - -// EmailPusher 极简 SMTP 邮件推送实现 (静态、解耦) -type EmailPusher struct{} - -// sanitizeEmailHeader removes CR/LF bytes so untrusted values cannot inject -// additional email headers (email header injection). -func sanitizeEmailHeader(v string) string { - v = strings.ReplaceAll(v, "\r", "") - v = strings.ReplaceAll(v, "\n", "") - return v -} - -// Send 发送邮件 -func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, ext map[string]any) (string, error) { - if cfg.URL == "" || cfg.Key == "" || cfg.Secret == "" { - return "", errors.New("email: SMTP configuration (url, key, secret) is incomplete") - } - if target == "" { - return "", errors.New("email: target email address is required") - } - - title := bodyTitle(body) - content := bodyContent(body, "

%s: %v

", "") - - // 邮件头和体 - from := cfg.Key - to := target - - // 如果 ext 中指定了 from_name,我们在 From 头部包含它 - fromName := "System Notification" - if ext != nil { - if fn, ok := ext["from_name"].(string); ok && fn != "" { - fromName = fn - } - } - - subjectHeader := fmt.Sprintf("Subject: %s\r\n", sanitizeEmailHeader(title)) - fromHeader := fmt.Sprintf("From: %s <%s>\r\n", sanitizeEmailHeader(fromName), sanitizeEmailHeader(from)) - toHeader := fmt.Sprintf("To: %s\r\n", sanitizeEmailHeader(to)) - mimeHeader := "MIME-version: 1.0;\r\nContent-Type: text/html; charset=\"UTF-8\";\r\n\r\n" - - // 拼装完整的邮件报文 - // 简单的 HTML 正文渲染 - htmlBody := fmt.Sprintf(`

%s

%s
`, title, content) - msg := []byte(fromHeader + toHeader + subjectHeader + mimeHeader + htmlBody + "\r\n") - - // 解析 Host 和 Port - host, port, err := net.SplitHostPort(cfg.URL) - if err != nil { - host = cfg.URL - port = "25" // 默认 SMTP 端口 - } - - auth := smtp.PlainAuth("", cfg.Key, cfg.Secret, host) - - // 异步超时处理 - errChan := make(chan error, 1) - util.Go(func() { - errChan <- smtp.SendMail(host+":"+port, auth, from, []string{to}, msg) - }) - - select { - case <-ctx.Done(): - return "", ctx.Err() - case err := <-errChan: - if err != nil { - return "", fmt.Errorf("email: send smtp mail failed: %w", err) - } - } - - return "", nil -} - -// ValidateConfig 校验邮件 SMTP 配置 -func (p *EmailPusher) ValidateConfig(cfg Config) error { - if cfg.URL == "" { - return errors.New("SMTP host:port is required") - } - if cfg.Key == "" { - return errors.New("SMTP username is required") - } - if cfg.Secret == "" { - return errors.New("SMTP password is required") - } - return nil -} diff --git a/backend/plugins/domain/message_gateway/push/email_test.go b/backend/plugins/domain/message_gateway/push/email_test.go deleted file mode 100644 index fb7367cc..00000000 --- a/backend/plugins/domain/message_gateway/push/email_test.go +++ /dev/null @@ -1,26 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import "testing" - -func TestSanitizeEmailHeader(t *testing.T) { - tests := []struct { - name string - input string - want string - }{ - {"plain", "System Notification", "System Notification"}, - {"crlf stripped", "alert\r\nBcc: attacker@example.com", "alertBcc: attacker@example.com"}, - {"cr stripped", "a\rb", "ab"}, - {"lf stripped", "a\nb", "ab"}, - } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - if got := sanitizeEmailHeader(tt.input); got != tt.want { - t.Errorf("sanitizeEmailHeader(%q) = %q, want %q", tt.input, got, tt.want) - } - }) - } -} diff --git a/backend/plugins/domain/message_gateway/push/lark.go b/backend/plugins/domain/message_gateway/push/lark.go deleted file mode 100644 index afa2cd58..00000000 --- a/backend/plugins/domain/message_gateway/push/lark.go +++ /dev/null @@ -1,256 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import ( - "Wavelet/pkg/httppool" - "bytes" - "context" - "crypto/hmac" - "crypto/sha256" - "encoding/base64" - "encoding/json" - "errors" - "fmt" - "net/http" - "strconv" - "strings" - "time" -) - -func init() { - Register("lark", &LarkPusher{}) -} - -const ( - msgTypeInteractive = "interactive" -) - -// LarkPusher 飞书 Webhook 机器人推送实现 -type LarkPusher struct{} - -type larkTextContent struct { - Text string `json:"text"` -} - -type larkCardHeaderTitle struct { - Content string `json:"content"` - Tag string `json:"tag"` -} - -type larkCardHeader struct { - Template string `json:"template"` // "blue", "orange", "red" etc. - Title larkCardHeaderTitle `json:"title"` -} - -type larkCardElementText struct { - Content string `json:"content"` - Tag string `json:"tag"` // "lark_md" -} - -type larkCardElement struct { - Tag string `json:"tag"` // "div" - Text larkCardElementText `json:"text"` -} - -type larkCardContent struct { - Header larkCardHeader `json:"header"` - Elements []larkCardElement `json:"elements"` -} - -type larkMessageRequest struct { - MessageType string `json:"msg_type"` - Timestamp string `json:"timestamp,omitempty"` - Sign string `json:"sign,omitempty"` - Content larkTextContent `json:"content,omitempty"` - Card *larkCardContent `json:"card,omitempty"` -} - -type larkMessageResponse struct { - Code int `json:"code"` - Msg string `json:"msg"` -} - -// Send 执行飞书消息发送 -// -//nolint:nestif,cyclop -func (p *LarkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, template string, _ map[string]any) (string, error) { - if cfg.URL == "" { - return "", errors.New("lark: URL is required") - } - - var req larkMessageRequest - - // 1. 如果有自定义模板,我们尝试进行解析 - if template != "" { - rendered := ParseTemplate(template, body) - - // 尝试解析原生的 Lark Card - var customCard larkCardContent - var rawMap map[string]any - _ = json.Unmarshal([]byte(rendered), &rawMap) - - if rawMap != nil && rawMap["elements"] != nil { - // 如果包含 elements 字段,说明是用户定制的原生飞书卡片 JSON - if err := json.Unmarshal([]byte(rendered), &customCard); err == nil { - req.MessageType = msgTypeInteractive - req.Card = &customCard - } else { - req.MessageType = "text" - req.Content.Text = rendered - } - } else { - // 说明配置的是系统统一通知消息 of JSON 模板:{"title": "...", "content": "...", "level": "..."} - type larkNotificationMessage struct { - Title string `json:"title"` - Content string `json:"content"` - Level string `json:"level"` - } - var msg larkNotificationMessage - if err := json.Unmarshal([]byte(rendered), &msg); err == nil && (msg.Title != "" || msg.Content != "") { - title := msg.Title - if title == "" { - title = defaultTitle - } - content := msg.Content - level := strings.ToUpper(msg.Level) - if level == "" { - level = levelInfo - } - - headerColor := "blue" - switch level { - case "IMPORTANT": - headerColor = "orange" - case "CRITICAL": - headerColor = "red" - } - - req.MessageType = msgTypeInteractive - req.Card = &larkCardContent{ - Header: larkCardHeader{ - Template: headerColor, - Title: larkCardHeaderTitle{ - Content: title, - Tag: "plain_text", - }, - }, - Elements: []larkCardElement{ - { - Tag: "div", - Text: larkCardElementText{ - Content: content, - Tag: "lark_md", - }, - }, - }, - } - } else { - // 兜底:如果无法按 JSON 解析出结构化字段,当做普通文本发送 - req.MessageType = "text" - req.Content.Text = rendered - } - } - } else { - // 2. 如果无模板,默认生成一个精美的飞书互动卡片 - title := bodyTitle(body) - content := bodyContent(body, "**%s**: %v", "\n") - level := bodyLevel(body) - - // 根据级别确定飞书卡片头部的背景色模板 - headerColor := "blue" - switch level { - case "IMPORTANT": - headerColor = "orange" - case "CRITICAL": - headerColor = "red" - } - - req.MessageType = msgTypeInteractive - req.Card = &larkCardContent{ - Header: larkCardHeader{ - Template: headerColor, - Title: larkCardHeaderTitle{ - Content: title, - Tag: "plain_text", - }, - }, - Elements: []larkCardElement{ - { - Tag: "div", - Text: larkCardElementText{ - Content: content, - Tag: "lark_md", - }, - }, - }, - } - } - - // 3. 计算签名 (如果配置了 secret) - if cfg.Secret != "" { - timestamp := time.Now().Unix() - sign, err := larkSign(cfg.Secret, timestamp) - if err != nil { - return "", fmt.Errorf("lark: sign failed: %w", err) - } - req.Timestamp = strconv.FormatInt(timestamp, 10) - req.Sign = sign - } - - jsonData, err := json.Marshal(req) - if err != nil { - return "", fmt.Errorf("lark: marshal request failed: %w", err) - } - - // 4. 发送 POST 请求 - httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, cfg.URL, bytes.NewBuffer(jsonData)) - if err != nil { - return "", fmt.Errorf("lark: create http request failed: %w", err) - } - httpReq.Header.Set("Content-Type", "application/json") - - client := httppool.NewClient(defaultHTTPClientTimeout) - resp, err := client.Do(httpReq) - if err != nil { - return "", fmt.Errorf("lark: http request failed: %w", err) - } - defer func() { _ = resp.Body.Close() }() - - if resp.StatusCode != http.StatusOK { - return "", fmt.Errorf("lark: http status %s", resp.Status) - } - - var res larkMessageResponse - if err := json.NewDecoder(resp.Body).Decode(&res); err != nil { - return "", fmt.Errorf("lark: decode response failed: %w", err) - } - - if res.Code != 0 { - return "", fmt.Errorf("lark: send message failed, code %d: %s", res.Code, res.Msg) - } - - return "", nil -} - -// ValidateConfig 校验飞书配置 -func (p *LarkPusher) ValidateConfig(cfg Config) error { - if cfg.URL == "" { - return errors.New("webhook URL is required") - } - if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") { - return errors.New("webhook URL must start with http:// or https://") - } - return nil -} - -func larkSign(secret string, timestamp int64) (string, error) { - stringToSign := fmt.Sprintf("%v", timestamp) + "\n" + secret - h := hmac.New(sha256.New, []byte(stringToSign)) - _, err := h.Write(nil) - if err != nil { - return "", err - } - return base64.StdEncoding.EncodeToString(h.Sum(nil)), nil -} diff --git a/backend/plugins/domain/message_gateway/push/telegram.go b/backend/plugins/domain/message_gateway/push/telegram.go deleted file mode 100644 index 68f6ba12..00000000 --- a/backend/plugins/domain/message_gateway/push/telegram.go +++ /dev/null @@ -1,140 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import ( - "Wavelet/pkg/httppool" - "bytes" - "context" - "encoding/json" - "errors" - "fmt" - "net/http" - "strings" -) - -func init() { - Register("telegram", &TelegramPusher{}) -} - -// TelegramPusher Telegram 机器人推送实现 -type TelegramPusher struct{} - -type telegramMessageRequest struct { - ChatID string `json:"chat_id"` - Text string `json:"text"` - ParseMode string `json:"parse_mode,omitempty"` -} - -type telegramErrorResponse struct { - Ok bool `json:"ok"` - ErrorCode int `json:"error_code"` - Description string `json:"description"` -} - -// Send 执行 Telegram 消息发送 -func (p *TelegramPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, template string, _ map[string]any) (string, error) { - if cfg.Secret == "" { - return "", errors.New("telegram: Bot Token (Secret) is required") - } - - chatID := target - if chatID == "" { - chatID = cfg.Key // Use default chat ID (Key) if target is blank - } - if chatID == "" { - return "", errors.New("telegram: chat_id (target or default Key) is required") - } - - baseURL := cfg.URL - if baseURL == "" { - baseURL = "https://api.telegram.org" - } - baseURL = strings.TrimSuffix(baseURL, "/") - - title := bodyTitle(body) - content := bodyContent(body, "%s: %v", "\n") - level := bodyLevel(body) - - var text string - if template != "" { - text = ParseTemplate(template, body) - } else { - text = fmt.Sprintf("[%s] %s\n\n%s", escapeHTML(level), escapeHTML(title), escapeHTML(content)) - } - - // Try sending with HTML parse mode - err := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, text, "HTML") - if err != nil { - // Fallback: send as plain text without parse mode - plainText := text - if template == "" { - plainText = fmt.Sprintf("[%s] %s\n\n%s", level, title, content) - } - fallbackErr := p.sendMessage(ctx, baseURL, cfg.Secret, chatID, plainText, "") - if fallbackErr != nil { - return "", fmt.Errorf("telegram: send message failed (fallback also failed): %w (original HTML error: %w)", fallbackErr, err) - } - } - - return "", nil -} - -// ValidateConfig 校验 Telegram 配置 -func (p *TelegramPusher) ValidateConfig(cfg Config) error { - if cfg.Secret == "" { - return errors.New("bot Token (Secret) is required") - } - if cfg.URL != "" { - if !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") { - return errors.New("API base URL must start with http:// or https://") - } - } - return nil -} - -func (p *TelegramPusher) sendMessage(ctx context.Context, baseURL, token, chatID, text, parseMode string) error { - apiURL := fmt.Sprintf("%s/bot%s/sendMessage", baseURL, token) - - reqPayload := telegramMessageRequest{ - ChatID: chatID, - Text: text, - ParseMode: parseMode, - } - - jsonData, err := json.Marshal(reqPayload) - if err != nil { - return fmt.Errorf("marshal request failed: %w", err) - } - - httpReq, err := http.NewRequestWithContext(ctx, http.MethodPost, apiURL, bytes.NewBuffer(jsonData)) - if err != nil { - return fmt.Errorf("create http request failed: %w", err) - } - httpReq.Header.Set("Content-Type", "application/json") - - client := httppool.NewClient(defaultHTTPClientTimeout) - resp, err := client.Do(httpReq) - if err != nil { - return fmt.Errorf("http request failed: %w", err) - } - defer func() { _ = resp.Body.Close() }() - - if resp.StatusCode != http.StatusOK { - var errRes telegramErrorResponse - if decodeErr := json.NewDecoder(resp.Body).Decode(&errRes); decodeErr == nil { - return fmt.Errorf("http status %d: %s", resp.StatusCode, errRes.Description) - } - return fmt.Errorf("http status %s", resp.Status) - } - - return nil -} - -func escapeHTML(s string) string { - s = strings.ReplaceAll(s, "&", "&") - s = strings.ReplaceAll(s, "<", "<") - s = strings.ReplaceAll(s, ">", ">") - return s -} diff --git a/backend/plugins/domain/message_gateway/push/telegram_test.go b/backend/plugins/domain/message_gateway/push/telegram_test.go deleted file mode 100644 index 74fcf728..00000000 --- a/backend/plugins/domain/message_gateway/push/telegram_test.go +++ /dev/null @@ -1,116 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import ( - "context" - "encoding/json" - "net/http" - "net/http/httptest" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestTelegramPusher_Send(t *testing.T) { - t.Run("successful send with HTML parse mode", func(t *testing.T) { - var receivedReq telegramMessageRequest - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - assert.Equal(t, "/botmy-token/sendMessage", r.URL.Path) - assert.Equal(t, http.MethodPost, r.Method) - assert.Equal(t, "application/json", r.Header.Get("Content-Type")) - - err := json.NewDecoder(r.Body).Decode(&receivedReq) - require.NoError(t, err) - - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"ok": true}`)) - })) - defer server.Close() - - pusher := &TelegramPusher{} - cfg := Config{ - Channel: "telegram", - URL: server.URL, - Secret: "my-token", - } - body := map[string]any{ - "title": "Alert", - "content": "Host down", - "level": "CRITICAL", - } - _, err := pusher.Send(context.Background(), cfg, "123456", body, "", nil) - require.NoError(t, err) - - assert.Equal(t, "123456", receivedReq.ChatID) - assert.Contains(t, receivedReq.Text, "[CRITICAL] Alert") - assert.Contains(t, receivedReq.Text, "Host down") - assert.Equal(t, "HTML", receivedReq.ParseMode) - }) - - t.Run("fallback to plain text on HTML error", func(t *testing.T) { - var requests []*telegramMessageRequest - server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - var req telegramMessageRequest - err := json.NewDecoder(r.Body).Decode(&req) - require.NoError(t, err) - requests = append(requests, &req) - - if len(requests) == 1 { - w.WriteHeader(http.StatusBadRequest) - _, _ = w.Write([]byte(`{"ok": false, "error_code": 400, "description": "Bad Request: can't parse entities"}`)) - } else { - w.WriteHeader(http.StatusOK) - _, _ = w.Write([]byte(`{"ok": true}`)) - } - })) - defer server.Close() - - pusher := &TelegramPusher{} - cfg := Config{ - Channel: "telegram", - URL: server.URL, - Secret: "my-token", - } - body := map[string]any{ - "title": "Alert & Info", - "content": "A < B comparison", - "level": "INFO", - } - _, err := pusher.Send(context.Background(), cfg, "123456", body, "", nil) - require.NoError(t, err) - - require.Len(t, requests, 2) - assert.Equal(t, "HTML", requests[0].ParseMode) - assert.Equal(t, "", requests[1].ParseMode) - assert.Contains(t, requests[1].Text, "[INFO] Alert & Info") - assert.Contains(t, requests[1].Text, "A < B comparison") - }) - - t.Run("validation error", func(t *testing.T) { - pusher := &TelegramPusher{} - cfg := Config{ - Channel: "telegram", - URL: "https://api.telegram.org", - } - err := pusher.ValidateConfig(cfg) - assert.Error(t, err) - - cfg = Config{ - Channel: "telegram", - URL: "ftp://api.telegram.org", - Secret: "token", - } - err = pusher.ValidateConfig(cfg) - assert.Error(t, err) - - cfg = Config{ - Channel: "telegram", - Secret: "token", - } - err = pusher.ValidateConfig(cfg) - assert.NoError(t, err) - }) -} diff --git a/backend/plugins/domain/message_gateway/push/template.go b/backend/plugins/domain/message_gateway/push/template.go deleted file mode 100644 index 9e472bd3..00000000 --- a/backend/plugins/domain/message_gateway/push/template.go +++ /dev/null @@ -1,111 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package push - -import ( - "encoding/json" - "fmt" - "maps" - "slices" - "strconv" - "strings" -) - -// ParseTemplate parses template strings by replacing {{placeholder}} structures with values from body. -// It is a single-pass parser designed for high performance and low allocations. -func ParseTemplate(template string, body map[string]any) string { - var buf strings.Builder - buf.Grow(len(template)) - - i := 0 - for { - pos := strings.Index(template[i:], "{{") - if pos == -1 { - buf.WriteString(template[i:]) - break - } - // Write prefix - buf.WriteString(template[i : i+pos]) - i += pos + 2 // skip "{{" - - endPos := strings.Index(template[i:], "}}") - if endPos == -1 { - // Unbalanced "{{" - buf.WriteString("{{") - buf.WriteString(template[i:]) - break - } - key := template[i : i+endPos] - if val, ok := body[key]; ok { - buf.WriteString(formatValue(val)) - } else { - // Keep the placeholder if key not found - buf.WriteString("{{") - buf.WriteString(key) - buf.WriteString("}}") - } - i += endPos + 2 // skip "}}" - } - return buf.String() -} - -func formatValue(v any) string { - if v == nil { - return "" - } - switch val := v.(type) { - case string: - return val - case []byte: - return string(val) - case int: - return strconv.Itoa(val) - case int32: - return strconv.FormatInt(int64(val), 10) - case int64: - return strconv.FormatInt(val, 10) - case float64: - return strconv.FormatFloat(val, 'f', -1, 64) - case bool: - return strconv.FormatBool(val) - default: - // If it's a map, slice, or struct, marshal it to JSON. - b, err := json.Marshal(v) - if err == nil { - return string(b) - } - return fmt.Sprintf("%v", v) - } -} - -// bodyTitle returns the notification title, falling back to the default. -func bodyTitle(body map[string]any) string { - if t, ok := body["title"].(string); ok && t != "" { - return t - } - return defaultTitle -} - -// bodyContent returns the notification body, rendering every entry with format -// (a "%s … %v" pair) and joining them with sep when no content field is given. -// Entries render in sorted key order so identical bodies always produce -// identical text. -func bodyContent(body map[string]any, format, sep string) string { - if c, ok := body["content"].(string); ok && c != "" { - return c - } - parts := make([]string, 0, len(body)) - for _, k := range slices.Sorted(maps.Keys(body)) { - parts = append(parts, fmt.Sprintf(format, k, body[k])) - } - return strings.Join(parts, sep) -} - -// bodyLevel returns the upper-cased notification level, falling back to INFO. -func bodyLevel(body map[string]any) string { - if l, ok := body["level"].(string); ok && l != "" { - return strings.ToUpper(l) - } - return levelInfo -} diff --git a/backend/plugins/domain/message_gateway/registry_test.go b/backend/plugins/domain/message_gateway/registry_test.go deleted file mode 100644 index ad86d322..00000000 --- a/backend/plugins/domain/message_gateway/registry_test.go +++ /dev/null @@ -1,38 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package message_gateway - -import ( - "context" - "testing" -) - -type stubChannel struct{} - -func (stubChannel) Type() string { return "stub" } -func (stubChannel) Connect(context.Context) error { - return nil -} -func (stubChannel) Disconnect(context.Context) error { return nil } -func (stubChannel) Send(context.Context, Recipient, OutboundMessage) error { - return nil -} -func (stubChannel) Capabilities() Capability { return Capability{Text: true} } - -func TestRegisterLookup(t *testing.T) { - Register("stub", func(ChannelConfig, Handler) (Channel, error) { - return stubChannel{}, nil - }) - fn, ok := Lookup("stub") - if !ok { - t.Fatal("expected factory") - } - ch, err := fn(ChannelConfig{}, nil) - if err != nil { - t.Fatal(err) - } - if ch.Type() != "stub" { - t.Fatalf("type=%s", ch.Type()) - } -} diff --git a/backend/plugins/domain/message_gateway/repository/push_test.go b/backend/plugins/domain/message_gateway/repository/push_test.go deleted file mode 100644 index 18376ccd..00000000 --- a/backend/plugins/domain/message_gateway/repository/push_test.go +++ /dev/null @@ -1,131 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package repository_test - -import ( - "Wavelet/pkg/testhelper" - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/repository" - "context" - "errors" - "testing" - - "github.com/glebarez/sqlite" - "gorm.io/gorm" -) - -// stubDBService satisfies contracts.DBService over a test database handle. -type stubDBService struct{ db *gorm.DB } - -func (s stubDBService) GORM() *gorm.DB { return s.db } - -func (s stubDBService) DB(_ context.Context) *gorm.DB { return s.db } - -func (s stubDBService) Named(_ string) *gorm.DB { return s.db } - -// TestFindUserByFieldRecordRejectsUnlistedColumns pins the column allow-list. The -// lookup column is interpolated into SQL, so an unlisted name must be refused before -// any query is built rather than trusted because call sites happen to pass literals. -func TestFindUserByFieldRecordRejectsUnlistedColumns(t *testing.T) { - db, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - if err := db.Table("w_users").Create(map[string]any{"id": 77, "username": "seeded"}).Error; err != nil { - t.Fatalf("seed user failed: %v", err) - } - - repository.SetDBServiceForTest(stubDBService{db: db}) - t.Cleanup(func() { repository.SetDBServiceForTest(nil) }) - - ctx := context.Background() - - user, err := repository.FindUserByFieldRecord(ctx, "username", "seeded") - if err != nil { - t.Fatalf("allowlisted lookup by username failed: %v", err) - } - if user.ID != 77 { - t.Errorf("allowlisted lookup returned ID %d, want 77", user.ID) - } - if _, err := repository.FindUserByFieldRecord(ctx, "id", uint64(77)); err != nil { - t.Errorf("allowlisted lookup by id failed: %v", err) - } - - cases := []struct { - name string - field string - }{ - {"tautology injection", `username = '' OR 1=1 --`}, - {"stacked statement", "id; DROP TABLE w_users"}, - {"column outside allow-list", "password"}, - {"empty field", ""}, - } - for _, tc := range cases { - if _, err := repository.FindUserByFieldRecord(ctx, tc.field, "seeded"); !errors.Is(err, errs.ErrUnsupportedUserLookupField) { - t.Errorf("%s: got err %v, want ErrUnsupportedUserLookupField", tc.name, err) - } - } - - var remaining int64 - if err := db.Table("w_users").Count(&remaining).Error; err != nil || remaining != 1 { - t.Fatalf("w_users damaged by rejected lookups: count=%d err=%v", remaining, err) - } -} - -// smtpTestValues are the four system-config rows the built-in email channel reads. -var smtpTestValues = map[string]string{ - "smtp_host": "mail.example.test", - "smtp_port": "465", - "smtp_username": "notify@example.test", - "smtp_password": "s3cret-value", -} - -// TestLoadSMTPConfigRecordMapsEveryKey guards the single-query rewrite: every field -// must still be filled from its own row. -func TestLoadSMTPConfigRecordMapsEveryKey(t *testing.T) { - db, _, cleanup := testhelper.SetupTestEnvironment(t) - defer cleanup() - - keys := make([]string, 0, len(smtpTestValues)) - for key := range smtpTestValues { - keys = append(keys, key) - } - if err := db.Table("w_system_configs").Where("key IN ?", keys).Delete(map[string]any{}).Error; err != nil { - t.Fatalf("clear smtp rows: %v", err) - } - for _, key := range keys { - row := map[string]any{"key": key, "value": smtpTestValues[key], "type": "system"} - if err := db.Table("w_system_configs").Create(row).Error; err != nil { - t.Fatalf("seed %s: %v", key, err) - } - } - - repository.SetDBServiceForTest(stubDBService{db: db}) - t.Cleanup(func() { repository.SetDBServiceForTest(nil) }) - - cfg, err := repository.LoadSMTPConfigRecord(context.Background()) - if err != nil { - t.Fatalf("LoadSMTPConfigRecord: %v", err) - } - if cfg.Host != smtpTestValues["smtp_host"] || cfg.Port != smtpTestValues["smtp_port"] || - cfg.Username != smtpTestValues["smtp_username"] || cfg.Password != smtpTestValues["smtp_password"] { - t.Errorf("got %+v, want every SMTP field mapped from its own row", cfg) - } -} - -// TestLoadSMTPConfigRecordSurfacesReadFailure pins the actual defect: a read that -// fails used to be discarded, returning four blank strings that callers could only -// interpret as "SMTP was never configured", so the notification was dropped silently. -func TestLoadSMTPConfigRecordSurfacesReadFailure(t *testing.T) { - bare, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{}) - if err != nil { - t.Fatalf("open bare sqlite: %v", err) - } - - repository.SetDBServiceForTest(stubDBService{db: bare}) - t.Cleanup(func() { repository.SetDBServiceForTest(nil) }) - - if _, err := repository.LoadSMTPConfigRecord(context.Background()); err == nil { - t.Fatal("LoadSMTPConfigRecord returned nil error although the config table cannot be read") - } -} diff --git a/backend/plugins/domain/message_gateway/service/push.go b/backend/plugins/domain/message_gateway/service/push.go deleted file mode 100644 index 0ea7081f..00000000 --- a/backend/plugins/domain/message_gateway/service/push.go +++ /dev/null @@ -1,1153 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package service - -import ( - "Wavelet/core/contracts" - "Wavelet/pkg/logger" - "Wavelet/pkg/util" - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/model" - pkgpush "Wavelet/plugins/domain/message_gateway/push" - "Wavelet/plugins/domain/message_gateway/repository" - "context" - "encoding/json" - "errors" - "fmt" - "strconv" - "strings" - "sync" - "time" -) - -var ( - builtInEventsMu sync.RWMutex - // BuiltInEvents lists all built-in events defined in custom_events. - BuiltInEvents []model.EventMetadata -) - -// RegisterBuiltInEvent registers a built-in event definition. -func RegisterBuiltInEvent(meta model.EventMetadata) { - builtInEventsMu.Lock() - defer builtInEventsMu.Unlock() - for i, e := range BuiltInEvents { - if e.Key == meta.Key { - BuiltInEvents[i] = meta - return - } - } - BuiltInEvents = append(BuiltInEvents, meta) -} - -// GetBuiltInEvents returns a copy of registered built-in events. -func GetBuiltInEvents() []model.EventMetadata { - builtInEventsMu.RLock() - defer builtInEventsMu.RUnlock() - out := make([]model.EventMetadata, len(BuiltInEvents)) - copy(out, BuiltInEvents) - return out -} - -// PushRegistryAdapter adapts contracts.PushRegistry onto the built-in event store. -type PushRegistryAdapter struct{} - -func (PushRegistryAdapter) RegisterBuiltInEvent(meta contracts.PushEventMeta) { - RegisterBuiltInEvent(eventMetadataFromContract(meta)) -} - -func (PushRegistryAdapter) SyncEvents(ctx context.Context) error { - return SyncEvents(ctx) -} - -func eventMetadataFromContract(meta contracts.PushEventMeta) model.EventMetadata { - return model.EventMetadata{ - Key: meta.Key, - Name: meta.Name, - Description: meta.Description, - DefaultTemplate: model.NotificationMessage{ - Title: meta.DefaultTemplate.Title, - Content: meta.DefaultTemplate.Content, - Level: meta.DefaultTemplate.Level, - Ext: meta.DefaultTemplate.Ext, - }, - } -} - -// SyncBuiltInEvents seeds a database row for every registered built-in event. -func SyncBuiltInEvents(ctx context.Context) error { - for _, meta := range GetBuiltInEvents() { - _, err := repository.GetPushEventByKeyRecord(ctx, meta.Key) - if errors.Is(err, errs.ErrRecordNotFound) { - var defaultTemplateStr string - if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil { - defaultTemplateStr = string(defaultTemplateBytes) - } - event := model.PushEvent{ - EventKey: meta.Key, - Name: meta.Name, - Channels: []string{}, - Targets: []string{}, - Template: defaultTemplateStr, - Enabled: false, - } - if err := repository.CreatePushEventRecord(ctx, &event); err != nil { - return err - } - } else if err != nil { - return err - } - } - return nil -} - -// ListPushEvents lists all configured push events. -func ListPushEvents(ctx context.Context) ([]model.PushEvent, error) { - return repository.ListPushEventsRecord(ctx) -} - -// CreatePushEvent stores a push event configuration for a built-in event or task type. -func CreatePushEvent(ctx context.Context, req model.CreatePushEventRequest) (model.PushEvent, error) { - eventKey, eventName, defaultTemplateBytes, err := GetEventInfo(req) - if err != nil { - return model.PushEvent{}, err - } - - count, err := repository.CountPushEventsByKeyRecord(ctx, eventKey) - if err != nil { - return model.PushEvent{}, err - } - if count > 0 { - return model.PushEvent{}, errors.New(errs.ErrEventAlreadyConfigured) - } - - templateStr := strings.TrimSpace(req.Template) - if templateStr == "" { - templateStr = string(defaultTemplateBytes) - } else { - var tempMap map[string]any - if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil { - return model.PushEvent{}, errors.New(errs.ErrTemplateInvalidJSON) - } - } - - channels := req.Channels - if channels == nil { - channels = []string{} - } - targets := req.Targets - if targets == nil { - targets = []string{} - } - - event := model.PushEvent{ - EventKey: eventKey, - Name: eventName, - TaskType: req.TaskType, - Channels: channels, - Targets: targets, - Template: templateStr, - Enabled: req.Enabled, - } - if err := event.Validate(); err != nil { - return model.PushEvent{}, err - } - if err := repository.CreatePushEventRecord(ctx, &event); err != nil { - return model.PushEvent{}, err - } - return event, nil -} - -// DeletePushEvent deletes a push event configuration by id. -func DeletePushEvent(ctx context.Context, id uint64) error { - event, err := repository.GetPushEventByIDRecord(ctx, id) - if err != nil { - return err - } - return repository.DeletePushEventRecord(ctx, &event) -} - -// UpdatePushEvent replaces mutable push event fields. -func UpdatePushEvent(ctx context.Context, id uint64, req model.UpdatePushEventRequest) error { - event, err := repository.GetPushEventByIDRecord(ctx, id) - if err != nil { - return err - } - - event.Channels = req.Channels - event.Targets = req.Targets - event.Template = req.Template - event.Enabled = req.Enabled - if err := event.Validate(); err != nil { - return err - } - return repository.SavePushEventRecord(ctx, &event) -} - -// TogglePushEvent flips the enabled flag of a push event. -func TogglePushEvent(ctx context.Context, id uint64) (bool, error) { - event, err := repository.GetPushEventByIDRecord(ctx, id) - if err != nil { - return false, err - } - - enabled := !event.Enabled - if enabled && len(event.Channels) == 0 { - return false, errors.New(errs.ErrEnableWithoutChannels) - } - if err := repository.UpdatePushEventEnabledRecord(ctx, &event, enabled); err != nil { - return false, err - } - return enabled, nil -} - -// ListPushHistories returns a paginated push delivery audit page. -func ListPushHistories(ctx context.Context, filter model.PushHistoryListFilter) (int64, []model.PushHistory, error) { - return repository.ListPushHistoriesRecord(ctx, filter) -} - -// ApplySMTPFallbackToPushConfig fills an email config from the system SMTP settings. -func ApplySMTPFallbackToPushConfig(ctx context.Context, cfg *pkgpush.Config) { - if cfg.Channel != model.ChannelEmail || (cfg.URL != "" && cfg.Key != "") { - return - } - smtp, err := repository.LoadSMTPConfigRecord(ctx) - if err != nil { - logger.ErrorF(ctx, "[Push] 读取 SMTP 系统配置失败: %v", err) - return - } - if smtp.Host == "" || smtp.Username == "" { - return - } - port := smtp.Port - if port == "" { - port = "587" - } - cfg.URL = smtp.Host + ":" + port - cfg.Key = smtp.Username - cfg.Secret = smtp.Password -} - -// RunPushTest validates an ad-hoc channel config and sends a connectivity probe. -func RunPushTest(ctx context.Context, cfg pkgpush.Config, target string) error { - pusher, err := pkgpush.GetPusher(cfg.Channel) - if err != nil { - return err - } - if err := pusher.ValidateConfig(cfg); err != nil { - return fmt.Errorf("%s: %w", errs.ErrValidationFailed, err) - } - - ApplySMTPFallbackToPushConfig(ctx, &cfg) - - testBody := map[string]any{ - model.KeyTitle: "测试通道推送", - model.KeyContent: "当您收到这条消息,说明当前渠道连通性测试通过。", - model.KeyLevel: model.DefaultLevelInfo, - } - if _, err := pusher.Send(ctx, cfg, target, testBody, "", nil); err != nil { - return err - } - return nil -} - -// ListPushChannels returns every configured push channel. -func ListPushChannels(ctx context.Context) ([]model.PushChannel, error) { - return repository.ListPushChannelsRecord(ctx) -} - -// CreatePushChannel validates uniqueness and persists a new push channel. -func CreatePushChannel(ctx context.Context, req model.CreatePushChannelRequest) (model.PushChannel, error) { - count, err := repository.CountPushChannelsByNameRecord(ctx, req.Name) - if err != nil { - return model.PushChannel{}, err - } - if count > 0 { - return model.PushChannel{}, errors.New(errs.ErrChannelNameExists) - } - - channel := model.PushChannel{ - Name: req.Name, - Description: req.Description, - Type: req.Type, - Token: req.Token, - URL: req.URL, - Other: req.Other, - Enabled: req.Enabled, - } - if err := channel.Validate(); err != nil { - return model.PushChannel{}, err - } - if err := repository.CreatePushChannelRecord(ctx, &channel); err != nil { - return model.PushChannel{}, err - } - return channel, nil -} - -// UpdatePushChannel replaces the mutable fields of an existing push channel. -func UpdatePushChannel(ctx context.Context, id uint64, req model.UpdatePushChannelRequest) (model.PushChannel, error) { - channel, err := repository.GetPushChannelByIDRecord(ctx, id) - if err != nil { - return model.PushChannel{}, err - } - - channel.Description = req.Description - channel.Type = req.Type - channel.Token = req.Token - channel.URL = req.URL - channel.Other = req.Other - channel.Enabled = req.Enabled - if err := channel.Validate(); err != nil { - return model.PushChannel{}, err - } - if err := repository.SavePushChannelRecord(ctx, &channel); err != nil { - return model.PushChannel{}, err - } - return channel, nil -} - -// DeletePushChannel removes a push channel by id. -func DeletePushChannel(ctx context.Context, id uint64) error { - channel, err := repository.GetPushChannelByIDRecord(ctx, id) - if err != nil { - return err - } - return repository.DeletePushChannelRecord(ctx, &channel) -} - -// LoadChannelForTest resolves the credentials under test, either from a stored -// channel name or from the ad-hoc values sent by the caller. -func LoadChannelForTest(ctx context.Context, req model.TestPushChannelRequest) (string, string, string, string, error) { - if req.Name != "" { - channel, err := repository.GetPushChannelByNameRecord(ctx, req.Name) - if err != nil { - return "", "", "", "", errors.New(errs.ErrChannelNotFound) - } - return channel.URL, channel.Token, channel.Other, channel.Type, nil - } - return req.URL, req.Token, req.Other, req.Type, nil -} - -// PreparePushChannelTest builds the connectivity probe payload for a channel. -func PreparePushChannelTest(ctx context.Context, req model.TestPushChannelRequest) (model.SendPayload, error) { - url, token, other, channelType, err := LoadChannelForTest(ctx, req) - if err != nil { - return model.SendPayload{}, err - } - - if channelType == model.ChannelEmail { - url, token, other = ResolveSMTPConfig(ctx, url, token, other) - } - - tempChannel := model.PushChannel{ - Name: "test_temp", - URL: url, - Token: token, - Other: other, - Type: channelType, - Enabled: true, - } - if err := tempChannel.Validate(); err != nil { - return model.SendPayload{}, err - } - url = tempChannel.URL - - var config pkgpush.Config - var renderedJSON string - switch channelType { - case model.ChannelLark: - config = pkgpush.Config{Channel: model.ChannelLark, URL: url, Secret: token} - renderedJSON = other - case model.ChannelEmail: - config = pkgpush.Config{Channel: model.ChannelEmail, URL: url, Key: token, Secret: other} - case model.ChannelTelegram: - config = pkgpush.Config{Channel: model.ChannelTelegram, URL: url, Secret: token, Key: other} - default: - config = pkgpush.Config{Channel: model.ChannelCustom, URL: url} - customPushReq := model.CustomPushRequest{ - Title: "通道测试通知", - Content: "这是一条来自系统的消息通道连通性测试消息。", - Description: "系统通道测试", - URL: "https://example.com", - To: req.Target, - } - renderedJSON = RenderCustomPayload(other, customPushReq) - } - - return model.SendPayload{ - EventKey: "test_channel", - Config: config, - Target: req.Target, - Body: model.NotificationMessage{ - Title: "通道测试通知", - Content: "这是一条来自系统的消息通道连通性测试消息。", - Level: model.DefaultLevelInfo, - }, - Template: renderedJSON, - }, nil -} - -// RenderCustomPayload substitutes the supported template variables of a custom -// webhook body, JSON-escaping every injected value. -func RenderCustomPayload(template string, req model.CustomPushRequest) string { - result := template - result = strings.ReplaceAll(result, "$title", EscapeJSONString(req.Title)) - result = strings.ReplaceAll(result, "$description", EscapeJSONString(req.Description)) - result = strings.ReplaceAll(result, "$content", EscapeJSONString(req.Content)) - result = strings.ReplaceAll(result, "$url", EscapeJSONString(req.URL)) - result = strings.ReplaceAll(result, "$to", EscapeJSONString(req.To)) - return result -} - -// EscapeJSONString renders s as a JSON string body without the surrounding quotes. -func EscapeJSONString(s string) string { - b, _ := json.Marshal(s) - const minJSONLen = 2 - if len(b) >= minJSONLen { - return string(b[1 : len(b)-1]) - } - return s -} - -// ListActivePushEventsByTaskType returns enabled push events for a given task type. -func ListActivePushEventsByTaskType(ctx context.Context, taskType string) ([]model.PushEvent, error) { - return repository.ListActivePushEventsByTaskTypeRecord(ctx, taskType) -} - -// QueryUser resolves a user through the UserService contract, falling back to the -// repository read path while the contract is not wired yet. -func QueryUser(ctx context.Context, fromService func(contracts.UserService) (*contracts.UserDTO, error), dbField string, dbVal any) (*contracts.UserDTO, error) { - if userSvc := GetUserService(ctx); userSvc != nil { - return fromService(userSvc) - } - if user, err := repository.FindUserByFieldRecord(ctx, dbField, dbVal); err == nil && user != nil { - return user, nil - } - return nil, errors.New(errs.ErrUserNotFound) -} - -// FindUserByID resolves a user by primary key. -func FindUserByID(ctx context.Context, id uint64) (*contracts.UserDTO, error) { - return QueryUser(ctx, func(s contracts.UserService) (*contracts.UserDTO, error) { - return s.GetUserByID(ctx, id) - }, "id", id) -} - -// FindUserByUsername resolves a user by login name. -func FindUserByUsername(ctx context.Context, username string) (*contracts.UserDTO, error) { - return QueryUser(ctx, func(s contracts.UserService) (*contracts.UserDTO, error) { - return s.GetUserByUsername(ctx, username) - }, "username", username) -} - -// LoadUserFromPayload extracts user info from data. -func LoadUserFromPayload(ctx context.Context, data map[string]any) any { - if u, exists := data["user"]; exists && u != nil { - return u - } - - if userID, ok := ExtractUserID(data); ok && userID > 0 { - if user, err := FindUserByID(ctx, userID); err == nil && user != nil { - return user - } - } - - if username := ExtractUsername(data); username != "" { - if user, err := FindUserByUsername(ctx, username); err == nil && user != nil { - return user - } - } - return nil -} - -// RecordPushHistory creates a push history audit record. -func RecordPushHistory(ctx context.Context, req model.SendPayload, status, errMsg string) error { - title := req.Body.Title - content := req.Body.Content - level := req.Body.Level - if title == "" { - title = "系统通知" - } - if level == "" { - level = model.DefaultLevelInfo - } - - target := req.Target - if target == "" { - if req.Config.URL != "" { - target = req.Config.URL - const maxTargetLen = 50 - const truncatedLen = 47 - if len(target) > maxTargetLen { - target = target[:truncatedLen] + "..." - } - } else { - target = "default" - } - } - - history := model.PushHistory{ - EventKey: req.EventKey, - Channel: req.Config.Channel, - Target: target, - Title: title, - Content: content, - Level: level, - Status: status, - ErrorMsg: errMsg, - } - return repository.CreatePushHistoryRecord(ctx, &history) -} - -// ResolveTarget parses dynamic placeholders into concrete receiver targets. -func ResolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string { - target = strings.TrimSpace(target) - if target == "" { - return "" - } - - resolved := ResolveDynamicKeyword(target, flatBody) - if strings.Contains(resolved, "@") { - return resolved - } - if val, matched := ResolveSystemTarget(ctx, resolved, channel); matched { - return val - } - - user, found := ResolveTargetUser(ctx, resolved, channel) - if !found { - return resolved - } - if channel == model.ChannelEmail && user.Email != "" { - return user.Email - } - if channel != model.ChannelEmail && user.Username != "" { - return user.Username - } - return resolved -} - -// ResolveDynamicKeyword resolves user.id, username, email keywords. -func ResolveDynamicKeyword(target string, flatBody map[string]any) string { - switch target { - case "user.id", "id": - if val, ok := flatBody["user.id"]; ok { - return fmt.Sprintf("%v", val) - } - if val, ok := flatBody["id"]; ok { - return fmt.Sprintf("%v", val) - } - case "user.username", "username": - if val, ok := flatBody["user.username"]; ok { - return fmt.Sprintf("%v", val) - } - if val, ok := flatBody["username"]; ok { - return fmt.Sprintf("%v", val) - } - case "user.email", model.ChannelEmail: - if val, ok := flatBody["user.email"]; ok { - return fmt.Sprintf("%v", val) - } - if val, ok := flatBody["email"]; ok { - return fmt.Sprintf("%v", val) - } - } - return target -} - -// ResolveTargetUser resolves user by numeric ID or username string. -func ResolveTargetUser(ctx context.Context, resolved, _ string) (contracts.UserDTO, bool) { - if id, err := strconv.ParseUint(resolved, 10, 64); err == nil { - if u, err := FindUserByID(ctx, id); err == nil && u != nil { - return *u, true - } - } - if u, err := FindUserByUsername(ctx, resolved); err == nil && u != nil { - return *u, true - } - return contracts.UserDTO{}, false -} - -// GetFirstAdminUser resolves the first administrator through the UserService -// contract, falling back to the repository read path when it is unavailable. -func GetFirstAdminUser(ctx context.Context) (*contracts.UserDTO, error) { - if userSvc := GetUserService(ctx); userSvc != nil { - return userSvc.GetFirstAdminUser(ctx) - } - if adminUser, err := repository.FindFirstAdminUserRecord(ctx); err == nil && adminUser != nil { - return adminUser, nil - } - return nil, errors.New(errs.ErrNoAdminUser) -} - -// ResolveSystemTarget maps system receiver aliases to administrator contact info. -func ResolveSystemTarget(ctx context.Context, resolved, channel string) (string, bool) { - if resolved != "系统" && resolved != "system" && resolved != "0" { - return "", false - } - adminUser, err := GetFirstAdminUser(ctx) - if err != nil || adminUser == nil { - return resolved, true - } - if channel == model.ChannelEmail && adminUser.Email != "" { - return adminUser.Email, true - } - if channel != model.ChannelEmail && adminUser.Username != "" { - return adminUser.Username, true - } - return resolved, true -} - -// ResolveSMTPConfig fills missing email endpoint fields from the system SMTP settings. -func ResolveSMTPConfig(ctx context.Context, url, token, other string) (string, string, string) { - if url != "" && token != "" { - return url, token, other - } - smtp, err := repository.LoadSMTPConfigRecord(ctx) - if err != nil { - logger.ErrorF(ctx, "[Push] 读取 SMTP 系统配置失败: %v", err) - return url, token, other - } - if smtp.Host == "" || smtp.Username == "" { - return url, token, other - } - port := smtp.Port - if port == "" { - port = "587" - } - if url == "" { - url = smtp.Host + ":" + port - } - if token == "" { - token = smtp.Username - } - if other == "" { - other = smtp.Password - } - return url, token, other -} - -// GetSystemUser gets a system user DTO. -func GetSystemUser(ctx context.Context) *contracts.UserDTO { - if adminUser, err := GetFirstAdminUser(ctx); err == nil && adminUser != nil { - return adminUser - } - return &contracts.UserDTO{ - Username: "system", - Nickname: "系统管理员", - } -} - -// FindBuiltInEvent finds a registered built-in event by key. -func FindBuiltInEvent(key string) (model.EventMetadata, bool) { - for _, meta := range GetBuiltInEvents() { - if meta.Key == key { - return meta, true - } - } - return model.EventMetadata{}, false -} - -// GetEventInfo derives the event key, display name and default template for a -// task-completion based event or a registered built-in event key. -func GetEventInfo(req model.CreatePushEventRequest) (string, string, []byte, error) { - if req.TaskType != "" { - taskName := req.TaskType - if taskSvc := GetTaskService(); taskSvc != nil { - if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok { - taskName = meta.DisplayName - } - } - eventKey := "task_completed:" + req.TaskType - eventName := "任务完成: " + taskName - defaultTemplate := model.NotificationMessage{ - Title: "任务完成: " + taskName, - Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。", - Level: model.DefaultLevelInfo, - } - defaultTemplateBytes, err := json.Marshal(defaultTemplate) - if err != nil { - return "", "", nil, err - } - return eventKey, eventName, defaultTemplateBytes, nil - } - - if req.EventKey == "" { - return "", "", nil, errors.New(errs.ErrEventKeyOrTaskType) - } - - meta, found := FindBuiltInEvent(req.EventKey) - if !found { - return "", "", nil, errors.New(errs.ErrUnsupportedEventKey) - } - - defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate) - if err != nil { - return "", "", nil, err - } - return req.EventKey, meta.Name, defaultTemplateBytes, nil -} - -// EnqueuePushTask dispatches a notification payload to the async push worker. -func EnqueuePushTask(ctx context.Context, payload model.SendPayload) error { - payloadBytes, err := json.Marshal(payload) - if err != nil { - return err - } - if taskSvc := GetTaskService(); taskSvc != nil { - _, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, "system") - return err - } - return errors.New(errs.ErrTaskServiceUnavailable) -} - -// GetFlatBody flattens nested body map. -func GetFlatBody(body map[string]any) map[string]any { - jsonBytes, err := json.Marshal(body) - if err != nil { - return body - } - var jsonMap map[string]any - if err := json.Unmarshal(jsonBytes, &jsonMap); err != nil { - return body - } - - flatResult := make(map[string]any) - FlattenMap("", jsonMap, flatResult) - return flatResult -} - -// FlattenMap recursively flattens map key-values. -func FlattenMap(prefix string, m, result map[string]any) { - for k, v := range m { - key := k - if prefix != "" { - key = prefix + "." + k - } - if nestedMap, ok := v.(map[string]any); ok { - FlattenMap(key, nestedMap, result) - } else { - result[key] = v - } - } -} - -const ( - // SendNotificationTask is the asynq task name for push notification. - SendNotificationTask = "push:send" - // TaskTypeSendNotification is the admin task manager type identifier. - TaskTypeSendNotification = "send_notification" -) - -// SendNotificationMeta represents the task metadata. -var SendNotificationMeta = contracts.TaskMetaDTO{ - Type: TaskTypeSendNotification, - AsynqTask: SendNotificationTask, - Name: "推送通知", - DisplayName: "推送通知", - Description: "异步执行系统通知的多渠道派发与推送", - Category: "push", - SupportsTime: false, - MaxRetry: 3, - Queue: "default", - Retryable: true, - Params: []contracts.TaskParamDTO{ - { - Name: "event_key", - Label: "事件标识", - Type: "string", - Required: true, - Placeholder: "admin_login", - Description: "事件标识 (如 admin_login)", - }, - { - Name: "target", - Label: "目标接收者", - Type: "string", - Required: false, - Description: "目标接收者", - }, - }, -} - -// PushHandler handles asynchronous notification sending. -type PushHandler struct{} - -// ValidatePayload validates and normalizes push parameters. -func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) { - if len(payload) == 0 { - return nil, errors.New(errs.ErrPayloadRequired) - } - - var req model.SendPayload - if err := json.Unmarshal(payload, &req); err != nil { - return nil, fmt.Errorf("%s: %w", errs.ErrInvalidJSONFormat, err) - } - - if req.Config.Channel == "" { - return nil, errors.New(errs.ErrChannelTypeRequired) - } - - return json.Marshal(req) -} - -// Execute performs the push send and logs delivery history audit. -func (h *PushHandler) Execute(ctx context.Context, payload []byte) error { - var req model.SendPayload - if err := json.Unmarshal(payload, &req); err != nil { - logger.ErrorF(ctx, "[Push] 解析推送参数失败: %v", err) - return fmt.Errorf("%s: %w", errs.ErrParsePayloadFailed, err) - } - - logger.InfoF(ctx, "[Push] 开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target) - - pusher, err := pkgpush.GetPusher(req.Config.Channel) - if err != nil { - errWrap := fmt.Errorf("%s: %w", errs.ErrGetPusherFailed, err) - logger.ErrorF(ctx, "[Push] 推送失败: %v", errWrap) - h.recordHistory(ctx, req, "failed", errWrap.Error()) - return errWrap - } - - flatBody := req.Body.Flatten() - upstreamResp, err := pusher.Send(ctx, req.Config, req.Target, flatBody, req.Template, nil) - - title := req.Body.Title - content := req.Body.Content - - if err != nil { - logger.ErrorF(ctx, "[Push] 消息推送失败 (标题: %s): %v, 上游返回: %s", title, err, upstreamResp) - h.recordHistory(ctx, req, "failed", err.Error()) - return fmt.Errorf("pusher.Send failed: %w", err) - } - - logger.InfoF(ctx, "[Push] 消息推送成功 (标题: %s, 内容摘要: %s), 上游返回: %s", title, content, upstreamResp) - h.recordHistory(ctx, req, "success", "") - - return nil -} - -func (h *PushHandler) recordHistory(ctx context.Context, req model.SendPayload, status, errMsg string) { - if dbErr := RecordPushHistory(ctx, req, status, errMsg); dbErr != nil { - logger.ErrorF(ctx, "[Push] 写入推送历史审计记录失败: %v", dbErr) - } -} - -// HandleTaskCompleted handles task completion notifications. -func HandleTaskCompleted(ctx context.Context, e contracts.TaskCompletedEvent) { - events, err := ListActivePushEventsByTaskType(ctx, e.TaskType) - if err != nil { - logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", e.TaskType, err) - return - } - if len(events) == 0 { - return - } - - body := map[string]any{ - "task_id": e.TaskID, - "task_name": e.TaskName, - "task_type": e.TaskType, - "task_status": e.Status, - "task_duration": e.Duration, - "time": time.Now().Format("2006-01-02 15:04:05"), - "task_error": e.ErrorMsg, - "task_result": e.ResultMsg, - } - - var payloadMap map[string]any - if e.Payload != "" { - if err := json.Unmarshal([]byte(e.Payload), &payloadMap); err == nil { - body["payload"] = payloadMap - ExtractUserFromMap(ctx, payloadMap, body) - } - } - if e.Detail != "" { - var detailMap map[string]any - if err := json.Unmarshal([]byte(e.Detail), &detailMap); err == nil { - body["detail"] = detailMap - ExtractUserFromMap(ctx, detailMap, body) - } - } - - for _, event := range events { - meta := model.EventMetadata{ - Key: event.EventKey, - Name: event.Name, - Description: "异步任务执行完毕触发的自动通知", - } - DefaultTrigger.Trigger(ctx, meta, body) - } -} - -// ExtractUserFromMap extracts user info from payload/detail into body map. -func ExtractUserFromMap(ctx context.Context, data, body map[string]any) { - if u, exists := body["user"]; exists && u != nil { - return - } - if user := LoadUserFromPayload(ctx, data); user != nil { - body["user"] = user - } -} - -// ExtractUserID extracts a user ID from map keys. -func ExtractUserID(data map[string]any) (uint64, bool) { - for _, k := range []string{"user_id", "userId", "uid"} { - val, ok := data[k] - if !ok || val == nil { - continue - } - switch v := val.(type) { - case float64: - if v >= 0 { - return uint64(v), true - } - case int: - if v >= 0 { - return uint64(v), true - } - case int64: - if v >= 0 { - return uint64(v), true - } - case uint64: - return v, true - case string: - if id, err := strconv.ParseUint(v, 10, 64); err == nil { - return id, true - } - } - } - return 0, false -} - -// ExtractUsername extracts a username string from map keys. -func ExtractUsername(data map[string]any) string { - for _, k := range []string{"username", "user_name"} { - if val, ok := data[k]; ok && val != nil { - if s, ok := val.(string); ok && s != "" { - return s - } - } - } - return "" -} - -// EventTrigger represents the unified event trigger class. -type EventTrigger struct{} - -// DefaultTrigger is the singleton instance of EventTrigger. -var DefaultTrigger = &EventTrigger{} - -// Trigger receives event metadata and processes the event notification dispatch asynchronously. -func (t *EventTrigger) Trigger(ctx context.Context, meta model.EventMetadata, body map[string]any) { - asyncCtx := context.WithoutCancel(ctx) - util.Go(func() { - if body == nil { - body = make(map[string]any) - } - if _, hasUser := body["user"]; !hasUser || body["user"] == nil { - body["user"] = GetSystemUser(asyncCtx) - } - - eventPtr, err := repository.GetActivePushEventByKey(asyncCtx, meta.Key) - if err != nil { - if errors.Is(err, errs.ErrRecordNotFound) { - return - } - logger.ErrorF(asyncCtx, "push_event_trigger: failed to get active event %s: %v", meta.Key, err) - return - } - event := *eventPtr - if len(event.Channels) == 0 { - return - } - - flatBody := GetFlatBody(body) - msg, _ := t.buildMessage(&event, meta, flatBody, body) - t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody) - }) -} - -func (t *EventTrigger) buildMessage(event *model.PushEvent, meta model.EventMetadata, flatBody, body map[string]any) (model.NotificationMessage, string) { - var msg model.NotificationMessage - renderedTemplate := "" - - templateSource := event.Template - if templateSource != "" { - var err error - msg, renderedTemplate, err = t.parseCustomTemplate(event, templateSource, flatBody) - if err != nil { - msg.Title = event.Name - msg.Content = renderedTemplate - msg.Level = model.DefaultLevelInfo - } - } else { - msg = t.parseDefaultTemplate(meta, flatBody) - } - - if msg.Ext == nil { - msg.Ext = make(map[string]any) - } - for k, v := range body { - if k == model.KeyTitle || k == model.KeyContent || k == model.KeyLevel { - continue - } - if _, exists := msg.Ext[k]; !exists { - msg.Ext[k] = v - } - } - - return msg, renderedTemplate -} - -func (t *EventTrigger) parseCustomTemplate(event *model.PushEvent, templateSource string, flatBody map[string]any) (model.NotificationMessage, string, error) { - var msg model.NotificationMessage - renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody) - - var tMap map[string]any - if err := json.Unmarshal([]byte(renderedTemplate), &tMap); err != nil { - return msg, renderedTemplate, err - } - - if title, ok := tMap[model.KeyTitle].(string); ok && title != "" { - msg.Title = title - } else { - msg.Title = event.Name - } - delete(tMap, model.KeyTitle) - - if content, ok := tMap[model.KeyContent].(string); ok && content != "" { - msg.Content = content - } else { - msg.Content = renderedTemplate - } - delete(tMap, model.KeyContent) - - if level, ok := tMap[model.KeyLevel].(string); ok && level != "" { - msg.Level = level - } else { - msg.Level = model.DefaultLevelInfo - } - delete(tMap, model.KeyLevel) - - msg.Ext = tMap - return msg, renderedTemplate, nil -} - -func (t *EventTrigger) parseDefaultTemplate(meta model.EventMetadata, flatBody map[string]any) model.NotificationMessage { - var msg model.NotificationMessage - msg.Title = pkgpush.ParseTemplate(meta.DefaultTemplate.Title, flatBody) - msg.Content = pkgpush.ParseTemplate(meta.DefaultTemplate.Content, flatBody) - msg.Level = pkgpush.ParseTemplate(meta.DefaultTemplate.Level, flatBody) - - if meta.DefaultTemplate.Ext != nil { - msg.Ext = make(map[string]any) - for k, v := range meta.DefaultTemplate.Ext { - if strVal, ok := v.(string); ok { - msg.Ext[k] = pkgpush.ParseTemplate(strVal, flatBody) - } else { - msg.Ext[k] = v - } - } - } - return msg -} - -func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta model.EventMetadata, event *model.PushEvent, msg model.NotificationMessage, flatBody map[string]any) { - for _, channelName := range event.Channels { - customChannel, err := repository.GetActivePushChannelByName(ctx, channelName) - if err == nil { - t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody) - continue - } - logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err) - } -} - -func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta model.EventMetadata, event *model.PushEvent, channel *model.PushChannel, msg model.NotificationMessage, flatBody map[string]any) { - if len(event.Targets) == 0 { - t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg) - return - } - - for _, target := range event.Targets { - resolvedTarget := ResolveTarget(ctx, target, flatBody, channel.Name) - t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, resolvedTarget, msg) - } -} - -func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta model.EventMetadata, channel *model.PushChannel, target string, msg model.NotificationMessage) { - var config pkgpush.Config - var renderedTemplate string - - switch channel.Type { - case model.ChannelLark: - config = pkgpush.Config{Channel: model.ChannelLark, URL: channel.URL, Secret: channel.Token} - renderedTemplate = channel.Other - case model.ChannelEmail: - url, token, other := ResolveSMTPConfig(ctx, channel.URL, channel.Token, channel.Other) - config = pkgpush.Config{Channel: model.ChannelEmail, URL: url, Key: token, Secret: other} - case model.ChannelTelegram: - config = pkgpush.Config{Channel: model.ChannelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other} - default: - config = pkgpush.Config{Channel: model.ChannelCustom, URL: channel.URL} - customPushReq := model.CustomPushRequest{ - Title: msg.Title, - Content: msg.Content, - Description: meta.Description, - To: target, - } - if urlVal, ok := msg.Ext["url"].(string); ok { - customPushReq.URL = urlVal - } - renderedTemplate = RenderCustomPayload(channel.Other, customPushReq) - } - - payload := model.SendPayload{ - EventKey: meta.Key, - Config: config, - Target: target, - Body: msg, - Template: renderedTemplate, - } - if err := EnqueuePushTask(ctx, payload); err != nil { - logger.ErrorF(ctx, "push_event_trigger: enqueuePushTask failed for %s channel %s -> %s: %v", channel.Type, channel.Name, target, err) - } -} - -// SyncEvents automatically registers/updates built-in events in the database. -func SyncEvents(ctx context.Context) error { - return SyncBuiltInEvents(ctx) -} - -// AdminLogin is the metadata definition for the admin login event. -var AdminLogin = model.EventMetadata{ - Key: "admin_login", - Name: "管理员登录", - DefaultTemplate: model.NotificationMessage{ - Title: "管理员登录提醒", - Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。", - Level: "INFO", - }, - Description: "当管理员成功登录系统时触发此通知", -} - -// HandleAdminLoggedIn 处理管理员登录事件并触发通知 -func HandleAdminLoggedIn(ctx context.Context, event contracts.AdminLoggedIn) { - if event.User == nil { - return - } - - body := map[string]any{ - "user": event.User, - "ip": event.IP, - "time": time.Now().Format("2006-01-02 15:04:05"), - } - DefaultTrigger.Trigger(ctx, AdminLogin, body) -} - -// RegisterCustomEvents registers default domain push notification events. -func RegisterCustomEvents() { - RegisterBuiltInEvent(AdminLogin) -} diff --git a/backend/plugins/domain/message_gateway/service/service.go b/backend/plugins/domain/message_gateway/service/service.go deleted file mode 100644 index dfd33db5..00000000 --- a/backend/plugins/domain/message_gateway/service/service.go +++ /dev/null @@ -1,405 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package service implements domain business logic and channel runners for message_gateway. -package service - -import ( - "Wavelet/core" - "Wavelet/core/contracts" - "Wavelet/pkg/logger" - "Wavelet/pkg/util" - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/model" - "Wavelet/plugins/domain/message_gateway/repository" - "context" - "crypto/rand" - "crypto/sha256" - "encoding/hex" - "encoding/json" - "errors" - "strconv" - "strings" - "sync" - "time" - "unicode" -) - -// Handler processes one inbound message. -type Handler func(ctx context.Context, msg model.InboundMessage) error - -// Factory constructs a Channel from decrypted config. -type Factory func(cfg model.ChannelConfig, onInbound Handler) (Channel, error) - -// Channel is one connected messaging adapter. -type Channel interface { - Type() string - Connect(ctx context.Context) error - Disconnect(ctx context.Context) error - Send(ctx context.Context, to model.Recipient, msg model.OutboundMessage) error - Capabilities() model.Capability -} - -var ( - factoriesMu sync.RWMutex - factories = map[string]Factory{} -) - -// Register stores a channel factory under typ. -func Register(typ string, fn Factory) { - factoriesMu.Lock() - defer factoriesMu.Unlock() - factories[typ] = fn -} - -// Lookup returns a previously registered factory. -func Lookup(typ string) (Factory, bool) { - factoriesMu.RLock() - defer factoriesMu.RUnlock() - fn, ok := factories[typ] - return fn, ok -} - -// CodeAlphabet excludes easily confused runes 0/O/1/I. -const CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" - -// CodeLength is the raw pairing code size. -const CodeLength = 8 - -// GenerateCode returns an 8-character pairing code. -func GenerateCode() (string, error) { - buf := make([]byte, CodeLength) - if _, err := rand.Read(buf); err != nil { - return "", err - } - out := make([]byte, CodeLength) - for i, b := range buf { - out[i] = CodeAlphabet[int(b)%len(CodeAlphabet)] - } - return string(out), nil -} - -// NormalizeCode strips separators and uppercases. -func NormalizeCode(s string) string { - var b strings.Builder - for _, r := range s { - if r == '-' || unicode.IsSpace(r) { - continue - } - b.WriteRune(unicode.ToUpper(r)) - } - return b.String() -} - -// FormatCode renders ABCD-EFGH. -func FormatCode(s string) string { - s = NormalizeCode(s) - if len(s) != CodeLength { - return s - } - return s[:4] + "-" + s[4:] -} - -var ( - credentialSecretMu sync.RWMutex - credentialSecret string -) - -// SetCredentialSecret sets the secret used to derive CredentialKey. -func SetCredentialSecret(secret string) { - credentialSecretMu.Lock() - defer credentialSecretMu.Unlock() - credentialSecret = secret -} - -// CredentialKey is AES-256 hex derived from the session secret. -func CredentialKey() string { - credentialSecretMu.RLock() - secret := credentialSecret - credentialSecretMu.RUnlock() - sum := sha256.Sum256([]byte(secret)) - return hex.EncodeToString(sum[:]) -} - -// EncryptCredentials encrypts a credential map as JSON. -func EncryptCredentials(creds map[string]string) (string, error) { - if creds == nil { - creds = map[string]string{} - } - raw, err := json.Marshal(creds) - if err != nil { - return "", err - } - return util.Encrypt(CredentialKey(), string(raw)) -} - -// DecryptCredentials decrypts a credential map. -func DecryptCredentials(ciphertext string) (map[string]string, error) { - if ciphertext == "" { - return map[string]string{}, nil - } - plain, err := util.Decrypt(CredentialKey(), ciphertext) - if err != nil { - return nil, err - } - var out map[string]string - if err := json.Unmarshal([]byte(plain), &out); err != nil { - return nil, err - } - if out == nil { - out = map[string]string{} - } - return out, nil -} - -// ParseExtra decodes optional extra JSON into a string map. -func ParseExtra(raw string) map[string]string { - if raw == "" { - return map[string]string{} - } - var out map[string]string - if err := json.Unmarshal([]byte(raw), &out); err != nil || out == nil { - return map[string]string{} - } - return out -} - -// EncodeExtra encodes extra fields as JSON. -func EncodeExtra(extra map[string]string) string { - if extra == nil { - return "" - } - raw, err := json.Marshal(extra) - if err != nil { - return "" - } - return string(raw) -} - -// Runner manages lifecycle for long-lived channel adapters (WebSocket, long-polling, etc.). -type Runner struct { - mu sync.Mutex - running bool - cancel context.CancelFunc -} - -// GlobalRunner is the default global runner instance. -var GlobalRunner = &Runner{} - -// Start starts all background long-lived channel runners. -func Start(ctx context.Context) error { - GlobalRunner.mu.Lock() - defer GlobalRunner.mu.Unlock() - - if GlobalRunner.running { - return nil - } - - runCtx, cancel := context.WithCancel(ctx) - GlobalRunner.cancel = cancel - GlobalRunner.running = true - - logger.InfoF(runCtx, "[MessageGateway] Starting bot channel runners...") - return nil -} - -// Stop stops the channel runner. -func Stop() { - GlobalRunner.mu.Lock() - defer GlobalRunner.mu.Unlock() - - if !GlobalRunner.running { - return - } - - if GlobalRunner.cancel != nil { - GlobalRunner.cancel() - } - GlobalRunner.running = false -} - -// Cordis contract singletons consumed by service layer. -var ( - cacheMu sync.RWMutex - cacheSvc contracts.CacheService - taskMu sync.RWMutex - taskSvc contracts.TaskService - userMu sync.RWMutex - userSvc contracts.UserService -) - -// SetCacheService sets the cache service. -func SetCacheService(s contracts.CacheService) { - cacheMu.Lock() - defer cacheMu.Unlock() - cacheSvc = s -} - -// SetTaskService sets the task service. -func SetTaskService(s contracts.TaskService) { - taskMu.Lock() - defer taskMu.Unlock() - taskSvc = s -} - -// SetUserService sets the user service. -func SetUserService(s contracts.UserService) { - userMu.Lock() - defer userMu.Unlock() - userSvc = s -} - -// GetCache resolves the cache service for the context. -func GetCache(ctx context.Context) contracts.CacheService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { - return s - } - } - cacheMu.RLock() - s := cacheSvc - cacheMu.RUnlock() - return s -} - -// GetTaskService returns the task service. -func GetTaskService() contracts.TaskService { - taskMu.RLock() - defer taskMu.RUnlock() - return taskSvc -} - -// GetUserService resolves the user service for the context. -func GetUserService(ctx context.Context) contracts.UserService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.UserService](c); err == nil && s != nil { - return s - } - } - userMu.RLock() - s := userSvc - userMu.RUnlock() - return s -} - -// BindChannel consumes a pairing code and binds the platform identity to the user. -func BindChannel(ctx context.Context, userID uint64, req model.BindRequest) (model.BindingDTO, error) { - channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64) - if err != nil || channelID == 0 { - return model.BindingDTO{}, errs.ErrChannelIDRequired - } - code := NormalizeCode(req.Code) - if code == "" { - return model.BindingDTO{}, errs.ErrCodeInvalid - } - pairing, err := repository.GetPairingCode(ctx, code) - if err != nil { - if errors.Is(err, errs.ErrRecordNotFound) { - return model.BindingDTO{}, errs.ErrCodeInvalid - } - return model.BindingDTO{}, err - } - if !pairing.ExpiresAt.After(time.Now()) { - return model.BindingDTO{}, errs.ErrCodeInvalid - } - if pairing.ChannelID != channelID { - return model.BindingDTO{}, errs.ErrChannelMismatch - } - ch, err := repository.GetMessageChannel(ctx, channelID) - if err != nil { - if errors.Is(err, errs.ErrRecordNotFound) { - return model.BindingDTO{}, errs.ErrCodeInvalid - } - return model.BindingDTO{}, err - } - if !ch.Enabled { - return model.BindingDTO{}, errs.ErrChannelDisabled - } - - existing, err := repository.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID) - if err != nil && !errors.Is(err, errs.ErrRecordNotFound) { - return model.BindingDTO{}, err - } - if err == nil && existing != nil { - if existing.UserID != userID { - return model.BindingDTO{}, errs.ErrPlatformAlreadyBound - } - _ = repository.DeletePairingCode(ctx, pairing.Code) - return ToBindingDTO(existing, ch), nil - } - - row := &model.MessageBinding{ - UserID: userID, - ChannelID: channelID, - PlatformUserID: pairing.PlatformUserID, - } - if err := repository.CreateMessageBinding(ctx, row); err != nil { - return model.BindingDTO{}, err - } - if err := repository.DeletePairingCode(ctx, pairing.Code); err != nil { - return model.BindingDTO{}, err - } - return ToBindingDTO(row, ch), nil -} - -// ListEnabledPublicChannels returns the channels a user may bind to. -func ListEnabledPublicChannels(ctx context.Context) ([]model.PublicChannelDTO, error) { - rows, err := repository.ListEnabledMessageChannels(ctx) - if err != nil { - return nil, err - } - out := make([]model.PublicChannelDTO, 0, len(rows)) - for _, row := range rows { - out = append(out, model.PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type}) - } - return out, nil -} - -// ListUserBindings returns the binding rows of one user enriched with channel info. -func ListUserBindings(ctx context.Context, userID uint64) ([]model.BindingDTO, error) { - rows, err := repository.ListBindingsByUser(ctx, userID) - if err != nil { - return nil, err - } - out := make([]model.BindingDTO, 0, len(rows)) - for i := range rows { - ch, err := repository.GetMessageChannel(ctx, rows[i].ChannelID) - if err != nil { - out = append(out, ToBindingDTO(&rows[i], nil)) - continue - } - out = append(out, ToBindingDTO(&rows[i], ch)) - } - return out, nil -} - -// UnbindChannel removes a binding owned by the given user. -func UnbindChannel(ctx context.Context, userID, bindingID uint64) error { - row, err := repository.GetMessageBinding(ctx, bindingID) - if err != nil { - if errors.Is(err, errs.ErrRecordNotFound) { - return errs.ErrBindingNotFound - } - return err - } - if row.UserID != userID { - return errs.ErrBindingForbidden - } - return repository.DeleteMessageBinding(ctx, bindingID) -} - -// ToBindingDTO projects a binding row and its optional channel onto the user DTO. -func ToBindingDTO(row *model.MessageBinding, ch *model.MessageChannel) model.BindingDTO { - dto := model.BindingDTO{ - ID: row.ID, - UserID: row.UserID, - ChannelID: row.ChannelID, - PlatformUserID: row.PlatformUserID, - CreatedAt: row.CreatedAt, - } - if ch != nil { - dto.ChannelName = ch.Name - dto.ChannelType = ch.Type - } - return dto -} diff --git a/backend/plugins/domain/message_gateway/channels/qq/adapter.go b/backend/plugins/domain/msg_gateway/channels/qq/adapter.go similarity index 85% rename from backend/plugins/domain/message_gateway/channels/qq/adapter.go rename to backend/plugins/domain/msg_gateway/channels/qq/adapter.go index 828f7d4d..a72899ab 100644 --- a/backend/plugins/domain/message_gateway/channels/qq/adapter.go +++ b/backend/plugins/domain/msg_gateway/channels/qq/adapter.go @@ -7,8 +7,9 @@ package qq import ( "Wavelet/pkg/logger" "Wavelet/pkg/util" - "Wavelet/plugins/domain/message_gateway/model" - "Wavelet/plugins/domain/message_gateway/service" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/service" "context" "fmt" "strings" @@ -33,7 +34,7 @@ type qqEvent struct { // Adapter is an official QQ Bot C2C channel. type Adapter struct { - cfg model.ChannelConfig + cfg do.ChannelConfig onInbound service.Handler api openapi.OpenAPI tokenSrc oauth2.TokenSource @@ -43,7 +44,7 @@ type Adapter struct { } // New constructs a QQ adapter. -func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, error) { +func New(cfg do.ChannelConfig, onInbound service.Handler) (service.Channel, error) { if strings.TrimSpace(cfg.Credentials["app_id"]) == "" || strings.TrimSpace(cfg.Credentials["app_secret"]) == "" { return nil, fmt.Errorf("qq: app_id and app_secret are required") } @@ -51,11 +52,11 @@ func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, e } // Type returns qq. -func (a *Adapter) Type() string { return model.ChannelTypeQQ } +func (a *Adapter) Type() string { return consts.ChannelTypeQQ } // Capabilities reports C2C text/media support. -func (a *Adapter) Capabilities() model.Capability { - return model.Capability{Text: true, Image: true, File: true, Reply: true} +func (a *Adapter) Capabilities() do.Capability { + return do.Capability{Text: true, Image: true, File: true, Reply: true} } // Connect starts the official WebSocket session (C2C intent). @@ -128,7 +129,7 @@ func (a *Adapter) Disconnect(_ context.Context) error { } // Send posts a C2C text reply. -func (a *Adapter) Send(ctx context.Context, to model.Recipient, msg model.OutboundMessage) error { +func (a *Adapter) Send(ctx context.Context, to do.Recipient, msg do.OutboundMessage) error { a.mu.Lock() api := a.api a.mu.Unlock() @@ -152,9 +153,9 @@ func (a *Adapter) handleEvent(ctx context.Context, ev qqEvent) { if disconnected || a.onInbound == nil { return } - _ = a.onInbound(ctx, model.InboundMessage{ + _ = a.onInbound(ctx, do.InboundMessage{ ChannelID: a.cfg.ID, - ChannelType: model.ChannelTypeQQ, + ChannelType: consts.ChannelTypeQQ, PlatformUserID: ev.UserID, ChatID: ev.UserID, MessageID: ev.MessageID, diff --git a/backend/plugins/domain/message_gateway/channels/qq/adapter_test.go b/backend/plugins/domain/msg_gateway/channels/qq/adapter_test.go similarity index 69% rename from backend/plugins/domain/message_gateway/channels/qq/adapter_test.go rename to backend/plugins/domain/msg_gateway/channels/qq/adapter_test.go index fe4b9f62..2d265a6a 100644 --- a/backend/plugins/domain/message_gateway/channels/qq/adapter_test.go +++ b/backend/plugins/domain/msg_gateway/channels/qq/adapter_test.go @@ -4,14 +4,14 @@ package qq import ( - "Wavelet/plugins/domain/message_gateway/model" + "Wavelet/plugins/domain/msg_gateway/model/do" "context" "testing" ) func TestHandleEvent_DropsNonC2C(t *testing.T) { var got int - a := &Adapter{onInbound: func(ctx context.Context, msg model.InboundMessage) error { + a := &Adapter{onInbound: func(_ context.Context, _ do.InboundMessage) error { got++ return nil }} @@ -22,8 +22,8 @@ func TestHandleEvent_DropsNonC2C(t *testing.T) { } func TestHandleEvent_C2CText(t *testing.T) { - var got model.InboundMessage - a := &Adapter{cfg: model.ChannelConfig{ID: 3}, onInbound: func(ctx context.Context, msg model.InboundMessage) error { + var got do.InboundMessage + a := &Adapter{cfg: do.ChannelConfig{ID: 3}, onInbound: func(_ context.Context, msg do.InboundMessage) error { got = msg return nil }} @@ -34,7 +34,7 @@ func TestHandleEvent_C2CText(t *testing.T) { } func TestNew_RequiresCreds(t *testing.T) { - _, err := New(model.ChannelConfig{}, nil) + _, err := New(do.ChannelConfig{}, nil) if err == nil { t.Fatal("expected error") } diff --git a/backend/plugins/domain/message_gateway/channels/telegram/adapter.go b/backend/plugins/domain/msg_gateway/channels/telegram/adapter.go similarity index 80% rename from backend/plugins/domain/message_gateway/channels/telegram/adapter.go rename to backend/plugins/domain/msg_gateway/channels/telegram/adapter.go index eb044e87..e4bf9a6a 100644 --- a/backend/plugins/domain/message_gateway/channels/telegram/adapter.go +++ b/backend/plugins/domain/msg_gateway/channels/telegram/adapter.go @@ -7,8 +7,9 @@ package telegram import ( "Wavelet/pkg/logger" "Wavelet/pkg/util" - "Wavelet/plugins/domain/message_gateway/model" - "Wavelet/plugins/domain/message_gateway/service" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/service" "context" "fmt" "os" @@ -22,13 +23,13 @@ import ( // Adapter is a Telegram private-chat channel. type Adapter struct { - cfg model.ChannelConfig + cfg do.ChannelConfig onInbound service.Handler bot *tele.Bot } // New constructs a Telegram adapter. Call service.Register from the runner. -func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, error) { +func New(cfg do.ChannelConfig, onInbound service.Handler) (service.Channel, error) { if strings.TrimSpace(cfg.Credentials["bot_token"]) == "" { return nil, fmt.Errorf("telegram: bot_token is required") } @@ -36,11 +37,11 @@ func New(cfg model.ChannelConfig, onInbound service.Handler) (service.Channel, e } // Type returns telegram. -func (a *Adapter) Type() string { return model.ChannelTypeTelegram } +func (a *Adapter) Type() string { return consts.ChannelTypeTelegram } // Capabilities reports private-chat media support. -func (a *Adapter) Capabilities() model.Capability { - return model.Capability{Text: true, Image: true, File: true, Reply: true} +func (a *Adapter) Capabilities() do.Capability { + return do.Capability{Text: true, Image: true, File: true, Reply: true} } // longPollWindow is how long Telegram may hold a getUpdates call open before @@ -50,7 +51,7 @@ func (a *Adapter) Capabilities() model.Capability { const longPollWindow = 10 * time.Second // buildTeleSettings assembles the telebot settings. -func buildTeleSettings(cfg model.ChannelConfig) tele.Settings { +func buildTeleSettings(cfg do.ChannelConfig) tele.Settings { pref := tele.Settings{ Token: cfg.Credentials["bot_token"], Poller: &tele.LongPoller{Timeout: longPollWindow}, @@ -99,7 +100,7 @@ func (a *Adapter) Disconnect(_ context.Context) error { } // Send replies to a private chat. -func (a *Adapter) Send(_ context.Context, to model.Recipient, msg model.OutboundMessage) error { +func (a *Adapter) Send(_ context.Context, to do.Recipient, msg do.OutboundMessage) error { if a.bot == nil { return fmt.Errorf("telegram: not connected") } @@ -118,9 +119,9 @@ func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) { if a.onInbound == nil { return } - msg := model.InboundMessage{ + msg := do.InboundMessage{ ChannelID: a.cfg.ID, - ChannelType: model.ChannelTypeTelegram, + ChannelType: consts.ChannelTypeTelegram, PlatformUserID: strconv.FormatInt(m.Sender.ID, 10), ChatID: strconv.FormatInt(m.Chat.ID, 10), MessageID: strconv.Itoa(m.ID), @@ -146,7 +147,7 @@ func (a *Adapter) handleTeleMessage(ctx context.Context, m *tele.Message) { // downloadMedia fetches message media into a scratch directory, returned so the // caller can remove it once the inbound handler no longer needs the paths. // An empty dir means nothing was downloaded. -func (a *Adapter) downloadMedia(m *tele.Message) (string, []model.Attachment) { +func (a *Adapter) downloadMedia(m *tele.Message) (string, []do.Attachment) { var files []*tele.File var names []string if m.Photo != nil { @@ -166,16 +167,16 @@ func (a *Adapter) downloadMedia(m *tele.Message) (string, []model.Attachment) { } dir, err := os.MkdirTemp("", "wg-tg-*") if err != nil { - return "", []model.Attachment{{Error: err.Error()}} + return "", []do.Attachment{{Error: err.Error()}} } - out := make([]model.Attachment, 0, len(files)) + out := make([]do.Attachment, 0, len(files)) for i, f := range files { path := filepath.Join(dir, names[i]) if err := a.bot.Download(f, path); err != nil { - out = append(out, model.Attachment{FileName: names[i], Error: err.Error()}) + out = append(out, do.Attachment{FileName: names[i], Error: err.Error()}) continue } - out = append(out, model.Attachment{Path: path, FileName: names[i]}) + out = append(out, do.Attachment{Path: path, FileName: names[i]}) } return dir, out } diff --git a/backend/plugins/domain/message_gateway/channels/telegram/adapter_test.go b/backend/plugins/domain/msg_gateway/channels/telegram/adapter_test.go similarity index 82% rename from backend/plugins/domain/message_gateway/channels/telegram/adapter_test.go rename to backend/plugins/domain/msg_gateway/channels/telegram/adapter_test.go index cd700e05..c2c75ef8 100644 --- a/backend/plugins/domain/message_gateway/channels/telegram/adapter_test.go +++ b/backend/plugins/domain/msg_gateway/channels/telegram/adapter_test.go @@ -4,7 +4,7 @@ package telegram import ( - "Wavelet/plugins/domain/message_gateway/model" + "Wavelet/plugins/domain/msg_gateway/model/do" "context" "testing" "time" @@ -16,7 +16,7 @@ import ( // telebot 以 int(timeout/time.Second) 下发给 getUpdates。写成裸整数会被解释为 // 纳秒,令 timeout=0,长轮询退化为对 Bot API 的空转轮询。 func TestBuildTeleSettingsLongPollWindow(t *testing.T) { - pref := buildTeleSettings(model.ChannelConfig{ + pref := buildTeleSettings(do.ChannelConfig{ Credentials: map[string]string{"bot_token": "token"}, Extra: map[string]string{"base_url": "https://tg.example.com/api/"}, }) @@ -35,7 +35,7 @@ func TestBuildTeleSettingsLongPollWindow(t *testing.T) { func TestHandleUpdate_DropsGroups(t *testing.T) { var got int - a := &Adapter{onInbound: func(ctx context.Context, msg model.InboundMessage) error { + a := &Adapter{onInbound: func(_ context.Context, _ do.InboundMessage) error { got++ return nil }} @@ -51,10 +51,10 @@ func TestHandleUpdate_DropsGroups(t *testing.T) { } func TestHandleUpdate_PrivateText(t *testing.T) { - var got model.InboundMessage + var got do.InboundMessage a := &Adapter{ - cfg: model.ChannelConfig{ID: 7, Type: "telegram"}, - onInbound: func(ctx context.Context, msg model.InboundMessage) error { + cfg: do.ChannelConfig{ID: 7, Type: "telegram"}, + onInbound: func(_ context.Context, msg do.InboundMessage) error { got = msg return nil }, @@ -71,7 +71,7 @@ func TestHandleUpdate_PrivateText(t *testing.T) { } func TestNew_RequiresToken(t *testing.T) { - _, err := New(model.ChannelConfig{}, nil) + _, err := New(do.ChannelConfig{}, nil) if err == nil { t.Fatal("expected error") } diff --git a/backend/plugins/domain/msg_gateway/consts/bot.go b/backend/plugins/domain/msg_gateway/consts/bot.go new file mode 100644 index 00000000..9c2fb67c --- /dev/null +++ b/backend/plugins/domain/msg_gateway/consts/bot.go @@ -0,0 +1,27 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package consts defines constants, error types, and identifiers for msg_gateway. +package consts + +// Bot Channel type and scope constants. +const ( + ChannelTypeTelegram = "telegram" + ChannelTypeQQ = "qq" + MessageChannelTypeTelegram = "telegram" + MessageChannelTypeQQ = "qq" + MessageOwnerScopeSystem = "system" +) + +// Bot Task and Schedule identifier constants. +const ( + TaskCleanupPairingCodes = "msg_gateway:cleanup_pairing_codes" + TaskDispatchBotMsg = "msg_gateway:dispatch_bot_msg" + TaskTypeDispatchBotMsg = "dispatch_bot_msg" +) + +// Pairing code constants. +const ( + CodeAlphabet = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" + CodeLength = 8 +) diff --git a/backend/plugins/domain/message_gateway/errs/errs.go b/backend/plugins/domain/msg_gateway/consts/errs.go similarity index 70% rename from backend/plugins/domain/message_gateway/errs/errs.go rename to backend/plugins/domain/msg_gateway/consts/errs.go index e6fa4280..e8c32c9b 100644 --- a/backend/plugins/domain/message_gateway/errs/errs.go +++ b/backend/plugins/domain/msg_gateway/consts/errs.go @@ -1,9 +1,7 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -// Package errs defines error sentinels and user-facing error message constants -// for the message_gateway plugin. -package errs +package consts import "errors" @@ -16,31 +14,33 @@ var ( ErrBindingForbidden = errors.New("cannot unbind another user's binding") ErrChannelIDRequired = errors.New("channel_id is required") ErrChannelDisabled = errors.New("channel is not enabled") + ErrChannelNotFound = errors.New("channel not found") + ErrEventNotFound = errors.New("notification event not found") + ErrUserNotFound = errors.New("user not found") + ErrNoAdminUser = errors.New("no admin user found") + ErrTaskServiceNotAvail = errors.New("task service not available") - // ErrRecordNotFound maps GORM's missing-row sentinel at the repository boundary so + // ErrRecordNotFound maps GORM's missing-row sentinel at the DAO boundary so // upper layers never import gorm. Its text matches gorm.ErrRecordNotFound verbatim. ErrRecordNotFound = errors.New("record not found") - - // ErrUnsupportedUserLookupField rejects a column name that the repository is not - // allowed to interpolate into a WHERE clause. - ErrUnsupportedUserLookupField = errors.New("unsupported user lookup field") ) // User-facing validation and error message constants. const ( - ErrNameRequired = "name is required" - ErrTypeInvalid = "type must be telegram or qq" - ErrTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text - ErrQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text - ErrChannelNotFound = "channel not found" - ErrChannelProbeFailed = "channel probe failed" - MaskedSecret = "********" + ErrNameRequired = "name is required" + ErrTypeInvalid = "type must be telegram or qq" + ErrTelegramTokenRequired = "telegram bot secret is required" //nolint:gosec // user-facing validation text + ErrQQCredentialsRequired = "qq app id and secret are required" //nolint:gosec // user-facing validation text + ErrChannelNotFoundText = "channel not found" + ErrChannelProbeFailed = "channel probe failed" + ErrBotDispatchTextRequired = "message text is required" + ErrBotChannelNotRegistered = "channel adapter is not registered" + MaskedSecret = "********" ErrLoginRequired = "login required" ErrInvalidBindingID = "invalid binding id" ErrInvalidChannelID = "invalid channel id" ErrInvalidEventID = "invalid event id" - ErrEventNotFound = "notification event not found" ErrValidationFailed = "validation failed" ErrMissingTelegramToken = "missing telegram bot token" @@ -61,8 +61,6 @@ const ( ErrEventKeyOrTaskType = "either event_key or task_type must be provided" ErrUnsupportedEventKey = "unsupported built-in event key" ErrTaskServiceUnavailable = "task service not available" - ErrUserNotFound = "user not found" - ErrNoAdminUser = "no admin user found" ErrPayloadRequired = "payload is required" ErrInvalidJSONFormat = "invalid json format" diff --git a/backend/plugins/domain/msg_gateway/consts/push.go b/backend/plugins/domain/msg_gateway/consts/push.go new file mode 100644 index 00000000..806950e9 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/consts/push.go @@ -0,0 +1,43 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package consts + +// Push notification channel type constants. +const ( + TypeCustom = "custom" + TypeEmail = "email" + TypeTelegram = "telegram" + ChannelCustom = "custom" + ChannelEmail = "email" + ChannelLark = "lark" + ChannelDingTalk = "dingtalk" + ChannelTelegram = "telegram" + ChannelBark = "bark" + ChannelDiscord = "discord" + ChannelSlack = "slack" + ChannelPushover = "pushover" +) + +// Push message template and payload keys. +const ( + DefaultLevelInfo = "INFO" + KeyTitle = "title" + KeyContent = "content" + KeyLevel = "level" + + KeyURL = "url" + KeyToken = "token" + KeyOther = "other" + + TypeText = "text" + TypePassword = "password" + TypeTextarea = "textarea" +) + +// Push task identifier constants. +const ( + TaskPushNotification = "msg_gateway:push_notification" + SendNotificationTask = "push:send" + TaskTypeSendNotification = "send_notification" +) diff --git a/backend/plugins/domain/message_gateway/handler/admin.go b/backend/plugins/domain/msg_gateway/controller/admin.go similarity index 83% rename from backend/plugins/domain/message_gateway/handler/admin.go rename to backend/plugins/domain/msg_gateway/controller/admin.go index f90ccacd..98488d48 100644 --- a/backend/plugins/domain/message_gateway/handler/admin.go +++ b/backend/plugins/domain/msg_gateway/controller/admin.go @@ -1,14 +1,14 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package handler +package controller import ( "Wavelet/pkg/response" - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/service" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/service" + "errors" "net/http" - "strconv" "github.com/gin-gonic/gin" ) @@ -19,7 +19,7 @@ import ( // @Tags admin-message-gateway // @Produce json // @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.Definition} +// @Success 200 {object} response.Any{data=[]do.Definition} // @Router /api/v1/admin/message-gateway/channels/definitions [get] func ListAdminChannelDefinitions(c *gin.Context) { c.JSON(http.StatusOK, response.OK(service.ListDefinitions())) @@ -31,7 +31,7 @@ func ListAdminChannelDefinitions(c *gin.Context) { // @Tags admin-message-gateway // @Produce json // @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.ChannelDTO} +// @Success 200 {object} response.Any{data=[]do.ChannelDTO} // @Router /api/v1/admin/message-gateway/channels [get] func ListAdminChannels(c *gin.Context) { rows, err := service.ListChannels(c.Request.Context()) @@ -43,17 +43,12 @@ func ListAdminChannels(c *gin.Context) { } func parseAdminChannelID(c *gin.Context) (uint64, bool) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, errs.ErrInvalidChannelID) - return 0, false - } - return id, true + return parseUint64Param(c, "id", consts.ErrInvalidChannelID) } func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) { - if err.Error() == errs.ErrChannelNotFound { - response.AbortNotFound(c, err.Error()) + if errors.Is(err, consts.ErrChannelNotFound) || err.Error() == consts.ErrChannelNotFoundText { + response.AbortNotFound(c, consts.ErrChannelNotFoundText) return } fallback(c, err.Error()) @@ -66,8 +61,8 @@ func handleAdminChannelError(c *gin.Context, err error, fallback func(c *gin.Con // @Accept json // @Produce json // @Security SessionCookie -// @Param request body model.CreateChannelRequest true "create body" -// @Success 200 {object} response.Any{data=model.ChannelDTO} +// @Param request body do.CreateChannelRequest true "create body" +// @Success 200 {object} response.Any{data=do.ChannelDTO} // @Failure 400 {object} response.Any // @Router /api/v1/admin/message-gateway/channels [post] func CreateAdminChannel(c *gin.Context) { @@ -82,8 +77,8 @@ func CreateAdminChannel(c *gin.Context) { // @Produce json // @Security SessionCookie // @Param id path int true "channel id" -// @Param request body model.UpdateChannelRequest true "update body" -// @Success 200 {object} response.Any{data=model.ChannelDTO} +// @Param request body do.UpdateChannelRequest true "update body" +// @Success 200 {object} response.Any{data=do.ChannelDTO} // @Failure 400 {object} response.Any // @Failure 404 {object} response.Any // @Router /api/v1/admin/message-gateway/channels/{id} [patch] diff --git a/backend/plugins/domain/msg_gateway/controller/base.go b/backend/plugins/domain/msg_gateway/controller/base.go new file mode 100644 index 00000000..f255d3d2 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/controller/base.go @@ -0,0 +1,69 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package controller + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/ginutil" + "Wavelet/pkg/response" + "context" + "net/http" + "strconv" + + "github.com/gin-gonic/gin" +) + +// currentUser extracts the authenticated UserDTO from gin.Context. +func currentUser(c *gin.Context) (*contracts.UserDTO, bool) { + return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) +} + +// parseUint64Param parses a uint64 URL path parameter. +func parseUint64Param(c *gin.Context, paramName, errInvalid string) (uint64, bool) { + id, err := strconv.ParseUint(c.Param(paramName), 10, 64) + if err != nil { + response.AbortBadRequest(c, errInvalid) + return 0, false + } + return id, true +} + +// handleJSONRequest binds a JSON body, executes the service handler, and writes the standard success envelope. +func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) { + var req Req + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + res, err := handler(c.Request.Context(), req) + if err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + c.JSON(http.StatusOK, response.OK(res)) +} + +// handleEntityUpdate resolves a path identifier and JSON body, executes the updater, and handles errors with onErr. +func handleEntityUpdate[Req any, Res any]( + c *gin.Context, + parseID func(*gin.Context) (uint64, bool), + updater func(ctx context.Context, id uint64, req Req) (Res, error), + onErr func(*gin.Context, error), +) { + id, ok := parseID(c) + if !ok { + return + } + var req Req + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } + dto, err := updater(c.Request.Context(), id, req) + if err != nil { + onErr(c, err) + return + } + c.JSON(http.StatusOK, response.OK(dto)) +} diff --git a/backend/plugins/domain/message_gateway/handler/push_channel.go b/backend/plugins/domain/msg_gateway/controller/push_channel.go similarity index 79% rename from backend/plugins/domain/message_gateway/handler/push_channel.go rename to backend/plugins/domain/msg_gateway/controller/push_channel.go index c3c07a72..89d440a0 100644 --- a/backend/plugins/domain/message_gateway/handler/push_channel.go +++ b/backend/plugins/domain/msg_gateway/controller/push_channel.go @@ -1,16 +1,15 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package handler +package controller import ( "Wavelet/pkg/response" - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/model" - "Wavelet/plugins/domain/message_gateway/service" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/service" "errors" "net/http" - "strconv" "github.com/gin-gonic/gin" ) @@ -24,7 +23,7 @@ import ( // @Success 200 {object} response.Any "通道配置定义列表" // @Router /api/v1/admin/push/channels/definitions [get] func ListPushChannelDefinitions(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(model.ListPushDefinitions())) + c.JSON(http.StatusOK, response.OK(do.ListPushDefinitions())) } // ListPushChannels 获取消息通道列表 @@ -33,7 +32,7 @@ func ListPushChannelDefinitions(c *gin.Context) { // @Tags admin-push // @Produce json // @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.PushChannel} "消息通道列表" +// @Success 200 {object} response.Any{data=[]entity.PushChannel} "消息通道列表" // @Router /api/v1/admin/push/channels [get] func ListPushChannels(c *gin.Context) { channels, err := service.ListPushChannels(c.Request.Context()) @@ -46,18 +45,13 @@ func ListPushChannels(c *gin.Context) { // parsePushChannelID reads the path identifier of a push channel. func parsePushChannelID(c *gin.Context) (uint64, bool) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, errs.ErrInvalidChannelID) - return 0, false - } - return id, true + return parseUint64Param(c, "id", consts.ErrInvalidChannelID) } // handlePushChannelNotFoundError maps a missing channel row to 404, others to fallback. func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) { - if errors.Is(err, errs.ErrRecordNotFound) { - response.AbortNotFound(c, errs.ErrChannelNotFound) + if errors.Is(err, consts.ErrRecordNotFound) || errors.Is(err, consts.ErrChannelNotFound) || err.Error() == consts.ErrChannelNotFoundText { + response.AbortNotFound(c, consts.ErrChannelNotFoundText) return } fallback(c, err.Error()) @@ -70,8 +64,8 @@ func handlePushChannelNotFoundError(c *gin.Context, err error, fallback func(c * // @Accept json // @Produce json // @Security SessionCookie -// @Param request body model.CreatePushChannelRequest true "创建参数" -// @Success 200 {object} response.Any{data=model.PushChannel} "创建成功" +// @Param request body do.CreatePushChannelRequest true "创建参数" +// @Success 200 {object} response.Any{data=entity.PushChannel} "创建成功" // @Router /api/v1/admin/push/channels [post] func CreatePushChannel(c *gin.Context) { handleJSONRequest(c, service.CreatePushChannel) @@ -85,8 +79,8 @@ func CreatePushChannel(c *gin.Context) { // @Produce json // @Security SessionCookie // @Param id path uint64 true "通道ID" -// @Param request body model.UpdatePushChannelRequest true "更新参数" -// @Success 200 {object} response.Any{data=model.PushChannel} "更新成功" +// @Param request body do.UpdatePushChannelRequest true "更新参数" +// @Success 200 {object} response.Any{data=entity.PushChannel} "更新成功" // @Router /api/v1/admin/push/channels/{id} [put] func UpdatePushChannel(c *gin.Context) { handleEntityUpdate(c, parsePushChannelID, service.UpdatePushChannel, func(c *gin.Context, err error) { @@ -123,11 +117,11 @@ func DeletePushChannel(c *gin.Context) { // @Accept json // @Produce json // @Security SessionCookie -// @Param request body model.TestPushChannelRequest true "测试参数" +// @Param request body do.TestPushChannelRequest true "测试参数" // @Success 200 {object} response.Any "测试触发成功" // @Router /api/v1/admin/push/channels/test [post] func TestPushChannel(c *gin.Context) { - var req model.TestPushChannelRequest + var req do.TestPushChannelRequest if err := c.ShouldBindJSON(&req); err != nil { response.AbortBadRequest(c, err.Error()) return diff --git a/backend/plugins/domain/message_gateway/handler/push_event.go b/backend/plugins/domain/msg_gateway/controller/push_event.go similarity index 86% rename from backend/plugins/domain/message_gateway/handler/push_event.go rename to backend/plugins/domain/msg_gateway/controller/push_event.go index 04c6e717..616eb5d8 100644 --- a/backend/plugins/domain/message_gateway/handler/push_event.go +++ b/backend/plugins/domain/msg_gateway/controller/push_event.go @@ -1,13 +1,13 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package handler +package controller import ( "Wavelet/pkg/response" - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/model" - "Wavelet/plugins/domain/message_gateway/service" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/service" "errors" "net/http" "strconv" @@ -21,7 +21,7 @@ import ( // @Tags admin-push // @Produce json // @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.PushEvent} "通知事件列表" +// @Success 200 {object} response.Any{data=[]entity.PushEvent} "通知事件列表" // @Router /api/v1/admin/push/events [get] func ListPushEvents(c *gin.Context) { ctx := c.Request.Context() @@ -47,18 +47,13 @@ func ListBuiltInPushEvents(c *gin.Context) { // parsePushEventID reads the path identifier of a push event. func parsePushEventID(c *gin.Context) (uint64, bool) { - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, errs.ErrInvalidEventID) - return 0, false - } - return id, true + return parseUint64Param(c, "id", consts.ErrInvalidEventID) } // handlePushEventNotFoundError maps a missing event row to 404, others to fallback. func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gin.Context, msg string)) { - if errors.Is(err, errs.ErrRecordNotFound) { - response.AbortNotFound(c, errs.ErrEventNotFound) + if errors.Is(err, consts.ErrRecordNotFound) || errors.Is(err, consts.ErrEventNotFound) || err.Error() == consts.ErrEventNotFound.Error() { + response.AbortNotFound(c, consts.ErrEventNotFound.Error()) return } fallback(c, err.Error()) @@ -71,8 +66,8 @@ func handlePushEventNotFoundError(c *gin.Context, err error, fallback func(c *gi // @Accept json // @Produce json // @Security SessionCookie -// @Param request body model.CreatePushEventRequest true "创建参数" -// @Success 200 {object} response.Any{data=model.PushEvent} "创建成功" +// @Param request body do.CreatePushEventRequest true "创建参数" +// @Success 200 {object} response.Any{data=entity.PushEvent} "创建成功" // @Router /api/v1/admin/push/events [post] func CreatePushEvent(c *gin.Context) { handleJSONRequest(c, service.CreatePushEvent) @@ -108,7 +103,7 @@ func DeletePushEvent(c *gin.Context) { // @Produce json // @Security SessionCookie // @Param id path int true "事件 ID" -// @Param request body model.UpdatePushEventRequest true "更新参数" +// @Param request body do.UpdatePushEventRequest true "更新参数" // @Success 200 {object} response.Any{data=string} "修改成功" // @Router /api/v1/admin/push/events/{id} [put] func UpdatePushEvent(c *gin.Context) { @@ -117,7 +112,7 @@ func UpdatePushEvent(c *gin.Context) { return } - var req model.UpdatePushEventRequest + var req do.UpdatePushEventRequest if err := c.ShouldBindJSON(&req); err != nil { response.AbortBadRequest(c, err.Error()) return @@ -175,8 +170,9 @@ func ListPushHistories(c *gin.Context) { pageSize = 20 } - total, results, err := service.ListPushHistories(c.Request.Context(), model.PushHistoryListFilter{ + total, results, err := service.ListPushHistories(c.Request.Context(), do.PushHistoryListFilter{ EventKey: c.Query("event_key"), + Channel: c.Query("channel"), Status: c.Query("status"), Page: page, PageSize: pageSize, @@ -199,11 +195,11 @@ func ListPushHistories(c *gin.Context) { // @Accept json // @Produce json // @Security SessionCookie -// @Param request body model.TestPushRequest true "测试请求体" +// @Param request body do.TestPushRequest true "测试请求体" // @Success 200 {object} response.Any{data=string} "测试成功" // @Router /api/v1/admin/push/test [post] func TestPush(c *gin.Context) { - var req model.TestPushRequest + var req do.TestPushRequest if err := c.ShouldBindJSON(&req); err != nil { response.AbortBadRequest(c, err.Error()) return diff --git a/backend/plugins/domain/message_gateway/handler/router.go b/backend/plugins/domain/msg_gateway/controller/router.go similarity index 96% rename from backend/plugins/domain/message_gateway/handler/router.go rename to backend/plugins/domain/msg_gateway/controller/router.go index 837505e6..cc8f1ec7 100644 --- a/backend/plugins/domain/message_gateway/handler/router.go +++ b/backend/plugins/domain/msg_gateway/controller/router.go @@ -1,7 +1,8 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package handler +// Package controller provides HTTP endpoints for msg_gateway. +package controller import ( "Wavelet/core/extpoints" diff --git a/backend/plugins/domain/message_gateway/handler/handlers.go b/backend/plugins/domain/msg_gateway/controller/user.go similarity index 55% rename from backend/plugins/domain/message_gateway/handler/handlers.go rename to backend/plugins/domain/msg_gateway/controller/user.go index 877ade86..6323f7d7 100644 --- a/backend/plugins/domain/message_gateway/handler/handlers.go +++ b/backend/plugins/domain/msg_gateway/controller/user.go @@ -1,81 +1,31 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -// Package handler provides HTTP endpoints for message_gateway. -package handler +package controller import ( - "Wavelet/core/contracts" - "Wavelet/pkg/ginutil" "Wavelet/pkg/response" - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/model" - "Wavelet/plugins/domain/message_gateway/service" - "context" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/service" "errors" "net/http" - "strconv" "github.com/gin-gonic/gin" ) -func currentUser(c *gin.Context) (*contracts.UserDTO, bool) { - return ginutil.GetFromContext[*contracts.UserDTO](c, contracts.AuthUserObjKey) -} - -// handleJSONRequest binds a JSON body, runs the service use case and writes the -// standard success envelope; any service error surfaces as a bad request. -func handleJSONRequest[Req any, Res any](c *gin.Context, handler func(ctx context.Context, req Req) (Res, error)) { - var req Req - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - res, err := handler(c.Request.Context(), req) - if err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - c.JSON(http.StatusOK, response.OK(res)) -} - -// handleEntityUpdate resolves a path identifier plus JSON body, runs the updater -// use case and writes the success envelope; error classification is delegated to onErr. -func handleEntityUpdate[Req any, Res any]( - c *gin.Context, - parseID func(*gin.Context) (uint64, bool), - updater func(ctx context.Context, id uint64, req Req) (Res, error), - onErr func(*gin.Context, error), -) { - id, ok := parseID(c) - if !ok { - return - } - var req Req - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, err.Error()) - return - } - dto, err := updater(c.Request.Context(), id, req) - if err != nil { - onErr(c, err) - return - } - c.JSON(http.StatusOK, response.OK(dto)) -} - // ListChannels lists enabled channels a user can bind. // @Summary List enabled messaging channels // @Description Returns enabled system bots the current user can pair with // @Tags message-gateway // @Produce json // @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.PublicChannelDTO} +// @Success 200 {object} response.Any{data=[]do.PublicChannelDTO} // @Failure 401 {object} response.Any // @Router /api/v1/message-gateway/channels [get] func ListChannels(c *gin.Context) { if user, ok := currentUser(c); !ok || user == nil { - response.AbortUnauthorized(c, errs.ErrLoginRequired) + response.AbortUnauthorized(c, consts.ErrLoginRequired) return } rows, err := service.ListEnabledPublicChannels(c.Request.Context()) @@ -92,13 +42,13 @@ func ListChannels(c *gin.Context) { // @Tags message-gateway // @Produce json // @Security SessionCookie -// @Success 200 {object} response.Any{data=[]model.BindingDTO} +// @Success 200 {object} response.Any{data=[]do.BindingDTO} // @Failure 401 {object} response.Any // @Router /api/v1/message-gateway/bindings [get] func ListBindings(c *gin.Context) { user, ok := currentUser(c) if !ok || user == nil { - response.AbortUnauthorized(c, errs.ErrLoginRequired) + response.AbortUnauthorized(c, consts.ErrLoginRequired) return } rows, err := service.ListUserBindings(c.Request.Context(), user.ID) @@ -116,25 +66,25 @@ func ListBindings(c *gin.Context) { // @Accept json // @Produce json // @Security SessionCookie -// @Param request body model.BindRequest true "bind body" -// @Success 200 {object} response.Any{data=model.BindingDTO} +// @Param request body do.BindRequest true "bind body" +// @Success 200 {object} response.Any{data=do.BindingDTO} // @Failure 400 {object} response.Any // @Failure 409 {object} response.Any // @Router /api/v1/message-gateway/bindings [post] func BindBinding(c *gin.Context) { user, ok := currentUser(c) if !ok || user == nil { - response.AbortUnauthorized(c, errs.ErrLoginRequired) + response.AbortUnauthorized(c, consts.ErrLoginRequired) return } - var req model.BindRequest + var req do.BindRequest if err := c.ShouldBindJSON(&req); err != nil { response.AbortBadRequest(c, err.Error()) return } dto, err := service.BindChannel(c.Request.Context(), user.ID, req) if err != nil { - if errors.Is(err, errs.ErrPlatformAlreadyBound) { + if errors.Is(err, consts.ErrPlatformAlreadyBound) { response.AbortConflict(c, err.Error()) return } @@ -158,20 +108,19 @@ func BindBinding(c *gin.Context) { func UnbindBinding(c *gin.Context) { user, ok := currentUser(c) if !ok || user == nil { - response.AbortUnauthorized(c, errs.ErrLoginRequired) + response.AbortUnauthorized(c, consts.ErrLoginRequired) return } - id, err := strconv.ParseUint(c.Param("id"), 10, 64) - if err != nil { - response.AbortBadRequest(c, errs.ErrInvalidBindingID) + id, ok := parseUint64Param(c, "id", consts.ErrInvalidBindingID) + if !ok { return } if err := service.UnbindChannel(c.Request.Context(), user.ID, id); err != nil { - if errors.Is(err, errs.ErrBindingNotFound) { + if errors.Is(err, consts.ErrBindingNotFound) { response.AbortNotFound(c, err.Error()) return } - if errors.Is(err, errs.ErrBindingForbidden) { + if errors.Is(err, consts.ErrBindingForbidden) { response.AbortForbidden(c, err.Error()) return } diff --git a/backend/plugins/domain/message_gateway/repository/repository.go b/backend/plugins/domain/msg_gateway/dao/bot.go similarity index 52% rename from backend/plugins/domain/message_gateway/repository/repository.go rename to backend/plugins/domain/msg_gateway/dao/bot.go index b6e5302a..630208c4 100644 --- a/backend/plugins/domain/message_gateway/repository/repository.go +++ b/backend/plugins/domain/msg_gateway/dao/bot.go @@ -1,68 +1,20 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -// Package repository provides data persistence for the message_gateway plugin. -package repository +package dao import ( - "Wavelet/core" - "Wavelet/core/contracts" "Wavelet/pkg/idgen" - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/model" + "Wavelet/plugins/domain/msg_gateway/model/entity" "context" "errors" - "sync" "time" "gorm.io/gorm" ) -var ( - dbMu sync.RWMutex - dbSvc contracts.DBService -) - -// SetDBServiceForTest injects a DBService for tests. Production wiring must use Apply. -func SetDBServiceForTest(s contracts.DBService) { - SetDBService(s) -} - -// SetDBService sets the database service singleton. -func SetDBService(s contracts.DBService) { - dbMu.Lock() - defer dbMu.Unlock() - dbSvc = s -} - -// GetDB resolves the persistence handle for the current call, preferring an -// explicitly injected *core.Context before falling back to the plugin singleton. -func GetDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } - } - dbMu.RLock() - s := dbSvc - dbMu.RUnlock() - if s != nil { - return s.DB(ctx) - } - return nil -} - -// mapNotFound translates GORM's missing-row sentinel into the plugin-level -// errs.ErrRecordNotFound so the service and handler layers stay free of gorm imports. -func mapNotFound(err error) error { - if errors.Is(err, gorm.ErrRecordNotFound) { - return errs.ErrRecordNotFound - } - return err -} - // CreateMessageChannel inserts a channel row. -func CreateMessageChannel(ctx context.Context, ch *model.MessageChannel) error { +func CreateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error { if ch.ID == 0 { ch.ID = idgen.NextUint64ID() } @@ -70,13 +22,13 @@ func CreateMessageChannel(ctx context.Context, ch *model.MessageChannel) error { } // UpdateMessageChannel saves a channel row. -func UpdateMessageChannel(ctx context.Context, ch *model.MessageChannel) error { +func UpdateMessageChannel(ctx context.Context, ch *entity.MessageChannel) error { return GetDB(ctx).Save(ch).Error } // GetMessageChannel loads a channel by id. -func GetMessageChannel(ctx context.Context, id uint64) (*model.MessageChannel, error) { - var ch model.MessageChannel +func GetMessageChannel(ctx context.Context, id uint64) (*entity.MessageChannel, error) { + var ch entity.MessageChannel if err := GetDB(ctx).Where("id = ?", id).First(&ch).Error; err != nil { return nil, mapNotFound(err) } @@ -84,8 +36,8 @@ func GetMessageChannel(ctx context.Context, id uint64) (*model.MessageChannel, e } // ListMessageChannels returns all channels newest first. -func ListMessageChannels(ctx context.Context) ([]model.MessageChannel, error) { - var rows []model.MessageChannel +func ListMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) { + var rows []entity.MessageChannel if err := GetDB(ctx).Order("id DESC").Find(&rows).Error; err != nil { return nil, err } @@ -95,18 +47,18 @@ func ListMessageChannels(ctx context.Context) ([]model.MessageChannel, error) { // DeleteMessageChannel removes pairings, bindings, then the channel. func DeleteMessageChannel(ctx context.Context, id uint64) error { return GetDB(ctx).Transaction(func(tx *gorm.DB) error { - if err := tx.Where("channel_id = ?", id).Delete(&model.MessagePairingCode{}).Error; err != nil { + if err := tx.Where("channel_id = ?", id).Delete(&entity.MessagePairingCode{}).Error; err != nil { return err } - if err := tx.Where("channel_id = ?", id).Delete(&model.MessageBinding{}).Error; err != nil { + if err := tx.Where("channel_id = ?", id).Delete(&entity.MessageBinding{}).Error; err != nil { return err } - return tx.Delete(&model.MessageChannel{}, id).Error + return tx.Delete(&entity.MessageChannel{}, id).Error }) } // CreateMessageBinding inserts a binding. -func CreateMessageBinding(ctx context.Context, b *model.MessageBinding) error { +func CreateMessageBinding(ctx context.Context, b *entity.MessageBinding) error { if b.ID == 0 { b.ID = idgen.NextUint64ID() } @@ -114,8 +66,8 @@ func CreateMessageBinding(ctx context.Context, b *model.MessageBinding) error { } // GetBindingByChannelPlatform finds a binding for a platform user on a channel. -func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*model.MessageBinding, error) { - var b model.MessageBinding +func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platformUserID string) (*entity.MessageBinding, error) { + var b entity.MessageBinding err := GetDB(ctx).Where("channel_id = ? AND platform_user_id = ?", channelID, platformUserID).First(&b).Error if err != nil { return nil, mapNotFound(err) @@ -124,17 +76,26 @@ func GetBindingByChannelPlatform(ctx context.Context, channelID uint64, platform } // ListBindingsByUser lists bindings for a Wavelet user. -func ListBindingsByUser(ctx context.Context, userID uint64) ([]model.MessageBinding, error) { - var rows []model.MessageBinding +func ListBindingsByUser(ctx context.Context, userID uint64) ([]entity.MessageBinding, error) { + var rows []entity.MessageBinding if err := GetDB(ctx).Where("user_id = ?", userID).Order("id DESC").Find(&rows).Error; err != nil { return nil, err } return rows, nil } +// ListBindingsByChannel lists bindings on one messaging channel. +func ListBindingsByChannel(ctx context.Context, channelID uint64) ([]entity.MessageBinding, error) { + var rows []entity.MessageBinding + if err := GetDB(ctx).Where("channel_id = ?", channelID).Order("id DESC").Find(&rows).Error; err != nil { + return nil, err + } + return rows, nil +} + // GetMessageBinding loads a binding by id. -func GetMessageBinding(ctx context.Context, id uint64) (*model.MessageBinding, error) { - var b model.MessageBinding +func GetMessageBinding(ctx context.Context, id uint64) (*entity.MessageBinding, error) { + var b entity.MessageBinding if err := GetDB(ctx).Where("id = ?", id).First(&b).Error; err != nil { return nil, mapNotFound(err) } @@ -143,12 +104,12 @@ func GetMessageBinding(ctx context.Context, id uint64) (*model.MessageBinding, e // DeleteMessageBinding deletes a binding by id. func DeleteMessageBinding(ctx context.Context, id uint64) error { - return GetDB(ctx).Delete(&model.MessageBinding{}, id).Error + return GetDB(ctx).Delete(&entity.MessageBinding{}, id).Error } // UpsertPairingCode reuses an unexpired code for the same channel+platform user. -func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*model.MessagePairingCode, error) { - var existing model.MessagePairingCode +func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, code string, expiresAt time.Time) (*entity.MessagePairingCode, error) { + var existing entity.MessagePairingCode err := GetDB(ctx). Where("channel_id = ? AND platform_user_id = ? AND expires_at > ?", channelID, platformUserID, time.Now()). First(&existing).Error @@ -158,7 +119,7 @@ func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, co if !errors.Is(err, gorm.ErrRecordNotFound) { return nil, err } - row := &model.MessagePairingCode{ + row := &entity.MessagePairingCode{ Code: code, ChannelID: channelID, PlatformUserID: platformUserID, @@ -171,8 +132,8 @@ func UpsertPairingCode(ctx context.Context, channelID uint64, platformUserID, co } // GetPairingCode loads a pairing code by normalized code string. -func GetPairingCode(ctx context.Context, code string) (*model.MessagePairingCode, error) { - var row model.MessagePairingCode +func GetPairingCode(ctx context.Context, code string) (*entity.MessagePairingCode, error) { + var row entity.MessagePairingCode if err := GetDB(ctx).Where("code = ?", code).First(&row).Error; err != nil { return nil, mapNotFound(err) } @@ -181,17 +142,17 @@ func GetPairingCode(ctx context.Context, code string) (*model.MessagePairingCode // DeletePairingCode removes a pairing code. func DeletePairingCode(ctx context.Context, code string) error { - return GetDB(ctx).Where("code = ?", code).Delete(&model.MessagePairingCode{}).Error + return GetDB(ctx).Where("code = ?", code).Delete(&entity.MessagePairingCode{}).Error } // DeleteExpiredPairingCodes removes expired pairing rows. func DeleteExpiredPairingCodes(ctx context.Context) error { - return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&model.MessagePairingCode{}).Error + return GetDB(ctx).Where("expires_at <= ?", time.Now()).Delete(&entity.MessagePairingCode{}).Error } // ListEnabledMessageChannels returns enabled channels. -func ListEnabledMessageChannels(ctx context.Context) ([]model.MessageChannel, error) { - var rows []model.MessageChannel +func ListEnabledMessageChannels(ctx context.Context) ([]entity.MessageChannel, error) { + var rows []entity.MessageChannel if err := GetDB(ctx).Where("enabled = ?", true).Find(&rows).Error; err != nil { return nil, err } diff --git a/backend/plugins/domain/msg_gateway/dao/bot_test.go b/backend/plugins/domain/msg_gateway/dao/bot_test.go new file mode 100644 index 00000000..aa66b500 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/dao/bot_test.go @@ -0,0 +1,64 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package dao_test + +import ( + "Wavelet/pkg/idgen" + "Wavelet/pkg/testhelper" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestBotDAO_ChannelAndBinding(t *testing.T) { + _ = idgen.Init(1) + db, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + require.NoError(t, db.AutoMigrate(&entity.MessageChannel{}, &entity.MessageBinding{}, &entity.MessagePairingCode{})) + + dao.SetDBServiceForTest(stubDBService{db: db}) + t.Cleanup(func() { dao.SetDBServiceForTest(nil) }) + + ctx := context.Background() + + ch := entity.MessageChannel{ + Name: "tg_bot", + Type: "telegram", + OwnerScope: "system", + Credentials: "encrypted_token", + Enabled: true, + } + require.NoError(t, dao.CreateMessageChannel(ctx, &ch)) + assert.NotZero(t, ch.ID) + + code, err := dao.UpsertPairingCode(ctx, ch.ID, "tg_user_1", "ABCD1234", time.Now().Add(10*time.Minute)) + require.NoError(t, err) + assert.Equal(t, "ABCD1234", code.Code) + + // Reusing pairing code + code2, err := dao.UpsertPairingCode(ctx, ch.ID, "tg_user_1", "XYZ9999", time.Now().Add(10*time.Minute)) + require.NoError(t, err) + assert.Equal(t, "ABCD1234", code2.Code) + + binding := entity.MessageBinding{ + UserID: 42, + ChannelID: ch.ID, + PlatformUserID: "tg_user_1", + } + require.NoError(t, dao.CreateMessageBinding(ctx, &binding)) + assert.NotZero(t, binding.ID) + + bindings, err := dao.ListBindingsByUser(ctx, 42) + require.NoError(t, err) + assert.Len(t, bindings, 1) + + require.NoError(t, dao.DeleteMessageChannel(ctx, ch.ID)) + _, err = dao.GetMessageChannel(ctx, ch.ID) + assert.Error(t, err) +} diff --git a/backend/plugins/domain/msg_gateway/dao/dao.go b/backend/plugins/domain/msg_gateway/dao/dao.go new file mode 100644 index 00000000..f48ac18a --- /dev/null +++ b/backend/plugins/domain/msg_gateway/dao/dao.go @@ -0,0 +1,80 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package dao provides database persistence and caching for the msg_gateway plugin. +package dao + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "Wavelet/plugins/domain/msg_gateway/consts" + "context" + "errors" + "sync" + + "gorm.io/gorm" +) + +var ( + dbMu sync.RWMutex + dbSvc contracts.DBService + cacheMu sync.RWMutex + cacheSvc contracts.CacheService +) + +// SetDBService sets the database service singleton. +func SetDBService(s contracts.DBService) { + dbMu.Lock() + defer dbMu.Unlock() + dbSvc = s +} + +// SetDBServiceForTest injects a DBService for tests. +func SetDBServiceForTest(s contracts.DBService) { + SetDBService(s) +} + +// SetCacheService sets the cache service singleton. +func SetCacheService(s contracts.CacheService) { + cacheMu.Lock() + defer cacheMu.Unlock() + cacheSvc = s +} + +// GetDB resolves the persistence handle for the current call. +func GetDB(ctx context.Context) *gorm.DB { + if c, ok := ctx.(*core.Context); ok && c != nil { + if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { + return s.DB(ctx) + } + } + dbMu.RLock() + s := dbSvc + dbMu.RUnlock() + if s != nil { + return s.DB(ctx) + } + return nil +} + +// GetCache resolves the cache service for the current call. +func GetCache(ctx context.Context) contracts.CacheService { + if c, ok := ctx.(*core.Context); ok && c != nil { + if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { + return s + } + } + cacheMu.RLock() + s := cacheSvc + cacheMu.RUnlock() + return s +} + +// mapNotFound translates GORM's missing-row sentinel into the plugin-level +// consts.ErrRecordNotFound so the service and controller layers stay free of gorm imports. +func mapNotFound(err error) error { + if errors.Is(err, gorm.ErrRecordNotFound) { + return consts.ErrRecordNotFound + } + return err +} diff --git a/backend/plugins/domain/message_gateway/repository/push.go b/backend/plugins/domain/msg_gateway/dao/push.go similarity index 52% rename from backend/plugins/domain/message_gateway/repository/push.go rename to backend/plugins/domain/msg_gateway/dao/push.go index a944c3c9..f3a22d03 100644 --- a/backend/plugins/domain/message_gateway/repository/push.go +++ b/backend/plugins/domain/msg_gateway/dao/push.go @@ -1,17 +1,12 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package repository +package dao import ( - "Wavelet/core" - "Wavelet/core/contracts" - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/model" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" "context" - "errors" - "fmt" - "sync" "time" "gorm.io/gorm" @@ -22,34 +17,9 @@ const ( activePushEventCacheTTL = 24 * time.Hour ) -var ( - cacheMu sync.RWMutex - cacheSvc contracts.CacheService -) - -// SetCacheService sets the cache service singleton. -func SetCacheService(s contracts.CacheService) { - cacheMu.Lock() - defer cacheMu.Unlock() - cacheSvc = s -} - -// GetCache resolves the cache service for the current call. -func GetCache(ctx context.Context) contracts.CacheService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { - return s - } - } - cacheMu.RLock() - s := cacheSvc - cacheMu.RUnlock() - return s -} - // ListPushChannelsRecord returns all push channels ordered by creation time descending. -func ListPushChannelsRecord(ctx context.Context) ([]model.PushChannel, error) { - var channels []model.PushChannel +func ListPushChannelsRecord(ctx context.Context) ([]entity.PushChannel, error) { + var channels []entity.PushChannel if err := GetDB(ctx).Order("created_at DESC").Find(&channels).Error; err != nil { return nil, err } @@ -57,17 +27,17 @@ func ListPushChannelsRecord(ctx context.Context) ([]model.PushChannel, error) { } // GetPushChannelByIDRecord loads a push channel by primary key. -func GetPushChannelByIDRecord(ctx context.Context, id uint64) (model.PushChannel, error) { - var channel model.PushChannel +func GetPushChannelByIDRecord(ctx context.Context, id uint64) (entity.PushChannel, error) { + var channel entity.PushChannel if err := GetDB(ctx).Where("id = ?", id).First(&channel).Error; err != nil { - return model.PushChannel{}, mapNotFound(err) + return entity.PushChannel{}, mapNotFound(err) } return channel, nil } // GetPushChannelByNameRecord loads a push channel by its unique name. -func GetPushChannelByNameRecord(ctx context.Context, name string) (*model.PushChannel, error) { - var channel model.PushChannel +func GetPushChannelByNameRecord(ctx context.Context, name string) (*entity.PushChannel, error) { + var channel entity.PushChannel if err := GetDB(ctx).Where("name = ?", name).First(&channel).Error; err != nil { return nil, mapNotFound(err) } @@ -77,14 +47,14 @@ func GetPushChannelByNameRecord(ctx context.Context, name string) (*model.PushCh // CountPushChannelsByNameRecord returns how many channels share the given name. func CountPushChannelsByNameRecord(ctx context.Context, name string) (int64, error) { var count int64 - if err := GetDB(ctx).Model(&model.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil { + if err := GetDB(ctx).Model(&entity.PushChannel{}).Where("name = ?", name).Count(&count).Error; err != nil { return 0, err } return count, nil } // CreatePushChannelRecord persists a new channel and invalidates cache. -func CreatePushChannelRecord(ctx context.Context, channel *model.PushChannel) error { +func CreatePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error { if err := GetDB(ctx).Create(channel).Error; err != nil { return err } @@ -93,7 +63,7 @@ func CreatePushChannelRecord(ctx context.Context, channel *model.PushChannel) er } // SavePushChannelRecord updates a channel and invalidates cache. -func SavePushChannelRecord(ctx context.Context, channel *model.PushChannel) error { +func SavePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error { if err := GetDB(ctx).Save(channel).Error; err != nil { return err } @@ -102,7 +72,7 @@ func SavePushChannelRecord(ctx context.Context, channel *model.PushChannel) erro } // DeletePushChannelRecord removes a channel and invalidates cache. -func DeletePushChannelRecord(ctx context.Context, channel *model.PushChannel) error { +func DeletePushChannelRecord(ctx context.Context, channel *entity.PushChannel) error { if err := GetDB(ctx).Delete(channel).Error; err != nil { return err } @@ -111,14 +81,15 @@ func DeletePushChannelRecord(ctx context.Context, channel *model.PushChannel) er } func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Duration, query func(db *gorm.DB, dest *T) error) (*T, error) { - var val T if cache := GetCache(ctx); cache != nil { + var val T if err := cache.Get(ctx, cacheKey, &val); err == nil { return &val, nil } } db := GetDB(ctx) + var val T if err := query(db, &val); err != nil { return nil, err } @@ -131,8 +102,8 @@ func getCachedOrQuery[T any](ctx context.Context, cacheKey string, ttl time.Dura } // GetActivePushChannelByName loads an enabled push channel, preferring the cache layer. -func GetActivePushChannelByName(ctx context.Context, name string) (*model.PushChannel, error) { - channel, err := getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *model.PushChannel) error { +func GetActivePushChannelByName(ctx context.Context, name string) (*entity.PushChannel, error) { + channel, err := getCachedOrQuery(ctx, "push:channel:active:"+name, activePushChannelCacheTTL, func(db *gorm.DB, dest *entity.PushChannel) error { return db.Where("name = ? AND enabled = ?", name, true).First(dest).Error }) if err != nil { @@ -149,8 +120,8 @@ func DeleteActivePushChannelCache(ctx context.Context, name string) { } // ListPushEventsRecord returns all push events ordered by creation time descending. -func ListPushEventsRecord(ctx context.Context) ([]model.PushEvent, error) { - var events []model.PushEvent +func ListPushEventsRecord(ctx context.Context) ([]entity.PushEvent, error) { + var events []entity.PushEvent if err := GetDB(ctx).Order("created_at DESC").Find(&events).Error; err != nil { return nil, err } @@ -158,19 +129,19 @@ func ListPushEventsRecord(ctx context.Context) ([]model.PushEvent, error) { } // GetPushEventByIDRecord loads a push event by primary key. -func GetPushEventByIDRecord(ctx context.Context, id uint64) (model.PushEvent, error) { - var event model.PushEvent +func GetPushEventByIDRecord(ctx context.Context, id uint64) (entity.PushEvent, error) { + var event entity.PushEvent if err := GetDB(ctx).First(&event, id).Error; err != nil { - return model.PushEvent{}, mapNotFound(err) + return entity.PushEvent{}, mapNotFound(err) } return event, nil } // GetPushEventByKeyRecord loads a push event by event key. -func GetPushEventByKeyRecord(ctx context.Context, key string) (model.PushEvent, error) { - var event model.PushEvent +func GetPushEventByKeyRecord(ctx context.Context, key string) (entity.PushEvent, error) { + var event entity.PushEvent if err := GetDB(ctx).Where("event_key = ?", key).First(&event).Error; err != nil { - return model.PushEvent{}, mapNotFound(err) + return entity.PushEvent{}, mapNotFound(err) } return event, nil } @@ -178,14 +149,14 @@ func GetPushEventByKeyRecord(ctx context.Context, key string) (model.PushEvent, // CountPushEventsByKeyRecord returns how many events use the given event key. func CountPushEventsByKeyRecord(ctx context.Context, key string) (int64, error) { var count int64 - if err := GetDB(ctx).Model(&model.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil { + if err := GetDB(ctx).Model(&entity.PushEvent{}).Where("event_key = ?", key).Count(&count).Error; err != nil { return 0, err } return count, nil } // CreatePushEventRecord persists a new push event and invalidates cache. -func CreatePushEventRecord(ctx context.Context, event *model.PushEvent) error { +func CreatePushEventRecord(ctx context.Context, event *entity.PushEvent) error { if err := GetDB(ctx).Create(event).Error; err != nil { return err } @@ -194,7 +165,7 @@ func CreatePushEventRecord(ctx context.Context, event *model.PushEvent) error { } // SavePushEventRecord updates a push event and invalidates cache. -func SavePushEventRecord(ctx context.Context, event *model.PushEvent) error { +func SavePushEventRecord(ctx context.Context, event *entity.PushEvent) error { if err := GetDB(ctx).Save(event).Error; err != nil { return err } @@ -203,7 +174,7 @@ func SavePushEventRecord(ctx context.Context, event *model.PushEvent) error { } // UpdatePushEventEnabledRecord toggles the enabled flag for a push event. -func UpdatePushEventEnabledRecord(ctx context.Context, event *model.PushEvent, enabled bool) error { +func UpdatePushEventEnabledRecord(ctx context.Context, event *entity.PushEvent, enabled bool) error { event.Enabled = enabled if err := GetDB(ctx).Model(event).Update("enabled", enabled).Error; err != nil { return err @@ -213,7 +184,7 @@ func UpdatePushEventEnabledRecord(ctx context.Context, event *model.PushEvent, e } // DeletePushEventRecord removes a push event and invalidates cache. -func DeletePushEventRecord(ctx context.Context, event *model.PushEvent) error { +func DeletePushEventRecord(ctx context.Context, event *entity.PushEvent) error { if err := GetDB(ctx).Delete(event).Error; err != nil { return err } @@ -222,8 +193,8 @@ func DeletePushEventRecord(ctx context.Context, event *model.PushEvent) error { } // ListActivePushEventsByTaskTypeRecord returns enabled events bound to a task type. -func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]model.PushEvent, error) { - var events []model.PushEvent +func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) ([]entity.PushEvent, error) { + var events []entity.PushEvent if err := GetDB(ctx).Where("task_type = ? AND enabled = ?", taskType, true).Find(&events).Error; err != nil { return nil, err } @@ -231,8 +202,8 @@ func ListActivePushEventsByTaskTypeRecord(ctx context.Context, taskType string) } // GetActivePushEventByKey loads an enabled push event, preferring the cache layer. -func GetActivePushEventByKey(ctx context.Context, key string) (*model.PushEvent, error) { - event, err := getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *model.PushEvent) error { +func GetActivePushEventByKey(ctx context.Context, key string) (*entity.PushEvent, error) { + event, err := getCachedOrQuery(ctx, "push:event:active:"+key, activePushEventCacheTTL, func(db *gorm.DB, dest *entity.PushEvent) error { return db.Where("event_key = ? AND enabled = ?", key, true).First(dest).Error }) if err != nil { @@ -249,11 +220,14 @@ func DeleteActivePushEventCache(ctx context.Context, key string) { } // ListPushHistoriesRecord returns paginated push history records. -func ListPushHistoriesRecord(ctx context.Context, filter model.PushHistoryListFilter) (int64, []model.PushHistory, error) { - query := GetDB(ctx).Model(&model.PushHistory{}).Order("created_at DESC") +func ListPushHistoriesRecord(ctx context.Context, filter do.PushHistoryListFilter) (int64, []entity.PushHistory, error) { + query := GetDB(ctx).Model(&entity.PushHistory{}).Order("created_at DESC") if filter.EventKey != "" { query = query.Where("event_key = ?", filter.EventKey) } + if filter.Channel != "" { + query = query.Where("channel = ?", filter.Channel) + } if filter.Status != "" { query = query.Where("status = ?", filter.Status) } @@ -263,7 +237,7 @@ func ListPushHistoriesRecord(ctx context.Context, filter model.PushHistoryListFi return 0, nil, err } - var results []model.PushHistory + var results []entity.PushHistory offset := (filter.Page - 1) * filter.PageSize if err := query.Offset(offset).Limit(filter.PageSize).Find(&results).Error; err != nil { return 0, nil, err @@ -273,92 +247,21 @@ func ListPushHistoriesRecord(ctx context.Context, filter model.PushHistoryListFi } // CreatePushHistoryRecord persists a push history audit record. -func CreatePushHistoryRecord(ctx context.Context, history *model.PushHistory) error { +func CreatePushHistoryRecord(ctx context.Context, history *entity.PushHistory) error { return GetDB(ctx).Create(history).Error } // PushHistoryQuery returns a scoped query builder for push histories. func PushHistoryQuery(ctx context.Context) *gorm.DB { - return GetDB(ctx).Model(&model.PushHistory{}) + return GetDB(ctx).Model(&entity.PushHistory{}) } -// smtpConfigKeys are the system-config rows backing the built-in email channel. -var smtpConfigKeys = []string{"smtp_host", "smtp_port", "smtp_username", "smtp_password"} - -// LoadSMTPConfigRecord reads the SMTP settings in one query. -// -// A key that is simply absent leaves its field empty, which is how an unconfigured -// mailer is represented. A read that fails is returned as an error, so callers -// cannot mistake an unhealthy database for "no SMTP configured" and silently drop -// the notification. -func LoadSMTPConfigRecord(ctx context.Context) (model.SMTPConfig, error) { +// DeletePushHistoriesBeforeRecord deletes push history records created before cutoff time. +func DeletePushHistoriesBeforeRecord(ctx context.Context, cutoff time.Time) (int64, error) { db := GetDB(ctx) if db == nil { - return model.SMTPConfig{}, errors.New("database not available") + return 0, nil } - - var rows []struct { - Key string - Value string - } - if err := db.Table("w_system_configs"). - Select("key", "value"). - Where("key IN ?", smtpConfigKeys). - Find(&rows).Error; err != nil { - return model.SMTPConfig{}, fmt.Errorf("read smtp system configs: %w", err) - } - - var cfg model.SMTPConfig - for _, row := range rows { - switch row.Key { - case "smtp_host": - cfg.Host = row.Value - case "smtp_port": - cfg.Port = row.Value - case "smtp_username": - cfg.Username = row.Value - case "smtp_password": - cfg.Password = row.Value - } - } - return cfg, nil -} - -// userLookupColumns allow-lists the columns FindUserByFieldRecord may filter on. -// The column name is concatenated into the WHERE clause, so anything not listed -// here must never reach the database. -var userLookupColumns = map[string]struct{}{ - "id": {}, - "username": {}, -} - -// FindUserByFieldRecord is the user lookup fallback for when the UserService -// contract is not wired yet. field must be one of userLookupColumns. -func FindUserByFieldRecord(ctx context.Context, field string, value any) (*contracts.UserDTO, error) { - if _, ok := userLookupColumns[field]; !ok { - return nil, errs.ErrUnsupportedUserLookupField - } - db := GetDB(ctx) - if db == nil { - return nil, errs.ErrRecordNotFound - } - var user contracts.UserDTO - if err := db.Table("w_users").Where(field+" = ?", value).First(&user).Error; err != nil { - return nil, err - } - return &user, nil -} - -// FindFirstAdminUserRecord is the admin lookup fallback for when the UserService -// contract is not wired yet. -func FindFirstAdminUserRecord(ctx context.Context) (*contracts.UserDTO, error) { - db := GetDB(ctx) - if db == nil { - return nil, errs.ErrRecordNotFound - } - var adminUser contracts.UserDTO - if err := db.Table("w_users").Where("is_admin = ?", true).Order("id ASC").First(&adminUser).Error; err != nil { - return nil, err - } - return &adminUser, nil + result := db.Where("created_at < ?", cutoff).Delete(&entity.PushHistory{}) + return result.RowsAffected, result.Error } diff --git a/backend/plugins/domain/msg_gateway/dao/push_test.go b/backend/plugins/domain/msg_gateway/dao/push_test.go new file mode 100644 index 00000000..19712ba5 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/dao/push_test.go @@ -0,0 +1,145 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package dao_test + +import ( + "Wavelet/pkg/idgen" + "Wavelet/pkg/testhelper" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +type stubDBService struct{ db *gorm.DB } + +func (s stubDBService) GORM() *gorm.DB { return s.db } +func (s stubDBService) DB(_ context.Context) *gorm.DB { return s.db } +func (s stubDBService) Named(_ string) *gorm.DB { return s.db } + +func TestPushChannelDAO_CRUD(t *testing.T) { + _ = idgen.Init(1) + db, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{})) + + dao.SetDBServiceForTest(stubDBService{db: db}) + t.Cleanup(func() { dao.SetDBServiceForTest(nil) }) + + ctx := context.Background() + + ch := entity.PushChannel{ + Name: "test_webhook", + Type: "custom", + URL: "https://example.com/hook", + Enabled: true, + } + require.NoError(t, dao.CreatePushChannelRecord(ctx, &ch)) + assert.NotZero(t, ch.ID) + + loaded, err := dao.GetPushChannelByIDRecord(ctx, ch.ID) + require.NoError(t, err) + assert.Equal(t, "test_webhook", loaded.Name) + + active, err := dao.GetActivePushChannelByName(ctx, "test_webhook") + require.NoError(t, err) + assert.Equal(t, ch.ID, active.ID) + + channels, err := dao.ListPushChannelsRecord(ctx) + require.NoError(t, err) + assert.NotEmpty(t, channels) + + require.NoError(t, dao.DeletePushChannelRecord(ctx, &ch)) + _, err = dao.GetPushChannelByIDRecord(ctx, ch.ID) + assert.Error(t, err) +} + +func TestPushEventDAO_CRUD(t *testing.T) { + _ = idgen.Init(1) + db, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + require.NoError(t, db.AutoMigrate(&entity.PushChannel{}, &entity.PushEvent{}, &entity.PushHistory{})) + + dao.SetDBServiceForTest(stubDBService{db: db}) + t.Cleanup(func() { dao.SetDBServiceForTest(nil) }) + + ctx := context.Background() + + ev := entity.PushEvent{ + EventKey: "test_event", + Name: "测试事件", + Channels: []string{"test_webhook"}, + Targets: []string{"admin"}, + Template: `{"title":"Hello"}`, + Enabled: true, + } + require.NoError(t, dao.CreatePushEventRecord(ctx, &ev)) + assert.NotZero(t, ev.ID) + + loaded, err := dao.GetPushEventByKeyRecord(ctx, "test_event") + require.NoError(t, err) + assert.Equal(t, "测试事件", loaded.Name) + + require.NoError(t, dao.UpdatePushEventEnabledRecord(ctx, &ev, false)) + loadedDisabled, err := dao.GetPushEventByIDRecord(ctx, ev.ID) + require.NoError(t, err) + assert.False(t, loadedDisabled.Enabled) + + require.NoError(t, dao.DeletePushEventRecord(ctx, &ev)) + _, err = dao.GetPushEventByIDRecord(ctx, ev.ID) + assert.Error(t, err) +} + +func TestPushHistoryDAO_Cleanup(t *testing.T) { + _ = idgen.Init(1) + db, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + require.NoError(t, db.AutoMigrate(&entity.PushHistory{})) + + dao.SetDBServiceForTest(stubDBService{db: db}) + t.Cleanup(func() { dao.SetDBServiceForTest(nil) }) + + ctx := context.Background() + + now := time.Now() + oldTime := now.Add(-40 * 24 * time.Hour) + recentTime := now.Add(-5 * 24 * time.Hour) + + oldHistory := entity.PushHistory{ + EventKey: "login", + Channel: "telegram", + Target: "123", + Title: "Old login", + Content: "Old content", + Level: "info", + Status: "success", + CreatedAt: oldTime, + } + recentHistory := entity.PushHistory{ + EventKey: "login", + Channel: "telegram", + Target: "123", + Title: "Recent login", + Content: "Recent content", + Level: "info", + Status: "success", + CreatedAt: recentTime, + } + require.NoError(t, db.Create(&oldHistory).Error) + require.NoError(t, db.Create(&recentHistory).Error) + + cutoff := now.Add(-30 * 24 * time.Hour) + deleted, err := dao.DeletePushHistoriesBeforeRecord(ctx, cutoff) + require.NoError(t, err) + assert.Equal(t, int64(1), deleted) + + var count int64 + db.Model(&entity.PushHistory{}).Count(&count) + assert.Equal(t, int64(1), count) +} diff --git a/backend/plugins/domain/message_gateway/migrations/postgres/00001_initial.sql b/backend/plugins/domain/msg_gateway/migrations/postgres/00001_initial.sql similarity index 100% rename from backend/plugins/domain/message_gateway/migrations/postgres/00001_initial.sql rename to backend/plugins/domain/msg_gateway/migrations/postgres/00001_initial.sql diff --git a/backend/plugins/domain/message_gateway/migrations/sqlite/00001_initial.sql b/backend/plugins/domain/msg_gateway/migrations/sqlite/00001_initial.sql similarity index 100% rename from backend/plugins/domain/message_gateway/migrations/sqlite/00001_initial.sql rename to backend/plugins/domain/msg_gateway/migrations/sqlite/00001_initial.sql diff --git a/backend/plugins/domain/msg_gateway/model/do/bot.go b/backend/plugins/domain/msg_gateway/model/do/bot.go new file mode 100644 index 00000000..3315cb46 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/model/do/bot.go @@ -0,0 +1,124 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package do defines domain objects, DTOs, and request/response payloads for msg_gateway. +package do + +import "time" + +// Capability describes what an adapter can send and receive. +type Capability struct { + Text bool + Image bool + File bool + Reply bool + Group bool +} + +// ChannelConfig is the decrypted runtime config passed to a factory. +type ChannelConfig struct { + ID uint64 + Type string + Name string + Credentials map[string]string + Extra map[string]string +} + +// Recipient is the outbound destination on a platform. +type Recipient struct { + ChatID string + PlatformUserID string +} + +// Attachment is a downloaded inbound file sitting on local disk. +type Attachment struct { + Path string + FileName string + MIME string + Error string +} + +// InboundMessage is a normalized private-chat message. +type InboundMessage struct { + ChannelID uint64 + ChannelType string + PlatformUserID string + ChatID string + MessageID string + Text string + Attachments []Attachment + BindingUserID *uint64 +} + +// OutboundMessage is a reply or probe send. +type OutboundMessage struct { + Text string + ReplyToID string + Attachments []Attachment +} + +// BindRequest is the user bind body. +type BindRequest struct { + ChannelID string `json:"channel_id"` + Code string `json:"code"` +} + +// BindingDTO is a user-facing binding row. +type BindingDTO struct { + ID uint64 `json:"id,string"` + UserID uint64 `json:"user_id,string"` + ChannelID uint64 `json:"channel_id,string"` + ChannelName string `json:"channel_name"` + ChannelType string `json:"channel_type"` + PlatformUserID string `json:"platform_user_id"` + CreatedAt time.Time `json:"created_at"` +} + +// PublicChannelDTO is an enabled channel a user can bind to. +type PublicChannelDTO struct { + ID uint64 `json:"id,string"` + Name string `json:"name"` + Type string `json:"type"` +} + +// Field is one admin form field. +type Field struct { + Key string `json:"key"` + Type string `json:"type"` + Required bool `json:"required"` +} + +// Definition describes a channel type form. +type Definition struct { + Type string `json:"type"` + Fields []Field `json:"fields"` +} + +// ChannelDTO represents a channel for admin consumption. +type ChannelDTO struct { + ID uint64 `json:"id,string"` + Name string `json:"name"` + Type string `json:"type"` + OwnerScope string `json:"owner_scope"` + OwnerID *uint64 `json:"owner_id,string,omitempty"` + Enabled bool `json:"enabled"` + Credentials map[string]string `json:"credentials"` + Extra map[string]string `json:"extra"` +} + +// CreateChannelRequest is admin create payload. +type CreateChannelRequest struct { + Name string `json:"name"` + Type string `json:"type"` + Enabled *bool `json:"enabled"` + Credentials map[string]string `json:"credentials"` + Extra map[string]string `json:"extra"` +} + +// UpdateChannelRequest is admin update payload. +type UpdateChannelRequest struct { + Name string `json:"name"` + Enabled *bool `json:"enabled"` + Credentials map[string]string `json:"credentials"` + Extra map[string]string `json:"extra"` +} diff --git a/backend/plugins/domain/message_gateway/model/push.go b/backend/plugins/domain/msg_gateway/model/do/push.go similarity index 58% rename from backend/plugins/domain/message_gateway/model/push.go rename to backend/plugins/domain/msg_gateway/model/do/push.go index b16c9e38..a990a286 100644 --- a/backend/plugins/domain/message_gateway/model/push.go +++ b/backend/plugins/domain/msg_gateway/model/do/push.go @@ -1,39 +1,12 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package model +package do import ( - "sync" + "Wavelet/plugins/domain/msg_gateway/consts" + pkgpush "Wavelet/plugins/domain/msg_gateway/push" "time" - - pkgpush "Wavelet/plugins/domain/message_gateway/push" -) - -// Push channel and payload constants. -const ( - ChannelCustom = "custom" - ChannelEmail = "email" - ChannelLark = "lark" - ChannelTelegram = "telegram" - DefaultLevelInfo = "INFO" - KeyTitle = "title" - KeyContent = "content" - KeyLevel = "level" - - // KeyURL represents the URL field key. - KeyURL = "url" - // KeyToken represents the Token field key. - KeyToken = "token" - // KeyOther represents the Other field key. - KeyOther = "other" - - // TypeText represents standard text input type. - TypeText = "text" - // TypePassword represents password input type. - TypePassword = "password" - // TypeTextarea represents textarea input type. - TypeTextarea = "textarea" ) // SMTPConfig mirrors the system SMTP settings consumed by the push service. @@ -135,12 +108,12 @@ type NotificationMessage struct { Ext map[string]any `json:"ext,omitempty"` } -// Flatten converts the structured NotificationMessage back to a flat map (original json structure). +// Flatten converts the structured NotificationMessage back to a flat map. func (m NotificationMessage) Flatten() map[string]any { res := map[string]any{ - KeyTitle: m.Title, - KeyContent: m.Content, - KeyLevel: m.Level, + consts.KeyTitle: m.Title, + consts.KeyContent: m.Content, + consts.KeyLevel: m.Level, } for k, v := range m.Ext { res[k] = v @@ -156,7 +129,7 @@ type EventMetadata struct { Description string `json:"description"` } -// SendPayload is the async push dispatch载荷 consumed by the notification worker. +// SendPayload is the async push dispatch payload consumed by the notification worker. type SendPayload struct { EventKey string `json:"event_key"` Config pkgpush.Config `json:"config"` @@ -176,138 +149,220 @@ type PushHistoryListFilter struct { PageSize int } -var ( - pushDefMu sync.RWMutex - pushDefinitions = make(map[string]PushDefinition) -) - -// RegisterPushChannelDefinition registers a channel definition. -func RegisterPushChannelDefinition(def PushDefinition) { - pushDefMu.Lock() - defer pushDefMu.Unlock() - pushDefinitions[def.Type] = def +// PushNotificationEvent defines the payload for eventbus notification trigger. +type PushNotificationEvent struct { + UserID uint64 `json:"user_id"` + Channel string `json:"channel"` + Title string `json:"title"` + Content string `json:"content"` + Metadata map[string]any `json:"metadata,omitempty"` } -// ListPushDefinitions returns all registered channel definitions. -func ListPushDefinitions() []PushDefinition { - pushDefMu.RLock() - defer pushDefMu.RUnlock() - - order := []string{ChannelCustom, ChannelLark, ChannelTelegram, ChannelEmail} - res := make([]PushDefinition, 0, len(pushDefinitions)) - for _, t := range order { - if d, ok := pushDefinitions[t]; ok { - res = append(res, d) - } - } - for t, d := range pushDefinitions { - found := false - for _, o := range order { - if o == t { - found = true - break - } - } - if !found { - res = append(res, d) - } - } - return res -} - -func init() { - RegisterPushChannelDefinition(PushDefinition{ - Type: ChannelCustom, +//nolint:goconst,dupl // Static push channel form definitions table +var defaultPushDefinitions = []PushDefinition{ + { + Type: consts.ChannelCustom, Name: "自定义消息通道", Description: "使用自定义 HTTP POST 请求向外部 Webhook 发送数据。", Fields: []PushField{ { - Key: KeyURL, + Key: consts.KeyURL, Label: "请求地址", - Type: TypeText, + Type: consts.TypeText, Required: true, Placeholder: "在此填写完整的请求地址,必须使用 HTTPS 协议", Description: "接口请求的完整 HTTPS URL,例如 https://api.example.com/webhook", }, { - Key: KeyOther, + Key: consts.KeyOther, Label: "请求体 (JSON)", - Type: TypeTextarea, + Type: consts.TypeTextarea, Required: true, Placeholder: "在此输入请求体,支持模板变量,必须为合法的 JSON 格式", Description: "可使用的变量:$title, $description, $content, $url, $to。例如 {\"text\": \"$content\"}", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ - Type: ChannelLark, + }, + { + Type: consts.ChannelLark, Name: "飞书群机器人", Description: "配置飞书群自定义机器人的 Webhook 接口投递。", Fields: []PushField{ { - Key: KeyURL, + Key: consts.KeyURL, Label: "Webhook 地址", - Type: TypeText, + Type: consts.TypeText, Required: true, Placeholder: "https://open.feishu.cn/open-apis/bot/v2/hook/YOUR_TOKEN", Description: "从飞书群机器人设置中复制的 Webhook URL", }, { - Key: KeyToken, + Key: consts.KeyToken, Label: "签名校验密钥 (Secret) (可选)", - Type: TypeText, + Type: consts.TypeText, Required: false, Placeholder: "可选,若机器人启用了安全设置中的签名校验,请在此输入", Description: "飞书群机器人安全设置中的签名校验 Key", }, { - Key: KeyOther, + Key: consts.KeyOther, Label: "自定义卡片 JSON 模版 (可选)", - Type: TypeTextarea, + Type: consts.TypeTextarea, Required: false, Placeholder: "可选,留空则默认使用系统内置的精美互动卡片", Description: "若填写,必须是合法的飞书卡片 JSON 格式", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ - Type: ChannelTelegram, + }, + { + Type: consts.ChannelDingTalk, + Name: "钉钉群机器人", + Description: "配置钉钉群自定义机器人的 Webhook 接口投递。", + Fields: []PushField{ + { + Key: consts.KeyURL, + Label: "Webhook 地址", + Type: consts.TypeText, + Required: true, + Placeholder: "https://oapi.dingtalk.com/robot/send?access_token=YOUR_TOKEN", + Description: "从钉钉群机器人设置中获取的完整 Webhook URL", + }, + { + Key: consts.KeyToken, + Label: "加签密钥 (Secret) (可选)", + Type: consts.TypeText, + Required: false, + Placeholder: "可选,若机器人启用了安全设置中的加签校验,请在此输入 SEC 开头的密钥", + Description: "钉钉群机器人安全设置中的加签 Secret", + }, + }, + }, + { + Type: consts.ChannelTelegram, Name: "Telegram 机器人", Description: "配置 Telegram 机器人推送消息。", Fields: []PushField{ { - Key: KeyURL, + Key: consts.KeyURL, Label: "API 基础地址 (可选)", - Type: TypeText, + Type: consts.TypeText, Required: false, Placeholder: "https://api.telegram.org", Description: "接口请求的 HTTPS 基础地址,留空默认为 https://api.telegram.org", }, { - Key: KeyToken, + Key: consts.KeyToken, Label: "机器人 Token (Bot Token)", - Type: TypePassword, + Type: consts.TypePassword, Required: true, Placeholder: "在此输入 Telegram 机器人的 Bot Token", Description: "通过 BotFather 申请到的机器人 Access Token", }, { - Key: KeyOther, + Key: consts.KeyOther, Label: "默认会话 ID (Chat ID) (可选)", - Type: TypeText, + Type: consts.TypeText, Required: false, Placeholder: "例如 -100123456789 或 @channel_name", Description: "默认的消息接收 Chat ID。如果通知事件中未配置 targets,将推送到此 ID", }, }, - }) - - RegisterPushChannelDefinition(PushDefinition{ - Type: ChannelEmail, + }, + { + Type: consts.ChannelBark, + Name: "Bark (iOS 推送)", + Description: "配置 Bark 推送通知至 iPhone / iPad 客户端。", + Fields: []PushField{ + { + Key: consts.KeyToken, + Label: "设备 Key (Device Key)", + Type: consts.TypeText, + Required: true, + Placeholder: "Bark App 首页显示的 Device Key", + Description: "从 Bark App 复制的设备专属 Key", + }, + { + Key: consts.KeyURL, + Label: "Bark 服务器地址 (可选)", + Type: consts.TypeText, + Required: false, + Placeholder: "https://api.day.app", + Description: "Bark 服务器地址,留空默认使用官方公共服务器 https://api.day.app", + }, + { + Key: consts.KeyOther, + Label: "额外配置 JSON (可选)", + Type: consts.TypeTextarea, + Required: false, + Placeholder: "{\"group\": \"Wavelet\", \"sound\": \"minuet\", \"icon\": \"https://...\"}", + Description: "可选的 JSON 配置,支持 group (分组)、sound (铃声)、icon (自定义图标)", + }, + }, + }, + { + Type: consts.ChannelDiscord, + Name: "Discord 频道", + Description: "配置 Discord 频道的 Incoming Webhook 消息推送。", + Fields: []PushField{ + { + Key: consts.KeyURL, + Label: "Webhook 地址", + Type: consts.TypeText, + Required: true, + Placeholder: "https://discord.com/api/webhooks/...", + Description: "从 Discord 频道集成设置中复制的 Webhook URL", + }, + }, + }, + { + Type: consts.ChannelSlack, + Name: "Slack 频道", + Description: "配置 Slack 频道的 Incoming Webhook 消息推送。", + Fields: []PushField{ + { + Key: consts.KeyURL, + Label: "Webhook 地址", + Type: consts.TypeText, + Required: true, + Placeholder: "https://hooks.slack.com/services/...", + Description: "从 Slack 应用配置中复制的 Incoming Webhook URL", + }, + }, + }, + { + Type: consts.ChannelPushover, + Name: "Pushover 推送", + Description: "配置 Pushover 即时推送到手机/桌面客户端。", + Fields: []PushField{ + { + Key: consts.KeyToken, + Label: "应用 Token (App Token)", + Type: consts.TypePassword, + Required: true, + Placeholder: "Pushover 创建应用生成的 API Token / Key", + Description: "从 Pushover 控制台创建的 Application API Token", + }, + { + Key: consts.KeyURL, + Label: "用户 Key (User Key)", + Type: consts.TypeText, + Required: true, + Placeholder: "Pushover 账号主页的 User Key", + Description: "Pushover 个人账号的 User Key", + }, + }, + }, + { + Type: consts.ChannelEmail, Name: "邮件推送通道", Description: "邮件推送通道直接使用系统全局 SMTP 设置进行发送,无需在此填写服务器配置。", Fields: []PushField{}, - }) + }, +} + +// ListPushDefinitions returns all registered channel definitions. +func ListPushDefinitions() []PushDefinition { + res := make([]PushDefinition, len(defaultPushDefinitions)) + copy(res, defaultPushDefinitions) + return res } diff --git a/backend/plugins/domain/msg_gateway/model/entity/bot.go b/backend/plugins/domain/msg_gateway/model/entity/bot.go new file mode 100644 index 00000000..0ab41d1f --- /dev/null +++ b/backend/plugins/domain/msg_gateway/model/entity/bot.go @@ -0,0 +1,56 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package entity defines GORM table mapping entities for msg_gateway. +package entity + +import "time" + +// MessageChannel is an admin-configured messaging adapter entity. +type MessageChannel struct { + ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` + Type string `json:"type" gorm:"size:32;not null"` + Name string `json:"name" gorm:"size:64;not null"` + OwnerScope string `json:"owner_scope" gorm:"size:32;not null;default:'system'"` + OwnerID *uint64 `json:"owner_id,omitempty"` + Credentials string `json:"credentials" gorm:"type:text;not null"` + Extra string `json:"extra" gorm:"type:text"` + Enabled bool `json:"enabled" gorm:"default:false;not null"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"` +} + +// TableName 表名 +func (MessageChannel) TableName() string { + return "w_message_channels" +} + +// MessageBinding maps a platform user to a Wavelet user on one channel. +type MessageBinding struct { + ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` + ChannelID uint64 `json:"channel_id" gorm:"not null;index"` + PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"` + UserID uint64 `json:"user_id" gorm:"not null;index"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName 表名 +func (MessageBinding) TableName() string { + return "w_message_bindings" +} + +// MessagePairingCode is a one-time bind code. +type MessagePairingCode struct { + ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` + Code string `json:"code" gorm:"size:32;uniqueIndex;not null"` + ChannelID uint64 `json:"channel_id" gorm:"not null;index"` + PlatformUserID string `json:"platform_user_id" gorm:"size:128;not null;index"` + UserID uint64 `json:"user_id" gorm:"not null;index"` + ExpiresAt time.Time `json:"expires_at" gorm:"not null;index"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"` +} + +// TableName 表名 +func (MessagePairingCode) TableName() string { + return "w_message_pairing_codes" +} diff --git a/backend/plugins/domain/msg_gateway/model/entity/push.go b/backend/plugins/domain/msg_gateway/model/entity/push.go new file mode 100644 index 00000000..8fb478dd --- /dev/null +++ b/backend/plugins/domain/msg_gateway/model/entity/push.go @@ -0,0 +1,95 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package entity + +import ( + "Wavelet/plugins/domain/msg_gateway/consts" + "errors" + "strings" + "time" +) + +// PushChannel 消息通道实体 +type PushChannel struct { + ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` + Name string `json:"name" gorm:"size:100;not null"` + Description string `json:"description" gorm:"size:255"` + Type string `json:"type" gorm:"size:50;not null;index"` + URL string `json:"url" gorm:"type:text"` + Token string `json:"token" gorm:"type:text"` + Other string `json:"other" gorm:"type:text"` + Enabled bool `json:"enabled" gorm:"index;not null;default:true"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"` +} + +// TableName 指定 GORM 表名 +func (PushChannel) TableName() string { + return "w_push_channels" +} + +// Validate 验证与标准化字段 +func (c *PushChannel) Validate() error { + c.Name = strings.TrimSpace(c.Name) + if c.Name == "" { + return errors.New(consts.ErrChannelNameRequired) + } + c.Type = strings.TrimSpace(c.Type) + if c.Type == "" { + return errors.New(consts.ErrChannelTypeRequired) + } + return nil +} + +// PushEvent 系统通知事件实体 +type PushEvent struct { + ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` + EventKey string `json:"event_key" gorm:"uniqueIndex;size:80;not null"` + Name string `json:"name" gorm:"size:100;not null"` + TaskType string `json:"task_type" gorm:"size:100;index;not null;default:''"` + Channels []string `json:"channels" gorm:"type:text;serializer:json;not null"` + Targets []string `json:"targets" gorm:"type:text;serializer:json;not null"` + Template string `json:"template" gorm:"type:text;not null"` + Enabled bool `json:"enabled" gorm:"index;not null;default:false"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"` + UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime;index"` +} + +// TableName 指定 GORM 表名 +func (PushEvent) TableName() string { + return "w_push_events" +} + +// Validate 验证 PushEvent 实体字段 +func (e *PushEvent) Validate() error { + e.EventKey = strings.TrimSpace(e.EventKey) + if e.EventKey == "" { + return errors.New(consts.ErrEventKeyRequired) + } + e.Name = strings.TrimSpace(e.Name) + if e.Name == "" { + return errors.New(consts.ErrNameRequired) + } + return nil +} + +// PushHistory 推送日志/历史实体 +type PushHistory struct { + ID uint64 `json:"id" gorm:"primaryKey;autoIncrement"` + EventKey string `json:"event_key" gorm:"size:80;not null;index"` + Channel string `json:"channel" gorm:"size:50;not null;index"` + Target string `json:"target" gorm:"size:255;not null"` + Title string `json:"title" gorm:"size:255;not null"` + Content string `json:"content" gorm:"type:text;not null"` + Level string `json:"level" gorm:"size:20;not null;default:'INFO'"` + Status string `json:"status" gorm:"size:20;not null;index"` + ErrorMsg string `json:"error_msg" gorm:"type:text"` + Payload string `json:"payload" gorm:"type:text"` + CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime;index"` +} + +// TableName 指定 GORM 表名 +func (PushHistory) TableName() string { + return "w_push_histories" +} diff --git a/backend/plugins/domain/msg_gateway/plugin.go b/backend/plugins/domain/msg_gateway/plugin.go new file mode 100644 index 00000000..85e1d58b --- /dev/null +++ b/backend/plugins/domain/msg_gateway/plugin.go @@ -0,0 +1,270 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package msg_gateway provides the Bot gateway, multi-channel notification dispatching, and asynchronous push worker domain plugin for Cordis. +package msg_gateway + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "Wavelet/core/extpoints" + "Wavelet/pkg/ginutil" + "Wavelet/pkg/util" + "Wavelet/plugins/domain/msg_gateway/channels/qq" + "Wavelet/plugins/domain/msg_gateway/channels/telegram" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/controller" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "Wavelet/plugins/domain/msg_gateway/service" + "context" + "embed" + "reflect" + "time" + + "github.com/gin-gonic/gin" +) + +//go:embed migrations/*/*.sql +var mgMigrations embed.FS + +// Option configures the msg_gateway plugin. +type Option func(*Plugin) + +// WithAutoStartRunner enables automatic bot runner startup in the background. +func WithAutoStartRunner(enable bool) Option { + return func(p *Plugin) { + p.autoStartRunner = enable + } +} + +// Plugin implements core.Plugin to provide Bot gateway and notification dispatch domain services. +type Plugin struct { + autoStartRunner bool + cancelRunner context.CancelFunc +} + +// New creates a new msg_gateway domain plugin. +func New(opts ...Option) *Plugin { + p := &Plugin{} + for _, opt := range opts { + if opt != nil { + opt(p) + } + } + return p +} + +// Name returns the unique identifier for the msg_gateway domain plugin. +func (p *Plugin) Name() string { + return "msg_gateway" +} + +// Inject declares required dependencies for the msg_gateway domain plugin. +func (p *Plugin) Inject() []reflect.Type { + return []reflect.Type{ + reflect.TypeFor[contracts.DBService](), + reflect.TypeFor[contracts.AuthService](), + } +} + +// Manifest returns the plugin metadata. +func (p *Plugin) Manifest() core.Manifest { + return core.Manifest{ + Name: "msg_gateway", + Version: "1.0.0", + Description: "Bot gateway, multi-channel notification push, and async worker dispatch plugin", + Author: "Wavelet Team", + } +} + +type mgAppConfig struct { + SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"` +} + +// DeclareConfig declares configuration bindings for the msg_gateway plugin. +func (p *Plugin) DeclareConfig() []core.ConfigBinding { + return []core.ConfigBinding{ + {Prefix: "app", Target: &mgAppConfig{}}, + } +} + +// Apply registers msg_gateway migrations, routes, tasks, schedules, events, and settings into the Context. +func (p *Plugin) Apply(ctx *core.Context) error { + var cfg mgAppConfig + if err := ctx.Config().Bind("app", &cfg); err == nil && cfg.SessionSecret != "" { + service.SetCredentialSecret(cfg.SessionSecret) + } + core.Bind[contracts.DBService](ctx, dao.SetDBService) + core.Bind[contracts.CacheService](ctx, func(cache contracts.CacheService) { + dao.SetCacheService(cache) + service.SetCacheService(cache) + }) + core.Bind[contracts.TaskService](ctx, service.SetTaskService) + core.Bind[contracts.UserService](ctx, service.SetUserService) + ctx.OnDispose(func() error { + dao.SetDBService(nil) + dao.SetCacheService(nil) + service.SetCacheService(nil) + service.SetTaskService(nil) + service.SetUserService(nil) + return nil + }) + + // 0. Resolve auth service for middleware (via IoC, not direct import) + denyAuth := ginutil.AuthUnavailable() + loginMW := denyAuth + adminMW := denyAuth + if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil { + if mw, ok := authSvc.RequireAuthMiddleware().(gin.HandlerFunc); ok { + loginMW = mw + } + if mw, ok := authSvc.RequireAdminMiddleware().(gin.HandlerFunc); ok { + adminMW = mw + } + } + + // 1. Register migrations + ctx.Migrations().Register("msg_gateway", mgMigrations) + + // 2. Register User HTTP Routes + controller.RegisterUserRoutes(ctx.Router().Group("/api/v1"), loginMW) + + // 3. Register Admin Message Gateway HTTP Routes + controller.RegisterAdminRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW) + + // 4. Register Admin Push HTTP Routes + controller.RegisterAdminPushRoutes(ctx.Router().Group("/api/v1/admin"), loginMW, adminMW) + + service.Register(consts.MessageChannelTypeTelegram, telegram.New) + service.Register(consts.MessageChannelTypeQQ, qq.New) + + const defaultTaskRetry = 3 + pushHandler := &service.PushHandler{} + + // 5. Register background tasks + ctx.Task().Register(consts.TaskPushNotification, func(c context.Context, payload []byte) error { + return pushHandler.Execute(c, payload) + }, + extpoints.WithTaskType("push_notification"), + extpoints.WithTaskName("消息网关推送通知"), + extpoints.WithTaskDescription("异步执行系统通知的多渠道派发与推送"), + extpoints.WithTaskCategory("push"), + extpoints.WithTaskRetry(defaultTaskRetry), + extpoints.WithTaskQueue("default"), + extpoints.WithTaskRetryable(true), + ) + + ctx.Task().Register(service.SendNotificationTask, func(c context.Context, payload []byte) error { + return pushHandler.Execute(c, payload) + }, extpoints.WithTaskMeta(service.SendNotificationMeta), extpoints.WithTaskRetry(defaultTaskRetry)) + + ctx.Task().Register(service.TaskDispatchBotMsg, &service.BotDispatchHandler{}, + extpoints.WithTaskMeta(service.BotDispatchMeta)) + + ctx.Task().Register(consts.TaskCleanupPairingCodes, func(c context.Context, _ []byte) error { + return dao.DeleteExpiredPairingCodes(c) + }, + extpoints.WithTaskType("cleanup_pairing_codes"), + extpoints.WithTaskName("清理过期配对码"), + extpoints.WithTaskDescription("定时清理已过期的平台 Bot 配对码"), + extpoints.WithTaskCategory("messaging"), + extpoints.WithTaskRetry(defaultTaskRetry), + extpoints.WithTaskQueue("default"), + extpoints.WithTaskRetryable(true), + ) + + // 6. Register Cron Schedules + ctx.Schedule().RegisterCron("*/10 * * * *", consts.TaskCleanupPairingCodes, map[string]any{"action": "cleanup"}) + + // 7. Register EventBus listeners for decoupled push triggers + ctx.Events().On("notification:push", func(c context.Context, e do.PushNotificationEvent) error { + meta := do.EventMetadata{ + Key: "eventbus:" + e.Channel, + Name: e.Title, + DefaultTemplate: do.NotificationMessage{ + Title: e.Title, + Content: e.Content, + Level: consts.DefaultLevelInfo, + Ext: e.Metadata, + }, + Description: "EventBus triggered notification", + } + service.DefaultTrigger.Trigger(c, meta, map[string]any{ + "user.id": e.UserID, + "title": e.Title, + "content": e.Content, + }) + return nil + }) + + // 8. Register task completed event listener + ctx.Events().On(contracts.EventTopicTaskCompleted, func(c context.Context, e contracts.TaskCompletedEvent) error { + service.HandleTaskCompleted(c, e) + return nil + }) + + // 8.1 Register system cleanup event listener + ctx.Events().On(contracts.EventTopicSystemCleanup, func(c context.Context, _ contracts.SystemCleanupEvent) error { + const defaultPushHistoryRetention = 30 * 24 * time.Hour + _, err := service.CleanupPushHistories(c, defaultPushHistoryRetention) + return err + }) + + // 9. Register built-in domain events and provide PushRegistry + service.RegisterCustomEvents() + core.Provide[contracts.PushRegistry](ctx, service.PushRegistryAdapter{}) + + // 10. Register Settings Schemas + ctx.Settings().Register(extpoints.SettingSchema{ + Key: "msg_gateway.pairing_code_expiry_minutes", + Default: 15, + Description: "Expiry duration for bot pairing codes in minutes", + Type: "integer", + Category: "messaging", + }) + ctx.Settings().Register(extpoints.SettingSchema{ + Key: "msg_gateway.max_bindings_per_user", + Default: 5, + Description: "Maximum platform bot bindings per user", + Type: "integer", + Category: "messaging", + }) + + // 11. Optional runner start & lifecycle + if p.autoStartRunner { + runnerCtx, cancel := context.WithCancel(ctx.GoContext()) + p.cancelRunner = cancel + util.Go(func() { + _ = service.Start(runnerCtx) + }) + } + + ctx.OnDispose(func() error { + if p.cancelRunner != nil { + p.cancelRunner() + } + return nil + }) + + return nil +} + +// Entity and DO aliases exported for integration test compatibility. +type ( + // MessageChannel is an alias for entity.MessageChannel. + MessageChannel = entity.MessageChannel + // MessageBinding is an alias for entity.MessageBinding. + MessageBinding = entity.MessageBinding + // MessagePairingCode is an alias for entity.MessagePairingCode. + MessagePairingCode = entity.MessagePairingCode + // PushChannel is an alias for entity.PushChannel. + PushChannel = entity.PushChannel + // PushEvent is an alias for entity.PushEvent. + PushEvent = entity.PushEvent + // PushHistory is an alias for entity.PushHistory. + PushHistory = entity.PushHistory + // PushNotificationEvent is an alias for do.PushNotificationEvent. + PushNotificationEvent = do.PushNotificationEvent +) diff --git a/backend/plugins/domain/message_gateway/plugin_registry_test.go b/backend/plugins/domain/msg_gateway/plugin_registry_test.go similarity index 89% rename from backend/plugins/domain/message_gateway/plugin_registry_test.go rename to backend/plugins/domain/msg_gateway/plugin_registry_test.go index 57cd9205..67a8cfe5 100644 --- a/backend/plugins/domain/message_gateway/plugin_registry_test.go +++ b/backend/plugins/domain/msg_gateway/plugin_registry_test.go @@ -1,13 +1,13 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package message_gateway_test +package msg_gateway_test import ( "Wavelet/core" "Wavelet/core/contracts" - "Wavelet/plugins/domain/message_gateway" - "Wavelet/plugins/domain/message_gateway/service" + "Wavelet/plugins/domain/msg_gateway" + "Wavelet/plugins/domain/msg_gateway/service" "context" "testing" @@ -17,7 +17,7 @@ import ( func TestPushRegistry(t *testing.T) { ctx := core.NewContext(context.Background()) - require.NoError(t, message_gateway.New().Apply(ctx)) + require.NoError(t, msg_gateway.New().Apply(ctx)) registry, err := core.Inject[contracts.PushRegistry](ctx) require.NoError(t, err) diff --git a/backend/plugins/domain/message_gateway/plugin_test.go b/backend/plugins/domain/msg_gateway/plugin_test.go similarity index 71% rename from backend/plugins/domain/message_gateway/plugin_test.go rename to backend/plugins/domain/msg_gateway/plugin_test.go index 3a6f8ef2..14bbc5f9 100644 --- a/backend/plugins/domain/message_gateway/plugin_test.go +++ b/backend/plugins/domain/msg_gateway/plugin_test.go @@ -1,11 +1,11 @@ // Copyright 2026 Arctel.net // SPDX-License-Identifier: Apache-2.0 -package message_gateway_test +package msg_gateway_test import ( "Wavelet/core" - "Wavelet/plugins/domain/message_gateway" + "Wavelet/plugins/domain/msg_gateway" "context" "io/fs" "testing" @@ -14,32 +14,32 @@ import ( "github.com/stretchr/testify/require" ) -func TestMessageGatewayPluginUnit(t *testing.T) { +func TestMsgGatewayPluginUnit(t *testing.T) { ctx := core.NewContext(context.Background()) - p := message_gateway.New() - assert.Equal(t, "message_gateway", p.Name()) + p := msg_gateway.New() + assert.Equal(t, "msg_gateway", p.Name()) assert.Equal(t, "1.0.0", p.Manifest().Version) require.NoError(t, p.Apply(ctx)) // Verify migrations - entry, ok := ctx.Migrations().Get("message_gateway") + entry, ok := ctx.Migrations().Get("msg_gateway") require.True(t, ok) entries, err := fs.ReadDir(entry.FS, entry.Dir) require.NoError(t, err) assert.NotEmpty(t, entries) // Verify tasks - task, ok := ctx.Tasks().Get("message_gateway:push_notification") + task, ok := ctx.Tasks().Get("msg_gateway:push_notification") require.True(t, ok) assert.Equal(t, 3, task.Retry) // Verify schedules - sched, ok := ctx.Schedules().Get("message_gateway:cleanup_pairing_codes") + sched, ok := ctx.Schedules().Get("msg_gateway:cleanup_pairing_codes") require.True(t, ok) assert.Equal(t, "*/10 * * * *", sched.Spec) // Verify settings - setting, ok := ctx.Settings().Get("message_gateway.max_bindings_per_user") + setting, ok := ctx.Settings().Get("msg_gateway.max_bindings_per_user") require.True(t, ok) assert.Equal(t, 5, setting.Default) } @@ -48,7 +48,7 @@ func TestMessageGatewayPluginUnit(t *testing.T) { // Register,则每次触发都投递到无人处理的任务类型,清理逻辑静默失效。 func TestEveryScheduleHasTaskHandler(t *testing.T) { ctx := core.NewContext(context.Background()) - require.NoError(t, message_gateway.New().Apply(ctx)) + require.NoError(t, msg_gateway.New().Apply(ctx)) schedules := ctx.Schedules().Schedules() require.NotEmpty(t, schedules) diff --git a/backend/plugins/domain/msg_gateway/push/bark.go b/backend/plugins/domain/msg_gateway/push/bark.go new file mode 100644 index 00000000..10048daa --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/bark.go @@ -0,0 +1,68 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "errors" + "fmt" + "strings" + + "github.com/nikoksr/notify" + "github.com/nikoksr/notify/service/bark" +) + +func init() { + Register("bark", &BarkPusher{}) +} + +// BarkPusher 基于 nikoksr/notify 的 Bark iOS 客户端通知推送实现 +type BarkPusher struct{} + +// Send 发送 Bark 通知 +func (p *BarkPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) { + deviceKey := cfg.Key + if deviceKey == "" { + deviceKey = cfg.Secret + } + if target != "" { + deviceKey = target + } + if deviceKey == "" { + return "", errors.New("bark: device key is required") + } + + serverURL := strings.TrimRight(cfg.URL, "/") + if serverURL == "" { + serverURL = bark.DefaultServerURL + } + + title := bodyTitle(body) + content := bodyContent(body, "%s: %v", "\n") + + barkService := bark.NewWithServers(deviceKey, serverURL) + notifier := notify.New() + notifier.UseServices(barkService) + + if err := notifier.Send(ctx, title, content); err != nil { + return "", fmt.Errorf("bark: notify send failed: %w", err) + } + + return "ok", nil +} + +// ValidateConfig 校验 Bark 配置 +func (p *BarkPusher) ValidateConfig(cfg Config) error { + deviceKey := cfg.Key + if deviceKey == "" { + deviceKey = cfg.Secret + } + if deviceKey == "" { + return errors.New("device key is required") + } + if cfg.URL != "" && !strings.HasPrefix(cfg.URL, "http://") && !strings.HasPrefix(cfg.URL, "https://") { + return errors.New("server URL must start with http:// or https://") + } + return nil +} diff --git a/backend/plugins/domain/msg_gateway/push/channels_test.go b/backend/plugins/domain/msg_gateway/push/channels_test.go new file mode 100644 index 00000000..d85d4684 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/channels_test.go @@ -0,0 +1,125 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "net/http" + "net/http/httptest" + "testing" +) + +func TestDingTalkPusher(t *testing.T) { + pusher, err := GetPusher("dingtalk") + if err != nil { + t.Fatalf("failed to get dingtalk pusher: %v", err) + } + + err = pusher.ValidateConfig(Config{URL: "https://oapi.dingtalk.com/robot/send?access_token=test"}) + if err != nil { + t.Errorf("ValidateConfig failed: %v", err) + } + + err = pusher.ValidateConfig(Config{}) + if err == nil { + t.Errorf("expected error for empty config, got nil") + } +} + +func TestBarkPusher(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"code":200,"message":"success"}`)) + })) + defer server.Close() + + pusher, err := GetPusher("bark") + if err != nil { + t.Fatalf("failed to get bark pusher: %v", err) + } + + err = pusher.ValidateConfig(Config{Key: "device_key_123"}) + if err != nil { + t.Errorf("ValidateConfig failed: %v", err) + } + + _, err = pusher.Send(context.Background(), Config{ + URL: server.URL, + Key: "device_key_123", + }, "", map[string]any{ + "title": "Alert", + "content": "Bark notification", + }, "", nil) + if err != nil { + t.Errorf("Send failed: %v", err) + } +} + +func TestDiscordPusher(t *testing.T) { + pusher, err := GetPusher("discord") + if err != nil { + t.Fatalf("failed to get discord pusher: %v", err) + } + + err = pusher.ValidateConfig(Config{Key: "bot_token_123"}) + if err != nil { + t.Errorf("ValidateConfig failed: %v", err) + } + + err = pusher.ValidateConfig(Config{}) + if err == nil { + t.Errorf("expected error for empty config, got nil") + } +} + +func TestSlackPusher(t *testing.T) { + pusher, err := GetPusher("slack") + if err != nil { + t.Fatalf("failed to get slack pusher: %v", err) + } + + err = pusher.ValidateConfig(Config{Key: "xoxb-123456"}) + if err != nil { + t.Errorf("ValidateConfig failed: %v", err) + } + + err = pusher.ValidateConfig(Config{}) + if err == nil { + t.Errorf("expected error for empty config, got nil") + } +} + +func TestPushoverPusher(t *testing.T) { + pusher, err := GetPusher("pushover") + if err != nil { + t.Fatalf("failed to get pushover pusher: %v", err) + } + + err = pusher.ValidateConfig(Config{Key: "app_token_123"}) + if err != nil { + t.Errorf("ValidateConfig failed: %v", err) + } + + err = pusher.ValidateConfig(Config{}) + if err == nil { + t.Errorf("expected error for empty config, got nil") + } +} + +func TestLarkPusher(t *testing.T) { + pusher, err := GetPusher("lark") + if err != nil { + t.Fatalf("failed to get lark pusher: %v", err) + } + + err = pusher.ValidateConfig(Config{URL: "https://open.feishu.cn/open-apis/bot/v2/hook/xxx"}) + if err != nil { + t.Errorf("ValidateConfig failed: %v", err) + } + + err = pusher.ValidateConfig(Config{}) + if err == nil { + t.Errorf("expected error for empty config, got nil") + } +} diff --git a/backend/plugins/domain/message_gateway/push/custom.go b/backend/plugins/domain/msg_gateway/push/custom.go similarity index 100% rename from backend/plugins/domain/message_gateway/push/custom.go rename to backend/plugins/domain/msg_gateway/push/custom.go diff --git a/backend/plugins/domain/message_gateway/push/custom_test.go b/backend/plugins/domain/msg_gateway/push/custom_test.go similarity index 100% rename from backend/plugins/domain/message_gateway/push/custom_test.go rename to backend/plugins/domain/msg_gateway/push/custom_test.go diff --git a/backend/plugins/domain/msg_gateway/push/dingtalk.go b/backend/plugins/domain/msg_gateway/push/dingtalk.go new file mode 100644 index 00000000..39d4c37b --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/dingtalk.go @@ -0,0 +1,66 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "errors" + "fmt" + "net/url" + "strings" + + "github.com/nikoksr/notify" + "github.com/nikoksr/notify/service/dingding" +) + +func init() { + Register("dingtalk", &DingTalkPusher{}) +} + +// DingTalkPusher 基于 nikoksr/notify 的钉钉机器人推送实现 +type DingTalkPusher struct{} + +// Send 发送钉钉通知 +func (p *DingTalkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, _ string, _ map[string]any) (string, error) { + token := cfg.Key + if token == "" { + if u, err := url.Parse(cfg.URL); err == nil { + token = u.Query().Get("access_token") + } + } + if token == "" { + token = cfg.URL + } + if token == "" { + return "", errors.New("dingtalk: access token or webhook URL is required") + } + + title := bodyTitle(body) + content := bodyContent(body, "**%s**: %v", "\n\n") + + dingService := dingding.New(&dingding.Config{ + Token: token, + Secret: cfg.Secret, + }) + + notifier := notify.New() + notifier.UseServices(dingService) + + if err := notifier.Send(ctx, title, content); err != nil { + return "", fmt.Errorf("dingtalk: notify send failed: %w", err) + } + + return "ok", nil +} + +// ValidateConfig 校验钉钉配置 +func (p *DingTalkPusher) ValidateConfig(cfg Config) error { + if cfg.URL == "" && cfg.Key == "" { + return errors.New("webhook URL or access token is required") + } + if cfg.URL != "" && !strings.HasPrefix(cfg.URL, "https://") { + return errors.New("webhook URL must use https:// protocol") + } + return nil +} diff --git a/backend/plugins/domain/msg_gateway/push/discord.go b/backend/plugins/domain/msg_gateway/push/discord.go new file mode 100644 index 00000000..a18d2635 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/discord.go @@ -0,0 +1,72 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "errors" + "fmt" + + "github.com/nikoksr/notify" + "github.com/nikoksr/notify/service/discord" +) + +func init() { + Register("discord", &DiscordPusher{}) +} + +// DiscordPusher 基于 nikoksr/notify 的 Discord 推送实现 +type DiscordPusher struct{} + +// Send 发送 Discord 通知 +func (p *DiscordPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) { + botToken := cfg.Key + if botToken == "" { + botToken = cfg.Secret + } + channelID := cfg.URL + if target != "" { + channelID = target + } + if channelID == "" { + channelID = cfg.Other + } + + if botToken == "" { + return "", errors.New("discord: bot token is required") + } + if channelID == "" { + return "", errors.New("discord: channel ID is required") + } + + title := bodyTitle(body) + content := bodyContent(body, "**%s**: %v", "\n") + + discordService := discord.New() + if err := discordService.AuthenticateWithBotToken(botToken); err != nil { + return "", fmt.Errorf("discord: auth failed: %w", err) + } + discordService.AddReceivers(channelID) + + notifier := notify.New() + notifier.UseServices(discordService) + + if err := notifier.Send(ctx, title, content); err != nil { + return "", fmt.Errorf("discord: notify send failed: %w", err) + } + + return "ok", nil +} + +// ValidateConfig 校验 Discord 配置 +func (p *DiscordPusher) ValidateConfig(cfg Config) error { + botToken := cfg.Key + if botToken == "" { + botToken = cfg.Secret + } + if botToken == "" { + return errors.New("bot token is required") + } + return nil +} diff --git a/backend/plugins/domain/msg_gateway/push/email.go b/backend/plugins/domain/msg_gateway/push/email.go new file mode 100644 index 00000000..6c672ae6 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/email.go @@ -0,0 +1,78 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + pkgmail "Wavelet/pkg/mail" + "context" + "errors" + "fmt" + "net" + "strconv" +) + +func init() { + Register("email", &EmailPusher{}) +} + +// EmailPusher 基于 pkg/mail 的 SMTP 邮件推送实现 +type EmailPusher struct{} + +// Send 发送邮件 +func (p *EmailPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, ext map[string]any) (string, error) { + if cfg.URL == "" || cfg.Key == "" || cfg.Secret == "" { + return "", errors.New("email: SMTP configuration (url, key, secret) is incomplete") + } + if target == "" { + return "", errors.New("email: target email address is required") + } + + title := bodyTitle(body) + content := bodyContent(body, "

%s: %v

", "") + + fromName := "System Notification" + if ext != nil { + if fn, ok := ext["from_name"].(string); ok && fn != "" { + fromName = fn + } + } + + htmlBody := fmt.Sprintf(`

%s

%s
`, title, content) + + host, portStr, err := net.SplitHostPort(cfg.URL) + port := 25 + if err != nil { + host = cfg.URL + } else if p, err := strconv.Atoi(portStr); err == nil && p > 0 { + port = p + } + + mailCfg := pkgmail.Config{ + Host: host, + Port: port, + Username: cfg.Key, + Password: cfg.Secret, + FromName: fromName, + } + + if err := pkgmail.SendMail(ctx, mailCfg, target, title, htmlBody); err != nil { + return "", fmt.Errorf("email: send smtp mail failed: %w", err) + } + + return "", nil +} + +// ValidateConfig 校验邮件 SMTP 配置 +func (p *EmailPusher) ValidateConfig(cfg Config) error { + if cfg.URL == "" { + return errors.New("SMTP host:port is required") + } + if cfg.Key == "" { + return errors.New("SMTP username is required") + } + if cfg.Secret == "" { + return errors.New("SMTP password is required") + } + return nil +} diff --git a/backend/plugins/domain/msg_gateway/push/email_test.go b/backend/plugins/domain/msg_gateway/push/email_test.go new file mode 100644 index 00000000..3d97aac4 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/email_test.go @@ -0,0 +1,65 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "testing" +) + +func TestEmailPusherValidateConfig(t *testing.T) { + pusher := &EmailPusher{} + + tests := []struct { + name string + cfg Config + wantErr bool + }{ + { + name: "empty url", + cfg: Config{URL: "", Key: "user", Secret: "pass"}, + wantErr: true, + }, + { + name: "empty key", + cfg: Config{URL: "smtp.example.com:587", Key: "", Secret: "pass"}, + wantErr: true, + }, + { + name: "empty secret", + cfg: Config{URL: "smtp.example.com:587", Key: "user", Secret: ""}, + wantErr: true, + }, + { + name: "valid config", + cfg: Config{URL: "smtp.example.com:587", Key: "user", Secret: "pass"}, + wantErr: false, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + err := pusher.ValidateConfig(tt.cfg) + if (err != nil) != tt.wantErr { + t.Errorf("ValidateConfig() error = %v, wantErr %v", err, tt.wantErr) + } + }) + } +} + +func TestEmailPusherSendValidation(t *testing.T) { + pusher := &EmailPusher{} + + // Missing target + _, err := pusher.Send(context.Background(), Config{URL: "127.0.0.1:25", Key: "u", Secret: "p"}, "", map[string]any{"title": "hi"}, "", nil) + if err == nil { + t.Errorf("expected error for empty target, got nil") + } + + // Missing config + _, err = pusher.Send(context.Background(), Config{}, "test@example.com", map[string]any{"title": "hi"}, "", nil) + if err == nil { + t.Errorf("expected error for empty config, got nil") + } +} diff --git a/backend/plugins/domain/msg_gateway/push/lark.go b/backend/plugins/domain/msg_gateway/push/lark.go new file mode 100644 index 00000000..c476eb4c --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/lark.go @@ -0,0 +1,52 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "errors" + "fmt" + "strings" + + "github.com/nikoksr/notify" + "github.com/nikoksr/notify/service/lark" +) + +func init() { + Register("lark", &LarkPusher{}) +} + +// LarkPusher 基于 nikoksr/notify 的飞书 Webhook 机器人推送实现 +type LarkPusher struct{} + +// Send 发送飞书通知 +func (p *LarkPusher) Send(ctx context.Context, cfg Config, _ string, body map[string]any, _ string, _ map[string]any) (string, error) { + if cfg.URL == "" { + return "", errors.New("lark: webhook URL is required") + } + + title := bodyTitle(body) + content := bodyContent(body, "**%s**: %v", "\n") + + larkService := lark.NewWebhookService(cfg.URL) + notifier := notify.New() + notifier.UseServices(larkService) + + if err := notifier.Send(ctx, title, content); err != nil { + return "", fmt.Errorf("lark: notify send failed: %w", err) + } + + return "ok", nil +} + +// ValidateConfig 校验飞书机器人配置 +func (p *LarkPusher) ValidateConfig(cfg Config) error { + if cfg.URL == "" { + return errors.New("webhook URL is required") + } + if !strings.HasPrefix(cfg.URL, "https://") { + return errors.New("webhook URL must use https:// protocol") + } + return nil +} diff --git a/backend/plugins/domain/message_gateway/push/push.go b/backend/plugins/domain/msg_gateway/push/push.go similarity index 94% rename from backend/plugins/domain/message_gateway/push/push.go rename to backend/plugins/domain/msg_gateway/push/push.go index e035d62a..3c37dc96 100644 --- a/backend/plugins/domain/message_gateway/push/push.go +++ b/backend/plugins/domain/msg_gateway/push/push.go @@ -15,6 +15,7 @@ const ( defaultTitle = "系统通知" levelInfo = "INFO" defaultHTTPClientTimeout = 10 * time.Second + maxResponseBodyBytes = 4096 ) // Config 基础通知渠道配置 @@ -23,6 +24,7 @@ type Config struct { URL string `json:"url,omitempty"` // Webhook 地址或 SMTP 地址 Secret string `json:"secret,omitempty"` // 签名密钥或 SMTP 密码/Token Key string `json:"key,omitempty"` // AppID 或 SMTP 用户名 + Other string `json:"other,omitempty"` // 附加配置 (如 ChatID / UserKey / 扩展 JSON) Ext map[string]any `json:"ext,omitempty"` // 预留拓展 JSON 配置 } diff --git a/backend/plugins/domain/msg_gateway/push/pushover.go b/backend/plugins/domain/msg_gateway/push/pushover.go new file mode 100644 index 00000000..bb1662f7 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/pushover.go @@ -0,0 +1,69 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "errors" + "fmt" + + "github.com/nikoksr/notify" + "github.com/nikoksr/notify/service/pushover" +) + +func init() { + Register("pushover", &PushoverPusher{}) +} + +// PushoverPusher 基于 nikoksr/notify 的 Pushover 移动端推送实现 +type PushoverPusher struct{} + +// Send 发送 Pushover 通知 +func (p *PushoverPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) { + appToken := cfg.Key + if appToken == "" { + appToken = cfg.Secret + } + if appToken == "" { + return "", errors.New("pushover: app token is required") + } + + userKey := cfg.URL + if userKey == "" { + userKey = cfg.Other + } + if target != "" { + userKey = target + } + if userKey == "" { + return "", errors.New("pushover: user key is required") + } + + title := bodyTitle(body) + content := bodyContent(body, "%s: %v", "\n") + + poService := pushover.New(appToken) + poService.AddReceivers(userKey) + + notifier := notify.New() + notifier.UseServices(poService) + + if err := notifier.Send(ctx, title, content); err != nil { + return "", fmt.Errorf("pushover: notify send failed: %w", err) + } + + return "ok", nil +} + +// ValidateConfig 校验 Pushover 配置 +func (p *PushoverPusher) ValidateConfig(cfg Config) error { + appToken := cfg.Key + if appToken == "" { + appToken = cfg.Secret + } + if appToken == "" { + return errors.New("app token is required") + } + return nil +} diff --git a/backend/plugins/domain/msg_gateway/push/slack.go b/backend/plugins/domain/msg_gateway/push/slack.go new file mode 100644 index 00000000..be162499 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/slack.go @@ -0,0 +1,69 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "errors" + "fmt" + + "github.com/nikoksr/notify" + "github.com/nikoksr/notify/service/slack" +) + +func init() { + Register("slack", &SlackPusher{}) +} + +// SlackPusher 基于 nikoksr/notify 的 Slack 推送实现 +type SlackPusher struct{} + +// Send 发送 Slack 通知 +func (p *SlackPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) { + token := cfg.Key + if token == "" { + token = cfg.Secret + } + channelID := cfg.URL + if target != "" { + channelID = target + } + if channelID == "" { + channelID = cfg.Other + } + + if token == "" { + return "", errors.New("slack: bot/api token is required") + } + if channelID == "" { + return "", errors.New("slack: channel ID is required") + } + + title := bodyTitle(body) + content := bodyContent(body, "*%s*: %v", "\n") + + slackService := slack.New(token) + slackService.AddReceivers(channelID) + + notifier := notify.New() + notifier.UseServices(slackService) + + if err := notifier.Send(ctx, title, content); err != nil { + return "", fmt.Errorf("slack: notify send failed: %w", err) + } + + return "ok", nil +} + +// ValidateConfig 校验 Slack 配置 +func (p *SlackPusher) ValidateConfig(cfg Config) error { + token := cfg.Key + if token == "" { + token = cfg.Secret + } + if token == "" { + return errors.New("slack token is required") + } + return nil +} diff --git a/backend/plugins/domain/msg_gateway/push/telegram.go b/backend/plugins/domain/msg_gateway/push/telegram.go new file mode 100644 index 00000000..590f1397 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/telegram.go @@ -0,0 +1,75 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "errors" + "fmt" + "strconv" + + "github.com/nikoksr/notify" + "github.com/nikoksr/notify/service/telegram" +) + +func init() { + Register("telegram", &TelegramPusher{}) +} + +// TelegramPusher 基于 nikoksr/notify 的 Telegram 机器人推送实现 +type TelegramPusher struct{} + +// Send 执行 Telegram 消息发送 +func (p *TelegramPusher) Send(ctx context.Context, cfg Config, target string, body map[string]any, _ string, _ map[string]any) (string, error) { + botToken := cfg.Secret + if botToken == "" { + botToken = cfg.Key + } + if botToken == "" { + return "", errors.New("telegram: bot token is required") + } + + chatIDStr := target + if chatIDStr == "" { + chatIDStr = cfg.Other + } + if chatIDStr == "" { + return "", errors.New("telegram: chat_id is required") + } + + chatID, err := strconv.ParseInt(chatIDStr, 10, 64) + if err != nil { + return "", fmt.Errorf("telegram: invalid chat_id %q: %w", chatIDStr, err) + } + + title := bodyTitle(body) + content := bodyContent(body, "%s: %v", "\n") + + tgService, err := telegram.New(botToken) + if err != nil { + return "", fmt.Errorf("telegram: init service failed: %w", err) + } + tgService.AddReceivers(chatID) + + notifier := notify.New() + notifier.UseServices(tgService) + + if err := notifier.Send(ctx, title, content); err != nil { + return "", fmt.Errorf("telegram: notify send failed: %w", err) + } + + return "ok", nil +} + +// ValidateConfig 校验 Telegram 机器人配置 +func (p *TelegramPusher) ValidateConfig(cfg Config) error { + token := cfg.Secret + if token == "" { + token = cfg.Key + } + if token == "" { + return errors.New("bot token is required") + } + return nil +} diff --git a/backend/plugins/domain/msg_gateway/push/telegram_test.go b/backend/plugins/domain/msg_gateway/push/telegram_test.go new file mode 100644 index 00000000..026ffbcd --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/telegram_test.go @@ -0,0 +1,28 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "context" + "testing" +) + +func TestTelegramPusherValidation(t *testing.T) { + pusher := &TelegramPusher{} + + err := pusher.ValidateConfig(Config{Secret: "123456:ABC-DEF1234ghIkl-zyx57W2v1u123ew11"}) + if err != nil { + t.Errorf("ValidateConfig failed: %v", err) + } + + err = pusher.ValidateConfig(Config{}) + if err == nil { + t.Errorf("expected error for empty config, got nil") + } + + _, err = pusher.Send(context.Background(), Config{Secret: "123:token"}, "not-a-number", map[string]any{"title": "test"}, "", nil) + if err == nil { + t.Errorf("expected error for invalid chat_id, got nil") + } +} diff --git a/backend/plugins/domain/msg_gateway/push/template.go b/backend/plugins/domain/msg_gateway/push/template.go new file mode 100644 index 00000000..5b6767d8 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/push/template.go @@ -0,0 +1,346 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package push + +import ( + "bytes" + "encoding/json" + "fmt" + "maps" + "regexp" + "slices" + "strconv" + "strings" + "text/template" + "time" +) + +var ( + // placeholderRegex matches {{ ... }} tags + placeholderRegex = regexp.MustCompile(`\{\{\s*([^}]+?)\s*\}\}`) + // identifierRegex matches simple identifiers like name or user.username + identifierRegex = regexp.MustCompile(`^[a-zA-Z_][a-zA-Z0-9_.]*$`) +) + +// jsonMap is a map that serializes to JSON when printed as a string in templates. +type jsonMap map[string]any + +func (m jsonMap) String() string { + b, err := json.Marshal(map[string]any(m)) + if err != nil { + return fmt.Sprintf("%v", map[string]any(m)) + } + return string(b) +} + +func (m jsonMap) MarshalJSON() ([]byte, error) { + return json.Marshal(map[string]any(m)) +} + +// jsonSlice is a slice that serializes to JSON when printed as a string in templates. +type jsonSlice []any + +func (s jsonSlice) String() string { + b, err := json.Marshal([]any(s)) + if err != nil { + return fmt.Sprintf("%v", []any(s)) + } + return string(b) +} + +func (s jsonSlice) MarshalJSON() ([]byte, error) { + return json.Marshal([]any(s)) +} + +var defaultFuncMap = template.FuncMap{ + "default": func(fallback any, val any) any { + if val == nil { + return fallback + } + switch v := val.(type) { + case string: + if v == "" { + return fallback + } + case bool: + if !v { + return fallback + } + case int: + if v == 0 { + return fallback + } + case int32: + if v == 0 { + return fallback + } + case int64: + if v == 0 { + return fallback + } + case float64: + if v == 0 { + return fallback + } + } + return val + }, + "toJson": func(v any) string { + b, err := json.Marshal(v) + if err != nil { + return fmt.Sprint(v) + } + return string(b) + }, + "upper": strings.ToUpper, + "lower": strings.ToLower, + "trim": strings.TrimSpace, + "dateFormat": func(format string, t any) string { + switch v := t.(type) { + case time.Time: + return v.Format(format) + case *time.Time: + if v != nil { + return v.Format(format) + } + } + return fmt.Sprint(t) + }, +} + +// hasKey checks if a dot-delimited or plain key exists in body +func hasKey(body map[string]any, key string) bool { + if _, ok := body[key]; ok { + return true + } + parts := strings.Split(key, ".") + var cur any = body + for _, part := range parts { + m, ok := cur.(map[string]any) + if !ok { + return false + } + val, exists := m[part] + if !exists { + return false + } + cur = val + } + return true +} + +// normalizeTemplate converts legacy {{key}} / {{user.name}} into Go template {{.user.name}} +// while preserving Go template keywords, dot expressions, pipelines, and missing placeholders. +func normalizeTemplate(tmpl string, body map[string]any) string { + return placeholderRegex.ReplaceAllStringFunc(tmpl, func(match string) string { + sub := strings.TrimSpace(match[2 : len(match)-2]) + if sub == "" { + return match + } + // If it's already a dot expression or special variable ($...) + if strings.HasPrefix(sub, ".") || strings.HasPrefix(sub, "$") { + return match + } + // If it's a known Go template keyword or block + firstWord := strings.Fields(sub)[0] + switch firstWord { + case "if", "else", "end", "range", "with", "template", "define", "block", "nil", "true", "false": + return match + } + // Check if it's a pipeline like `key | default "val"` + if strings.Contains(sub, "|") { + const pipelineSplitCount = 2 + parts := strings.SplitN(sub, "|", pipelineSplitCount) + left := strings.TrimSpace(parts[0]) + right := strings.TrimSpace(parts[1]) + if identifierRegex.MatchString(left) && !strings.HasPrefix(left, ".") && !strings.HasPrefix(left, "$") { + return fmt.Sprintf("{{ .%s | %s }}", left, right) + } + return match + } + // Simple identifier: if present in body, convert to dot expression; otherwise preserve as is for fallback + if identifierRegex.MatchString(sub) { + if hasKey(body, sub) { + return fmt.Sprintf("{{ .%s }}", sub) + } + // Missing key: keep original text so fallback or literal is preserved + return match + } + return match + }) +} + +// prepareContext pre-processes the body map so that: +// 1. Dotted keys like "user.name" are expanded to nested map structure. +// 2. Complex structs, slices, and maps have JSON-friendly string representations when directly interpolated. +func prepareContext(body map[string]any) jsonMap { + if body == nil { + return make(jsonMap) + } + ctx := make(jsonMap, len(body)) + for k, v := range body { + formatted := formatContextValue(v) + ctx[k] = formatted + // If key contains '.', expand into nested hierarchy + if strings.Contains(k, ".") { + parts := strings.Split(k, ".") + cur := ctx + for i := 0; i < len(parts)-1; i++ { + sub, ok := cur[parts[i]].(jsonMap) + if !ok { + sub = make(jsonMap) + cur[parts[i]] = sub + } + cur = sub + } + cur[parts[len(parts)-1]] = formatted + } + } + return ctx +} + +// formatContextValue formats slices and maps to JSON representation for direct string printing, +// while preserving basic scalar types for template functions. +func formatContextValue(v any) any { + if v == nil { + return "" + } + switch val := v.(type) { + case string, int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64, float32, float64, bool, time.Time: + return val + case []byte: + return string(val) + case map[string]any: + jm := make(jsonMap, len(val)) + for k, subVal := range val { + jm[k] = formatContextValue(subVal) + } + return jm + case []any: + js := make(jsonSlice, len(val)) + for i, subVal := range val { + js[i] = formatContextValue(subVal) + } + return js + case []string: + js := make(jsonSlice, len(val)) + for i, subVal := range val { + js[i] = subVal + } + return js + default: + return val + } +} + +// ParseTemplate parses template strings by replacing {{placeholder}} structures with values from body. +// It supports Go text/template expressions (e.g. if/else, pipelines, default, toJson) as well as legacy {{key}} placeholders. +func ParseTemplate(templateStr string, body map[string]any) string { + if templateStr == "" { + return "" + } + + normalized := normalizeTemplate(templateStr, body) + ctx := prepareContext(body) + + tmpl, err := template.New("push_tmpl"). + Funcs(defaultFuncMap). + Option("missingkey=zero"). + Parse(normalized) + if err != nil { + return fallbackReplace(templateStr, body) + } + + var buf bytes.Buffer + if err = tmpl.Execute(&buf, ctx); err != nil { + return fallbackReplace(templateStr, body) + } + + return buf.String() +} + +func fallbackReplace(template string, body map[string]any) string { + var buf strings.Builder + buf.Grow(len(template)) + + i := 0 + for { + pos := strings.Index(template[i:], "{{") + if pos == -1 { + buf.WriteString(template[i:]) + break + } + buf.WriteString(template[i : i+pos]) + i += pos + 2 + + endPos := strings.Index(template[i:], "}}") + if endPos == -1 { + buf.WriteString("{{") + buf.WriteString(template[i:]) + break + } + key := strings.TrimSpace(template[i : i+endPos]) + key = strings.TrimPrefix(key, ".") + if val, ok := body[key]; ok { + buf.WriteString(formatValue(val)) + } else { + buf.WriteString("{{") + buf.WriteString(template[i : i+endPos]) + buf.WriteString("}}") + } + i += endPos + 2 + } + return buf.String() +} + +func formatValue(v any) string { + if v == nil { + return "" + } + switch val := v.(type) { + case string: + return val + case []byte: + return string(val) + case int: + return strconv.Itoa(val) + case int32: + return strconv.FormatInt(int64(val), 10) + case int64: + return strconv.FormatInt(val, 10) + case float64: + return strconv.FormatFloat(val, 'f', -1, 64) + case bool: + return strconv.FormatBool(val) + default: + b, err := json.Marshal(v) + if err == nil { + return string(b) + } + return fmt.Sprintf("%v", v) + } +} + +// bodyTitle returns the notification title, falling back to the default. +func bodyTitle(body map[string]any) string { + if t, ok := body["title"].(string); ok && t != "" { + return t + } + return defaultTitle +} + +// bodyContent returns the notification body, rendering every entry with format +// (a "%s … %v" pair) and joining them with sep when no content field is given. +// Entries render in sorted key order so identical bodies always produce +// identical text. +func bodyContent(body map[string]any, format, sep string) string { + if c, ok := body["content"].(string); ok && c != "" { + return c + } + parts := make([]string, 0, len(body)) + for _, k := range slices.Sorted(maps.Keys(body)) { + parts = append(parts, fmt.Sprintf(format, k, body[k])) + } + return strings.Join(parts, sep) +} diff --git a/backend/plugins/domain/message_gateway/push/template_test.go b/backend/plugins/domain/msg_gateway/push/template_test.go similarity index 57% rename from backend/plugins/domain/message_gateway/push/template_test.go rename to backend/plugins/domain/msg_gateway/push/template_test.go index 5132212a..cf11d140 100644 --- a/backend/plugins/domain/message_gateway/push/template_test.go +++ b/backend/plugins/domain/msg_gateway/push/template_test.go @@ -64,6 +64,60 @@ func TestParseTemplate(t *testing.T) { body: map[string]any{"obj": map[string]any{"key": "value"}}, expected: `obj: {"key":"value"}`, }, + { + name: "nested property from flat map", + template: "hello {{user.username}}", + body: map[string]any{"user.username": "Alice"}, + expected: "hello Alice", + }, + { + name: "nested property from nested map", + template: "hello {{user.username}}", + body: map[string]any{"user": map[string]any{"username": "Bob"}}, + expected: "hello Bob", + }, + { + name: "go template dot syntax", + template: "hello {{.user.username}}", + body: map[string]any{"user": map[string]any{"username": "Charlie"}}, + expected: "hello Charlie", + }, + { + name: "default value helper fallback", + template: "hello {{.nickname | default \"Guest\"}}", + body: map[string]any{"nickname": ""}, + expected: "hello Guest", + }, + { + name: "default value helper provided", + template: "hello {{.nickname | default \"Guest\"}}", + body: map[string]any{"nickname": "David"}, + expected: "hello David", + }, + { + name: "conditional if else true", + template: "{{if .is_admin}}Admin: {{.name}}{{else}}User: {{.name}}{{end}}", + body: map[string]any{"is_admin": true, "name": "Eve"}, + expected: "Admin: Eve", + }, + { + name: "conditional if else false", + template: "{{if .is_admin}}Admin: {{.name}}{{else}}User: {{.name}}{{end}}", + body: map[string]any{"is_admin": false, "name": "Frank"}, + expected: "User: Frank", + }, + { + name: "upper and lower helper", + template: "{{.title | upper}} - {{.level | lower}}", + body: map[string]any{"title": "Warning", "level": "INFO"}, + expected: "WARNING - info", + }, + { + name: "toJson helper", + template: "payload: {{toJson .data}}", + body: map[string]any{"data": map[string]any{"status": "ok"}}, + expected: `payload: {"status":"ok"}`, + }, } for _, tt := range tests { diff --git a/backend/plugins/domain/message_gateway/service/admin.go b/backend/plugins/domain/msg_gateway/service/bot_channel.go similarity index 59% rename from backend/plugins/domain/message_gateway/service/admin.go rename to backend/plugins/domain/msg_gateway/service/bot_channel.go index 50633dea..3d64f14c 100644 --- a/backend/plugins/domain/message_gateway/service/admin.go +++ b/backend/plugins/domain/msg_gateway/service/bot_channel.go @@ -4,9 +4,10 @@ package service import ( - "Wavelet/plugins/domain/message_gateway/errs" - "Wavelet/plugins/domain/message_gateway/model" - "Wavelet/plugins/domain/message_gateway/repository" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" "context" "encoding/json" "errors" @@ -22,19 +23,19 @@ import ( const defaultTelegramAPI = "https://api.telegram.org" // ListDefinitions returns the admin form schema of every supported channel type. -func ListDefinitions() []model.Definition { - return []model.Definition{ +func ListDefinitions() []do.Definition { + return []do.Definition{ { - Type: model.MessageChannelTypeTelegram, - Fields: []model.Field{ - {Key: "token", Type: "password", Required: true}, - {Key: "api_base", Type: "text", Required: false}, + Type: consts.MessageChannelTypeTelegram, + Fields: []do.Field{ + {Key: "token", Type: consts.TypePassword, Required: true}, + {Key: "api_base", Type: consts.TypeText, Required: false}, }, }, { - Type: model.MessageChannelTypeQQ, - Fields: []model.Field{ - {Key: "app_id", Type: "text", Required: true}, + Type: consts.MessageChannelTypeQQ, + Fields: []do.Field{ + {Key: "app_id", Type: consts.TypeText, Required: true}, {Key: "client_secret", Type: "password", Required: true}, }, }, @@ -42,25 +43,25 @@ func ListDefinitions() []model.Definition { } // CreateChannel validates the admin payload and persists an encrypted channel. -func CreateChannel(ctx context.Context, req model.CreateChannelRequest) (model.ChannelDTO, error) { +func CreateChannel(ctx context.Context, req do.CreateChannelRequest) (do.ChannelDTO, error) { name := strings.TrimSpace(req.Name) if name == "" { - return model.ChannelDTO{}, errors.New(errs.ErrNameRequired) + return do.ChannelDTO{}, errors.New(consts.ErrNameRequired) } channelType := strings.TrimSpace(req.Type) - if channelType != model.MessageChannelTypeTelegram && channelType != model.MessageChannelTypeQQ { - return model.ChannelDTO{}, errors.New(errs.ErrTypeInvalid) + if channelType != consts.MessageChannelTypeTelegram && channelType != consts.MessageChannelTypeQQ { + return do.ChannelDTO{}, errors.New(consts.ErrTypeInvalid) } creds := req.Credentials if creds == nil { creds = map[string]string{} } if err := ValidateCredentials(channelType, creds, false); err != nil { - return model.ChannelDTO{}, err + return do.ChannelDTO{}, err } cipher, err := EncryptCredentials(creds) if err != nil { - return model.ChannelDTO{}, err + return do.ChannelDTO{}, err } extra := req.Extra if extra == nil { @@ -70,32 +71,32 @@ func CreateChannel(ctx context.Context, req model.CreateChannelRequest) (model.C if req.Enabled != nil { enabled = *req.Enabled } - row := &model.MessageChannel{ + row := &entity.MessageChannel{ Name: name, Type: channelType, - OwnerScope: model.MessageOwnerScopeSystem, + OwnerScope: consts.MessageOwnerScopeSystem, Enabled: enabled, Credentials: cipher, Extra: EncodeExtra(extra), } - if err := repository.CreateMessageChannel(ctx, row); err != nil { - return model.ChannelDTO{}, err + if err := dao.CreateMessageChannel(ctx, row); err != nil { + return do.ChannelDTO{}, err } return ToDTO(row, creds, extra), nil } // UpdateChannel patches a channel; empty secrets keep the stored ciphertext. -func UpdateChannel(ctx context.Context, id uint64, req model.UpdateChannelRequest) (model.ChannelDTO, error) { - row, err := repository.GetMessageChannel(ctx, id) +func UpdateChannel(ctx context.Context, id uint64, req do.UpdateChannelRequest) (do.ChannelDTO, error) { + row, err := dao.GetMessageChannel(ctx, id) if err != nil { - if errors.Is(err, errs.ErrRecordNotFound) { - return model.ChannelDTO{}, errors.New(errs.ErrChannelNotFound) + if errors.Is(err, consts.ErrRecordNotFound) { + return do.ChannelDTO{}, errors.New(consts.ErrChannelNotFoundText) } - return model.ChannelDTO{}, err + return do.ChannelDTO{}, err } creds, err := DecryptCredentials(row.Credentials) if err != nil { - return model.ChannelDTO{}, err + return do.ChannelDTO{}, err } extra := ParseExtra(row.Extra) @@ -120,30 +121,30 @@ func UpdateChannel(ctx context.Context, id uint64, req model.UpdateChannelReques merged[k] = v } if err := ValidateCredentials(row.Type, merged, true); err != nil { - return model.ChannelDTO{}, err + return do.ChannelDTO{}, err } creds = merged } cipher, err := EncryptCredentials(creds) if err != nil { - return model.ChannelDTO{}, err + return do.ChannelDTO{}, err } row.Credentials = cipher row.Extra = EncodeExtra(extra) - if err := repository.UpdateMessageChannel(ctx, row); err != nil { - return model.ChannelDTO{}, err + if err := dao.UpdateMessageChannel(ctx, row); err != nil { + return do.ChannelDTO{}, err } return ToDTO(row, creds, extra), nil } // ListChannels returns every channel with secrets masked. -func ListChannels(ctx context.Context) ([]model.ChannelDTO, error) { - rows, err := repository.ListMessageChannels(ctx) +func ListChannels(ctx context.Context) ([]do.ChannelDTO, error) { + rows, err := dao.ListMessageChannels(ctx) if err != nil { return nil, err } - out := make([]model.ChannelDTO, 0, len(rows)) + out := make([]do.ChannelDTO, 0, len(rows)) for i := range rows { creds, _ := DecryptCredentials(rows[i].Credentials) extra := ParseExtra(rows[i].Extra) @@ -154,21 +155,21 @@ func ListChannels(ctx context.Context) ([]model.ChannelDTO, error) { // DeleteChannel removes a channel together with its bindings and pairing codes. func DeleteChannel(ctx context.Context, id uint64) error { - if _, err := repository.GetMessageChannel(ctx, id); err != nil { - if errors.Is(err, errs.ErrRecordNotFound) { - return errors.New(errs.ErrChannelNotFound) + if _, err := dao.GetMessageChannel(ctx, id); err != nil { + if errors.Is(err, consts.ErrRecordNotFound) { + return errors.New(consts.ErrChannelNotFoundText) } return err } - return repository.DeleteMessageChannel(ctx, id) + return dao.DeleteMessageChannel(ctx, id) } // ProbeChannel verifies the stored credentials against the upstream platform. func ProbeChannel(ctx context.Context, id uint64) error { - row, err := repository.GetMessageChannel(ctx, id) + row, err := dao.GetMessageChannel(ctx, id) if err != nil { - if errors.Is(err, errs.ErrRecordNotFound) { - return errors.New(errs.ErrChannelNotFound) + if errors.Is(err, consts.ErrRecordNotFound) { + return errors.New(consts.ErrChannelNotFoundText) } return err } @@ -177,12 +178,12 @@ func ProbeChannel(ctx context.Context, id uint64) error { return err } switch row.Type { - case model.MessageChannelTypeTelegram: + case consts.MessageChannelTypeTelegram: return ProbeTelegram(ctx, creds) - case model.MessageChannelTypeQQ: + case consts.MessageChannelTypeQQ: return ProbeQQ(ctx, creds) default: - return errors.New(errs.ErrTypeInvalid) + return errors.New(consts.ErrTypeInvalid) } } @@ -190,7 +191,7 @@ func ProbeChannel(ctx context.Context, id uint64) error { func ProbeTelegram(ctx context.Context, creds map[string]string) error { tok := creds["token"] if strings.TrimSpace(tok) == "" { - return errors.New(errs.ErrMissingTelegramToken) + return errors.New(consts.ErrMissingTelegramToken) } base := creds["api_base"] base = strings.TrimRight(strings.TrimSpace(base), "/") @@ -210,7 +211,7 @@ func ProbeTelegram(ctx context.Context, creds map[string]string) error { defer func() { _ = resp.Body.Close() }() body, _ := io.ReadAll(resp.Body) if resp.StatusCode != http.StatusOK { - return fmt.Errorf("%s (%d): %s", errs.ErrTelegramGetMeFailed, resp.StatusCode, string(body)) + return fmt.Errorf("%s (%d): %s", consts.ErrTelegramGetMeFailed, resp.StatusCode, string(body)) } var res struct { OK bool `json:"ok"` @@ -219,7 +220,7 @@ func ProbeTelegram(ctx context.Context, creds map[string]string) error { return err } if !res.OK { - return fmt.Errorf("%s: %s", errs.ErrTelegramNotOK, string(body)) + return fmt.Errorf("%s: %s", consts.ErrTelegramNotOK, string(body)) } return nil } @@ -229,7 +230,7 @@ func ProbeQQ(_ context.Context, creds map[string]string) error { appID := strings.TrimSpace(creds["app_id"]) secret := strings.TrimSpace(creds["client_secret"]) if appID == "" || secret == "" { - return errors.New(errs.ErrMissingQQCredentials) + return errors.New(consts.ErrMissingQQCredentials) } credentials := &token.QQBotCredentials{ AppID: appID, @@ -238,10 +239,10 @@ func ProbeQQ(_ context.Context, creds map[string]string) error { tokSrc := token.NewQQBotTokenSource(credentials) tok, err := tokSrc.Token() if err != nil { - return fmt.Errorf("%s: %w", errs.ErrQQTokenFetchFailed, err) + return fmt.Errorf("%s: %w", consts.ErrQQTokenFetchFailed, err) } if tok == nil || tok.AccessToken == "" { - return errors.New(errs.ErrQQEmptyToken) + return errors.New(consts.ErrQQEmptyToken) } return nil } @@ -249,31 +250,31 @@ func ProbeQQ(_ context.Context, creds map[string]string) error { // ValidateCredentials checks the admin submitted credentials for a channel type. func ValidateCredentials(t string, creds map[string]string, isUpdate bool) error { switch t { - case model.MessageChannelTypeTelegram: + case consts.MessageChannelTypeTelegram: tok := creds["token"] if strings.TrimSpace(tok) == "" && !isUpdate { - return errors.New(errs.ErrTelegramTokenRequired) + return errors.New(consts.ErrTelegramTokenRequired) } if base, ok := creds["api_base"]; ok && strings.TrimSpace(base) != "" { if !strings.HasPrefix(base, "http://") && !strings.HasPrefix(base, "https://") { - return errors.New(errs.ErrAPIBaseInvalid) + return errors.New(consts.ErrAPIBaseInvalid) } } - case model.MessageChannelTypeQQ: + case consts.MessageChannelTypeQQ: appID := creds["app_id"] secret := creds["client_secret"] if (strings.TrimSpace(appID) == "" || strings.TrimSpace(secret) == "") && !isUpdate { - return errors.New(errs.ErrQQCredentialsRequired) + return errors.New(consts.ErrQQCredentialsRequired) } default: - return errors.New(errs.ErrTypeInvalid) + return errors.New(consts.ErrTypeInvalid) } return nil } // ToDTO projects a channel row onto the admin DTO with credentials masked. -func ToDTO(row *model.MessageChannel, creds, extra map[string]string) model.ChannelDTO { - return model.ChannelDTO{ +func ToDTO(row *entity.MessageChannel, creds, extra map[string]string) do.ChannelDTO { + return do.ChannelDTO{ ID: row.ID, Name: row.Name, Type: row.Type, @@ -284,27 +285,3 @@ func ToDTO(row *model.MessageChannel, creds, extra map[string]string) model.Chan Extra: extra, } } - -// MaskCredentials hides secret bearing credential entries. -func MaskCredentials(_ string, in map[string]string) map[string]string { - out := make(map[string]string, len(in)) - for k, v := range in { - if k == "token" || k == "client_secret" { - out[k] = MaskSecret(v) - } else { - out[k] = v - } - } - return out -} - -const minMaskSecretLength = 8 - -// MaskSecret keeps only a short visible prefix and suffix of a secret. -func MaskSecret(s string) string { - s = strings.TrimSpace(s) - if len(s) <= minMaskSecretLength { - return "******" - } - return s[:4] + "..." + s[len(s)-4:] -} diff --git a/backend/plugins/domain/msg_gateway/service/bot_crypto.go b/backend/plugins/domain/msg_gateway/service/bot_crypto.go new file mode 100644 index 00000000..ce9600aa --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_crypto.go @@ -0,0 +1,113 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/pkg/util" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "strings" + "sync" +) + +var ( + credentialSecretMu sync.RWMutex + credentialSecret string +) + +// SetCredentialSecret sets the secret used to derive CredentialKey. +func SetCredentialSecret(secret string) { + credentialSecretMu.Lock() + defer credentialSecretMu.Unlock() + credentialSecret = secret +} + +// CredentialKey is AES-256 hex derived from the session secret. +func CredentialKey() string { + credentialSecretMu.RLock() + secret := credentialSecret + credentialSecretMu.RUnlock() + sum := sha256.Sum256([]byte(secret)) + return hex.EncodeToString(sum[:]) +} + +// EncryptCredentials encrypts a credential map as JSON ciphertext. +func EncryptCredentials(creds map[string]string) (string, error) { + if creds == nil { + creds = map[string]string{} + } + raw, err := json.Marshal(creds) + if err != nil { + return "", err + } + return util.Encrypt(CredentialKey(), string(raw)) +} + +// DecryptCredentials decrypts a credential map from ciphertext. +func DecryptCredentials(ciphertext string) (map[string]string, error) { + if ciphertext == "" { + return map[string]string{}, nil + } + plain, err := util.Decrypt(CredentialKey(), ciphertext) + if err != nil { + return nil, err + } + var out map[string]string + if err := json.Unmarshal([]byte(plain), &out); err != nil { + return nil, err + } + if out == nil { + out = map[string]string{} + } + return out, nil +} + +// ParseExtra decodes optional extra JSON into a string map. +func ParseExtra(raw string) map[string]string { + if raw == "" { + return map[string]string{} + } + var out map[string]string + if err := json.Unmarshal([]byte(raw), &out); err != nil || out == nil { + return map[string]string{} + } + return out +} + +// EncodeExtra encodes extra fields as JSON string. +func EncodeExtra(extra map[string]string) string { + if extra == nil { + return "" + } + raw, err := json.Marshal(extra) + if err != nil { + return "" + } + return string(raw) +} + +// MaskCredentials hides secret-bearing credential entries. +func MaskCredentials(_ string, in map[string]string) map[string]string { + out := make(map[string]string, len(in)) + for k, v := range in { + if k == "token" || k == "client_secret" || k == "app_secret" || k == "bot_token" { + out[k] = MaskSecret(v) + } else { + out[k] = v + } + } + return out +} + +const minMaskSecretLength = 8 + +// MaskSecret keeps only a short visible prefix and suffix of a secret. +func MaskSecret(s string) string { + s = strings.TrimSpace(s) + if len(s) <= minMaskSecretLength { + return "******" + } + return s[:4] + "..." + s[len(s)-4:] +} diff --git a/backend/plugins/domain/msg_gateway/service/bot_dispatch.go b/backend/plugins/domain/msg_gateway/service/bot_dispatch.go new file mode 100644 index 00000000..448f1fa0 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_dispatch.go @@ -0,0 +1,191 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "context" + "encoding/json" + "errors" + "fmt" + "strings" +) + +const ( + // TaskDispatchBotMsg is the queue pattern for bot downlink dispatch. + TaskDispatchBotMsg = consts.TaskDispatchBotMsg + // TaskTypeDispatchBotMsg is the admin type identifier for bot downlink dispatch. + TaskTypeDispatchBotMsg = consts.TaskTypeDispatchBotMsg + + taskQueueDefault = "default" + taskParamTypeString = "string" + paramNameText = "text" +) + +// BotDispatchMeta describes the bot downlink dispatch task. +var BotDispatchMeta = contracts.TaskMetaDTO{ + Type: TaskTypeDispatchBotMsg, + AsynqTask: TaskDispatchBotMsg, + Name: "分发 Bot 消息", + DisplayName: "分发 Bot 消息", + Description: "向已绑定的平台账号异步下发 Bot 文本消息", + Category: "messaging", + Queue: taskQueueDefault, + Retryable: true, + Params: []contracts.TaskParamDTO{ + {Name: paramNameText, Label: "消息内容", Type: consts.TypeText, Required: true, Placeholder: "要发送的文本", Description: "下发给绑定用户的文本"}, + {Name: "channel_id", Label: "频道 ID", Type: taskParamTypeString, Required: false, Placeholder: "留空表示全部启用频道", Description: "仅向指定频道的绑定发送"}, + {Name: "user_id", Label: "用户 ID", Type: taskParamTypeString, Required: false, Placeholder: "留空表示频道下全部绑定", Description: "仅向指定 Wavelet 用户的绑定发送"}, + }, +} + +type botDispatchPayload struct { + Text string `json:"text"` + ChannelID uint64 `json:"channel_id,string"` + UserID uint64 `json:"user_id,string"` +} + +// BotDispatchHandler sends a text message through enabled bot channels. +type BotDispatchHandler struct{} + +// ValidatePayload requires a non-empty message body. +func (h *BotDispatchHandler) ValidatePayload(payload []byte) ([]byte, error) { + p, err := parseBotDispatchPayload(payload) + if err != nil { + return nil, err + } + return json.Marshal(p) +} + +// Execute delivers the text to matching channel bindings. +func (h *BotDispatchHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) { + p, err := parseBotDispatchPayload(payload) + if err != nil { + return nil, err + } + + channels, err := dao.ListEnabledMessageChannels(ctx) + if err != nil { + return nil, err + } + if p.ChannelID != 0 { + filtered := channels[:0] + for i := range channels { + if channels[i].ID == p.ChannelID { + filtered = append(filtered, channels[i]) + } + } + channels = filtered + if len(channels) == 0 { + return nil, consts.ErrChannelNotFound + } + } + + sent := 0 + failed := 0 + for i := range channels { + n, ferr := dispatchOnChannel(ctx, &channels[i], p.UserID, p.Text) + sent += n + failed += ferr + } + msg := fmt.Sprintf("Bot 消息已尝试发送,成功 %d,失败 %d", sent, failed) + if svc := GetTaskService(ctx); svc != nil { + svc.AppendLog(ctx, "%s", msg) + } + if sent == 0 && failed > 0 { + return nil, errors.New(msg) + } + return &contracts.TaskResultDTO{Message: msg}, nil +} + +func parseBotDispatchPayload(payload []byte) (botDispatchPayload, error) { + var p botDispatchPayload + if len(payload) > 0 { + if err := json.Unmarshal(payload, &p); err != nil { + return p, fmt.Errorf("%s: %w", consts.ErrInvalidJSONFormat, err) + } + } + p.Text = strings.TrimSpace(p.Text) + if p.Text == "" { + return p, errors.New(consts.ErrBotDispatchTextRequired) + } + return p, nil +} + +func dispatchOnChannel(ctx context.Context, row *entity.MessageChannel, userID uint64, text string) (sent, failed int) { + factory, ok := Lookup(row.Type) + if !ok { + logger.ErrorF(ctx, "bot dispatch: %s type=%s", consts.ErrBotChannelNotRegistered, row.Type) + return 0, 1 + } + cfg, err := channelConfigFromRow(row) + if err != nil { + logger.ErrorF(ctx, "bot dispatch: decode channel %d: %v", row.ID, err) + return 0, 1 + } + ch, err := factory(cfg, nil) + if err != nil { + logger.ErrorF(ctx, "bot dispatch: create adapter %d: %v", row.ID, err) + return 0, 1 + } + if err := ch.Connect(ctx); err != nil { + logger.ErrorF(ctx, "bot dispatch: connect channel %d: %v", row.ID, err) + return 0, 1 + } + defer func() { _ = ch.Disconnect(ctx) }() + + bindings, err := dao.ListBindingsByChannel(ctx, row.ID) + if err != nil { + logger.ErrorF(ctx, "bot dispatch: list bindings %d: %v", row.ID, err) + return 0, 1 + } + for i := range bindings { + if userID != 0 && bindings[i].UserID != userID { + continue + } + to := do.Recipient{ + ChatID: bindings[i].PlatformUserID, + PlatformUserID: bindings[i].PlatformUserID, + } + if err := ch.Send(ctx, to, do.OutboundMessage{Text: text}); err != nil { + logger.ErrorF(ctx, "bot dispatch: send channel=%d user=%d: %v", row.ID, bindings[i].UserID, err) + failed++ + continue + } + sent++ + } + return sent, failed +} + +func channelConfigFromRow(row *entity.MessageChannel) (do.ChannelConfig, error) { + creds, err := DecryptCredentials(row.Credentials) + if err != nil { + return do.ChannelConfig{}, err + } + if creds == nil { + creds = map[string]string{} + } + if creds["bot_token"] == "" && creds["token"] != "" { + creds["bot_token"] = creds["token"] + } + if creds["app_secret"] == "" && creds["client_secret"] != "" { + creds["app_secret"] = creds["client_secret"] + } + extra := ParseExtra(row.Extra) + if extra["base_url"] == "" && creds["api_base"] != "" { + extra["base_url"] = creds["api_base"] + } + return do.ChannelConfig{ + ID: row.ID, + Type: row.Type, + Name: row.Name, + Credentials: creds, + Extra: extra, + }, nil +} diff --git a/backend/plugins/domain/msg_gateway/service/bot_dispatch_test.go b/backend/plugins/domain/msg_gateway/service/bot_dispatch_test.go new file mode 100644 index 00000000..e4ec9850 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_dispatch_test.go @@ -0,0 +1,46 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service_test + +import ( + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "Wavelet/plugins/domain/msg_gateway/service" + "context" + "path/filepath" + "testing" + + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +type dispatchTestDB struct{ db *gorm.DB } + +func (m *dispatchTestDB) GORM() *gorm.DB { return m.db } +func (m *dispatchTestDB) DB(ctx context.Context) *gorm.DB { return m.db.WithContext(ctx) } +func (m *dispatchTestDB) Named(_ string) *gorm.DB { return m.db } + +func TestBotDispatchValidatePayload(t *testing.T) { + h := &service.BotDispatchHandler{} + _, err := h.ValidatePayload([]byte(`{}`)) + require.Error(t, err) + _, err = h.ValidatePayload([]byte(`{"text":"hello"}`)) + require.NoError(t, err) +} + +func TestBotDispatchNoChannels(t *testing.T) { + testDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "dispatch.db")), &gorm.Config{}) + require.NoError(t, err) + require.NoError(t, testDB.AutoMigrate(&entity.MessageChannel{}, &entity.MessageBinding{})) + dao.SetDBServiceForTest(&dispatchTestDB{db: testDB}) + t.Cleanup(func() { dao.SetDBServiceForTest(nil) }) + + h := &service.BotDispatchHandler{} + res, err := h.Execute(context.Background(), []byte(`{"text":"hello"}`)) + require.NoError(t, err) + require.NotNil(t, res) + assert.Contains(t, res.Message, "成功 0") +} diff --git a/backend/plugins/domain/msg_gateway/service/bot_pairing.go b/backend/plugins/domain/msg_gateway/service/bot_pairing.go new file mode 100644 index 00000000..365be507 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_pairing.go @@ -0,0 +1,174 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "context" + "crypto/rand" + "errors" + "strconv" + "strings" + "time" + "unicode" +) + +// GenerateCode returns an 8-character pairing code using crypto/rand. +func GenerateCode() (string, error) { + buf := make([]byte, consts.CodeLength) + if _, err := rand.Read(buf); err != nil { + return "", err + } + out := make([]byte, consts.CodeLength) + for i, b := range buf { + out[i] = consts.CodeAlphabet[int(b)%len(consts.CodeAlphabet)] + } + return string(out), nil +} + +// NormalizeCode strips separators and uppercases. +func NormalizeCode(s string) string { + var b strings.Builder + for _, r := range s { + if r == '-' || unicode.IsSpace(r) { + continue + } + b.WriteRune(unicode.ToUpper(r)) + } + return b.String() +} + +// FormatCode renders ABCD-EFGH format. +func FormatCode(s string) string { + s = NormalizeCode(s) + if len(s) != consts.CodeLength { + return s + } + return s[:4] + "-" + s[4:] +} + +// BindChannel consumes a pairing code and binds the platform identity to the user. +func BindChannel(ctx context.Context, userID uint64, req do.BindRequest) (do.BindingDTO, error) { + channelID, err := strconv.ParseUint(strings.TrimSpace(req.ChannelID), 10, 64) + if err != nil || channelID == 0 { + return do.BindingDTO{}, consts.ErrChannelIDRequired + } + code := NormalizeCode(req.Code) + if code == "" { + return do.BindingDTO{}, consts.ErrCodeInvalid + } + pairing, err := dao.GetPairingCode(ctx, code) + if err != nil { + if errors.Is(err, consts.ErrRecordNotFound) { + return do.BindingDTO{}, consts.ErrCodeInvalid + } + return do.BindingDTO{}, err + } + if !pairing.ExpiresAt.After(time.Now()) { + return do.BindingDTO{}, consts.ErrCodeInvalid + } + if pairing.ChannelID != channelID { + return do.BindingDTO{}, consts.ErrChannelMismatch + } + ch, err := dao.GetMessageChannel(ctx, channelID) + if err != nil { + if errors.Is(err, consts.ErrRecordNotFound) { + return do.BindingDTO{}, consts.ErrCodeInvalid + } + return do.BindingDTO{}, err + } + if !ch.Enabled { + return do.BindingDTO{}, consts.ErrChannelDisabled + } + + existing, err := dao.GetBindingByChannelPlatform(ctx, channelID, pairing.PlatformUserID) + if err != nil && !errors.Is(err, consts.ErrRecordNotFound) { + return do.BindingDTO{}, err + } + if err == nil && existing != nil { + if existing.UserID != userID { + return do.BindingDTO{}, consts.ErrPlatformAlreadyBound + } + _ = dao.DeletePairingCode(ctx, pairing.Code) + return ToBindingDTO(existing, ch), nil + } + + row := &entity.MessageBinding{ + UserID: userID, + ChannelID: channelID, + PlatformUserID: pairing.PlatformUserID, + } + if err := dao.CreateMessageBinding(ctx, row); err != nil { + return do.BindingDTO{}, err + } + if err := dao.DeletePairingCode(ctx, pairing.Code); err != nil { + return do.BindingDTO{}, err + } + return ToBindingDTO(row, ch), nil +} + +// ListEnabledPublicChannels returns the channels a user may bind to. +func ListEnabledPublicChannels(ctx context.Context) ([]do.PublicChannelDTO, error) { + rows, err := dao.ListEnabledMessageChannels(ctx) + if err != nil { + return nil, err + } + out := make([]do.PublicChannelDTO, 0, len(rows)) + for _, row := range rows { + out = append(out, do.PublicChannelDTO{ID: row.ID, Name: row.Name, Type: row.Type}) + } + return out, nil +} + +// ListUserBindings returns the binding rows of one user enriched with channel info. +func ListUserBindings(ctx context.Context, userID uint64) ([]do.BindingDTO, error) { + rows, err := dao.ListBindingsByUser(ctx, userID) + if err != nil { + return nil, err + } + out := make([]do.BindingDTO, 0, len(rows)) + for i := range rows { + ch, err := dao.GetMessageChannel(ctx, rows[i].ChannelID) + if err != nil { + out = append(out, ToBindingDTO(&rows[i], nil)) + continue + } + out = append(out, ToBindingDTO(&rows[i], ch)) + } + return out, nil +} + +// UnbindChannel removes a binding owned by the given user. +func UnbindChannel(ctx context.Context, userID, bindingID uint64) error { + row, err := dao.GetMessageBinding(ctx, bindingID) + if err != nil { + if errors.Is(err, consts.ErrRecordNotFound) { + return consts.ErrBindingNotFound + } + return err + } + if row.UserID != userID { + return consts.ErrBindingForbidden + } + return dao.DeleteMessageBinding(ctx, bindingID) +} + +// ToBindingDTO projects a binding row and its optional channel onto the user DTO. +func ToBindingDTO(row *entity.MessageBinding, ch *entity.MessageChannel) do.BindingDTO { + dto := do.BindingDTO{ + ID: row.ID, + UserID: row.UserID, + ChannelID: row.ChannelID, + PlatformUserID: row.PlatformUserID, + CreatedAt: row.CreatedAt, + } + if ch != nil { + dto.ChannelName = ch.Name + dto.ChannelType = ch.Type + } + return dto +} diff --git a/backend/plugins/domain/msg_gateway/service/bot_pairing_test.go b/backend/plugins/domain/msg_gateway/service/bot_pairing_test.go new file mode 100644 index 00000000..b6774127 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_pairing_test.go @@ -0,0 +1,27 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service_test + +import ( + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/service" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestGenerateCode_AlphabetAndLength(t *testing.T) { + code, err := service.GenerateCode() + require.NoError(t, err) + assert.Len(t, code, consts.CodeLength) + for _, r := range code { + assert.Contains(t, consts.CodeAlphabet, string(r)) + } +} + +func TestNormalizeAndFormat(t *testing.T) { + assert.Equal(t, "ABCDEFGH", service.NormalizeCode("ab-cd-ef-gh")) + assert.Equal(t, "ABCD-EFGH", service.FormatCode("ABCDEFGH")) +} diff --git a/backend/plugins/domain/msg_gateway/service/bot_runner.go b/backend/plugins/domain/msg_gateway/service/bot_runner.go new file mode 100644 index 00000000..712fa58f --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_runner.go @@ -0,0 +1,88 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/pkg/logger" + "Wavelet/plugins/domain/msg_gateway/model/do" + "context" + "sync" +) + +// Handler processes one inbound message. +type Handler func(ctx context.Context, msg do.InboundMessage) error + +// Factory constructs a Channel from decrypted config. +type Factory func(cfg do.ChannelConfig, onInbound Handler) (Channel, error) + +// Channel is one connected messaging adapter. +type Channel interface { + Type() string + Connect(ctx context.Context) error + Disconnect(ctx context.Context) error + Send(ctx context.Context, to do.Recipient, msg do.OutboundMessage) error + Capabilities() do.Capability +} + +var ( + factoriesMu sync.RWMutex + factories = map[string]Factory{} +) + +// Register stores a channel factory under typ. +func Register(typ string, fn Factory) { + factoriesMu.Lock() + defer factoriesMu.Unlock() + factories[typ] = fn +} + +// Lookup returns a previously registered factory. +func Lookup(typ string) (Factory, bool) { + factoriesMu.RLock() + defer factoriesMu.RUnlock() + fn, ok := factories[typ] + return fn, ok +} + +// Runner manages lifecycle for long-lived channel adapters (WebSocket, long-polling, etc.). +type Runner struct { + mu sync.Mutex + running bool + cancel context.CancelFunc +} + +// GlobalRunner is the default global runner instance. +var GlobalRunner = &Runner{} + +// Start starts all background long-lived channel runners. +func Start(ctx context.Context) error { + GlobalRunner.mu.Lock() + defer GlobalRunner.mu.Unlock() + + if GlobalRunner.running { + return nil + } + + runCtx, cancel := context.WithCancel(ctx) + GlobalRunner.cancel = cancel + GlobalRunner.running = true + + logger.InfoF(runCtx, "[MessageGateway] Starting bot channel runners...") + return nil +} + +// Stop stops the channel runner. +func Stop() { + GlobalRunner.mu.Lock() + defer GlobalRunner.mu.Unlock() + + if !GlobalRunner.running { + return + } + + if GlobalRunner.cancel != nil { + GlobalRunner.cancel() + } + GlobalRunner.running = false +} diff --git a/backend/plugins/domain/msg_gateway/service/bot_runner_test.go b/backend/plugins/domain/msg_gateway/service/bot_runner_test.go new file mode 100644 index 00000000..3caa01da --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/bot_runner_test.go @@ -0,0 +1,34 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service_test + +import ( + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/service" + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type stubChannel struct{} + +func (stubChannel) Type() string { return "stub" } +func (stubChannel) Connect(context.Context) error { return nil } +func (stubChannel) Disconnect(context.Context) error { return nil } +func (stubChannel) Send(context.Context, do.Recipient, do.OutboundMessage) error { return nil } +func (stubChannel) Capabilities() do.Capability { return do.Capability{Text: true} } + +func TestRegisterLookup(t *testing.T) { + service.Register("stub", func(do.ChannelConfig, service.Handler) (service.Channel, error) { + return stubChannel{}, nil + }) + fn, ok := service.Lookup("stub") + require.True(t, ok) + + ch, err := fn(do.ChannelConfig{}, nil) + require.NoError(t, err) + assert.Equal(t, "stub", ch.Type()) +} diff --git a/backend/plugins/domain/msg_gateway/service/push_channel.go b/backend/plugins/domain/msg_gateway/service/push_channel.go new file mode 100644 index 00000000..23a407cf --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_channel.go @@ -0,0 +1,168 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + pkgpush "Wavelet/plugins/domain/msg_gateway/push" + "context" + "errors" + "fmt" +) + +// ListPushChannels returns every configured push channel. +func ListPushChannels(ctx context.Context) ([]entity.PushChannel, error) { + return dao.ListPushChannelsRecord(ctx) +} + +// CreatePushChannel validates uniqueness and persists a new push channel. +func CreatePushChannel(ctx context.Context, req do.CreatePushChannelRequest) (entity.PushChannel, error) { + count, err := dao.CountPushChannelsByNameRecord(ctx, req.Name) + if err != nil { + return entity.PushChannel{}, err + } + if count > 0 { + return entity.PushChannel{}, errors.New(consts.ErrChannelNameExists) + } + + channel := entity.PushChannel{ + Name: req.Name, + Description: req.Description, + Type: req.Type, + Token: req.Token, + URL: req.URL, + Other: req.Other, + Enabled: req.Enabled, + } + if err := channel.Validate(); err != nil { + return entity.PushChannel{}, err + } + if err := dao.CreatePushChannelRecord(ctx, &channel); err != nil { + return entity.PushChannel{}, err + } + return channel, nil +} + +// UpdatePushChannel replaces the mutable fields of an existing push channel. +func UpdatePushChannel(ctx context.Context, id uint64, req do.UpdatePushChannelRequest) (entity.PushChannel, error) { + channel, err := dao.GetPushChannelByIDRecord(ctx, id) + if err != nil { + return entity.PushChannel{}, err + } + + channel.Description = req.Description + channel.Type = req.Type + channel.Token = req.Token + channel.URL = req.URL + channel.Other = req.Other + channel.Enabled = req.Enabled + if err := channel.Validate(); err != nil { + return entity.PushChannel{}, err + } + if err := dao.SavePushChannelRecord(ctx, &channel); err != nil { + return entity.PushChannel{}, err + } + return channel, nil +} + +// DeletePushChannel removes a push channel by id. +func DeletePushChannel(ctx context.Context, id uint64) error { + channel, err := dao.GetPushChannelByIDRecord(ctx, id) + if err != nil { + return err + } + return dao.DeletePushChannelRecord(ctx, &channel) +} + +// LoadChannelForTest resolves the credentials under test, either from a stored +// channel name or from the ad-hoc values sent by the caller. +func LoadChannelForTest(ctx context.Context, req do.TestPushChannelRequest) (string, string, string, string, error) { + if req.Name != "" { + channel, err := dao.GetPushChannelByNameRecord(ctx, req.Name) + if err != nil { + return "", "", "", "", errors.New(consts.ErrChannelNotFoundText) + } + return channel.URL, channel.Token, channel.Other, channel.Type, nil + } + return req.URL, req.Token, req.Other, req.Type, nil +} + +// PreparePushChannelTest builds the connectivity probe payload for a channel. +func PreparePushChannelTest(ctx context.Context, req do.TestPushChannelRequest) (do.SendPayload, error) { + url, token, other, channelType, err := LoadChannelForTest(ctx, req) + if err != nil { + return do.SendPayload{}, err + } + + tempChannel := entity.PushChannel{ + Name: "test_temp", + URL: url, + Token: token, + Other: other, + Type: channelType, + Enabled: true, + } + if err := tempChannel.Validate(); err != nil { + return do.SendPayload{}, err + } + url = tempChannel.URL + + var config pkgpush.Config + var renderedJSON string + switch channelType { + case consts.ChannelLark: + config = pkgpush.Config{Channel: consts.ChannelLark, URL: url, Secret: token} + renderedJSON = other + case consts.ChannelEmail: + config = pkgpush.Config{Channel: consts.ChannelEmail, URL: url, Key: token, Secret: other} + case consts.ChannelTelegram: + config = pkgpush.Config{Channel: consts.ChannelTelegram, URL: url, Secret: token, Key: other} + default: + config = pkgpush.Config{Channel: consts.ChannelCustom, URL: url} + customPushReq := do.CustomPushRequest{ + Title: "通道测试通知", + Content: "这是一条来自系统的消息通道连通性测试消息。", + Description: "系统通道测试", + URL: "https://example.com", + To: req.Target, + } + renderedJSON = RenderCustomPayload(other, customPushReq) + } + + return do.SendPayload{ + EventKey: "test_channel", + Config: config, + Target: req.Target, + Body: do.NotificationMessage{ + Title: "通道测试通知", + Content: "这是一条来自系统的消息通道连通性测试消息。", + Level: consts.DefaultLevelInfo, + }, + Template: renderedJSON, + }, nil +} + +// RunPushTest validates an ad-hoc channel config and sends a connectivity probe. +func RunPushTest(ctx context.Context, cfg pkgpush.Config, target string) error { + pusher, err := pkgpush.GetPusher(cfg.Channel) + if err != nil { + return err + } + if err := pusher.ValidateConfig(cfg); err != nil { + return fmt.Errorf("%s: %w", consts.ErrValidationFailed, err) + } + + testBody := map[string]any{ + consts.KeyTitle: "测试通道推送", + consts.KeyContent: "当您收到这条消息,说明当前渠道连通性测试通过。", + consts.KeyLevel: consts.DefaultLevelInfo, + } + if _, err := pusher.Send(ctx, cfg, target, testBody, "", nil); err != nil { + return err + } + return nil +} diff --git a/backend/plugins/domain/msg_gateway/service/push_event.go b/backend/plugins/domain/msg_gateway/service/push_event.go new file mode 100644 index 00000000..90173ab3 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_event.go @@ -0,0 +1,339 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + "context" + "encoding/json" + "errors" + "strings" + "sync" + "time" +) + +var ( + builtInEventsMu sync.RWMutex + // BuiltInEvents lists all built-in events defined across the domain. + BuiltInEvents []do.EventMetadata +) + +// RegisterBuiltInEvent registers a built-in event definition. +func RegisterBuiltInEvent(meta do.EventMetadata) { + builtInEventsMu.Lock() + defer builtInEventsMu.Unlock() + for i, e := range BuiltInEvents { + if e.Key == meta.Key { + BuiltInEvents[i] = meta + return + } + } + BuiltInEvents = append(BuiltInEvents, meta) +} + +// GetBuiltInEvents returns a copy of registered built-in events. +func GetBuiltInEvents() []do.EventMetadata { + builtInEventsMu.RLock() + defer builtInEventsMu.RUnlock() + out := make([]do.EventMetadata, len(BuiltInEvents)) + copy(out, BuiltInEvents) + return out +} + +// FindBuiltInEvent finds a registered built-in event by key. +func FindBuiltInEvent(key string) (do.EventMetadata, bool) { + for _, meta := range GetBuiltInEvents() { + if meta.Key == key { + return meta, true + } + } + return do.EventMetadata{}, false +} + +// PushRegistryAdapter adapts contracts.PushRegistry onto the built-in event store. +type PushRegistryAdapter struct{} + +// RegisterBuiltInEvent records a built-in push event definition from cross-plugin contract. +func (PushRegistryAdapter) RegisterBuiltInEvent(meta contracts.PushEventMeta) { + RegisterBuiltInEvent(eventMetadataFromContract(meta)) +} + +// SyncEvents persists registered built-in events into the database. +func (PushRegistryAdapter) SyncEvents(ctx context.Context) error { + return SyncEvents(ctx) +} + +func eventMetadataFromContract(meta contracts.PushEventMeta) do.EventMetadata { + return do.EventMetadata{ + Key: meta.Key, + Name: meta.Name, + Description: meta.Description, + DefaultTemplate: do.NotificationMessage{ + Title: meta.DefaultTemplate.Title, + Content: meta.DefaultTemplate.Content, + Level: meta.DefaultTemplate.Level, + Ext: meta.DefaultTemplate.Ext, + }, + } +} + +// SyncBuiltInEvents seeds a database row for every registered built-in event. +func SyncBuiltInEvents(ctx context.Context) error { + for _, meta := range GetBuiltInEvents() { + _, err := dao.GetPushEventByKeyRecord(ctx, meta.Key) + if errors.Is(err, consts.ErrRecordNotFound) { + var defaultTemplateStr string + if defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate); err == nil { + defaultTemplateStr = string(defaultTemplateBytes) + } + event := entity.PushEvent{ + EventKey: meta.Key, + Name: meta.Name, + Channels: []string{}, + Targets: []string{}, + Template: defaultTemplateStr, + Enabled: false, + } + if err := dao.CreatePushEventRecord(ctx, &event); err != nil { + return err + } + } else if err != nil { + return err + } + } + return nil +} + +// SyncEvents automatically registers/updates built-in events in the database. +func SyncEvents(ctx context.Context) error { + return SyncBuiltInEvents(ctx) +} + +// ListPushEvents lists all configured push events. +func ListPushEvents(ctx context.Context) ([]entity.PushEvent, error) { + return dao.ListPushEventsRecord(ctx) +} + +// CreatePushEvent stores a push event configuration for a built-in event or task type. +func CreatePushEvent(ctx context.Context, req do.CreatePushEventRequest) (entity.PushEvent, error) { + eventKey, eventName, defaultTemplateBytes, err := GetEventInfo(ctx, req) + if err != nil { + return entity.PushEvent{}, err + } + + count, err := dao.CountPushEventsByKeyRecord(ctx, eventKey) + if err != nil { + return entity.PushEvent{}, err + } + if count > 0 { + return entity.PushEvent{}, errors.New(consts.ErrEventAlreadyConfigured) + } + + templateStr := strings.TrimSpace(req.Template) + if templateStr == "" { + templateStr = string(defaultTemplateBytes) + } else { + var tempMap map[string]any + if err := json.Unmarshal([]byte(templateStr), &tempMap); err != nil { + return entity.PushEvent{}, errors.New(consts.ErrTemplateInvalidJSON) + } + } + + channels := req.Channels + if channels == nil { + channels = []string{} + } + targets := req.Targets + if targets == nil { + targets = []string{} + } + + event := entity.PushEvent{ + EventKey: eventKey, + Name: eventName, + TaskType: req.TaskType, + Channels: channels, + Targets: targets, + Template: templateStr, + Enabled: req.Enabled, + } + if err := event.Validate(); err != nil { + return entity.PushEvent{}, err + } + if err := dao.CreatePushEventRecord(ctx, &event); err != nil { + return entity.PushEvent{}, err + } + return event, nil +} + +// DeletePushEvent deletes a push event configuration by id. +func DeletePushEvent(ctx context.Context, id uint64) error { + event, err := dao.GetPushEventByIDRecord(ctx, id) + if err != nil { + return err + } + return dao.DeletePushEventRecord(ctx, &event) +} + +// UpdatePushEvent replaces mutable push event fields. +func UpdatePushEvent(ctx context.Context, id uint64, req do.UpdatePushEventRequest) error { + event, err := dao.GetPushEventByIDRecord(ctx, id) + if err != nil { + return err + } + + event.Channels = req.Channels + event.Targets = req.Targets + event.Template = req.Template + event.Enabled = req.Enabled + if err := event.Validate(); err != nil { + return err + } + return dao.SavePushEventRecord(ctx, &event) +} + +// TogglePushEvent flips the enabled flag of a push event. +func TogglePushEvent(ctx context.Context, id uint64) (bool, error) { + event, err := dao.GetPushEventByIDRecord(ctx, id) + if err != nil { + return false, err + } + + enabled := !event.Enabled + if enabled && len(event.Channels) == 0 { + return false, errors.New(consts.ErrEnableWithoutChannels) + } + if err := dao.UpdatePushEventEnabledRecord(ctx, &event, enabled); err != nil { + return false, err + } + return enabled, nil +} + +// ListActivePushEventsByTaskType returns enabled push events for a given task type. +func ListActivePushEventsByTaskType(ctx context.Context, taskType string) ([]entity.PushEvent, error) { + return dao.ListActivePushEventsByTaskTypeRecord(ctx, taskType) +} + +// GetEventInfo derives the event key, display name and default template for a +// task-completion based event or a registered built-in event key. +func GetEventInfo(ctx context.Context, req do.CreatePushEventRequest) (string, string, []byte, error) { + if req.TaskType != "" { + taskName := req.TaskType + if taskSvc := GetTaskService(ctx); taskSvc != nil { + if meta, ok := taskSvc.GetTaskMeta(req.TaskType); ok { + taskName = meta.DisplayName + } + } + eventKey := "task_completed:" + req.TaskType + eventName := "任务完成: " + taskName + defaultTemplate := do.NotificationMessage{ + Title: "任务完成: " + taskName, + Content: "异步任务 {{task_name}} (ID: {{task_id}}) 已完成。状态: {{task_status}},耗时: {{task_duration}} ms。", + Level: consts.DefaultLevelInfo, + } + defaultTemplateBytes, err := json.Marshal(defaultTemplate) + if err != nil { + return "", "", nil, err + } + return eventKey, eventName, defaultTemplateBytes, nil + } + + if req.EventKey == "" { + return "", "", nil, errors.New(consts.ErrEventKeyOrTaskType) + } + + meta, found := FindBuiltInEvent(req.EventKey) + if !found { + return "", "", nil, errors.New(consts.ErrUnsupportedEventKey) + } + + defaultTemplateBytes, err := json.Marshal(meta.DefaultTemplate) + if err != nil { + return "", "", nil, err + } + return req.EventKey, meta.Name, defaultTemplateBytes, nil +} + +// AdminLogin is the metadata definition for the admin login event. +var AdminLogin = do.EventMetadata{ + Key: "admin_login", + Name: "管理员登录", + DefaultTemplate: do.NotificationMessage{ + Title: "管理员登录提醒", + Content: "管理员 {{user.username}} 于 {{time}} 从 IP {{ip}} 登录系统。", + Level: consts.DefaultLevelInfo, + }, + Description: "当管理员成功登录系统时触发此通知", +} + +// HandleAdminLoggedIn 处理管理员登录事件并触发通知 +func HandleAdminLoggedIn(ctx context.Context, event contracts.AdminLoggedIn) { + if event.User == nil { + return + } + + body := map[string]any{ + "user": event.User, + "ip": event.IP, + "time": time.Now().Format("2006-01-02 15:04:05"), + } + DefaultTrigger.Trigger(ctx, AdminLogin, body) +} + +// HandleTaskCompleted handles task completion notifications. +func HandleTaskCompleted(ctx context.Context, e contracts.TaskCompletedEvent) { + events, err := ListActivePushEventsByTaskType(ctx, e.TaskType) + if err != nil { + logger.ErrorF(ctx, "push_task_completed_listener: failed to query push events for task type %s: %v", e.TaskType, err) + return + } + if len(events) == 0 { + return + } + + body := map[string]any{ + "task_id": e.TaskID, + "task_name": e.TaskName, + "task_type": e.TaskType, + "task_status": e.Status, + "task_duration": e.Duration, + "time": time.Now().Format("2006-01-02 15:04:05"), + "task_error": e.ErrorMsg, + "task_result": e.ResultMsg, + } + + var payloadMap map[string]any + if e.Payload != "" { + if err := json.Unmarshal([]byte(e.Payload), &payloadMap); err == nil { + body["payload"] = payloadMap + ExtractUserFromMap(ctx, payloadMap, body) + } + } + if e.Detail != "" { + var detailMap map[string]any + if err := json.Unmarshal([]byte(e.Detail), &detailMap); err == nil { + body["detail"] = detailMap + ExtractUserFromMap(ctx, detailMap, body) + } + } + + for _, event := range events { + meta := do.EventMetadata{ + Key: event.EventKey, + Name: event.Name, + Description: "异步任务执行完毕触发的自动通知", + } + DefaultTrigger.Trigger(ctx, meta, body) + } +} + +// RegisterCustomEvents registers default domain push notification events. +func RegisterCustomEvents() { + RegisterBuiltInEvent(AdminLogin) +} diff --git a/backend/plugins/domain/msg_gateway/service/push_template.go b/backend/plugins/domain/msg_gateway/service/push_template.go new file mode 100644 index 00000000..34330ce7 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_template.go @@ -0,0 +1,63 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/plugins/domain/msg_gateway/model/do" + "encoding/json" + "strings" +) + +// GetFlatBody flattens nested body map into dot-separated key-value map. +func GetFlatBody(body map[string]any) map[string]any { + jsonBytes, err := json.Marshal(body) + if err != nil { + return body + } + var jsonMap map[string]any + if err := json.Unmarshal(jsonBytes, &jsonMap); err != nil { + return body + } + + flatResult := make(map[string]any) + FlattenMap("", jsonMap, flatResult) + return flatResult +} + +// FlattenMap recursively flattens map key-values with dot notation. +func FlattenMap(prefix string, m, result map[string]any) { + for k, v := range m { + key := k + if prefix != "" { + key = prefix + "." + k + } + if nestedMap, ok := v.(map[string]any); ok { + FlattenMap(key, nestedMap, result) + } else { + result[key] = v + } + } +} + +// RenderCustomPayload substitutes the supported template variables of a custom +// webhook body, JSON-escaping every injected value. +func RenderCustomPayload(template string, req do.CustomPushRequest) string { + result := template + result = strings.ReplaceAll(result, "$title", EscapeJSONString(req.Title)) + result = strings.ReplaceAll(result, "$description", EscapeJSONString(req.Description)) + result = strings.ReplaceAll(result, "$content", EscapeJSONString(req.Content)) + result = strings.ReplaceAll(result, "$url", EscapeJSONString(req.URL)) + result = strings.ReplaceAll(result, "$to", EscapeJSONString(req.To)) + return result +} + +// EscapeJSONString renders s as a JSON string body without the surrounding quotes. +func EscapeJSONString(s string) string { + b, _ := json.Marshal(s) + const minJSONLen = 2 + if len(b) >= minJSONLen { + return string(b[1 : len(b)-1]) + } + return s +} diff --git a/backend/plugins/domain/msg_gateway/service/push_trigger.go b/backend/plugins/domain/msg_gateway/service/push_trigger.go new file mode 100644 index 00000000..f3904b70 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_trigger.go @@ -0,0 +1,398 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/pkg/util" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + pkgpush "Wavelet/plugins/domain/msg_gateway/push" + "context" + "encoding/json" + "errors" + "fmt" + "strconv" + "strings" +) + +// EventTrigger represents the unified event trigger class. +type EventTrigger struct{} + +// DefaultTrigger is the singleton instance of EventTrigger. +var DefaultTrigger = &EventTrigger{} + +// Trigger receives event metadata and processes the event notification dispatch asynchronously. +func (t *EventTrigger) Trigger(ctx context.Context, meta do.EventMetadata, body map[string]any) { + asyncCtx := context.WithoutCancel(ctx) + util.Go(func() { + if body == nil { + body = make(map[string]any) + } + if _, hasUser := body["user"]; !hasUser || body["user"] == nil { + body["user"] = GetSystemUser(asyncCtx) + } + + eventPtr, err := dao.GetActivePushEventByKey(asyncCtx, meta.Key) + if err != nil { + if errors.Is(err, consts.ErrRecordNotFound) { + return + } + logger.ErrorF(asyncCtx, "push_event_trigger: failed to get active event %s: %v", meta.Key, err) + return + } + event := *eventPtr + if len(event.Channels) == 0 { + return + } + + flatBody := GetFlatBody(body) + msg, _ := t.buildMessage(&event, meta, flatBody, body) + t.enqueuePushTasks(asyncCtx, meta, &event, msg, flatBody) + }) +} + +func (t *EventTrigger) buildMessage(event *entity.PushEvent, meta do.EventMetadata, flatBody, body map[string]any) (do.NotificationMessage, string) { + var msg do.NotificationMessage + renderedTemplate := "" + + templateSource := event.Template + if templateSource != "" { + var err error + msg, renderedTemplate, err = t.parseCustomTemplate(event, templateSource, flatBody) + if err != nil { + msg.Title = event.Name + msg.Content = renderedTemplate + msg.Level = consts.DefaultLevelInfo + } + } else { + msg = t.parseDefaultTemplate(meta, flatBody) + } + + if msg.Ext == nil { + msg.Ext = make(map[string]any) + } + for k, v := range body { + if k == consts.KeyTitle || k == consts.KeyContent || k == consts.KeyLevel { + continue + } + if _, exists := msg.Ext[k]; !exists { + msg.Ext[k] = v + } + } + + return msg, renderedTemplate +} + +func (t *EventTrigger) parseCustomTemplate(event *entity.PushEvent, templateSource string, flatBody map[string]any) (do.NotificationMessage, string, error) { + var msg do.NotificationMessage + renderedTemplate := pkgpush.ParseTemplate(templateSource, flatBody) + + var tMap map[string]any + if err := json.Unmarshal([]byte(renderedTemplate), &tMap); err != nil { + return msg, renderedTemplate, err + } + + if title, ok := tMap[consts.KeyTitle].(string); ok && title != "" { + msg.Title = title + } else { + msg.Title = event.Name + } + delete(tMap, consts.KeyTitle) + + if content, ok := tMap[consts.KeyContent].(string); ok && content != "" { + msg.Content = content + } else { + msg.Content = renderedTemplate + } + delete(tMap, consts.KeyContent) + + if level, ok := tMap[consts.KeyLevel].(string); ok && level != "" { + msg.Level = level + } else { + msg.Level = consts.DefaultLevelInfo + } + delete(tMap, consts.KeyLevel) + + msg.Ext = tMap + return msg, renderedTemplate, nil +} + +func (t *EventTrigger) parseDefaultTemplate(meta do.EventMetadata, flatBody map[string]any) do.NotificationMessage { + var msg do.NotificationMessage + msg.Title = pkgpush.ParseTemplate(meta.DefaultTemplate.Title, flatBody) + msg.Content = pkgpush.ParseTemplate(meta.DefaultTemplate.Content, flatBody) + msg.Level = pkgpush.ParseTemplate(meta.DefaultTemplate.Level, flatBody) + + if meta.DefaultTemplate.Ext != nil { + msg.Ext = make(map[string]any) + for k, v := range meta.DefaultTemplate.Ext { + if strVal, ok := v.(string); ok { + msg.Ext[k] = pkgpush.ParseTemplate(strVal, flatBody) + } else { + msg.Ext[k] = v + } + } + } + return msg +} + +func (t *EventTrigger) enqueuePushTasks(ctx context.Context, meta do.EventMetadata, event *entity.PushEvent, msg do.NotificationMessage, flatBody map[string]any) { + for _, channelName := range event.Channels { + customChannel, err := dao.GetActivePushChannelByName(ctx, channelName) + if err == nil { + t.enqueueCustomPushChannelTasks(ctx, meta, event, customChannel, msg, flatBody) + continue + } + logger.WarnF(ctx, "push_event_trigger: channel %q not found in DB or disabled: %v", channelName, err) + } +} + +func (t *EventTrigger) enqueueCustomPushChannelTasks(ctx context.Context, meta do.EventMetadata, event *entity.PushEvent, channel *entity.PushChannel, msg do.NotificationMessage, flatBody map[string]any) { + if len(event.Targets) == 0 { + t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, "", msg) + return + } + + for _, target := range event.Targets { + resolvedTarget := ResolveTarget(ctx, target, flatBody, channel.Name) + t.enqueueSingleCustomPushChannelTask(ctx, meta, channel, resolvedTarget, msg) + } +} + +func (t *EventTrigger) enqueueSingleCustomPushChannelTask(ctx context.Context, meta do.EventMetadata, channel *entity.PushChannel, target string, msg do.NotificationMessage) { + var config pkgpush.Config + var renderedTemplate string + + switch channel.Type { + case consts.ChannelLark: + config = pkgpush.Config{Channel: consts.ChannelLark, URL: channel.URL, Secret: channel.Token} + renderedTemplate = channel.Other + case consts.ChannelEmail: + config = pkgpush.Config{Channel: consts.ChannelEmail, URL: channel.URL, Key: channel.Token, Secret: channel.Other} + case consts.ChannelTelegram: + config = pkgpush.Config{Channel: consts.ChannelTelegram, URL: channel.URL, Secret: channel.Token, Key: channel.Other} + default: + config = pkgpush.Config{Channel: consts.ChannelCustom, URL: channel.URL} + customPushReq := do.CustomPushRequest{ + Title: msg.Title, + Content: msg.Content, + Description: meta.Description, + To: target, + } + if urlVal, ok := msg.Ext["url"].(string); ok { + customPushReq.URL = urlVal + } + renderedTemplate = RenderCustomPayload(channel.Other, customPushReq) + } + + payload := do.SendPayload{ + EventKey: meta.Key, + Config: config, + Target: target, + Body: msg, + Template: renderedTemplate, + } + if err := EnqueuePushTask(ctx, payload); err != nil { + logger.ErrorF(ctx, "push_event_trigger: enqueuePushTask failed for %s channel %s -> %s: %v", channel.Type, channel.Name, target, err) + } +} + +// ResolveTarget parses dynamic placeholders into concrete receiver targets. +func ResolveTarget(ctx context.Context, target string, flatBody map[string]any, channel string) string { + target = strings.TrimSpace(target) + if target == "" { + return "" + } + + resolved := ResolveDynamicKeyword(target, flatBody) + if strings.Contains(resolved, "@") { + return resolved + } + if val, matched := ResolveSystemTarget(ctx, resolved, channel); matched { + return val + } + + user, found := ResolveTargetUser(ctx, resolved, channel) + if !found { + return resolved + } + if channel == consts.ChannelEmail && user.Email != "" { + return user.Email + } + if channel != consts.ChannelEmail && user.Username != "" { + return user.Username + } + return resolved +} + +// ResolveDynamicKeyword resolves user.id, username, email keywords. +func ResolveDynamicKeyword(target string, flatBody map[string]any) string { + switch target { + case "user.id", "id": + if val, ok := flatBody["user.id"]; ok { + return fmt.Sprintf("%v", val) + } + if val, ok := flatBody["id"]; ok { + return fmt.Sprintf("%v", val) + } + case "user.username", "username": + if val, ok := flatBody["user.username"]; ok { + return fmt.Sprintf("%v", val) + } + if val, ok := flatBody["username"]; ok { + return fmt.Sprintf("%v", val) + } + case "user.email", consts.ChannelEmail: + if val, ok := flatBody["user.email"]; ok { + return fmt.Sprintf("%v", val) + } + if val, ok := flatBody["email"]; ok { + return fmt.Sprintf("%v", val) + } + } + return target +} + +// ResolveTargetUser resolves user by numeric ID or username string via UserService contract. +func ResolveTargetUser(ctx context.Context, resolved, _ string) (contracts.UserDTO, bool) { + if id, err := strconv.ParseUint(resolved, 10, 64); err == nil { + if u, err := FindUserByID(ctx, id); err == nil && u != nil { + return *u, true + } + } + if u, err := FindUserByUsername(ctx, resolved); err == nil && u != nil { + return *u, true + } + return contracts.UserDTO{}, false +} + +// FindUserByID resolves a user by primary key through UserService. +func FindUserByID(ctx context.Context, id uint64) (*contracts.UserDTO, error) { + if userSvc := GetUserService(ctx); userSvc != nil { + return userSvc.GetUserByID(ctx, id) + } + return nil, consts.ErrUserNotFound +} + +// FindUserByUsername resolves a user by login name through UserService. +func FindUserByUsername(ctx context.Context, username string) (*contracts.UserDTO, error) { + if userSvc := GetUserService(ctx); userSvc != nil { + return userSvc.GetUserByUsername(ctx, username) + } + return nil, consts.ErrUserNotFound +} + +// GetFirstAdminUser resolves the first administrator through the UserService contract. +func GetFirstAdminUser(ctx context.Context) (*contracts.UserDTO, error) { + if userSvc := GetUserService(ctx); userSvc != nil { + return userSvc.GetFirstAdminUser(ctx) + } + return nil, consts.ErrNoAdminUser +} + +// ResolveSystemTarget maps system receiver aliases to administrator contact info. +func ResolveSystemTarget(ctx context.Context, resolved, channel string) (string, bool) { + if resolved != "系统" && resolved != "system" && resolved != "0" { + return "", false + } + adminUser, err := GetFirstAdminUser(ctx) + if err != nil || adminUser == nil { + return resolved, true + } + if channel == consts.ChannelEmail && adminUser.Email != "" { + return adminUser.Email, true + } + if channel != consts.ChannelEmail && adminUser.Username != "" { + return adminUser.Username, true + } + return resolved, true +} + +// GetSystemUser gets a system user DTO. +func GetSystemUser(ctx context.Context) *contracts.UserDTO { + if adminUser, err := GetFirstAdminUser(ctx); err == nil && adminUser != nil { + return adminUser + } + return &contracts.UserDTO{ + Username: "system", + Nickname: "系统管理员", + } +} + +// LoadUserFromPayload extracts user info from data. +func LoadUserFromPayload(ctx context.Context, data map[string]any) any { + if u, exists := data["user"]; exists && u != nil { + return u + } + + if userID, ok := ExtractUserID(data); ok && userID > 0 { + if user, err := FindUserByID(ctx, userID); err == nil && user != nil { + return user + } + } + + if username := ExtractUsername(data); username != "" { + if user, err := FindUserByUsername(ctx, username); err == nil && user != nil { + return user + } + } + return nil +} + +// ExtractUserFromMap extracts user info from payload/detail into body map. +func ExtractUserFromMap(ctx context.Context, data, body map[string]any) { + if u, exists := body["user"]; exists && u != nil { + return + } + if user := LoadUserFromPayload(ctx, data); user != nil { + body["user"] = user + } +} + +// ExtractUserID extracts a user ID from map keys. +func ExtractUserID(data map[string]any) (uint64, bool) { + for _, k := range []string{"user_id", "userId", "uid"} { + val, ok := data[k] + if !ok || val == nil { + continue + } + switch v := val.(type) { + case float64: + if v >= 0 { + return uint64(v), true + } + case int: + if v >= 0 { + return uint64(v), true + } + case int64: + if v >= 0 { + return uint64(v), true + } + case uint64: + return v, true + case string: + if id, err := strconv.ParseUint(v, 10, 64); err == nil { + return id, true + } + } + } + return 0, false +} + +// ExtractUsername extracts a username string from map keys. +func ExtractUsername(data map[string]any) string { + for _, k := range []string{"username", "user_name"} { + if val, ok := data[k]; ok && val != nil { + if s, ok := val.(string); ok && s != "" { + return s + } + } + } + return "" +} diff --git a/backend/plugins/domain/msg_gateway/service/push_worker.go b/backend/plugins/domain/msg_gateway/service/push_worker.go new file mode 100644 index 00000000..a29ec3f6 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/push_worker.go @@ -0,0 +1,189 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package service + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + "Wavelet/plugins/domain/msg_gateway/consts" + "Wavelet/plugins/domain/msg_gateway/dao" + "Wavelet/plugins/domain/msg_gateway/model/do" + "Wavelet/plugins/domain/msg_gateway/model/entity" + pkgpush "Wavelet/plugins/domain/msg_gateway/push" + "context" + "encoding/json" + "errors" + "fmt" + "time" +) + +const ( + // SendNotificationTask is the asynq task name for push notification. + SendNotificationTask = consts.SendNotificationTask + // TaskTypeSendNotification is the admin task manager type identifier. + TaskTypeSendNotification = consts.TaskTypeSendNotification +) + +// SendNotificationMeta represents the task metadata. +var SendNotificationMeta = contracts.TaskMetaDTO{ + Type: TaskTypeSendNotification, + AsynqTask: SendNotificationTask, + Name: "推送通知", + DisplayName: "推送通知", + Description: "异步执行系统通知的多渠道派发与推送", + Category: "push", + SupportsTime: false, + MaxRetry: 3, + Queue: "default", + Retryable: true, + Params: []contracts.TaskParamDTO{ + { + Name: "event_key", + Label: "事件标识", + Type: "string", + Required: true, + Placeholder: "admin_login", + Description: "事件标识 (如 admin_login)", + }, + { + Name: "target", + Label: "目标接收者", + Type: "string", + Required: false, + Description: "目标接收者", + }, + }, +} + +// PushHandler handles asynchronous notification sending. +type PushHandler struct{} + +// ValidatePayload validates and normalizes push parameters. +func (h *PushHandler) ValidatePayload(payload []byte) ([]byte, error) { + if len(payload) == 0 { + return nil, errors.New(consts.ErrPayloadRequired) + } + + var req do.SendPayload + if err := json.Unmarshal(payload, &req); err != nil { + return nil, fmt.Errorf("%s: %w", consts.ErrInvalidJSONFormat, err) + } + + if req.Config.Channel == "" { + return nil, errors.New(consts.ErrChannelTypeRequired) + } + + return json.Marshal(req) +} + +// Execute performs the push send and logs delivery history audit. +func (h *PushHandler) Execute(ctx context.Context, payload []byte) error { + var req do.SendPayload + if err := json.Unmarshal(payload, &req); err != nil { + logger.ErrorF(ctx, "[Push] 解析推送参数失败: %v", err) + return fmt.Errorf("%s: %w", consts.ErrParsePayloadFailed, err) + } + + logger.InfoF(ctx, "[Push] 开始推送通知: 事件 = %s, 渠道 = %s, 接收目标 = %s", req.EventKey, req.Config.Channel, req.Target) + + pusher, err := pkgpush.GetPusher(req.Config.Channel) + if err != nil { + errWrap := fmt.Errorf("%s: %w", consts.ErrGetPusherFailed, err) + logger.ErrorF(ctx, "[Push] 推送失败: %v", errWrap) + h.recordHistory(ctx, req, "failed", errWrap.Error()) + return errWrap + } + + flatBody := req.Body.Flatten() + upstreamResp, err := pusher.Send(ctx, req.Config, req.Target, flatBody, req.Template, nil) + + title := req.Body.Title + content := req.Body.Content + + if err != nil { + logger.ErrorF(ctx, "[Push] 消息推送失败 (标题: %s): %v, 上游返回: %s", title, err, upstreamResp) + h.recordHistory(ctx, req, "failed", err.Error()) + return fmt.Errorf("pusher.Send failed: %w", err) + } + + logger.InfoF(ctx, "[Push] 消息推送成功 (标题: %s, 内容摘要: %s), 上游返回: %s", title, content, upstreamResp) + h.recordHistory(ctx, req, "success", "") + + return nil +} + +func (h *PushHandler) recordHistory(ctx context.Context, req do.SendPayload, status, errMsg string) { + if dbErr := RecordPushHistory(ctx, req, status, errMsg); dbErr != nil { + logger.ErrorF(ctx, "[Push] 写入推送历史审计记录失败: %v", dbErr) + } +} + +// EnqueuePushTask dispatches a notification payload to the async push worker. +func EnqueuePushTask(ctx context.Context, payload do.SendPayload) error { + payloadBytes, err := json.Marshal(payload) + if err != nil { + return err + } + if taskSvc := GetTaskService(ctx); taskSvc != nil { + _, err = taskSvc.Dispatch(ctx, "send_notification", payloadBytes, contracts.TaskTriggerSystem) + return err + } + return errors.New(consts.ErrTaskServiceUnavailable) +} + +// RecordPushHistory creates a push history audit record. +func RecordPushHistory(ctx context.Context, req do.SendPayload, status, errMsg string) error { + title := req.Body.Title + content := req.Body.Content + level := req.Body.Level + if title == "" { + title = "系统通知" + } + if level == "" { + level = consts.DefaultLevelInfo + } + + target := req.Target + if target == "" { + if req.Config.URL != "" { + target = req.Config.URL + const maxTargetLen = 50 + const truncatedLen = 47 + if len(target) > maxTargetLen { + target = target[:truncatedLen] + "..." + } + } else { + target = "default" + } + } + + history := entity.PushHistory{ + EventKey: req.EventKey, + Channel: req.Config.Channel, + Target: target, + Title: title, + Content: content, + Level: level, + Status: status, + ErrorMsg: errMsg, + } + return dao.CreatePushHistoryRecord(ctx, &history) +} + +// ListPushHistories returns a paginated push delivery audit page. +func ListPushHistories(ctx context.Context, filter do.PushHistoryListFilter) (int64, []entity.PushHistory, error) { + return dao.ListPushHistoriesRecord(ctx, filter) +} + +// CleanupPushHistories removes push delivery audit records older than the retention duration. +func CleanupPushHistories(ctx context.Context, retention time.Duration) (int64, error) { + cutoff := time.Now().Add(-retention) + deleted, err := dao.DeletePushHistoriesBeforeRecord(ctx, cutoff) + if err != nil { + logger.WarnF(ctx, "[Push] 清理历史推送日志失败: %v", err) + return 0, err + } + logger.InfoF(ctx, "[Push] 已清理 %s 前推送历史日志,共 %d 条", cutoff.Format(time.RFC3339), deleted) + return deleted, nil +} diff --git a/backend/plugins/domain/msg_gateway/service/service.go b/backend/plugins/domain/msg_gateway/service/service.go new file mode 100644 index 00000000..dc6f7959 --- /dev/null +++ b/backend/plugins/domain/msg_gateway/service/service.go @@ -0,0 +1,75 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +// Package service implements domain business logic, bot gateway adapters, and push notification services for msg_gateway. +package service + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "context" + "sync" +) + +// Platform service dependencies resolved from Cordis context or global fallbacks. +var ( + cacheMu sync.RWMutex + cacheSvc contracts.CacheService + taskMu sync.RWMutex + taskSvc contracts.TaskService + userMu sync.RWMutex + userSvc contracts.UserService +) + +// SetCacheService sets the cache service. +func SetCacheService(s contracts.CacheService) { + cacheMu.Lock() + defer cacheMu.Unlock() + cacheSvc = s +} + +// SetTaskService sets the task service. +func SetTaskService(s contracts.TaskService) { + taskMu.Lock() + defer taskMu.Unlock() + taskSvc = s +} + +// SetUserService sets the user service. +func SetUserService(s contracts.UserService) { + userMu.Lock() + defer userMu.Unlock() + userSvc = s +} + +// GetCache resolves the cache service for the context. +func GetCache(ctx context.Context) contracts.CacheService { + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s + } + cacheMu.RLock() + s := cacheSvc + cacheMu.RUnlock() + return s +} + +// GetTaskService resolves the task service for the context. +func GetTaskService(ctx context.Context) contracts.TaskService { + if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil { + return s + } + taskMu.RLock() + defer taskMu.RUnlock() + return taskSvc +} + +// GetUserService resolves the user service for the context. +func GetUserService(ctx context.Context) contracts.UserService { + if s, err := core.InjectFrom[contracts.UserService](ctx); err == nil && s != nil { + return s + } + userMu.RLock() + s := userSvc + userMu.RUnlock() + return s +} diff --git a/backend/plugins/domain/risk_control/logstore/db_helper.go b/backend/plugins/domain/risk_control/logstore/db_helper.go index af540ef9..ed33c591 100644 --- a/backend/plugins/domain/risk_control/logstore/db_helper.go +++ b/backend/plugins/domain/risk_control/logstore/db_helper.go @@ -42,10 +42,8 @@ func SetChDBForTest(db *gorm.DB) { } func getDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) } dbMu.RLock() s := dbSvc diff --git a/backend/plugins/domain/risk_control/plugin.go b/backend/plugins/domain/risk_control/plugin.go index 3976fc3a..dbe9e0ec 100644 --- a/backend/plugins/domain/risk_control/plugin.go +++ b/backend/plugins/domain/risk_control/plugin.go @@ -91,17 +91,12 @@ func (p *Plugin) Apply(ctx *core.Context) error { var dbCfg rcDBConfig _ = ctx.Config().Bind("database", &dbCfg) - SetAccessLogEnabled(chCfg.Enabled) + // Access logs persist on the active log database (SQLite / Postgres / ClickHouse). + // Collection is independent of ClickHouse being enabled. + SetAccessLogEnabled(true) logstore.SetDefaultDatabases(dbCfg.Enabled, chCfg.Enabled) - // 0. Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - logstore.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - logstore.SetDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, logstore.SetDBService) ctx.OnDispose(func() error { logstore.SetDBService(nil) return nil diff --git a/backend/plugins/domain/risk_control/plugin_test.go b/backend/plugins/domain/risk_control/plugin_test.go index fbc37a6d..aa986b66 100644 --- a/backend/plugins/domain/risk_control/plugin_test.go +++ b/backend/plugins/domain/risk_control/plugin_test.go @@ -35,6 +35,8 @@ func TestRiskControlPluginUnit(t *testing.T) { setting, ok := ctx.Settings().Get("risk_control.enable_access_log") require.True(t, ok) assert.Equal(t, true, setting.Default) + assert.True(t, risk_control.IsAccessLogEnabled(), + "access log collection must be on even when ClickHouse is disabled") require.NoError(t, ctx.Dispose()) assert.False(t, customMWCalled) // not dispatched via gin engine here diff --git a/backend/plugins/domain/risk_control/service.go b/backend/plugins/domain/risk_control/service.go index d25c84fd..c502027d 100644 --- a/backend/plugins/domain/risk_control/service.go +++ b/backend/plugins/domain/risk_control/service.go @@ -13,8 +13,12 @@ import ( "time" ) -// fallbackLogEngine 是日志库状态不可得时对外暴露的引擎标识。 -const fallbackLogEngine = "sqlite" +const ( + // fallbackLogEngine 是日志库状态不可得时对外暴露的引擎标识。 + fallbackLogEngine = "sqlite" + // accessLogMaxFlushWait 强制把不足 MinBatchSize 的访问日志刷盘,避免管理台低频访问永远看不到记录。 + accessLogMaxFlushWait = 2 * time.Second +) var ( logWriterMu sync.RWMutex @@ -30,6 +34,7 @@ func InitLogWriter(ctx context.Context) { } cfg := batchwriter.DefaultConfig() + cfg.MaxFlushWait = accessLogMaxFlushWait writer, err := batchwriter.New[*logstore.UserAccessLog](cfg, writeAccessLogBatch, batchwriter.WithDropHandler[*logstore.UserAccessLog](func(item *logstore.UserAccessLog) { path := "" diff --git a/backend/plugins/domain/system/plugin.go b/backend/plugins/domain/system/plugin.go index b03f98e8..1a45bcc9 100644 --- a/backend/plugins/domain/system/plugin.go +++ b/backend/plugins/domain/system/plugin.go @@ -9,8 +9,8 @@ import ( "Wavelet/core/contracts" "Wavelet/pkg/logger" "Wavelet/pkg/response" + "context" "net/http" - "reflect" "github.com/gin-gonic/gin" ) @@ -28,13 +28,6 @@ func (p *Plugin) Name() string { return "system" } -// Inject declares required dependencies for the system domain plugin. -func (p *Plugin) Inject() []reflect.Type { - return []reflect.Type{ - reflect.TypeFor[contracts.DBService](), - } -} - // Manifest returns the plugin metadata. func (p *Plugin) Manifest() core.Manifest { return core.Manifest{ @@ -47,42 +40,43 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers system routes. func (p *Plugin) Apply(ctx *core.Context) error { - appName := ctx.Config().String("app.app_name", "Wavelet") - // 1. Health check ctx.Router().GET("/api/healthz", func(c *gin.Context) { c.JSON(http.StatusOK, gin.H{"status": "ok"}) }) ctx.Router().RegisterWhitelist("/api/healthz") - // 2. Public config + // 2. Public config — owned data comes from PublicConfigProvider (admin). ctx.Router().GET("/api/v1/config/public", func(c *gin.Context) { - if p, err := core.Inject[contracts.PublicConfigProvider](ctx); err == nil && p != nil { - data, err := p.PublicConfig(c.Request.Context()) - if err != nil { - logger.ErrorF(c.Request.Context(), "[System] public config provider failed: %v", err) - response.AbortInternal(c, "public config unavailable") - return - } - c.JSON(http.StatusOK, response.OK(data)) + provider := resolvePublicConfigProvider(c.Request.Context(), ctx) + if provider == nil { + c.JSON(http.StatusOK, response.OK(map[string]string{})) return } - configs, err := listPublicSystemConfigs(c.Request.Context(), ctx) + data, err := provider.PublicConfig(c.Request.Context()) if err != nil { - logger.ErrorF(c.Request.Context(), "[System] query public system configs failed: %v", err) + logger.ErrorF(c.Request.Context(), "[System] public config provider failed: %v", err) + response.AbortInternal(c, "public config unavailable") + return } - c.JSON(http.StatusOK, response.OK(gin.H{ - "configs": configs, - "app": gin.H{ - "name": appName, - }, - })) - }) - - // 3. Custom injection - ctx.Router().GET("/custom", func(c *gin.Context) { - c.JSON(http.StatusOK, response.OK(gin.H{"custom": true})) + if data == nil { + data = map[string]string{} + } + c.JSON(http.StatusOK, response.OK(data)) }) + ctx.Router().RegisterWhitelist("/api/v1/config/public") return nil } + +func resolvePublicConfigProvider(reqCtx context.Context, appCtx *core.Context) contracts.PublicConfigProvider { + if p, err := core.InjectFrom[contracts.PublicConfigProvider](reqCtx); err == nil && p != nil { + return p + } + if appCtx != nil { + if p, err := core.Inject[contracts.PublicConfigProvider](appCtx); err == nil && p != nil { + return p + } + } + return nil +} diff --git a/backend/plugins/domain/system/public_config_test.go b/backend/plugins/domain/system/public_config_test.go index 47d8840f..d2065556 100644 --- a/backend/plugins/domain/system/public_config_test.go +++ b/backend/plugins/domain/system/public_config_test.go @@ -18,13 +18,17 @@ import ( "github.com/gin-gonic/gin" ) -type stubPublic struct{ payload any } +type stubPublic struct{ payload map[string]string } -func (s stubPublic) PublicConfig(context.Context) (any, error) { return s.payload, nil } +func (s stubPublic) PublicConfig(context.Context) (map[string]string, error) { + return s.payload, nil +} type errPublic struct{ err error } -func (s errPublic) PublicConfig(context.Context) (any, error) { return nil, s.err } +func (s errPublic) PublicConfig(context.Context) (map[string]string, error) { + return nil, s.err +} func TestPublicConfigUsesProviderWhenPresent(t *testing.T) { gin.SetMode(gin.TestMode) @@ -54,16 +58,22 @@ func TestPublicConfigDefaultWithoutProvider(t *testing.T) { if err := New().Apply(ctx); err != nil { t.Fatal(err) } + if !ctx.Router().IsWhitelisted("/api/v1/config/public") { + t.Fatal("GET /api/v1/config/public not whitelisted") + } raw := invokePublicConfig(t, publicConfigHandler(t, ctx)) var data map[string]any if err := json.Unmarshal(raw, &data); err != nil { t.Fatal(err) } - if _, ok := data["configs"]; !ok { - t.Fatalf("data = %s, want key configs", raw) + if len(data) != 0 { + t.Fatalf("data = %s, want empty flat map", raw) } - if _, ok := data["app"]; !ok { - t.Fatalf("data = %s, want key app", raw) + if _, ok := data["configs"]; ok { + t.Fatalf("data = %s, default payload must not wrap configs", raw) + } + if _, ok := data["app"]; ok { + t.Fatalf("data = %s, default payload must not wrap app", raw) } } diff --git a/backend/plugins/domain/system/repository.go b/backend/plugins/domain/system/repository.go deleted file mode 100644 index 612feb4c..00000000 --- a/backend/plugins/domain/system/repository.go +++ /dev/null @@ -1,35 +0,0 @@ -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package system - -import ( - "Wavelet/core" - "Wavelet/core/contracts" - "context" -) - -// publicSystemConfig 前端公共配置接口的只读投影。 -type publicSystemConfig struct { - Key string `json:"key"` - Value string `json:"value"` -} - -// listPublicSystemConfigs 读取对前端可见的系统配置项。 -// -// 注意:w_system_configs 的所有者插件是 admin,contracts 目前尚未暴露读取契约, -// 因此此处仍只能直连只读查询;待 admin 提供 SettingsService 契约后应改为调用契约。 -func listPublicSystemConfigs(ctx context.Context, appCtx *core.Context) ([]publicSystemConfig, error) { - dbSvc, err := core.Inject[contracts.DBService](appCtx) - if err != nil { - return nil, err - } - if dbSvc == nil { - return nil, nil - } - var configs []publicSystemConfig - err = dbSvc.DB(ctx).Table("w_system_configs"). - Where("visibility = ?", "visible"). - Find(&configs).Error - return configs, err -} diff --git a/backend/plugins/domain/upload/exports.go b/backend/plugins/domain/upload/exports.go index dc789f5f..596f7600 100644 --- a/backend/plugins/domain/upload/exports.go +++ b/backend/plugins/domain/upload/exports.go @@ -112,3 +112,6 @@ type RebuildUploadStatsHandler = uploadtask.RebuildUploadStatsHandler // WarmImageCachePayload is the payload for image cache warmup tasks. type WarmImageCachePayload = uploadtask.WarmImageCachePayload + +// CleanupOrphanUploads removes unconfirmed upload files. +var CleanupOrphanUploads = uploadtask.CleanupOrphanUploads diff --git a/backend/plugins/domain/upload/plugin.go b/backend/plugins/domain/upload/plugin.go index 3fa08eaf..39d5e91c 100644 --- a/backend/plugins/domain/upload/plugin.go +++ b/backend/plugins/domain/upload/plugin.go @@ -8,6 +8,7 @@ import ( "Wavelet/core" "Wavelet/core/contracts" "Wavelet/core/extpoints" + "Wavelet/pkg/ginutil" "Wavelet/plugins/domain/upload/filesrv" "Wavelet/plugins/domain/upload/handler" "Wavelet/plugins/domain/upload/shared" @@ -56,83 +57,57 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers upload routes, tasks, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { - // Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - shared.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - shared.SetDBService(db) - }) - } - - // Bind CacheService - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - shared.SetCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - shared.SetCacheService(cache) - }) - } - - // Bind StorageService - if storage, err := core.Inject[contracts.StorageService](ctx); err == nil && storage != nil { - shared.SetStorageService(storage) - } else { - core.When[contracts.StorageService](ctx, func(storage contracts.StorageService) { - shared.SetStorageService(storage) - }) - } - - // Bind TaskService - if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil { - shared.SetTaskService(taskSvc) - } else { - core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) { - shared.SetTaskService(taskSvc) - }) - } - - // Bind AuthService - if authSvc, err := core.Inject[contracts.AuthService](ctx); err == nil && authSvc != nil { - shared.SetAuthService(authSvc) - } else { - core.When[contracts.AuthService](ctx, func(authSvc contracts.AuthService) { - shared.SetAuthService(authSvc) - }) - } + core.Bind[contracts.DBService](ctx, shared.SetDBService) + core.Bind[contracts.CacheService](ctx, shared.SetCacheService) + core.Bind[contracts.StorageService](ctx, shared.SetStorageService) + core.Bind[contracts.TaskService](ctx, shared.SetTaskService) + core.Bind[contracts.AuthService](ctx, shared.SetAuthService) + core.Provide[contracts.UploadService](ctx, &uploadServiceImpl{}) ctx.OnDispose(func() error { shared.ResetServices() return nil }) - // 0. Resolve auth service for middleware - var authSvc contracts.AuthService - if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil { - return err + denyAuth := ginutil.AuthUnavailable() + loginMW := func(c *gin.Context) { + if svc := shared.GetAuthService(c.Request.Context()); svc != nil { + if mw, ok := svc.RequireAuthMiddleware().(gin.HandlerFunc); ok && mw != nil { + mw(c) + return + } + } + denyAuth(c) + } + adminMW := func(c *gin.Context) { + if svc := shared.GetAuthService(c.Request.Context()); svc != nil { + if mw, ok := svc.RequireAdminMiddleware().(gin.HandlerFunc); ok && mw != nil { + mw(c) + return + } + } + denyAuth(c) } - loginMW := authSvc.RequireAuthMiddleware().(gin.HandlerFunc) // 0a. Register migrations ctx.Migrations().Register("upload", uploadMigrations) // 1. Register File Server Routes - ctx.Router().GET("/f/:id", filesrv.ServeFileByID) + // TIP: loginMW populates AuthUserObjKey so private files can be checked for ownership. + ctx.Router().GET("/f/:id", loginMW, filesrv.ServeFileByID) // 2. Register User/Admin Upload HTTP Routes uploadGroup := ctx.Router().Group("/api/v1/upload", loginMW) { uploadGroup.POST("", handler.UploadFile) - uploadGroup.GET("", handler.ListFiles) - uploadGroup.DELETE("/:id", handler.DeleteFile) - uploadGroup.POST("/batch-download", handler.BatchDownloadFiles) + uploadGroup.DELETE("/:id", handler.DeleteMyFile) uploadGroup.GET("/my", handler.ListMyFiles) uploadGroup.PUT("/:id", handler.UpdateMyFile) uploadGroup.GET("/download/:id", handler.DownloadFile) uploadGroup.POST("/download/batch", handler.BatchDownloadFiles) } - adminUploadGroup := ctx.Router().Group("/api/v1/admin/uploads", loginMW) + adminUploadGroup := ctx.Router().Group("/api/v1/admin/uploads", loginMW, adminMW) { adminUploadGroup.GET("", handler.ListFiles) adminUploadGroup.GET("/stats", handler.GetFileStats) @@ -148,34 +123,16 @@ func (p *Plugin) Apply(ctx *core.Context) error { defaultSingleRetry = 1 ) - // 3. Register tasks. Handlers take raw payload bytes rather than a driver - // specific task type so they run under both the asynq and in-process workers. - cleanupHandler := &task.SystemCleanupHandler{} - ctx.Task().Register(task.SystemCleanupTask, func(c context.Context, payload []byte) error { - _, err := cleanupHandler.Execute(c, payload) - return err - }, extpoints.WithTaskMeta(task.SystemCleanupMeta), extpoints.WithTaskRetry(defaultCleanupRetry)) + ctx.Task().Register(task.SystemCleanupTask, &task.SystemCleanupHandler{}, extpoints.WithTaskMeta(task.SystemCleanupMeta), extpoints.WithTaskRetry(defaultCleanupRetry)) + ctx.Task().Register(task.RebuildUploadStatsTask, &task.RebuildUploadStatsHandler{}, extpoints.WithTaskMeta(task.RebuildUploadStatsMeta), extpoints.WithTaskRetry(defaultStatsRetry)) + ctx.Task().Register(task.StorageMigrationTask, &task.MigrationHandler{}, extpoints.WithTaskMeta(task.StorageMigrationMeta), extpoints.WithTaskRetry(defaultSingleRetry)) + ctx.Task().Register(task.WarmImageCacheTask, &task.WarmImageCacheHandler{}, extpoints.WithTaskMeta(task.WarmImageCacheMeta), extpoints.WithTaskRetry(1)) - rebuildStatsHandler := &task.RebuildUploadStatsHandler{} - ctx.Task().Register(task.RebuildUploadStatsTask, func(c context.Context, payload []byte) error { - _, err := rebuildStatsHandler.Execute(c, payload) + // 4. Register Event Listeners for domain events + ctx.Events().On(contracts.EventTopicSystemCleanup, func(c context.Context, _ contracts.SystemCleanupEvent) error { + _, _, err := task.CleanupOrphanUploads(c) return err - }, extpoints.WithTaskMeta(task.RebuildUploadStatsMeta), extpoints.WithTaskRetry(defaultStatsRetry)) - - migrationHandler := &task.MigrationHandler{} - ctx.Task().Register(task.StorageMigrationTask, func(c context.Context, payload []byte) error { - _, err := migrationHandler.Execute(c, payload) - return err - }, extpoints.WithTaskMeta(task.StorageMigrationMeta), extpoints.WithTaskRetry(defaultSingleRetry)) - - warmHandler := &task.WarmImageCacheHandler{} - ctx.Task().Register(task.WarmImageCacheTask, func(c context.Context, payload []byte) error { - _, err := warmHandler.Execute(c, payload) - return err - }, extpoints.WithTaskMeta(task.WarmImageCacheMeta), extpoints.WithTaskRetry(1)) - - // 4. Register Cron Schedule - ctx.Schedule().RegisterCron("0 3 * * *", task.SystemCleanupTask, nil) + }) // 5. Register Settings Schemas ctx.Settings().Register(extpoints.SettingSchema{ diff --git a/backend/plugins/domain/upload/plugin_test.go b/backend/plugins/domain/upload/plugin_test.go index 39de4d5a..b93059e4 100644 --- a/backend/plugins/domain/upload/plugin_test.go +++ b/backend/plugins/domain/upload/plugin_test.go @@ -23,6 +23,10 @@ func (stubAuthService) RequireAuthMiddleware() any { return gin.HandlerFunc(func(c *gin.Context) { c.Next() }) } +func (stubAuthService) RequireAdminMiddleware() any { + return gin.HandlerFunc(func(c *gin.Context) { c.Next() }) +} + func TestUserUploadRoutes(t *testing.T) { gin.SetMode(gin.TestMode) ctx := core.NewContext(context.Background()) @@ -34,12 +38,18 @@ func TestUserUploadRoutes(t *testing.T) { } want := []string{ + "POST /api/v1/upload", + "DELETE /api/v1/upload/:id", "GET /api/v1/upload/my", "PUT /api/v1/upload/:id", "GET /api/v1/upload/download/:id", "POST /api/v1/upload/download/batch", - "GET /api/v1/upload", - "POST /api/v1/upload/batch-download", + "GET /api/v1/admin/uploads", + "GET /api/v1/admin/uploads/stats", + "DELETE /api/v1/admin/uploads/:id", + "GET /api/v1/admin/uploads/download/:id", + "POST /api/v1/admin/uploads/download/batch", + "GET /api/v1/admin/uploads/types", } found := make(map[string]bool, len(want)) for _, rd := range ctx.Router().Routes() { diff --git a/backend/plugins/domain/upload/service.go b/backend/plugins/domain/upload/service.go new file mode 100644 index 00000000..9374eb00 --- /dev/null +++ b/backend/plugins/domain/upload/service.go @@ -0,0 +1,94 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package upload + +import ( + "Wavelet/core/contracts" + "Wavelet/plugins/domain/upload/models" + "Wavelet/plugins/domain/upload/repository" + "Wavelet/plugins/domain/upload/shared" + "context" + "errors" +) + +type uploadServiceImpl struct{} + +func (s *uploadServiceImpl) GetByID(ctx context.Context, id uint64) (*contracts.UploadDTO, error) { + u, err := repository.GetActiveUploadByID(ctx, id) + if err != nil { + return nil, err + } + dto := toUploadDTO(&u) + return &dto, nil +} + +func (s *uploadServiceImpl) OpenStoredUpload(ctx context.Context, id uint64) (*contracts.OpenedUploadDTO, error) { + u, err := repository.GetActiveUploadByID(ctx, id) + if err != nil { + return nil, err + } + storageSvc := shared.GetStorage(ctx) + if storageSvc == nil { + return nil, errors.New("storage service not available") + } + obj, err := storageSvc.Get(ctx, u.FilePath) + if err != nil { + return nil, err + } + return &contracts.OpenedUploadDTO{ + Upload: toUploadDTO(&u), + Body: obj.Body, + ContentType: obj.ContentType, + ContentLength: obj.ContentLength, + }, nil +} + +func (s *uploadServiceImpl) Remove(ctx context.Context, id uint64) error { + _, err := Remove(ctx, id) + return err +} + +func (s *uploadServiceImpl) RemoveOwned(ctx context.Context, id uint64, userID uint64) error { + _, err := RemoveOwned(ctx, userID, id) + return err +} + +func (s *uploadServiceImpl) FindByHash(ctx context.Context, hash string, size int64) (*contracts.UploadDTO, error) { + u, err := FindByHash(ctx, hash, size) + if err != nil { + return nil, err + } + dto := toUploadDTO(&u) + return &dto, nil +} + +func (s *uploadServiceImpl) RebuildStats(ctx context.Context) error { + return RebuildUploadStats(ctx) +} + +func toUploadDTO(u *models.Upload) contracts.UploadDTO { + return contracts.UploadDTO{ + ID: u.ID, + UserID: u.UserID, + FileName: u.FileName, + FilePath: u.FilePath, + MimeType: u.MimeType, + Size: u.FileSize, + Hash: u.Hash, + Status: string(u.Status), + Type: u.Type, + Metadata: contracts.UploadMetadataDTO{ + Width: u.Metadata.Width, + Height: u.Metadata.Height, + Duration: u.Metadata.Duration, + OriginalMime: u.Metadata.OriginalMime, + UserAgent: u.Metadata.UserAgent, + ClientIP: u.Metadata.ClientIP, + Bucket: u.Metadata.Bucket, + Extra: u.Metadata.Extra, + }, + CreatedAt: u.CreatedAt, + UpdatedAt: u.UpdatedAt, + } +} diff --git a/backend/plugins/domain/upload/shared/context_services.go b/backend/plugins/domain/upload/shared/context_services.go index cb6c2415..7e9494cd 100644 --- a/backend/plugins/domain/upload/shared/context_services.go +++ b/backend/plugins/domain/upload/shared/context_services.go @@ -69,10 +69,8 @@ func ResetServices() { // GetDB resolves the GORM DB instance. func GetDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) } svcMu.RLock() s := dbSvc @@ -85,10 +83,8 @@ func GetDB(ctx context.Context) *gorm.DB { // GetCache resolves the CacheService instance. func GetCache(ctx context.Context) contracts.CacheService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.CacheService](c); err == nil && s != nil { - return s - } + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s } svcMu.RLock() s := cacheSvc @@ -98,10 +94,8 @@ func GetCache(ctx context.Context) contracts.CacheService { // GetStorage resolves the StorageService instance. func GetStorage(ctx context.Context) contracts.StorageService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.StorageService](c); err == nil && s != nil { - return s - } + if s, err := core.InjectFrom[contracts.StorageService](ctx); err == nil && s != nil { + return s } svcMu.RLock() s := storageSvc @@ -110,7 +104,10 @@ func GetStorage(ctx context.Context) contracts.StorageService { } // GetTaskService resolves the TaskService instance. -func GetTaskService() contracts.TaskService { +func GetTaskService(ctx context.Context) contracts.TaskService { + if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil { + return s + } svcMu.RLock() defer svcMu.RUnlock() return taskSvc @@ -118,10 +115,8 @@ func GetTaskService() contracts.TaskService { // GetAuthService resolves the AuthService instance. func GetAuthService(ctx context.Context) contracts.AuthService { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.AuthService](c); err == nil && s != nil { - return s - } + if s, err := core.InjectFrom[contracts.AuthService](ctx); err == nil && s != nil { + return s } svcMu.RLock() s := authSvc diff --git a/backend/plugins/domain/upload/task/cleanup.go b/backend/plugins/domain/upload/task/cleanup.go index 7b09ee2c..743ba74e 100644 --- a/backend/plugins/domain/upload/task/cleanup.go +++ b/backend/plugins/domain/upload/task/cleanup.go @@ -33,9 +33,9 @@ const ( var SystemCleanupMeta = contracts.TaskMetaDTO{ Type: TaskTypeSystemCleanup, AsynqTask: SystemCleanupTask, - Name: "系统垃圾清理", - DisplayName: "系统垃圾清理", - Description: "定期清理未使用上传文件、历史推送记录和过期任务执行日志", + Name: "清理未确认上传文件", + DisplayName: "清理未确认上传文件", + Description: "定期清理超过1小时未使用的上传临时文件与底层存储资源", Category: "maintenance", SupportsTime: false, MaxRetry: 3, @@ -43,13 +43,29 @@ var SystemCleanupMeta = contracts.TaskMetaDTO{ Retryable: true, } -// SystemCleanupHandler 系统定期垃圾清理异步任务处理器 +// SystemCleanupHandler 未确认上传文件清理异步任务处理器 type SystemCleanupHandler struct{} -// Execute 执行系统清理(包含文件清理、历史推送日志和任务执行日志清理) +// Execute 执行系统清理(清理未使用上传文件) func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*contracts.TaskResultDTO, error) { + totalProcessed, totalDeleted, err := CleanupOrphanUploads(ctx) + if err != nil { + return nil, err + } + + msg := fmt.Sprintf( + "系统垃圾清理完成,处理未确认文件: %d 个,物理删除: %d 个", + totalProcessed, + totalDeleted, + ) + logger.InfoF(ctx, "%s", msg) + return &contracts.TaskResultDTO{Message: msg}, nil +} + +// CleanupOrphanUploads 扫描并清理超过1小时未确认的 pending 状态上传文件及物理存储 +func CleanupOrphanUploads(ctx context.Context) (int, int, error) { if uploadstorage.ReadOnly(ctx) { - return nil, errors.New(shared.ErrStorageReadOnly) + return 0, 0, errors.New(shared.ErrStorageReadOnly) } const batchSize = 100 var lastID uint64 @@ -62,14 +78,14 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*contract db := shared.GetDB(ctx) if db == nil { - return nil, errors.New("database service not available") + return 0, 0, errors.New("database service not available") } storageSvc := shared.GetStorage(ctx) for { if err := ctx.Err(); err != nil { - return nil, fmt.Errorf("system cleanup canceled: %w", err) + return totalProcessed, totalDeleted, fmt.Errorf("system cleanup canceled: %w", err) } var pendingUploads []models.Upload @@ -79,7 +95,7 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*contract Limit(batchSize). Find(&pendingUploads).Error; err != nil { logger.ErrorF(ctx, "查询过期待使用上传文件失败: %v", err) - return nil, fmt.Errorf("failed to query pending uploads: %w", err) + return totalProcessed, totalDeleted, fmt.Errorf("failed to query pending uploads: %w", err) } if len(pendingUploads) == 0 { @@ -88,7 +104,7 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*contract for i := range pendingUploads { if err := ctx.Err(); err != nil { - return nil, fmt.Errorf("system cleanup canceled: %w", err) + return totalProcessed, totalDeleted, fmt.Errorf("system cleanup canceled: %w", err) } upload := &pendingUploads[i] @@ -117,39 +133,5 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*contract } } - // 清理过期任务执行记录 - var deletedExecutions int64 - sevenDaysAgo := time.Now().Add(-7 * 24 * time.Hour) - if err := db. - Table("w_task_executions"). - Where("created_at < ?", sevenDaysAgo). - Delete(&struct{}{}).Error; err != nil { - logger.WarnF(ctx, "清理过期任务执行日志失败: %v", err) - } else { - deletedExecutions = db.RowsAffected - logger.InfoF(ctx, "已清理 7 天前任务执行日志,共 %d 条", deletedExecutions) - } - - // 清理已过期推送日志 - var deletedPushLogs int64 - thirtyDaysAgo := time.Now().Add(-30 * 24 * time.Hour) - if err := db. - Table("w_push_logs"). - Where("created_at < ?", thirtyDaysAgo). - Delete(&struct{}{}).Error; err != nil { - logger.WarnF(ctx, "清理历史推送日志失败: %v", err) - } else { - deletedPushLogs = db.RowsAffected - logger.InfoF(ctx, "已清理 30 天前推送日志,共 %d 条", deletedPushLogs) - } - - msg := fmt.Sprintf( - "系统垃圾清理完成,处理未确认文件: %d 个,物理删除: %d 个,清理过期任务日志: %d 条,清理历史推送日志: %d 条", - totalProcessed, - totalDeleted, - deletedExecutions, - deletedPushLogs, - ) - logger.InfoF(ctx, "%s", msg) - return &contracts.TaskResultDTO{Message: msg}, nil + return totalProcessed, totalDeleted, nil } diff --git a/backend/plugins/domain/user/errs.go b/backend/plugins/domain/user/errs.go index df4102f2..75c5073c 100644 --- a/backend/plugins/domain/user/errs.go +++ b/backend/plugins/domain/user/errs.go @@ -8,7 +8,9 @@ const ( errInvalidParams = "无效的请求参数" errUserNotFound = "用户不存在" //nolint:gosec // error message, not hardcoded credentials - errPasswordMismatch = "用户名或密码错误" + errPasswordMismatch = "用户名或密码错误" + errTooManyLoginAttempts = "登录尝试过于频繁,请稍后重试" + //nolint:gosec // error message, not hardcoded credentials //nolint:gosec // error message, not hardcoded credentials errOldPasswordIncorrect = "原密码不正确" //nolint:gosec // error message, not hardcoded credentials @@ -40,4 +42,12 @@ const ( //nolint:gosec // error message, not hardcoded credentials errServicePasswordTooShort = "密码长度至少为 8 位" errUniqueUsernameFailed = "failed to generate unique username" + errInvalidEmail = "邮箱地址无效" + errInvalidEmailCode = "验证码必须是 6 位数字" + errInvalidTaskPayload = "任务参数无效" + errMailSubjectRequired = "邮件主题不能为空" + errMailBodyRequired = "邮件内容不能为空" + errSMTPNotConfigured = "SMTP 未配置" + errEmailCacheUnavailable = "缓存服务不可用,无法保存验证码" + errSendEmailFailed = "邮件发送失败" ) diff --git a/backend/plugins/domain/user/handlers.go b/backend/plugins/domain/user/handlers.go index 8ddea90b..21512306 100644 --- a/backend/plugins/domain/user/handlers.go +++ b/backend/plugins/domain/user/handlers.go @@ -5,15 +5,19 @@ package user import ( "Wavelet/core/contracts" + "Wavelet/pkg/idgen" "Wavelet/pkg/logger" "Wavelet/pkg/response" + "Wavelet/pkg/util" "context" "crypto/rand" "crypto/sha256" "encoding/hex" + "encoding/json" "net/http" "strconv" "sync" + "time" "github.com/gin-contrib/sessions" "github.com/gin-gonic/gin" @@ -77,15 +81,16 @@ func invalidateTokenCache(ctx context.Context, tokenHash string) { } } -// Login 用户密码登录 -// @Summary 用户密码登录 -// @Description 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。 +// Login 用户登录 +// @Summary 用户登录 +// @Description 使用用户名和密码登录系统,验证通过后建立 Session 并返回用户信息。 // @Tags user // @Accept json // @Produce json // @Param request body user.loginRequest true "登录请求参数" // @Success 200 {object} response.Any "登录成功,返回用户信息" // @Failure 400 {object} response.Any "用户名或密码错误" +// @Failure 429 {object} response.Any "登录尝试过于频繁" // @Failure 500 {object} response.Any "服务内部错误" // @Router /api/v1/user/login [post] func Login(c *gin.Context) { @@ -95,8 +100,28 @@ func Login(c *gin.Context) { return } - user, err := GetUserByUsername(c.Request.Context(), req.Username) + ctx := c.Request.Context() + clientIP := c.ClientIP() + rateKey := "auth:login:ip:" + clientIP + if clientIP == "" { + rateKey = "auth:login:user:" + req.Username + } + + limiter := getLimiter(ctx) + if limiter != nil { + res, err := limiter.Allow(ctx, rateKey, contracts.Rate{ + Limit: 5, + Period: 1 * time.Minute, + }) + if err == nil && !res.Allowed { + response.AbortTooManyRequests(c, errTooManyLoginAttempts) + return + } + } + + user, err := GetUserByUsername(ctx, req.Username) if err != nil { + util.DummyCheckPassword(req.Password) response.AbortUnauthorized(c, errPasswordMismatch) return } @@ -106,14 +131,25 @@ func Login(c *gin.Context) { return } + if limiter != nil { + _ = limiter.Reset(ctx, rateKey) + } + + if user.ID == 0 { + newID := idgen.NextUint64ID() + if err := getDB(ctx).Model(&User{}).Where("username = ?", user.Username).Update("id", newID).Error; err == nil { + user.ID = newID + } + } + sess := sessions.Default(c) - sess.Set(contracts.AuthUserIDKey, user.ID) + sess.Set(contracts.AuthUserIDKey, strconv.FormatUint(user.ID, 10)) sess.Set(contracts.AuthUserNameKey, user.Username) needChange := user.NeedChangePassword || user.IsPlaintextPassword() user.NeedChangePassword = needChange sess.Set("need_change_password", needChange) if err := sess.Save(); err != nil { - logger.ErrorF(c.Request.Context(), "save session failed on login: %v", err) + logger.ErrorF(ctx, "save session failed on login: %v", err) } c.JSON(http.StatusOK, response.OK(user)) @@ -128,6 +164,7 @@ func Login(c *gin.Context) { // @Param request body user.registerRequest true "注册请求参数" // @Success 200 {object} response.Any "注册并登录成功,返回用户信息" // @Failure 400 {object} response.Any "参数错误、用户名已存在或注册已关闭" +// @Failure 429 {object} response.Any "注册尝试过于频繁" // @Failure 500 {object} response.Any "服务内部错误" // @Router /api/v1/user/register [post] func Register(c *gin.Context) { @@ -137,6 +174,19 @@ func Register(c *gin.Context) { return } + ctx := c.Request.Context() + clientIP := c.ClientIP() + if limiter := getLimiter(ctx); limiter != nil && clientIP != "" { + res, err := limiter.Allow(ctx, "auth:register:ip:"+clientIP, contracts.Rate{ + Limit: 10, + Period: 1 * time.Minute, + }) + if err == nil && !res.Allowed { + response.AbortTooManyRequests(c, errTooManyLoginAttempts) + return + } + } + newUser := &User{ Username: req.Username, Email: req.Email, @@ -153,7 +203,7 @@ func Register(c *gin.Context) { } sess := sessions.Default(c) - sess.Set(contracts.AuthUserIDKey, newUser.ID) + sess.Set(contracts.AuthUserIDKey, strconv.FormatUint(newUser.ID, 10)) sess.Set(contracts.AuthUserNameKey, newUser.Username) if err := sess.Save(); err != nil { logger.ErrorF(c.Request.Context(), "save session failed on register: %v", err) @@ -190,10 +240,36 @@ func Logout(c *gin.Context) { // @Tags user // @Accept json // @Produce json +// @Param request body user.sendEmailCodeRequest true "目标邮箱" // @Success 200 {object} response.Any "发送成功" // @Failure 400 {object} response.Any "参数错误" +// @Failure 500 {object} response.Any "发送失败" // @Router /api/v1/user/send-email-code [post] func SendEmailCode(c *gin.Context) { + var req sendEmailCodeRequest + if err := c.ShouldBindJSON(&req); err != nil { + response.AbortBadRequest(c, errInvalidParams) + return + } + payload, err := json.Marshal(sendEmailCodePayload{Email: req.Email}) + if err != nil { + response.AbortInternal(c, errSendEmailFailed) + return + } + ctx := c.Request.Context() + if taskSvc := getTaskService(ctx); taskSvc != nil { + if _, err := taskSvc.Dispatch(ctx, TaskTypeSendEmailCode, payload, contracts.TaskTriggerSystem); err != nil { + logger.ErrorF(ctx, "dispatch send_email_code failed: %v", err) + response.AbortInternal(c, errSendEmailFailed) + return + } + c.JSON(http.StatusOK, response.OK(gin.H{"sent": true})) + return + } + if _, err := (&SendEmailCodeHandler{}).Execute(ctx, payload); err != nil { + response.AbortBadRequest(c, err.Error()) + return + } c.JSON(http.StatusOK, response.OK(gin.H{"sent": true})) } diff --git a/backend/plugins/domain/user/login_rate_limit_test.go b/backend/plugins/domain/user/login_rate_limit_test.go new file mode 100644 index 00000000..4b5f1f89 --- /dev/null +++ b/backend/plugins/domain/user/login_rate_limit_test.go @@ -0,0 +1,118 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user_test + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "Wavelet/pkg/response" + "Wavelet/plugins/domain/auth" + "Wavelet/plugins/domain/user" + "Wavelet/plugins/infra/cache_memory" + database "Wavelet/plugins/infra/database" + "bytes" + "context" + "encoding/json" + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-contrib/sessions" + "github.com/gin-contrib/sessions/cookie" + "github.com/gin-gonic/gin" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestUserLoginRateLimiting(t *testing.T) { + gin.SetMode(gin.TestMode) + + ctx := core.NewContext(context.Background()) + ctx.Config().SetSource(core.NewMapSource(nil)) + require.NoError(t, ctx.Config().Resolve()) + + testDB := setupTestDB(t) + require.NoError(t, database.New(database.WithDB(testDB)).Apply(ctx)) + require.NoError(t, cache_memory.New().Apply(ctx)) + require.NoError(t, auth.New().Apply(ctx)) + require.NoError(t, user.New().Apply(ctx)) + + // Create a test user + userSvc, err := core.Inject[contracts.UserService](ctx) + require.NoError(t, err) + createdUser, err := userSvc.CreateUser(context.Background(), contracts.CreateUserRequest{ + Username: "ratelimit_user", + Password: "CorrectPassword123!", + }) + require.NoError(t, err) + require.NotNil(t, createdUser) + + router := gin.New() + router.Use(response.ErrorHandlerMiddleware()) + store := cookie.NewStore([]byte("test-secret-key-session")) + router.Use(sessions.Sessions("wavelet_session", store)) + router.Use(func(c *gin.Context) { + c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), ctx.Root())) + c.Next() + }) + + for _, rd := range ctx.Router().Routes() { + handlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers)) + for _, m := range rd.Middlewares { + if h, ok := m.(gin.HandlerFunc); ok { + handlers = append(handlers, h) + } else if fn, ok := m.(func(*gin.Context)); ok { + handlers = append(handlers, fn) + } + } + for _, raw := range rd.Handlers { + if h, ok := raw.(gin.HandlerFunc); ok { + handlers = append(handlers, h) + } else if fn, ok := raw.(func(*gin.Context)); ok { + handlers = append(handlers, fn) + } + } + router.Handle(rd.Method, rd.Path, handlers...) + } + + loginBody, _ := json.Marshal(map[string]string{ + "username": "ratelimit_user", + "password": "WrongPassword!", + }) + + // Make 5 failed login attempts (Limit is 5) + for i := 1; i <= 5; i++ { + req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewReader(loginBody)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = "192.168.1.100:12345" + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusUnauthorized, w.Code, "attempt %d should be 401 Unauthorized", i) + } + + // 6th attempt from the same IP should be blocked with 429 Too Many Requests + { + req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewReader(loginBody)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = "192.168.1.100:12345" + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusTooManyRequests, w.Code, "6th attempt should be 429 Too Many Requests") + + var resp map[string]any + err := json.Unmarshal(w.Body.Bytes(), &resp) + require.NoError(t, err) + assert.Equal(t, "登录尝试过于频繁,请稍后重试", resp["error_msg"]) + } + + // Another IP is not blocked + { + req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewReader(loginBody)) + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = "192.168.1.101:12345" + w := httptest.NewRecorder() + router.ServeHTTP(w, req) + assert.Equal(t, http.StatusUnauthorized, w.Code, "different IP should receive 401, not 429") + } +} diff --git a/backend/plugins/domain/user/models.go b/backend/plugins/domain/user/models.go index b742f9a9..9e367c57 100644 --- a/backend/plugins/domain/user/models.go +++ b/backend/plugins/domain/user/models.go @@ -109,6 +109,10 @@ type registerRequest struct { Email string `json:"email"` } +type sendEmailCodeRequest struct { + Email string `json:"email" binding:"required"` +} + // changePasswordRequest 修改密码请求参数 type changePasswordRequest struct { OldPassword string `json:"old_password" binding:"required"` diff --git a/backend/plugins/domain/user/plugin.go b/backend/plugins/domain/user/plugin.go index a46c37b4..6d41c35e 100644 --- a/backend/plugins/domain/user/plugin.go +++ b/backend/plugins/domain/user/plugin.go @@ -76,16 +76,15 @@ func (p *Plugin) Manifest() core.Manifest { // Apply registers user migrations, services, routes, tasks, schedules, and settings into the Context. func (p *Plugin) Apply(ctx *core.Context) error { - // 0. Bind DBService from Context - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - SetDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, SetDBService) + core.Bind[contracts.CacheService](ctx, SetCacheService) + core.Bind[contracts.TaskService](ctx, SetTaskService) + core.Bind[contracts.LimiterService](ctx, SetLimiterService) ctx.OnDispose(func() error { SetDBService(nil) + SetCacheService(nil) + SetTaskService(nil) + SetLimiterService(nil) return nil }) @@ -101,11 +100,8 @@ func (p *Plugin) Apply(ctx *core.Context) error { if mw, ok := authSvc.DisallowTokenAuthMiddleware().(gin.HandlerFunc); ok { noTokenMW = mw } - } else { - core.When[contracts.AuthService](ctx, func(svc contracts.AuthService) { - SetAuthService(svc) - }) } + core.Bind[contracts.AuthService](ctx, SetAuthService) ctx.OnDispose(func() error { SetAuthService(nil) return nil @@ -120,19 +116,12 @@ func (p *Plugin) Apply(ctx *core.Context) error { } core.Provide[contracts.UserService](ctx, p.userSvc) - passThrough := gin.HandlerFunc(func(c *gin.Context) { c.Next() }) - loginCap, registerCap, emailCap := passThrough, passThrough, passThrough - if capSvc, err := core.Inject[contracts.CaptchaService](ctx); err == nil && capSvc != nil { - if mw, ok := capSvc.VerifyMiddleware("login").(gin.HandlerFunc); ok { - loginCap = mw - } - if mw, ok := capSvc.VerifyMiddleware("register").(gin.HandlerFunc); ok { - registerCap = mw - } - if mw, ok := capSvc.VerifyMiddleware("send_email_code").(gin.HandlerFunc); ok { - emailCap = mw - } - } + // CAP middleware is resolved per request. user Apply runs before cap in + // the default plugin list; snapshotting CaptchaService here would leave + // login/register as a permanent pass-through. + loginCap := captchaGuard(ctx, "login") + registerCap := captchaGuard(ctx, "register") + emailCap := captchaGuard(ctx, "send_email_code") // 3. Register HTTP Routes userGroup := ctx.Router().Group("/api/v1/user") @@ -155,90 +144,19 @@ func (p *Plugin) Apply(ctx *core.Context) error { } } - const ( - defaultUserTaskRetry = 3 - paramTypeString = "string" - paramNameEmail = "email" - ) + ctx.Task().Register(TaskSendEmailCode, &SendEmailCodeHandler{}, + extpoints.WithTaskMeta(SendEmailCodeMeta), extpoints.WithTaskRetry(defaultUserTaskRetry)) + ctx.Task().Register(TaskSendMail, &SendMailHandler{}, + extpoints.WithTaskMeta(SendMailMeta), extpoints.WithTaskRetry(defaultUserTaskRetry)) + ctx.Task().Register(TaskCleanupInactive, &CleanupInactiveHandler{}, + extpoints.WithTaskMeta(CleanupInactiveMeta)) - // 4. Register background tasks - ctx.Task().Register("user:send_email_code", func(_ context.Context, _ []byte) error { - return nil - }, - extpoints.WithTaskType("send_email_code"), - extpoints.WithTaskName("发送邮箱验证码"), - extpoints.WithTaskDescription("异步发送用户注册与验证邮箱验证码"), - extpoints.WithTaskCategory("user"), - extpoints.WithTaskRetry(defaultUserTaskRetry), - extpoints.WithTaskQueue("default"), - extpoints.WithTaskRetryable(true), - extpoints.WithTaskParams( - contracts.TaskParamDTO{ - Name: paramNameEmail, - Label: "目标邮箱", - Type: paramTypeString, - Required: true, - Placeholder: "user@example.com", - Description: "接收验证码的目标邮箱", - }, - contracts.TaskParamDTO{ - Name: "code", - Label: "验证码", - Type: paramTypeString, - Required: true, - Placeholder: "123456", - Description: "6 位数字验证码", - }, - ), - ) - - ctx.Task().Register("mail:send", func(_ context.Context, _ []byte) error { - return nil - }, - extpoints.WithTaskType("send_email"), - extpoints.WithTaskName("发送邮件"), - extpoints.WithTaskDescription("异步发送系统邮件"), - extpoints.WithTaskCategory("mail"), - extpoints.WithTaskRetry(defaultUserTaskRetry), - extpoints.WithTaskQueue("default"), - extpoints.WithTaskRetryable(true), - extpoints.WithTaskParams( - contracts.TaskParamDTO{ - Name: "to", - Label: "接收邮箱 (To)", - Type: paramTypeString, - Required: true, - Placeholder: "receiver@example.com", - Description: "接收邮件的目标邮箱地址", - }, - contracts.TaskParamDTO{ - Name: "subject", - Label: "邮件主题 (Subject)", - Type: paramTypeString, - Required: true, - Placeholder: "请输入邮件主题", - Description: "发送邮件的主题标题", - }, - contracts.TaskParamDTO{ - Name: "body", - Label: "邮件内容 (Body)", - Type: "text", - Required: true, - Placeholder: "请输入邮件内容(支持 HTML格式)", - Description: "发送邮件的内容主体", - }, - ), - ) - - ctx.Task().Register("user:cleanup_inactive", func(_ context.Context, _ []byte) error { - return nil - }, - extpoints.WithTaskType("cleanup_inactive_users"), - extpoints.WithTaskName("清理未激活用户"), - extpoints.WithTaskDescription("清理长期未激活的注册用户与临时凭据"), - extpoints.WithTaskCategory("user"), - extpoints.WithTaskQueue("default"), - ) + // 4.1 Register Event Listeners for domain events + ctx.Events().On(contracts.EventTopicSystemCleanup, func(c context.Context, _ contracts.SystemCleanupEvent) error { + handler := &CleanupInactiveHandler{} + _, err := handler.Execute(c, nil) + return err + }) // 5. Register Settings Schemas ctx.Settings().Register(extpoints.SettingSchema{ @@ -267,3 +185,31 @@ func (p *Plugin) Apply(ctx *core.Context) error { return nil } + +func captchaGuard(appCtx *core.Context, scope string) gin.HandlerFunc { + return func(c *gin.Context) { + svc := resolveCaptchaService(c.Request.Context(), appCtx) + if svc == nil { + c.Next() + return + } + mw, ok := svc.VerifyMiddleware(scope).(gin.HandlerFunc) + if !ok || mw == nil { + c.Next() + return + } + mw(c) + } +} + +func resolveCaptchaService(reqCtx context.Context, appCtx *core.Context) contracts.CaptchaService { + if s, err := core.InjectFrom[contracts.CaptchaService](reqCtx); err == nil && s != nil { + return s + } + if appCtx != nil { + if s, err := core.Inject[contracts.CaptchaService](appCtx); err == nil && s != nil { + return s + } + } + return nil +} diff --git a/backend/plugins/domain/user/plugin_captcha_test.go b/backend/plugins/domain/user/plugin_captcha_test.go index 62d542a8..6fc31d60 100644 --- a/backend/plugins/domain/user/plugin_captcha_test.go +++ b/backend/plugins/domain/user/plugin_captcha_test.go @@ -5,6 +5,8 @@ package user_test import ( "context" + "net/http" + "net/http/httptest" "reflect" "testing" @@ -72,3 +74,77 @@ func TestApplyWithCaptchaServiceWrapsLogin(t *testing.T) { } t.Fatal("missing POST /api/v1/user/login") } + +type denyCaptchaService struct{} + +func (denyCaptchaService) VerifyMiddleware(string) any { + return gin.HandlerFunc(func(c *gin.Context) { + c.AbortWithStatus(http.StatusUnauthorized) + }) +} + +func (denyCaptchaService) ChallengeHandler() any { return gin.HandlerFunc(func(c *gin.Context) {}) } + +func (denyCaptchaService) RedeemHandler() any { return gin.HandlerFunc(func(c *gin.Context) {}) } + +func TestLoginCaptchaGuardResolvesServiceAfterApply(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx := core.NewContext(context.Background()) + if err := user.New().Apply(ctx); err != nil { + t.Fatal(err) + } + core.Provide[contracts.CaptchaService](ctx, denyCaptchaService{}) + + handler := loginCaptchaGuard(t, ctx) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/user/login", nil) + handler(c) + + if !c.IsAborted() { + t.Fatal("login captcha guard did not abort after late CaptchaService provide") + } + if w.Code != http.StatusUnauthorized { + t.Fatalf("status = %d, want %d", w.Code, http.StatusUnauthorized) + } +} + +func TestLoginCaptchaGuardPassesWithoutCaptchaService(t *testing.T) { + gin.SetMode(gin.TestMode) + ctx := core.NewContext(context.Background()) + if err := user.New().Apply(ctx); err != nil { + t.Fatal(err) + } + + handler := loginCaptchaGuard(t, ctx) + w := httptest.NewRecorder() + c, _ := gin.CreateTestContext(w) + c.Request = httptest.NewRequest(http.MethodPost, "/api/v1/user/login", nil) + handler(c) + + if c.IsAborted() { + t.Fatal("login captcha guard aborted without CaptchaService") + } +} + +func loginCaptchaGuard(t *testing.T, ctx *core.Context) gin.HandlerFunc { + t.Helper() + for _, rd := range ctx.Router().Routes() { + if rd.Method != "POST" || rd.Path != "/api/v1/user/login" { + continue + } + if len(rd.Handlers) == 0 { + t.Fatal("POST /api/v1/user/login has no handlers") + } + switch h := rd.Handlers[0].(type) { + case gin.HandlerFunc: + return h + case func(*gin.Context): + return h + default: + t.Fatalf("unexpected handler type %T", rd.Handlers[0]) + } + } + t.Fatal("missing POST /api/v1/user/login") + return nil +} diff --git a/backend/plugins/domain/user/plugin_test.go b/backend/plugins/domain/user/plugin_test.go index a5a03783..e4020e2a 100644 --- a/backend/plugins/domain/user/plugin_test.go +++ b/backend/plugins/domain/user/plugin_test.go @@ -166,6 +166,8 @@ func TestUserLoginHTTPHandler(t *testing.T) { IsActive: true, } require.NoError(t, user.CreateUser(context.Background(), plainUser)) + require.NotZero(t, plainUser.ID) + require.Greater(t, plainUser.ID, uint64(10000), "user id must be snowflake, not sqlite autoincrement") reqBodyPlain := `{"username":"plain_admin","password":"12345678"}` reqPlain, _ := http.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewBufferString(reqBodyPlain)) diff --git a/backend/plugins/domain/user/repository.go b/backend/plugins/domain/user/repository.go index 2bb02a9e..a40e68a3 100644 --- a/backend/plugins/domain/user/repository.go +++ b/backend/plugins/domain/user/repository.go @@ -6,17 +6,22 @@ package user import ( "Wavelet/core" "Wavelet/core/contracts" + "Wavelet/pkg/idgen" "Wavelet/pkg/util" "context" + "errors" "strings" "sync" + "time" "gorm.io/gorm" ) var ( - dbMu sync.RWMutex - dbSvc contracts.DBService + dbMu sync.RWMutex + dbSvc contracts.DBService + limiterMu sync.RWMutex + limiterSvc contracts.LimiterService ) // SetDBService sets the active DBService contract for the user domain plugin. @@ -26,11 +31,16 @@ func SetDBService(s contracts.DBService) { dbSvc = s } +// SetLimiterService sets the active LimiterService contract for the user domain plugin. +func SetLimiterService(s contracts.LimiterService) { + limiterMu.Lock() + defer limiterMu.Unlock() + limiterSvc = s +} + func getDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) } dbMu.RLock() @@ -43,6 +53,17 @@ func getDB(ctx context.Context) *gorm.DB { return nil } +func getLimiter(ctx context.Context) contracts.LimiterService { + if s, err := core.InjectFrom[contracts.LimiterService](ctx); err == nil && s != nil { + return s + } + + limiterMu.RLock() + s := limiterSvc + limiterMu.RUnlock() + return s +} + // GetUserByID 通过 ID 获取用户 func GetUserByID(ctx context.Context, id uint64) (*User, error) { var u User @@ -82,8 +103,11 @@ func GetUserByEmail(ctx context.Context, email string) (*User, error) { return &u, nil } -// CreateUser 创建用户 +// CreateUser 创建用户。ID 为空时用雪花算法分配,避免 SQLite/GORM 把 0 当成自增主键。 func CreateUser(ctx context.Context, u *User) error { + if u != nil && u.ID == 0 { + u.ID = idgen.NextUint64ID() + } return getDB(ctx).Create(u).Error } @@ -92,27 +116,6 @@ func UpdateUser(ctx context.Context, u *User) error { return getDB(ctx).Save(u).Error } -// ListUsers 分页查询用户 -func ListUsers(ctx context.Context, page, pageSize int, keyword string) ([]*User, int64, error) { - db := getDB(ctx).Model(&User{}) - if keyword != "" { - escaped := util.EscapeLike(keyword) - db = db.Where("username LIKE ? ESCAPE '\\' OR nickname LIKE ? ESCAPE '\\' OR email LIKE ? ESCAPE '\\'", "%"+escaped+"%", "%"+escaped+"%", "%"+escaped+"%") - } - - var total int64 - if err := db.Count(&total).Error; err != nil { - return nil, 0, err - } - - var users []*User - offset := (page - 1) * pageSize - if err := db.Offset(offset).Limit(pageSize).Order("id DESC").Find(&users).Error; err != nil { - return nil, 0, err - } - return users, total, nil -} - // GetAccessTokenByHash 通过 Hash 查询访问令牌 func GetAccessTokenByHash(ctx context.Context, tokenHash string) (*AccessToken, error) { var token AccessToken @@ -180,6 +183,26 @@ func DeleteUserWithRelations(ctx context.Context, id uint64) error { }) } +// ListInactiveNeverLoggedInUserIDs returns non-admin users created before cutoff +// who have never logged in. Seeded admin/system accounts are excluded. +func ListInactiveNeverLoggedInUserIDs(ctx context.Context, cutoff time.Time) ([]uint64, error) { + db := getDB(ctx) + if db == nil { + return nil, errors.New("database not available") + } + var ids []uint64 + unixEpoch := time.Unix(0, 0).UTC() + err := db.Model(&User{}). + Where("is_admin = ? AND username NOT IN ?", false, []string{"admin", "system"}). + Where("created_at < ?", cutoff). + Where("last_login_at IS NULL OR last_login_at < ?", unixEpoch). + Pluck("id", &ids).Error + if err != nil { + return nil, err + } + return ids, nil +} + // GetFirstAdminUser 获取第一个管理员用户 func GetFirstAdminUser(ctx context.Context) (*User, error) { var u User diff --git a/backend/plugins/domain/user/session_access_test.go b/backend/plugins/domain/user/session_access_test.go new file mode 100644 index 00000000..53ff5eec --- /dev/null +++ b/backend/plugins/domain/user/session_access_test.go @@ -0,0 +1,291 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user_test + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "Wavelet/pkg/response" + "Wavelet/plugins/domain/auth" + "Wavelet/plugins/domain/upload" + uploadmodels "Wavelet/plugins/domain/upload/models" + "Wavelet/plugins/domain/user" + "bytes" + "context" + "database/sql" + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "testing" + + database "Wavelet/plugins/infra/database" + + "github.com/gin-contrib/sessions" + "github.com/gin-contrib/sessions/cookie" + "github.com/gin-gonic/gin" +) + +type stubStorageService struct{ contracts.StorageService } + +type loginEnvelope struct { + ErrorMsg string `json:"error_msg"` + Data json.RawMessage `json:"data"` +} + +func mountUserAuthEngine(t *testing.T) (*gin.Engine, contracts.UserService) { + t.Helper() + gin.SetMode(gin.TestMode) + + ctx := core.NewContext(context.Background()) + ctx.Config().SetSource(core.NewMapSource(nil)) + if err := ctx.Config().Resolve(); err != nil { + t.Fatalf("Config.Resolve() error = %v", err) + } + testDB := setupTestDB(t) + if err := testDB.AutoMigrate(&uploadmodels.Upload{}); err != nil { + t.Fatalf("AutoMigrate(Upload) error = %v", err) + } + if err := database.New(database.WithDB(testDB)).Apply(ctx); err != nil { + t.Fatalf("database.Apply() error = %v", err) + } + if err := auth.New().Apply(ctx); err != nil { + t.Fatalf("auth.Apply() error = %v", err) + } + if err := user.New().Apply(ctx); err != nil { + t.Fatalf("user.Apply() error = %v", err) + } + core.Provide[contracts.StorageService](ctx, stubStorageService{}) + if err := upload.New().Apply(ctx); err != nil { + t.Fatalf("upload.Apply() error = %v", err) + } + + userSvc, err := core.Inject[contracts.UserService](ctx) + if err != nil || userSvc == nil { + t.Fatalf("Inject UserService: svc=%v err=%v", userSvc, err) + } + + engine := gin.New() + engine.Use(response.ErrorHandlerMiddleware()) + engine.Use(sessions.Sessions("wavelet_session_id", cookie.NewStore([]byte("test-session-secret")))) + engine.Use(func(c *gin.Context) { + c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), ctx.Root())) + c.Next() + }) + + for _, rd := range ctx.Router().Routes() { + handlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers)) + for _, m := range rd.Middlewares { + h, ok := m.(gin.HandlerFunc) + if !ok { + fn, ok := m.(func(*gin.Context)) + if !ok { + t.Fatalf("unsupported middleware type %T for %s %s", m, rd.Method, rd.Path) + } + h = fn + } + handlers = append(handlers, h) + } + for _, raw := range rd.Handlers { + h, ok := raw.(gin.HandlerFunc) + if !ok { + fn, ok := raw.(func(*gin.Context)) + if !ok { + t.Fatalf("unsupported handler type %T for %s %s", raw, rd.Method, rd.Path) + } + h = fn + } + handlers = append(handlers, h) + } + engine.Handle(rd.Method, rd.Path, handlers...) + } + + return engine, userSvc +} + +func loginAndCookie(t *testing.T, engine *gin.Engine, username, password string) []*http.Cookie { + t.Helper() + body := `{"username":"` + username + `","password":"` + password + `"}` + req := httptest.NewRequest(http.MethodPost, "/api/v1/user/login", bytes.NewBufferString(body)) + req.Header.Set("Content-Type", "application/json") + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("POST /api/v1/user/login username=%s status = %d, want 200 body=%s", username, rec.Code, rec.Body.String()) + } + cookies := rec.Result().Cookies() + if len(cookies) == 0 { + t.Fatalf("POST /api/v1/user/login username=%s Set-Cookie missing, headers=%v", username, rec.Header()) + } + return cookies +} + +func getWithCookies(engine *gin.Engine, path string, cookies []*http.Cookie) *httptest.ResponseRecorder { + req := httptest.NewRequest(http.MethodGet, path, nil) + for _, c := range cookies { + req.AddCookie(c) + } + rec := httptest.NewRecorder() + engine.ServeHTTP(rec, req) + return rec +} + +func TestNonAdminSessionCanAccessProtectedAPIs(t *testing.T) { + engine, userSvc := mountUserAuthEngine(t) + bg := context.Background() + + admin, err := userSvc.CreateUser(bg, contracts.CreateUserRequest{ + Username: "admin_user", + Password: "Password123!", + Email: "admin_user@example.com", + IsAdmin: true, + }) + if err != nil { + t.Fatalf("CreateUser(admin) error = %v", err) + } + if err := userSvc.SetUserAdmin(bg, admin.ID, true); err != nil { + t.Fatalf("SetUserAdmin() error = %v", err) + } + + member, err := userSvc.CreateUser(bg, contracts.CreateUserRequest{ + Username: "plain_user", + Password: "Password123!", + Email: "plain_user@example.com", + IsAdmin: false, + }) + if err != nil { + t.Fatalf("CreateUser(member) error = %v", err) + } + if member.IsAdmin { + t.Fatalf("CreateUser(member).IsAdmin = true, want false") + } + + cases := []struct { + name string + username string + wantID uint64 + }{ + {name: "admin", username: "admin_user", wantID: admin.ID}, + {name: "non-admin", username: "plain_user", wantID: member.ID}, + } + + protected := []string{ + "/api/v1/user/self", + "/api/v1/user-info", + "/api/v1/upload/my?page=1&page_size=12", + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + cookies := loginAndCookie(t, engine, tc.username, "Password123!") + for _, path := range protected { + rec := getWithCookies(engine, path, cookies) + if rec.Code != http.StatusOK { + t.Errorf("GET %s as %s status = %d, want 200 body=%s", path, tc.name, rec.Code, rec.Body.String()) + continue + } + var env loginEnvelope + if err := json.Unmarshal(rec.Body.Bytes(), &env); err != nil { + t.Errorf("GET %s as %s decode error = %v body=%s", path, tc.name, err, rec.Body.String()) + continue + } + if env.ErrorMsg != "" { + t.Errorf("GET %s as %s error_msg = %q, want empty", path, tc.name, env.ErrorMsg) + } + if bytes.Contains(env.Data, []byte(`"username"`)) { + var payload struct { + ID json.RawMessage `json:"id"` + Username string `json:"username"` + } + if err := json.Unmarshal(env.Data, &payload); err != nil { + t.Errorf("GET %s as %s data decode error = %v data=%s", path, tc.name, err, string(env.Data)) + continue + } + if payload.Username != tc.username { + t.Errorf("GET %s as %s username = %q, want %q", path, tc.name, payload.Username, tc.username) + } + if len(payload.ID) == 0 || payload.ID[0] != '"' { + t.Errorf("GET %s as %s id JSON = %s, want a string (snowflake ids exceed JS MAX_SAFE_INTEGER)", path, tc.name, payload.ID) + } + } + } + }) + } +} + +func TestLoginBackfillsNullUserIDSoProtectedAPIsSucceed(t *testing.T) { + engine, _ := mountUserAuthEngine(t) + db := database.DB(context.Background()) + if db == nil { + t.Fatal("database.DB() = nil, want the test database") + } + + legacy := user.User{ + Username: "legacy_zero", + Email: "legacy_zero@example.com", + IsActive: true, + } + if err := legacy.SetEncryptedPassword("Password123!"); err != nil { + t.Fatalf("SetEncryptedPassword() error = %v", err) + } + sqlDB, err := db.DB() + if err != nil { + t.Fatalf("db.DB() error = %v", err) + } + if _, err := sqlDB.Exec( + "INSERT INTO w_users (id, username, password, email, is_active, is_admin) VALUES (0, ?, ?, ?, 1, 0)", + legacy.Username, legacy.Password, legacy.Email, + ); err != nil { + t.Fatalf("INSERT legacy user error = %v", err) + } + var stored sql.NullInt64 + if err := sqlDB.QueryRow("SELECT id FROM w_users WHERE username = ?", legacy.Username).Scan(&stored); err != nil { + t.Fatalf("SELECT id error = %v", err) + } + if stored.Valid && stored.Int64 != 0 { + t.Fatalf("legacy user id = %d, want 0 or NULL to reproduce the 401", stored.Int64) + } + + cookies := loginAndCookie(t, engine, legacy.Username, "Password123!") + protected := []string{ + "/api/v1/user/self", + "/api/v1/user-info", + "/api/v1/upload/my?page=1&page_size=12", + } + for _, path := range protected { + rec := getWithCookies(engine, path, cookies) + if rec.Code != http.StatusOK { + t.Errorf("GET %s as legacy_zero status = %d, want 200 body=%s", path, rec.Code, rec.Body.String()) + continue + } + var env loginEnvelope + if err := json.Unmarshal(rec.Body.Bytes(), &env); err != nil { + t.Errorf("GET %s as legacy_zero decode error = %v body=%s", path, err, rec.Body.String()) + continue + } + if env.ErrorMsg != "" { + t.Errorf("GET %s as legacy_zero error_msg = %q, want empty", path, env.ErrorMsg) + } + } + + var backfilled uint64 + if err := db.Raw("SELECT id FROM w_users WHERE username = ?", legacy.Username).Scan(&backfilled).Error; err != nil { + t.Fatalf("SELECT backfilled id error = %v", err) + } + if backfilled == 0 { + t.Errorf("legacy user id after login = 0, want a snowflake id") + } +} + +func TestLoginRequiredRejectsMissingSession(t *testing.T) { + engine, _ := mountUserAuthEngine(t) + rec := getWithCookies(engine, "/api/v1/user/self", nil) + if rec.Code != http.StatusUnauthorized { + t.Errorf("GET /api/v1/user/self without cookie status = %d, want 401 body=%s", rec.Code, rec.Body.String()) + } + body, _ := io.ReadAll(rec.Body) + if !bytes.Contains(body, []byte("未登录")) && !bytes.Contains(body, []byte("用户不存在")) { + t.Errorf("GET /api/v1/user/self without cookie body = %s, want 未登录", body) + } +} diff --git a/backend/plugins/domain/user/task.go b/backend/plugins/domain/user/task.go new file mode 100644 index 00000000..439079aa --- /dev/null +++ b/backend/plugins/domain/user/task.go @@ -0,0 +1,385 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "Wavelet/pkg/logger" + pkgmail "Wavelet/pkg/mail" + "context" + "crypto/rand" + "encoding/binary" + "encoding/json" + "errors" + "fmt" + "net/mail" + "strconv" + "strings" + "sync" + "time" + "unicode" +) + +const ( + // TaskSendEmailCode is the queue pattern for email verification codes. + TaskSendEmailCode = "user:send_email_code" + // TaskTypeSendEmailCode is the admin type identifier for email verification codes. + TaskTypeSendEmailCode = "send_email_code" + // TaskSendMail is the queue pattern for generic outbound mail. + TaskSendMail = "mail:send" + // TaskTypeSendMail is the admin type identifier for generic outbound mail. + TaskTypeSendMail = "send_email" + // TaskCleanupInactive is the queue pattern for inactive-user cleanup. + TaskCleanupInactive = "user:cleanup_inactive" + // TaskTypeCleanupInactive is the admin type identifier for inactive-user cleanup. + TaskTypeCleanupInactive = "cleanup_inactive_users" + + defaultUserTaskRetry = 3 + emailCodeTTL = 10 * time.Minute + emailCodeCacheKeyPrefix = "user:email_code:" + inactiveRetentionDays = 30 + hoursPerDay = 24 + inactiveRetention = inactiveRetentionDays * hoursPerDay * time.Hour + smtpConfigKeyHost = "smtp_host" + smtpConfigKeyPort = "smtp_port" + smtpConfigKeyUsername = "smtp_username" + smtpConfigKeyPassword = "smtp_password" + defaultSMTPPort = 587 + emailCodeLength = 6 + emailCodeModulo = 1000000 + taskQueueDefault = "default" + taskParamTypeString = "string" + taskParamTypeText = "text" + paramNameEmail = "email" +) + +var smtpConfigKeys = []string{ + smtpConfigKeyHost, smtpConfigKeyPort, smtpConfigKeyUsername, smtpConfigKeyPassword, +} + +var ( + cacheMu sync.RWMutex + cacheSvc contracts.CacheService + taskMu sync.RWMutex + taskSvc contracts.TaskService +) + +// SetCacheService sets the cache contract used to store email verification codes. +func SetCacheService(s contracts.CacheService) { + cacheMu.Lock() + defer cacheMu.Unlock() + cacheSvc = s +} + +// SetTaskService sets the task contract used by HTTP handlers to enqueue mail jobs. +func SetTaskService(s contracts.TaskService) { + taskMu.Lock() + defer taskMu.Unlock() + taskSvc = s +} + +func getCache(ctx context.Context) contracts.CacheService { + if s, err := core.InjectFrom[contracts.CacheService](ctx); err == nil && s != nil { + return s + } + cacheMu.RLock() + defer cacheMu.RUnlock() + return cacheSvc +} + +func getTaskService(ctx context.Context) contracts.TaskService { + if s, err := core.InjectFrom[contracts.TaskService](ctx); err == nil && s != nil { + return s + } + taskMu.RLock() + defer taskMu.RUnlock() + return taskSvc +} + +func appendTaskLog(ctx context.Context, format string, args ...any) { + if svc := getTaskService(ctx); svc != nil { + svc.AppendLog(ctx, format, args...) + } +} + +// SendEmailCodeMeta describes the email verification-code task. +var SendEmailCodeMeta = contracts.TaskMetaDTO{ + Type: TaskTypeSendEmailCode, + AsynqTask: TaskSendEmailCode, + Name: "发送邮箱验证码", + DisplayName: "发送邮箱验证码", + Description: "异步发送用户注册与验证邮箱验证码", + Category: "user", + MaxRetry: defaultUserTaskRetry, + Queue: taskQueueDefault, + Retryable: true, + Params: []contracts.TaskParamDTO{ + {Name: paramNameEmail, Label: "目标邮箱", Type: taskParamTypeString, Required: true, Placeholder: "user@example.com", Description: "接收验证码的目标邮箱"}, + {Name: "code", Label: "验证码", Type: taskParamTypeString, Required: false, Placeholder: "123456", Description: "6 位数字验证码,留空则自动生成"}, + }, +} + +// SendMailMeta describes the generic outbound-mail task. +var SendMailMeta = contracts.TaskMetaDTO{ + Type: TaskTypeSendMail, + AsynqTask: TaskSendMail, + Name: "发送邮件", + DisplayName: "发送邮件", + Description: "异步发送系统邮件", + Category: "mail", + MaxRetry: defaultUserTaskRetry, + Queue: taskQueueDefault, + Retryable: true, + Params: []contracts.TaskParamDTO{ + {Name: "to", Label: "接收邮箱 (To)", Type: taskParamTypeString, Required: true, Placeholder: "receiver@example.com", Description: "接收邮件的目标邮箱地址"}, + {Name: "subject", Label: "邮件主题 (Subject)", Type: taskParamTypeString, Required: true, Placeholder: "请输入邮件主题", Description: "发送邮件的主题标题"}, + {Name: "body", Label: "邮件内容 (Body)", Type: taskParamTypeText, Required: true, Placeholder: "请输入邮件内容(支持 HTML格式)", Description: "发送邮件的内容主体"}, + }, +} + +// CleanupInactiveMeta describes the inactive-user cleanup task. +var CleanupInactiveMeta = contracts.TaskMetaDTO{ + Type: TaskTypeCleanupInactive, + AsynqTask: TaskCleanupInactive, + Name: "清理未激活用户", + DisplayName: "清理未激活用户", + Description: "清理长期未登录的注册用户及其访问令牌", + Category: "user", + Queue: taskQueueDefault, + Retryable: true, +} + +type sendEmailCodePayload struct { + Email string `json:"email"` + Code string `json:"code"` +} + +type sendMailPayload struct { + To string `json:"to"` + Subject string `json:"subject"` + Body string `json:"body"` +} + +// SendEmailCodeHandler sends a 6-digit email verification code and caches it. +type SendEmailCodeHandler struct{} + +// ValidatePayload checks the destination address and optional code. +func (h *SendEmailCodeHandler) ValidatePayload(payload []byte) ([]byte, error) { + p, err := parseSendEmailCodePayload(payload) + if err != nil { + return nil, err + } + return json.Marshal(p) +} + +// Execute generates (if needed), caches, and emails the verification code. +func (h *SendEmailCodeHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) { + p, err := parseSendEmailCodePayload(payload) + if err != nil { + return nil, err + } + if p.Code == "" { + p.Code, err = generateEmailCode() + if err != nil { + return nil, err + } + } + + cache := getCache(ctx) + if cache == nil { + return nil, errors.New(errEmailCacheUnavailable) + } + if err := cache.Set(ctx, emailCodeCacheKey(p.Email), p.Code, emailCodeTTL); err != nil { + return nil, fmt.Errorf("store email code: %w", err) + } + + cfg, err := loadSMTPConfig(ctx) + if err != nil { + return nil, err + } + subject := "邮箱验证码" + body := fmt.Sprintf("

您的验证码是 %s,%d 分钟内有效。

", p.Code, int(emailCodeTTL.Minutes())) + appendTaskLog(ctx, "发送邮箱验证码到 %s", maskEmail(p.Email)) + if err := pkgmail.SendMail(ctx, cfg, p.Email, subject, body); err != nil { + logger.ErrorF(ctx, "send email code failed: %v", err) + return nil, errors.New(errSendEmailFailed) + } + return &contracts.TaskResultDTO{Message: fmt.Sprintf("验证码已发送至 %s", maskEmail(p.Email))}, nil +} + +// SendMailHandler sends a generic HTML email through the configured SMTP server. +type SendMailHandler struct{} + +// ValidatePayload checks to/subject/body. +func (h *SendMailHandler) ValidatePayload(payload []byte) ([]byte, error) { + p, err := parseSendMailPayload(payload) + if err != nil { + return nil, err + } + return json.Marshal(p) +} + +// Execute sends the mail. +func (h *SendMailHandler) Execute(ctx context.Context, payload []byte) (*contracts.TaskResultDTO, error) { + p, err := parseSendMailPayload(payload) + if err != nil { + return nil, err + } + cfg, err := loadSMTPConfig(ctx) + if err != nil { + return nil, err + } + appendTaskLog(ctx, "发送邮件到 %s,主题: %s", maskEmail(p.To), p.Subject) + if err := pkgmail.SendMail(ctx, cfg, p.To, p.Subject, p.Body); err != nil { + logger.ErrorF(ctx, "send mail failed: %v", err) + return nil, errors.New(errSendEmailFailed) + } + return &contracts.TaskResultDTO{Message: fmt.Sprintf("邮件已发送至 %s", maskEmail(p.To))}, nil +} + +// CleanupInactiveHandler deletes users who registered long ago and never logged in. +type CleanupInactiveHandler struct{} + +// Execute removes stale never-logged-in non-admin users and their access tokens. +func (h *CleanupInactiveHandler) Execute(ctx context.Context, _ []byte) (*contracts.TaskResultDTO, error) { + cutoff := time.Now().Add(-inactiveRetention) + ids, err := ListInactiveNeverLoggedInUserIDs(ctx, cutoff) + if err != nil { + return nil, err + } + appendTaskLog(ctx, "扫描到 %d 个超过 %d 天未登录的注册用户", len(ids), int(inactiveRetention.Hours()/float64(hoursPerDay))) + deleted := 0 + for _, id := range ids { + if err := DeleteUserWithRelations(ctx, id); err != nil { + logger.ErrorF(ctx, "cleanup inactive user %d failed: %v", id, err) + continue + } + deleted++ + } + msg := fmt.Sprintf("已清理 %d 个长期未登录用户及其访问令牌", deleted) + appendTaskLog(ctx, "%s", msg) + return &contracts.TaskResultDTO{Message: msg}, nil +} + +func parseSendEmailCodePayload(payload []byte) (sendEmailCodePayload, error) { + var p sendEmailCodePayload + if len(payload) > 0 { + if err := json.Unmarshal(payload, &p); err != nil { + return p, errors.New(errInvalidTaskPayload) + } + } + p.Email = normalizeEmail(p.Email) + if err := validateEmail(p.Email); err != nil { + return p, err + } + p.Code = strings.TrimSpace(p.Code) + if p.Code != "" && !isSixDigitCode(p.Code) { + return p, errors.New(errInvalidEmailCode) + } + return p, nil +} + +func parseSendMailPayload(payload []byte) (sendMailPayload, error) { + var p sendMailPayload + if err := json.Unmarshal(payload, &p); err != nil { + return p, errors.New(errInvalidTaskPayload) + } + p.To = normalizeEmail(p.To) + p.Subject = strings.TrimSpace(p.Subject) + if err := validateEmail(p.To); err != nil { + return p, err + } + if p.Subject == "" { + return p, errors.New(errMailSubjectRequired) + } + if strings.TrimSpace(p.Body) == "" { + return p, errors.New(errMailBodyRequired) + } + return p, nil +} + +func loadSMTPConfig(ctx context.Context) (pkgmail.Config, error) { + db := getDB(ctx) + if db == nil { + return pkgmail.Config{}, errors.New(errSMTPNotConfigured) + } + var rows []struct { + Key string + Value string + } + if err := db.Table("w_system_configs"). + Select("key", "value"). + Where("key IN ?", smtpConfigKeys). + Find(&rows).Error; err != nil { + return pkgmail.Config{}, fmt.Errorf("read smtp config: %w", err) + } + cfg := pkgmail.Config{Port: defaultSMTPPort} + for _, row := range rows { + switch row.Key { + case smtpConfigKeyHost: + cfg.Host = strings.TrimSpace(row.Value) + case smtpConfigKeyPort: + if n, err := strconv.Atoi(strings.TrimSpace(row.Value)); err == nil && n > 0 { + cfg.Port = n + } + case smtpConfigKeyUsername: + cfg.Username = strings.TrimSpace(row.Value) + case smtpConfigKeyPassword: + cfg.Password = row.Value + } + } + if cfg.Host == "" || cfg.Username == "" { + return pkgmail.Config{}, errors.New(errSMTPNotConfigured) + } + return cfg, nil +} + +func generateEmailCode() (string, error) { + var buf [4]byte + if _, err := rand.Read(buf[:]); err != nil { + return "", err + } + n := binary.BigEndian.Uint32(buf[:]) % emailCodeModulo + return fmt.Sprintf("%06d", n), nil +} + +func emailCodeCacheKey(email string) string { + return emailCodeCacheKeyPrefix + normalizeEmail(email) +} + +func normalizeEmail(email string) string { + return strings.ToLower(strings.TrimSpace(email)) +} + +func validateEmail(email string) error { + if email == "" { + return errors.New(errEmailEmpty) + } + addr, err := mail.ParseAddress(email) + if err != nil || !strings.EqualFold(addr.Address, email) { + return errors.New(errInvalidEmail) + } + return nil +} + +func isSixDigitCode(code string) bool { + if len(code) != emailCodeLength { + return false + } + for _, r := range code { + if !unicode.IsDigit(r) { + return false + } + } + return true +} + +func maskEmail(email string) string { + at := strings.IndexByte(email, '@') + if at <= 1 { + return "***" + } + return email[:1] + "***" + email[at:] +} diff --git a/backend/plugins/domain/user/task_test.go b/backend/plugins/domain/user/task_test.go new file mode 100644 index 00000000..8c4ccb2b --- /dev/null +++ b/backend/plugins/domain/user/task_test.go @@ -0,0 +1,109 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package user_test + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/idgen" + "Wavelet/plugins/domain/user" + "context" + "path/filepath" + "testing" + "time" + + "github.com/glebarez/sqlite" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm" +) + +type taskTestDB struct{ db *gorm.DB } + +func (m *taskTestDB) GORM() *gorm.DB { return m.db } +func (m *taskTestDB) DB(ctx context.Context) *gorm.DB { return m.db.WithContext(ctx) } +func (m *taskTestDB) Named(_ string) *gorm.DB { return m.db } + +type sysConfigRow struct { + Key string `gorm:"primaryKey;size:64"` + Value string `gorm:"type:text"` +} + +func (sysConfigRow) TableName() string { return "w_system_configs" } + +func setupUserTaskDB(t *testing.T) *gorm.DB { + t.Helper() + _ = idgen.Init(1) + testDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "user_task.db")), &gorm.Config{}) + require.NoError(t, err) + require.NoError(t, testDB.AutoMigrate(&user.User{}, &user.AccessToken{}, &sysConfigRow{})) + user.SetDBService(&taskTestDB{db: testDB}) + t.Cleanup(func() { user.SetDBService(nil) }) + return testDB +} + +func TestSendEmailCodeValidatePayload(t *testing.T) { + h := &user.SendEmailCodeHandler{} + _, err := h.ValidatePayload([]byte(`{"email":"not-an-email"}`)) + require.Error(t, err) + + out, err := h.ValidatePayload([]byte(`{"email":"User@Example.com"}`)) + require.NoError(t, err) + assert.Contains(t, string(out), `"user@example.com"`) +} + +func TestSendMailValidatePayload(t *testing.T) { + h := &user.SendMailHandler{} + _, err := h.ValidatePayload([]byte(`{"to":"a@b.com","subject":"","body":"x"}`)) + require.Error(t, err) + _, err = h.ValidatePayload([]byte(`{"to":"a@b.com","subject":"Hi","body":"

ok

"}`)) + require.NoError(t, err) +} + +func TestSendMailRequiresSMTP(t *testing.T) { + setupUserTaskDB(t) + h := &user.SendMailHandler{} + _, err := h.Execute(context.Background(), []byte(`{"to":"a@b.com","subject":"Hi","body":"

ok

"}`)) + require.Error(t, err) + assert.Contains(t, err.Error(), "SMTP") +} + +func TestCleanupInactiveNeverLoggedInUsers(t *testing.T) { + db := setupUserTaskDB(t) + old := time.Now().Add(-40 * 24 * time.Hour) + stale := user.User{ID: 42, Username: "stale", Password: "x", IsActive: true, CreatedAt: old} + require.NoError(t, db.Create(&stale).Error) + require.NoError(t, db.Model(&stale).Updates(map[string]any{ + "created_at": old, + "last_login_at": time.Time{}, + }).Error) + + fresh := user.User{ID: 43, Username: "fresh", Password: "x", IsActive: true, LastLoginAt: time.Now()} + require.NoError(t, db.Create(&fresh).Error) + + admin := user.User{ID: 1, Username: "admin", Password: "x", IsAdmin: true, CreatedAt: old} + require.NoError(t, db.Create(&admin).Error) + require.NoError(t, db.Model(&admin).Updates(map[string]any{ + "created_at": old, + "last_login_at": time.Time{}, + }).Error) + + h := &user.CleanupInactiveHandler{} + res, err := h.Execute(context.Background(), nil) + require.NoError(t, err) + require.NotNil(t, res) + assert.Contains(t, res.Message, "1") + + _, err = user.GetUserByID(context.Background(), 42) + assert.Error(t, err) + _, err = user.GetUserByID(context.Background(), 43) + require.NoError(t, err) + _, err = user.GetUserByID(context.Background(), 1) + require.NoError(t, err) +} + +func TestSendEmailCodeMetaExported(t *testing.T) { + assert.Equal(t, "send_email_code", user.SendEmailCodeMeta.Type) + assert.Equal(t, "user:send_email_code", user.SendEmailCodeMeta.AsynqTask) + _ = contracts.TaskHandler(&user.SendEmailCodeHandler{}) +} diff --git a/backend/plugins/drivers/driver_asynq_cron/plugin.go b/backend/plugins/drivers/driver_asynq_cron/plugin.go index d822d894..5439583c 100644 --- a/backend/plugins/drivers/driver_asynq_cron/plugin.go +++ b/backend/plugins/drivers/driver_asynq_cron/plugin.go @@ -120,23 +120,8 @@ func (p *Plugin) Apply(ctx *core.Context) error { } p.mu.Unlock() - // Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } - - // Bind TaskService - if taskSvc, err := core.Inject[contracts.TaskService](ctx); err == nil && taskSvc != nil { - setTaskService(taskSvc) - } else { - core.When[contracts.TaskService](ctx, func(taskSvc contracts.TaskService) { - setTaskService(taskSvc) - }) - } + core.Bind[contracts.DBService](ctx, setDBService) + core.Bind[contracts.TaskService](ctx, setTaskService) ctx.OnDispose(func() error { setDBService(nil) diff --git a/backend/plugins/drivers/driver_asynq_worker/db_helper.go b/backend/plugins/drivers/driver_asynq_worker/db_helper.go index 025fbca1..a212de18 100644 --- a/backend/plugins/drivers/driver_asynq_worker/db_helper.go +++ b/backend/plugins/drivers/driver_asynq_worker/db_helper.go @@ -34,10 +34,8 @@ func SetRedisClient(c redis.UniversalClient) { } func getDB(ctx context.Context) *gorm.DB { - if c, ok := ctx.(*core.Context); ok && c != nil { - if s, err := core.Inject[contracts.DBService](c); err == nil && s != nil { - return s.DB(ctx) - } + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) } dbMu.RLock() s := dbSvc diff --git a/backend/plugins/drivers/driver_asynq_worker/executor.go b/backend/plugins/drivers/driver_asynq_worker/executor.go index fc8459f5..6d214bf1 100644 --- a/backend/plugins/drivers/driver_asynq_worker/executor.go +++ b/backend/plugins/drivers/driver_asynq_worker/executor.go @@ -4,6 +4,7 @@ package driver_asynq_worker import ( + "Wavelet/core/contracts" "Wavelet/pkg/idgen" "Wavelet/pkg/logger" "Wavelet/pkg/util" @@ -208,7 +209,7 @@ func RetryTask(ctx context.Context, id uint64) (string, error) { MaxRetry: execution.MaxRetry, RetryCount: execution.RetryCount + 1, Payload: execution.Payload, - TriggeredBy: "retry", + TriggeredBy: contracts.TaskTriggerRetry, } if err := createTaskExecution(ctx, newExecution); err != nil { @@ -265,8 +266,9 @@ func ProcessTask(ctx context.Context, t *asynq.Task) error { taskID := t.ResultWriter().TaskID() - // 注入 taskID 到 context + // 注入 taskID 到 context,并把原始 asynq.Task 交给适配器 Handler。 ctx = withTaskID(ctx, taskID) + ctx = context.WithValue(ctx, asynqTaskCtxKey{}, t) // 查找处理器 handler, ok := getHandler(t.Type()) @@ -277,18 +279,20 @@ func ProcessTask(ctx context.Context, t *asynq.Task) error { return err } - // 加载或动态创建执行记录 + // 加载或动态创建执行记录。无 DB 时仍执行业务 Handler,避免测试/精简拓扑 panic。 + var execution *TaskExecution now := time.Now() - execution, err := getOrCreateTaskExecution(ctx, taskID, t, taskPayload, now) - if err == nil { - updateExecutionOnStart(ctx, execution, now) + if getDB(ctx) != nil { + var err error + execution, err = getOrCreateTaskExecution(ctx, taskID, t, taskPayload, now) + if err == nil { + updateExecutionOnStart(ctx, execution, now) + } } if execution != nil { AppendLog(ctx, "[系统] 开始执行异步任务 [名称: %s, 类型: %s],重试次数: %d/%d", execution.TaskName, t.Type(), execution.RetryCount, execution.MaxRetry) - } else { - AppendLog(ctx, "[系统] 开始执行异步任务 [类型: %s]", t.Type()) } // 开始计时 @@ -390,7 +394,7 @@ func getOrCreateTaskExecution(ctx context.Context, taskID string, t *asynq.Task, MaxRetry: meta.MaxRetry, RetryCount: 0, Payload: string(payload), - TriggeredBy: "schedule", + TriggeredBy: contracts.TaskTriggerSchedule, StartedAt: &now, } diff --git a/backend/plugins/drivers/driver_asynq_worker/plugin.go b/backend/plugins/drivers/driver_asynq_worker/plugin.go index 05905182..4258c8e0 100644 --- a/backend/plugins/drivers/driver_asynq_worker/plugin.go +++ b/backend/plugins/drivers/driver_asynq_worker/plugin.go @@ -8,6 +8,7 @@ import ( "Wavelet/core" "Wavelet/core/contracts" "context" + "encoding/json" "errors" "fmt" "sync" @@ -144,14 +145,7 @@ func (p *Plugin) Apply(ctx *core.Context) error { ResetAsynqClient() p.mu.Unlock() - // 0. Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, setDBService) ctx.OnDispose(func() error { setDBService(nil) return nil @@ -191,12 +185,15 @@ func (p *Plugin) Start(_ context.Context) error { mux := asynq.NewServeMux() if p.coreCtx != nil && p.coreCtx.Tasks() != nil { + appCtx := p.coreCtx.Root() for _, td := range p.coreCtx.Tasks().Tasks() { - handler, err := toAsynqHandler(td.Handler) + handler, err := toAsynqHandler(td.Pattern, td.Handler) if err != nil { return fmt.Errorf("driver_asynq_worker: invalid handler for task pattern %q: %w", td.Pattern, err) } - mux.Handle(td.Pattern, handler) + mux.Handle(td.Pattern, asynq.HandlerFunc(func(c context.Context, t *asynq.Task) error { + return handler.ProcessTask(core.WithAppContext(c, appCtx), t) + })) } } @@ -276,21 +273,90 @@ func (p *Plugin) Mux() *asynq.ServeMux { return p.mux } -func toAsynqHandler(h any) (asynq.Handler, error) { +func toAsynqHandler(pattern string, h any) (asynq.Handler, error) { if h == nil { return nil, errors.New("nil handler") } + if th, ok := h.(TaskHandler); ok { + RegisterHandler(pattern, th) + return asynq.HandlerFunc(ProcessTask), nil + } + if th, ok := h.(contracts.TaskHandler); ok { + RegisterHandler(pattern, contractTaskAdapter{inner: th}) + return asynq.HandlerFunc(ProcessTask), nil + } + if fn, ok := h.(func(context.Context, []byte) (*contracts.TaskResultDTO, error)); ok { + RegisterHandler(pattern, contractFuncAdapter{fn: fn}) + return asynq.HandlerFunc(ProcessTask), nil + } + + inner, err := toRawAsynqHandler(h) + if err != nil { + return nil, err + } + RegisterHandler(pattern, &asynqHandlerAdapter{inner: inner}) + return asynq.HandlerFunc(ProcessTask), nil +} + +type contractTaskAdapter struct { + inner contracts.TaskHandler +} + +func (a contractTaskAdapter) Execute(ctx context.Context, payload []byte) (*TaskResult, error) { + res, err := a.inner.Execute(ctx, payload) + if err != nil { + return nil, err + } + return dtoToTaskResult(res), nil +} + +func (a contractTaskAdapter) ValidatePayload(payload []byte) ([]byte, error) { + if v, ok := a.inner.(PayloadValidator); ok { + return v.ValidatePayload(payload) + } + return payload, nil +} + +type contractFuncAdapter struct { + fn func(context.Context, []byte) (*contracts.TaskResultDTO, error) +} + +func (a contractFuncAdapter) Execute(ctx context.Context, payload []byte) (*TaskResult, error) { + res, err := a.fn(ctx, payload) + if err != nil { + return nil, err + } + return dtoToTaskResult(res), nil +} + +func dtoToTaskResult(res *contracts.TaskResultDTO) *TaskResult { + if res == nil { + return &TaskResult{Message: "ok"} + } + out := &TaskResult{Message: res.Message} + if res.Detail == nil { + return out + } + if s, ok := res.Detail.(string); ok { + out.Detail = s + return out + } + b, err := json.Marshal(res.Detail) + if err != nil { + out.Detail = fmt.Sprint(res.Detail) + return out + } + out.Detail = string(b) + return out +} + +func toRawAsynqHandler(h any) (asynq.Handler, error) { switch fn := h.(type) { case asynq.HandlerFunc: return fn, nil case asynq.Handler: return fn, nil - case TaskHandler: - return asynq.HandlerFunc(func(c context.Context, t *asynq.Task) error { - RegisterHandler(t.Type(), fn) - return ProcessTask(c, t) - }), nil case func(context.Context, *asynq.Task) error: return asynq.HandlerFunc(fn), nil case func(context.Context, []byte) error: @@ -316,6 +382,23 @@ func toAsynqHandler(h any) (asynq.Handler, error) { } } +type asynqTaskCtxKey struct{} + +type asynqHandlerAdapter struct { + inner asynq.Handler +} + +func (a *asynqHandlerAdapter) Execute(ctx context.Context, payload []byte) (*TaskResult, error) { + t, _ := ctx.Value(asynqTaskCtxKey{}).(*asynq.Task) + if t == nil { + t = asynq.NewTask("", payload) + } + if err := a.inner.ProcessTask(ctx, t); err != nil { + return nil, err + } + return &TaskResult{Message: "ok"}, nil +} + type taskServiceImpl struct{} func (s *taskServiceImpl) Dispatch(ctx context.Context, taskType string, payload []byte, triggeredBy string) (string, error) { @@ -472,6 +555,15 @@ func (s *taskServiceImpl) GetExecution(ctx context.Context, id uint64) (*contrac return &dto, nil } +func (s *taskServiceImpl) GetExecutionByTaskID(ctx context.Context, taskID string) (*contracts.TaskExecutionDTO, error) { + exec, err := GetTaskExecutionByTaskID(ctx, taskID) + if err != nil { + return nil, err + } + dto := toTaskExecutionDTO(exec) + return &dto, nil +} + func toTaskExecutionDTO(exec *TaskExecution) contracts.TaskExecutionDTO { return contracts.TaskExecutionDTO{ ID: exec.ID, diff --git a/backend/plugins/drivers/driver_http/engine_slash_test.go b/backend/plugins/drivers/driver_http/engine_slash_test.go index e58f588d..efbfa2f2 100644 --- a/backend/plugins/drivers/driver_http/engine_slash_test.go +++ b/backend/plugins/drivers/driver_http/engine_slash_test.go @@ -4,9 +4,12 @@ package driver_http import ( - "testing" - + "Wavelet/core" "Wavelet/core/extpoints" + "context" + "net/http" + "net/http/httptest" + "testing" ) func TestBuildEngineDefaultRedirectsTrailingSlash(t *testing.T) { @@ -111,3 +114,29 @@ func bindAppConfig(t *testing.T, values map[string]any, env map[string]string) h } func boolPtr(v bool) *bool { return &v } + +func TestDriverHTTPSwaggerMount(t *testing.T) { + ctx := core.NewContext(t.Context()) + ctx.Config().SetSource(core.NewMapSource(map[string]any{"app.env": "development"})) + if err := ctx.Config().Resolve(); err != nil { + t.Fatal(err) + } + p := New(WithAddr("127.0.0.1:0")) + if err := p.Apply(ctx); err != nil { + t.Fatal(err) + } + startCtx, cancel := context.WithCancel(t.Context()) + defer cancel() + if err := p.Start(startCtx); err != nil { + t.Fatal(err) + } + defer func() { _ = p.Stop(t.Context()) }() + + w := httptest.NewRecorder() + req, _ := http.NewRequestWithContext(t.Context(), http.MethodGet, "/swagger/index.html", nil) + req.RequestURI = "/swagger/index.html" + p.Engine().ServeHTTP(w, req) + if w.Code != http.StatusOK { + t.Fatalf("expected 200 for /swagger/index.html, got %d (body: %s)", w.Code, w.Body.String()) + } +} diff --git a/backend/plugins/drivers/driver_http/frontend.go b/backend/plugins/drivers/driver_http/frontend.go index 0aa99c33..73b4dfd4 100644 --- a/backend/plugins/drivers/driver_http/frontend.go +++ b/backend/plugins/drivers/driver_http/frontend.go @@ -17,7 +17,7 @@ const indexFile = "index.html" // serverOwnedPrefixes are backend-owned namespaces. A miss there must keep Gin's default // 404 instead of silently returning the frontend shell, which would mask broken API links. -var serverOwnedPrefixes = []string{"/api/", "/f/"} +var serverOwnedPrefixes = []string{"/api/", "/f/", "/swagger/"} // registerFrontend mounts assets as the NoRoute fallback so client-side routes resolve. // It is a no-op when assets is nil, i.e. the binary was built without the embed_frontend tag. diff --git a/backend/plugins/drivers/driver_http/plugin.go b/backend/plugins/drivers/driver_http/plugin.go index 83d933bc..b323786c 100644 --- a/backend/plugins/drivers/driver_http/plugin.go +++ b/backend/plugins/drivers/driver_http/plugin.go @@ -7,6 +7,7 @@ package driver_http import ( "Wavelet/core" "Wavelet/core/contracts" + _ "Wavelet/docs" // swagger documentation registration "Wavelet/pkg/util" "context" "errors" @@ -17,6 +18,8 @@ import ( "time" "github.com/gin-gonic/gin" + swaggerFiles "github.com/swaggo/files" + ginSwagger "github.com/swaggo/gin-swagger" ) const ( @@ -118,27 +121,13 @@ func (p *Plugin) Apply(ctx *core.Context) error { } p.mu.Unlock() - // Bind DBService from Context - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { - setDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - setDBService(db) - }) - } + core.Bind[contracts.DBService](ctx, setDBService) ctx.OnDispose(func() error { setDBService(nil) return nil }) - // Bind CacheService from Context - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - setCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - setCacheService(cache) - }) - } + core.Bind[contracts.CacheService](ctx, setCacheService) ctx.OnDispose(func() error { setCacheService(nil) return nil @@ -181,29 +170,20 @@ func (p *Plugin) Start(ctx context.Context) error { } } - // Mount routes collected in Context RouterExtension - if p.coreCtx != nil && p.coreCtx.Router() != nil { - SetWhitelist(p.coreCtx.Router().Whitelist()) - for _, rd := range p.coreCtx.Router().Routes() { - allHandlers := make([]gin.HandlerFunc, 0, len(rd.Middlewares)+len(rd.Handlers)) + if err := p.mountContextRoutes(ctx); err != nil { + return err + } - for _, m := range rd.Middlewares { - gh, err := toGinHandler(m) - if err != nil { - return fmt.Errorf("driver_http: invalid middleware for route %s %s: %w", rd.Method, rd.Path, err) - } - allHandlers = append(allHandlers, gh) + // Mount Swagger in non-production environments + if p.coreCtx != nil { + var appCfg httpAppConfig + _ = p.coreCtx.Config().Bind("app", &appCfg) + if appCfg.Env != "production" && appCfg.Env != "prod" { + swaggerHandler := ginSwagger.WrapHandler(swaggerFiles.Handler) + p.engine.GET("/swagger/*any", swaggerHandler) + if appCfg.APIPrefix != "" { + p.engine.GET(appCfg.APIPrefix+"/swagger/*any", swaggerHandler) } - - for _, h := range rd.Handlers { - gh, err := toGinHandler(h) - if err != nil { - return fmt.Errorf("driver_http: invalid handler for route %s %s: %w", rd.Method, rd.Path, err) - } - allHandlers = append(allHandlers, gh) - } - - p.engine.Handle(rd.Method, rd.Path, allHandlers...) } } @@ -266,6 +246,56 @@ func (p *Plugin) Stop(ctx context.Context) error { return err } +func (p *Plugin) mountContextRoutes(ctx context.Context) error { + if p.coreCtx == nil || p.coreCtx.Router() == nil || p.engine == nil { + return nil + } + p.engine.Use(appContextMiddleware(ctx, p.coreCtx.Root())) + SetWhitelist(p.coreCtx.Router().Whitelist()) + + globalMW, err := toGinHandlers(p.coreCtx.Router().Middlewares()) + if err != nil { + return fmt.Errorf("driver_http: invalid global middleware: %w", err) + } + + for _, rd := range p.coreCtx.Router().Routes() { + routeMW, convErr := toGinHandlers(rd.Middlewares) + if convErr != nil { + return fmt.Errorf("driver_http: invalid middleware for route %s %s: %w", rd.Method, rd.Path, convErr) + } + handlers, convErr := toGinHandlers(rd.Handlers) + if convErr != nil { + return fmt.Errorf("driver_http: invalid handler for route %s %s: %w", rd.Method, rd.Path, convErr) + } + allHandlers := make([]gin.HandlerFunc, 0, len(globalMW)+len(routeMW)+len(handlers)) + allHandlers = append(allHandlers, globalMW...) + allHandlers = append(allHandlers, routeMW...) + allHandlers = append(allHandlers, handlers...) + p.engine.Handle(rd.Method, rd.Path, allHandlers...) + } + return nil +} + +func toGinHandlers(hs []any) ([]gin.HandlerFunc, error) { + out := make([]gin.HandlerFunc, 0, len(hs)) + for _, h := range hs { + gh, err := toGinHandler(h) + if err != nil { + return nil, err + } + out = append(out, gh) + } + return out, nil +} + +//nolint:contextcheck // middleware must wrap the gin request context, not Start's ctx +func appContextMiddleware(_ context.Context, appCtx *core.Context) gin.HandlerFunc { + return func(c *gin.Context) { + c.Request = c.Request.WithContext(core.WithAppContext(c.Request.Context(), appCtx)) + c.Next() + } +} + // Addr returns the current listening address (or configured address if not yet started). func (p *Plugin) Addr() string { p.mu.RLock() diff --git a/backend/plugins/drivers/driver_inproc_cron/plugin.go b/backend/plugins/drivers/driver_inproc_cron/plugin.go index 49fbc3fe..40caf4d9 100644 --- a/backend/plugins/drivers/driver_inproc_cron/plugin.go +++ b/backend/plugins/drivers/driver_inproc_cron/plugin.go @@ -83,7 +83,11 @@ func (p *Plugin) Start(ctx context.Context) error { p.scheduler = newInprocScheduler(p.coreCtx.Schedules(), p.coreCtx.Tasks(), taskSvc) } - return p.scheduler.Start(ctx) + runCtx := ctx + if p.coreCtx != nil { + runCtx = core.WithAppContext(ctx, p.coreCtx.Root()) + } + return p.scheduler.Start(runCtx) } // Stop terminates the in-process cron scheduler. diff --git a/backend/plugins/drivers/driver_inproc_cron/scheduler.go b/backend/plugins/drivers/driver_inproc_cron/scheduler.go index 8b3b33f0..b0760c7e 100644 --- a/backend/plugins/drivers/driver_inproc_cron/scheduler.go +++ b/backend/plugins/drivers/driver_inproc_cron/scheduler.go @@ -99,7 +99,7 @@ func (s *inprocScheduler) registerJob(ctx context.Context, def extpoints.Schedul _, err := s.cronRunner.AddFunc(cronSpec, func() { if s.taskSvc != nil { - if _, dispatchErr := s.taskSvc.Dispatch(ctx, taskType, payloadBytes, "inproc_cron"); dispatchErr != nil { + if _, dispatchErr := s.taskSvc.Dispatch(ctx, taskType, payloadBytes, contracts.TaskTriggerSchedule); dispatchErr != nil { logger.ErrorF(ctx, "driver_inproc_cron: dispatch task %q failed: %v", taskType, dispatchErr) } return diff --git a/backend/plugins/drivers/driver_inproc_worker/db_helper.go b/backend/plugins/drivers/driver_inproc_worker/db_helper.go new file mode 100644 index 00000000..32077db9 --- /dev/null +++ b/backend/plugins/drivers/driver_inproc_worker/db_helper.go @@ -0,0 +1,37 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package driver_inproc_worker + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "context" + "sync" + + "gorm.io/gorm" +) + +var ( + dbMu sync.RWMutex + dbSvc contracts.DBService +) + +func setDBService(s contracts.DBService) { + dbMu.Lock() + defer dbMu.Unlock() + dbSvc = s +} + +func getDB(ctx context.Context) *gorm.DB { + if s, err := core.InjectFrom[contracts.DBService](ctx); err == nil && s != nil { + return s.DB(ctx) + } + dbMu.RLock() + s := dbSvc + dbMu.RUnlock() + if s != nil { + return s.DB(ctx) + } + return nil +} diff --git a/backend/plugins/drivers/driver_inproc_worker/executor.go b/backend/plugins/drivers/driver_inproc_worker/executor.go index 7fcc9f2f..1a57688b 100644 --- a/backend/plugins/drivers/driver_inproc_worker/executor.go +++ b/backend/plugins/drivers/driver_inproc_worker/executor.go @@ -4,10 +4,14 @@ package driver_inproc_worker import ( + "Wavelet/core" + "Wavelet/core/contracts" "Wavelet/core/extpoints" "Wavelet/pkg/idgen" + "Wavelet/pkg/logger" "Wavelet/pkg/util" "context" + "encoding/json" "errors" "fmt" "sync" @@ -39,6 +43,7 @@ type InprocQueue struct { // baseCtx is the app-lifetime context captured at Start; task handlers // derive their timeouts from it so shutdown cancellation propagates. baseCtx context.Context + appCtx *core.Context } // NewInprocQueue creates a new InprocQueue with a given concurrency and queue capacity. @@ -59,34 +64,56 @@ func NewInprocQueue(concurrency, queueCap int, taskReg extpoints.TaskExtension) } // Enqueue puts a new task into the in-process queue. -func (q *InprocQueue) Enqueue(taskType string, payload []byte, source string) (string, error) { +// taskType may be the registration pattern or the admin type identifier. +func (q *InprocQueue) Enqueue(ctx context.Context, taskType string, payload []byte, source string) (string, error) { if !q.running.Load() { return "", errors.New("driver_inproc_worker: queue is not running") } - taskID := fmt.Sprintf("inproc_%d", idgen.NextUint64ID()) + td, ok := q.lookupTask(taskType) + if !ok { + return "", fmt.Errorf("driver_inproc_worker: unknown task type %q", taskType) + } + + if source == "" { + source = contracts.TaskTriggerManual + } + idType := td.Type + if idType == "" { + idType = td.Pattern + } + taskID := fmt.Sprintf("%s_%s_%d", source, idType, idgen.NextUint64ID()) msg := TaskMessage{ ID: taskID, - TaskType: taskType, + TaskType: td.Pattern, Payload: payload, Source: source, CreatedAt: time.Now(), + RetryLeft: td.Retry, } - if q.taskReg != nil { - if td, ok := q.taskReg.Get(taskType); ok { - msg.RetryLeft = td.Retry - } + if err := q.createExecution(ctx, msg, td); err != nil { + return "", err } + q.appendExecutionLog(ctx, taskID, fmt.Sprintf("[系统] 任务已成功入队,等待调度执行 (最大重试次数: %d)", td.Retry)) + select { case q.queue <- msg: return taskID, nil default: + q.failExecution(ctx, msg, errors.New("queue is full"), 0) return "", errors.New("driver_inproc_worker: queue is full") } } +func (q *InprocQueue) lookupTask(taskType string) (extpoints.TaskDefinition, bool) { + if q.taskReg == nil { + return extpoints.TaskDefinition{}, false + } + return q.taskReg.Get(taskType) +} + // Start begins processing tasks with the worker pool. ctx is the app-lifetime // context used as the parent for per-task execution contexts. func (q *InprocQueue) Start(ctx context.Context) { @@ -145,12 +172,10 @@ func (q *InprocQueue) workerLoop(ctx context.Context) { } func (q *InprocQueue) executeTask(ctx context.Context, msg TaskMessage) { - if q.taskReg == nil { - return - } - - td, ok := q.taskReg.Get(msg.TaskType) + td, ok := q.lookupTask(msg.TaskType) if !ok { + logger.ErrorF(ctx, "driver_inproc_worker: no handler for task %q", msg.TaskType) + q.failExecution(ctx, msg, fmt.Errorf("unregistered task handler: %s", msg.TaskType), 0) return } @@ -161,44 +186,174 @@ func (q *InprocQueue) executeTask(ctx context.Context, msg TaskMessage) { taskCtx, cancel := context.WithTimeout(ctx, timeout) defer cancel() - - err := invokeHandler(taskCtx, td.Handler, msg.Payload) - if err != nil && msg.RetryLeft > 0 { - msg.RetryLeft-- - // Retry with backoff - util.Go(func() { - select { - case <-time.After(defaultRetryBackoff): - case <-q.stopCh: - return - case <-ctx.Done(): - return - } - if q.running.Load() { - select { - case q.queue <- msg: - default: - } - } - }) + if q.appCtx != nil { + taskCtx = core.WithAppContext(taskCtx, q.appCtx) + ctx = core.WithAppContext(ctx, q.appCtx) } + + q.markRunning(ctx, msg) + start := time.Now() + result, err := invokeHandler(taskCtx, td.Handler, msg.Payload) + duration := time.Since(start) + + if err != nil { + q.failExecution(ctx, msg, err, duration) + if msg.RetryLeft > 0 { + msg.RetryLeft-- + util.Go(func() { + select { + case <-time.After(defaultRetryBackoff): + case <-q.stopCh: + return + case <-ctx.Done(): + return + } + if q.running.Load() { + select { + case q.queue <- msg: + default: + } + } + }) + } + return + } + q.succeedExecution(ctx, msg, duration, result) } -func invokeHandler(ctx context.Context, handler any, payload []byte) error { +func invokeHandler(ctx context.Context, handler any, payload []byte) (*contracts.TaskResultDTO, error) { if handler == nil { - return errors.New("nil task handler") + return nil, errors.New("nil task handler") } switch fn := handler.(type) { - case func(context.Context, []byte) error: + case contracts.TaskHandler: + return fn.Execute(ctx, payload) + case func(context.Context, []byte) (*contracts.TaskResultDTO, error): return fn(ctx, payload) + case func(context.Context, []byte) error: + return nil, fn(ctx, payload) case func(context.Context) error: - return fn(ctx) + return nil, fn(ctx) case func([]byte) error: - return fn(payload) + return nil, fn(payload) case func() error: - return fn() + return nil, fn() default: - return fmt.Errorf("unsupported handler type: %T", handler) + return nil, fmt.Errorf("unsupported handler type: %T", handler) + } +} + +func (q *InprocQueue) createExecution(ctx context.Context, msg TaskMessage, td extpoints.TaskDefinition) error { + db := getDB(ctx) + if db == nil { + return nil + } + + name := td.Name + if name == "" { + name = td.DisplayName + } + if name == "" { + name = td.Pattern + } + exec := &taskExecution{ + ID: idgen.NextUint64ID(), + TaskID: msg.ID, + TaskType: td.Pattern, + TaskName: name, + Status: taskExecutionStatusPending, + Retryable: td.Retryable || td.Retry > 0, + MaxRetry: td.Retry, + RetryCount: 0, + Payload: string(msg.Payload), + TriggeredBy: msg.Source, + } + if err := db.Create(exec).Error; err != nil { + return fmt.Errorf("driver_inproc_worker: create task execution: %w", err) + } + return nil +} + +func (q *InprocQueue) markRunning(ctx context.Context, msg TaskMessage) { + db := getDB(ctx) + if db == nil { + return + } + now := time.Now() + updates := map[string]any{ + taskExecutionColStatus: taskExecutionStatusRunning, + "started_at": now, + } + if err := db.Model(&taskExecution{}).Where("task_id = ?", msg.ID).Updates(updates).Error; err != nil { + logger.ErrorF(ctx, "driver_inproc_worker: mark running failed taskID=%s: %v", msg.ID, err) + return + } + q.appendExecutionLog(ctx, msg.ID, fmt.Sprintf("[系统] 开始执行异步任务 [类型: %s]", msg.TaskType)) +} + +func (q *InprocQueue) succeedExecution(ctx context.Context, msg TaskMessage, duration time.Duration, result *contracts.TaskResultDTO) { + db := getDB(ctx) + if db == nil { + return + } + now := time.Now() + resultText := "ok" + if result != nil { + resultText = result.Message + if result.Detail != nil { + if s, ok := result.Detail.(string); ok && s != "" { + resultText = result.Message + "\n" + s + } else if b, err := json.Marshal(result.Detail); err == nil && len(b) > 0 && string(b) != "null" { + resultText = result.Message + "\n" + string(b) + } + } + } + updates := map[string]any{ + taskExecutionColStatus: taskExecutionStatusSucceeded, + "error_message": "", + "result": resultText, + "finished_at": now, + "duration": duration.Milliseconds(), + } + if err := db.Model(&taskExecution{}).Where("task_id = ?", msg.ID).Updates(updates).Error; err != nil { + logger.ErrorF(ctx, "driver_inproc_worker: mark succeeded failed taskID=%s: %v", msg.ID, err) + return + } + q.appendExecutionLog(ctx, msg.ID, fmt.Sprintf("[系统] 任务执行成功,耗时: %d ms", duration.Milliseconds())) +} + +func (q *InprocQueue) failExecution(ctx context.Context, msg TaskMessage, execErr error, duration time.Duration) { + db := getDB(ctx) + if db == nil { + return + } + now := time.Now() + updates := map[string]any{ + taskExecutionColStatus: taskExecutionStatusFailed, + "error_message": execErr.Error(), + "finished_at": now, + "duration": duration.Milliseconds(), + } + if err := db.Model(&taskExecution{}).Where("task_id = ?", msg.ID).Updates(updates).Error; err != nil { + logger.ErrorF(ctx, "driver_inproc_worker: mark failed failed taskID=%s: %v", msg.ID, err) + return + } + q.appendExecutionLog(ctx, msg.ID, fmt.Sprintf("[系统] 任务执行失败,耗时: %d ms,错误原因: %v", duration.Milliseconds(), execErr)) +} + +func (q *InprocQueue) appendExecutionLog(ctx context.Context, taskID, logLine string) { + db := getDB(ctx) + if db == nil { + return + } + now := time.Now().Format("15:04:05") + line := fmt.Sprintf("[%s] %s\n", now, logLine) + var exec taskExecution + if err := db.Where("task_id = ?", taskID).First(&exec).Error; err != nil { + return + } + if err := db.Model(&exec).Update("log", exec.Log+line).Error; err != nil { + logger.ErrorF(ctx, "driver_inproc_worker: append log failed taskID=%s: %v", taskID, err) } } diff --git a/backend/plugins/drivers/driver_inproc_worker/plugin.go b/backend/plugins/drivers/driver_inproc_worker/plugin.go index cd004818..119c3b1d 100644 --- a/backend/plugins/drivers/driver_inproc_worker/plugin.go +++ b/backend/plugins/drivers/driver_inproc_worker/plugin.go @@ -8,6 +8,7 @@ import ( "Wavelet/core" "Wavelet/core/contracts" "context" + "errors" "sync" "time" ) @@ -24,15 +25,15 @@ var ( ) // DispatchTask enqueues a background task to the active in-process worker queue. -func DispatchTask(_ context.Context, taskType string, payload []byte, source string) (string, error) { +func DispatchTask(ctx context.Context, taskType string, payload []byte, source string) (string, error) { globalMu.RLock() q := globalQueue globalMu.RUnlock() if q == nil { - return "", nil + return "", errors.New("driver_inproc_worker: queue is not running") } - return q.Enqueue(taskType, payload, source) + return q.Enqueue(ctx, taskType, payload, source) } // Option configures the in-process worker driver plugin. @@ -118,10 +119,13 @@ func (p *Plugin) ConfigEnabled(view core.ConfigView) bool { func (p *Plugin) Apply(ctx *core.Context) error { p.coreCtx = ctx + core.Bind[contracts.DBService](ctx, setDBService) + taskSvc := newInprocTaskService(ctx.Tasks()) core.Provide[contracts.TaskService](ctx, taskSvc) ctx.OnDispose(func() error { + setDBService(nil) shutdownCtx, cancel := context.WithTimeout(context.Background(), p.shutdownTimeout) defer cancel() return p.Stop(shutdownCtx) @@ -141,6 +145,9 @@ func (p *Plugin) Start(ctx context.Context) error { if p.queue == nil { p.queue = NewInprocQueue(p.concurrency, p.queueCapacity, p.coreCtx.Tasks()) } + if p.coreCtx != nil { + p.queue.appCtx = p.coreCtx.Root() + } globalMu.Lock() globalQueue = p.queue diff --git a/backend/plugins/drivers/driver_inproc_worker/plugin_test.go b/backend/plugins/drivers/driver_inproc_worker/plugin_test.go index 1107e433..85feaff7 100644 --- a/backend/plugins/drivers/driver_inproc_worker/plugin_test.go +++ b/backend/plugins/drivers/driver_inproc_worker/plugin_test.go @@ -5,8 +5,10 @@ package driver_inproc_worker_test import ( "Wavelet/core" + "Wavelet/core/contracts" "Wavelet/core/extpoints" "Wavelet/pkg/idgen" + "Wavelet/pkg/testhelper" "Wavelet/plugins/drivers/driver_inproc_worker" "context" "sync/atomic" @@ -15,8 +17,19 @@ import ( "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "gorm.io/gorm" ) +type testDBService struct { + db *gorm.DB +} + +func (m *testDBService) GORM() *gorm.DB { return m.db } + +func (m *testDBService) DB(ctx context.Context) *gorm.DB { return m.db.WithContext(ctx) } + +func (m *testDBService) Named(_ string) *gorm.DB { return m.db } + func TestInprocWorkerPlugin(t *testing.T) { require.NoError(t, idgen.Init(1)) ctx := core.NewContext(context.Background()) @@ -51,3 +64,53 @@ func TestInprocWorkerPlugin(t *testing.T) { // Stop driver require.NoError(t, p.Stop(context.Background())) } + +func TestInprocWorkerDispatchByTypeTracksExecution(t *testing.T) { + require.NoError(t, idgen.Init(1)) + testDB, _, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + ctx := core.NewContext(context.Background()) + core.Provide[contracts.DBService](ctx, &testDBService{db: testDB}) + + p := driver_inproc_worker.New( + driver_inproc_worker.WithConcurrency(2), + driver_inproc_worker.WithShutdownTimeout(time.Second), + ) + require.NoError(t, p.Apply(ctx)) + + var executedCount atomic.Int32 + ctx.Tasks().Register("system:cleanup", func(_ context.Context, _ []byte) (*contracts.TaskResultDTO, error) { + executedCount.Add(1) + return &contracts.TaskResultDTO{Message: "cleaned 3 files"}, nil + }, + extpoints.WithTaskType("system_cleanup"), + extpoints.WithTaskName("系统垃圾清理"), + extpoints.WithTaskRetry(1), + extpoints.WithTaskRetryable(true), + ) + + require.NoError(t, p.Start(context.Background())) + t.Cleanup(func() { + _ = p.Stop(context.Background()) + }) + + taskID, err := driver_inproc_worker.DispatchTask(context.Background(), "system_cleanup", []byte("payload"), "manual") + require.NoError(t, err) + assert.NotEmpty(t, taskID) + + require.Eventually(t, func() bool { + return executedCount.Load() == 1 + }, 2*time.Second, 20*time.Millisecond, "inproc worker should execute task dispatched by admin type") + + taskSvc, err := core.Inject[contracts.TaskService](ctx) + require.NoError(t, err) + + require.Eventually(t, func() bool { + execs, total, listErr := taskSvc.ListExecutions(context.Background(), "", "", 1, 10) + if listErr != nil || total == 0 || len(execs) == 0 { + return false + } + return execs[0].TaskID == taskID && execs[0].Status == "succeeded" && execs[0].TaskName == "系统垃圾清理" && execs[0].Result == "cleaned 3 files" + }, 2*time.Second, 20*time.Millisecond, "inproc worker should persist a succeeded execution record") +} diff --git a/backend/plugins/drivers/driver_inproc_worker/task_service.go b/backend/plugins/drivers/driver_inproc_worker/task_service.go index 0f3e7d83..41545179 100644 --- a/backend/plugins/drivers/driver_inproc_worker/task_service.go +++ b/backend/plugins/drivers/driver_inproc_worker/task_service.go @@ -25,8 +25,22 @@ func (s *inprocTaskService) Dispatch(ctx context.Context, taskType string, paylo return DispatchTask(ctx, taskType, payload, triggeredBy) } -func (s *inprocTaskService) Retry(_ context.Context, id uint64) (string, error) { - return fmt.Sprintf("inproc_retry_%d", id), nil +func (s *inprocTaskService) Retry(ctx context.Context, id uint64) (string, error) { + db := getDB(ctx) + if db == nil { + return "", errors.New("driver_inproc_worker: db not initialized") + } + var exec taskExecution + if err := db.Where("id = ?", id).First(&exec).Error; err != nil { + return "", fmt.Errorf("driver_inproc_worker: task execution not found: %w", err) + } + if exec.Status != taskExecutionStatusFailed { + return "", fmt.Errorf("driver_inproc_worker: only failed tasks can be retried, current status: %s", exec.Status) + } + if !exec.Retryable { + return "", errors.New("driver_inproc_worker: task is not retryable") + } + return DispatchTask(ctx, exec.TaskType, []byte(exec.Payload), "retry") } func (s *inprocTaskService) ListTasks() []contracts.TaskMetaDTO { @@ -64,10 +78,80 @@ func (s *inprocTaskService) ReloadScheduler() error { func (s *inprocTaskService) AppendLog(_ context.Context, _ string, _ ...any) { } -func (s *inprocTaskService) ListExecutions(_ context.Context, _, _ string, _, _ int) ([]contracts.TaskExecutionDTO, int64, error) { - return []contracts.TaskExecutionDTO{}, 0, nil +func (s *inprocTaskService) ListExecutions(ctx context.Context, taskType, status string, page, pageSize int) ([]contracts.TaskExecutionDTO, int64, error) { + db := getDB(ctx) + if db == nil { + return []contracts.TaskExecutionDTO{}, 0, nil + } + if page <= 0 { + page = 1 + } + if pageSize <= 0 { + pageSize = 20 + } + query := db.Model(&taskExecution{}) + if taskType != "" { + query = query.Where("task_type = ?", taskType) + } + if status != "" { + query = query.Where("status = ?", status) + } + var total int64 + if err := query.Count(&total).Error; err != nil { + return nil, 0, err + } + var rows []taskExecution + offset := (page - 1) * pageSize + if err := query.Order("id DESC").Offset(offset).Limit(pageSize).Find(&rows).Error; err != nil { + return nil, 0, err + } + res := make([]contracts.TaskExecutionDTO, 0, len(rows)) + for i := range rows { + res = append(res, toExecutionDTO(&rows[i])) + } + return res, total, nil } -func (s *inprocTaskService) GetExecution(_ context.Context, _ uint64) (*contracts.TaskExecutionDTO, error) { - return nil, errors.New("driver_inproc_worker: task executions are not tracked") +func findExecution(ctx context.Context, query any, args ...any) (*contracts.TaskExecutionDTO, error) { + db := getDB(ctx) + if db == nil { + return nil, errors.New("driver_inproc_worker: db not initialized") + } + var exec taskExecution + if err := db.Where(query, args...).First(&exec).Error; err != nil { + return nil, err + } + dto := toExecutionDTO(&exec) + return &dto, nil +} + +func (s *inprocTaskService) GetExecution(ctx context.Context, id uint64) (*contracts.TaskExecutionDTO, error) { + return findExecution(ctx, "id = ?", id) +} + +func (s *inprocTaskService) GetExecutionByTaskID(ctx context.Context, taskID string) (*contracts.TaskExecutionDTO, error) { + return findExecution(ctx, "task_id = ?", taskID) +} + +func toExecutionDTO(exec *taskExecution) contracts.TaskExecutionDTO { + return contracts.TaskExecutionDTO{ + ID: exec.ID, + TaskID: exec.TaskID, + TaskType: exec.TaskType, + TaskName: exec.TaskName, + Status: string(exec.Status), + Retryable: exec.Retryable, + MaxRetry: exec.MaxRetry, + RetryCount: exec.RetryCount, + Log: exec.Log, + ErrorMessage: exec.ErrorMessage, + Result: exec.Result, + StartedAt: exec.StartedAt, + FinishedAt: exec.FinishedAt, + Duration: exec.Duration, + Payload: exec.Payload, + TriggeredBy: exec.TriggeredBy, + CreatedAt: exec.CreatedAt, + UpdatedAt: exec.UpdatedAt, + } } diff --git a/backend/plugins/drivers/driver_inproc_worker/types.go b/backend/plugins/drivers/driver_inproc_worker/types.go new file mode 100644 index 00000000..c3e5f84f --- /dev/null +++ b/backend/plugins/drivers/driver_inproc_worker/types.go @@ -0,0 +1,44 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package driver_inproc_worker + +import "time" + +type taskExecutionStatus string + +const ( + taskExecutionStatusPending taskExecutionStatus = "pending" + taskExecutionStatusRunning taskExecutionStatus = "running" + taskExecutionStatusSucceeded taskExecutionStatus = "succeeded" + taskExecutionStatusFailed taskExecutionStatus = "failed" + + taskExecutionColStatus = "status" +) + +// taskExecution maps to the admin-owned w_task_executions table so the +// console can list in-process runs the same way it lists Asynq runs. +type taskExecution struct { + ID uint64 `gorm:"primaryKey"` + TaskID string `gorm:"size:128;uniqueIndex;not null"` + TaskType string `gorm:"size:64;index;not null"` + TaskName string `gorm:"size:128"` + Status taskExecutionStatus `gorm:"size:32;index;not null"` + Retryable bool `gorm:"not null;default:false"` + MaxRetry int `gorm:"not null;default:0"` + RetryCount int `gorm:"not null;default:0"` + Log string `gorm:"type:text"` + ErrorMessage string `gorm:"type:text"` + Result string `gorm:"type:text"` + StartedAt *time.Time `gorm:"index"` + FinishedAt *time.Time + Duration int64 `gorm:"comment:耗时毫秒"` + Payload string `gorm:"type:text"` + TriggeredBy string `gorm:"size:32;not null;default:system"` + CreatedAt time.Time + UpdatedAt time.Time +} + +func (taskExecution) TableName() string { + return "w_task_executions" +} diff --git a/backend/plugins/drivers/drivers_test.go b/backend/plugins/drivers/drivers_test.go index 2fc891d2..5b9c008a 100644 --- a/backend/plugins/drivers/drivers_test.go +++ b/backend/plugins/drivers/drivers_test.go @@ -5,7 +5,10 @@ package drivers_test import ( "Wavelet/core" + "Wavelet/core/contracts" "Wavelet/core/extpoints" + "Wavelet/pkg/idgen" + "Wavelet/pkg/testhelper" "Wavelet/plugins/drivers/driver_asynq_cron" "Wavelet/plugins/drivers/driver_asynq_worker" "Wavelet/plugins/drivers/driver_http" @@ -24,8 +27,19 @@ import ( "github.com/hibiken/asynq" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "gorm.io/gorm" ) +type testDBService struct { + db *gorm.DB +} + +func (m *testDBService) GORM() *gorm.DB { return m.db } + +func (m *testDBService) DB(ctx context.Context) *gorm.DB { return m.db.WithContext(ctx) } + +func (m *testDBService) Named(_ string) *gorm.DB { return m.db } + func init() { gin.SetMode(gin.TestMode) } @@ -138,6 +152,41 @@ func TestHTTPDriverLifecycle(t *testing.T) { require.NoError(t, err) } +func TestHTTPDriverAppliesGlobalMiddlewareRegisteredAfterRoutes(t *testing.T) { + ctx := core.NewContext(context.Background()) + ctx.Config().SetSource(core.NewMapSource(nil)) + require.NoError(t, ctx.Config().Resolve()) + + ctx.Router().GET("/early", func(c *gin.Context) { + c.Status(http.StatusOK) + }) + + var called atomic.Bool + ctx.Router().Use(func(c *gin.Context) { + called.Store(true) + c.Next() + }) + + httpPlugin := driver_http.New(driver_http.WithAddr("127.0.0.1:0")) + require.NoError(t, httpPlugin.Apply(ctx)) + d, ok := ctx.Driver(core.DriverTypeHTTP) + require.True(t, ok) + require.NoError(t, d.Start(context.Background())) + t.Cleanup(func() { + stopCtx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + _ = d.Stop(stopCtx) + }) + + resp, err := http.Get(fmt.Sprintf("http://%s/early", httpPlugin.Addr())) + require.NoError(t, err) + defer resp.Body.Close() + require.Equal(t, http.StatusOK, resp.StatusCode) + if !called.Load() { + t.Fatal("Router.Use middleware registered after the route must still run at HTTP Start") + } +} + func TestAsynqWorkerDriverLifecycle(t *testing.T) { mr, err := miniredis.Run() require.NoError(t, err) @@ -214,6 +263,59 @@ func TestAsynqWorkerDriverLifecycle(t *testing.T) { require.NoError(t, err) } +func TestAsynqWorkerDispatchTracksExecution(t *testing.T) { + _ = idgen.Init(1) + testDB, mr, cleanup := testhelper.SetupTestEnvironment(t) + defer cleanup() + + ctx := core.NewContext(context.Background()) + ctx.Config().SetSource(core.NewMapSource(nil)) + require.NoError(t, ctx.Config().Resolve()) + core.Provide[contracts.DBService](ctx, &testDBService{db: testDB}) + + var processed atomic.Bool + ctx.Tasks().Register("system:cleanup", func(_ context.Context, _ []byte) (*contracts.TaskResultDTO, error) { + processed.Store(true) + return &contracts.TaskResultDTO{Message: "cleaned 3 files"}, nil + }, + extpoints.WithTaskType("system_cleanup"), + extpoints.WithTaskName("系统垃圾清理"), + extpoints.WithTaskQueue("default"), + extpoints.WithTaskRetry(1), + extpoints.WithTaskRetryable(true), + ) + + workerPlugin := driver_asynq_worker.New( + driver_asynq_worker.WithRedisOpt(asynq.RedisClientOpt{Addr: mr.Addr()}), + driver_asynq_worker.WithConcurrency(2), + driver_asynq_worker.WithShutdownTimeout(2*time.Second), + ) + require.NoError(t, workerPlugin.Apply(ctx)) + require.NoError(t, workerPlugin.Start(context.Background())) + t.Cleanup(func() { + _ = workerPlugin.Stop(context.Background()) + }) + + taskSvc, err := core.Inject[contracts.TaskService](ctx) + require.NoError(t, err) + + taskID, err := taskSvc.Dispatch(context.Background(), "system_cleanup", []byte(`{}`), "manual") + require.NoError(t, err) + require.NotEmpty(t, taskID) + + require.Eventually(t, func() bool { + return processed.Load() + }, 5*time.Second, 50*time.Millisecond, "asynq worker should execute dispatched func handler") + + require.Eventually(t, func() bool { + execs, _, listErr := taskSvc.ListExecutions(context.Background(), "", "", 1, 10) + if listErr != nil || len(execs) == 0 { + return false + } + return execs[0].TaskID == taskID && execs[0].Status == "succeeded" && execs[0].Result == "cleaned 3 files" + }, 5*time.Second, 50*time.Millisecond, "task execution should become succeeded after worker runs") +} + func TestAsynqCronDriverLifecycle(t *testing.T) { mr, err := miniredis.Run() require.NoError(t, err) diff --git a/backend/plugins/infra/cache/limiter.go b/backend/plugins/infra/cache/limiter.go new file mode 100644 index 00000000..13a9c8a8 --- /dev/null +++ b/backend/plugins/infra/cache/limiter.go @@ -0,0 +1,100 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package cache + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/limiter" + "context" + + "github.com/go-redis/redis_rate/v10" + "github.com/redis/go-redis/v9" +) + +type redisLimiterImpl struct { + limiter *redis_rate.Limiter + keyPrefix string +} + +func newRedisLimiter(client redis.UniversalClient, keyPrefix string) contracts.LimiterService { + return &redisLimiterImpl{ + limiter: redis_rate.NewLimiter(client), + keyPrefix: keyPrefix, + } +} + +func (r *redisLimiterImpl) prefixedKey(key string) string { + if r.keyPrefix != "" { + return r.keyPrefix + "limiter:" + key + } + return PrefixedKey("limiter:" + key) +} + +func (r *redisLimiterImpl) Allow(ctx context.Context, key string, rate contracts.Rate) (*contracts.RateLimitResult, error) { + return r.AllowN(ctx, key, rate, 1) +} + +func (r *redisLimiterImpl) AllowN(ctx context.Context, key string, rate contracts.Rate, n int) (*contracts.RateLimitResult, error) { + limit := redis_rate.Limit{ + Rate: rate.Limit, + Period: rate.Period, + Burst: rate.Limit, + } + + res, err := r.limiter.AllowN(ctx, r.prefixedKey(key), limit, n) + if err != nil { + return nil, err + } + + return &contracts.RateLimitResult{ + Allowed: res.Allowed > 0, + Remaining: res.Remaining, + ResetAfter: res.ResetAfter, + RetryAfter: res.RetryAfter, + }, nil +} + +func (r *redisLimiterImpl) Reset(ctx context.Context, key string) error { + return r.limiter.Reset(ctx, r.prefixedKey(key)) +} + +type memoryLimiterFallback struct { + limiter *limiter.MemoryLimiter +} + +func newMemoryLimiterFallback() contracts.LimiterService { + return &memoryLimiterFallback{ + limiter: limiter.NewMemoryLimiter(), + } +} + +func (m *memoryLimiterFallback) Allow(ctx context.Context, key string, rate contracts.Rate) (*contracts.RateLimitResult, error) { + res, err := m.limiter.Allow(ctx, key, limiter.Rate{Limit: rate.Limit, Period: rate.Period}) + if err != nil { + return nil, err + } + return &contracts.RateLimitResult{ + Allowed: res.Allowed, + Remaining: res.Remaining, + ResetAfter: res.ResetAfter, + RetryAfter: res.RetryAfter, + }, nil +} + +func (m *memoryLimiterFallback) AllowN(ctx context.Context, key string, rate contracts.Rate, n int) (*contracts.RateLimitResult, error) { + res, err := m.limiter.AllowN(ctx, key, limiter.Rate{Limit: rate.Limit, Period: rate.Period}, n) + if err != nil { + return nil, err + } + return &contracts.RateLimitResult{ + Allowed: res.Allowed, + Remaining: res.Remaining, + ResetAfter: res.ResetAfter, + RetryAfter: res.RetryAfter, + }, nil +} + +func (m *memoryLimiterFallback) Reset(ctx context.Context, key string) error { + return m.limiter.Reset(ctx, key) +} diff --git a/backend/plugins/infra/cache/limiter_test.go b/backend/plugins/infra/cache/limiter_test.go new file mode 100644 index 00000000..26c30aad --- /dev/null +++ b/backend/plugins/infra/cache/limiter_test.go @@ -0,0 +1,84 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package cache_test + +import ( + "Wavelet/core" + "Wavelet/core/contracts" + "Wavelet/plugins/infra/cache" + "context" + "testing" + "time" + + "github.com/alicebob/miniredis/v2" + "github.com/redis/go-redis/v9" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestRedisLimiterService(t *testing.T) { + mr, err := miniredis.Run() + require.NoError(t, err) + defer mr.Close() + + rdb := redis.NewClient(&redis.Options{ + Addr: mr.Addr(), + }) + defer func() { _ = rdb.Close() }() + + p := cache.New( + cache.WithRedis(rdb), + cache.WithKeyPrefix("test:"), + ) + ctx := core.NewContext(context.Background()) + ctx.Config().SetSource(core.NewMapSource(map[string]any{ + "redis.enabled": true, + })) + require.NoError(t, ctx.Config().Resolve()) + require.NoError(t, p.Apply(ctx)) + + limiter, err := core.Inject[contracts.LimiterService](ctx) + require.NoError(t, err) + require.NotNil(t, limiter) + + testCtx := context.Background() + rate := contracts.Rate{ + Limit: 3, + Period: time.Minute, + } + + // 1st request + res, err := limiter.Allow(testCtx, "user:123", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) + assert.Equal(t, 2, res.Remaining) + + // 2nd request + res, err = limiter.Allow(testCtx, "user:123", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) + assert.Equal(t, 1, res.Remaining) + + // 3rd request + res, err = limiter.Allow(testCtx, "user:123", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) + assert.Equal(t, 0, res.Remaining) + + // 4th request - blocked + res, err = limiter.Allow(testCtx, "user:123", rate) + require.NoError(t, err) + assert.False(t, res.Allowed) + assert.Equal(t, 0, res.Remaining) + assert.Greater(t, res.RetryAfter, time.Duration(0)) + + // Reset + err = limiter.Reset(testCtx, "user:123") + require.NoError(t, err) + + // Allowed again after reset + res, err = limiter.Allow(testCtx, "user:123", rate) + require.NoError(t, err) + assert.True(t, res.Allowed) +} diff --git a/backend/plugins/infra/cache/plugin.go b/backend/plugins/infra/cache/plugin.go index 0155da2d..b2f4f9d7 100644 --- a/backend/plugins/infra/cache/plugin.go +++ b/backend/plugins/infra/cache/plugin.go @@ -139,6 +139,15 @@ func (p *Plugin) Apply(ctx *core.Context) error { } core.Provide[contracts.CacheService](ctx, svc) + + var limiterSvc contracts.LimiterService + if redisClient != nil { + limiterSvc = newRedisLimiter(redisClient, p.keyPrefix) + } else { + limiterSvc = newMemoryLimiterFallback() + } + core.Provide[contracts.LimiterService](ctx, limiterSvc) + return nil } diff --git a/backend/plugins/infra/cache_memory/limiter.go b/backend/plugins/infra/cache_memory/limiter.go new file mode 100644 index 00000000..e847d0b3 --- /dev/null +++ b/backend/plugins/infra/cache_memory/limiter.go @@ -0,0 +1,56 @@ +// Copyright 2026 Arctel.net +// SPDX-License-Identifier: Apache-2.0 + +package cache_memory + +import ( + "Wavelet/core/contracts" + "Wavelet/pkg/limiter" + "context" +) + +type memoryLimiterAdapter struct { + limiter *limiter.MemoryLimiter +} + +func newMemoryLimiterService() contracts.LimiterService { + return &memoryLimiterAdapter{ + limiter: limiter.NewMemoryLimiter(), + } +} + +func (a *memoryLimiterAdapter) Allow(ctx context.Context, key string, rate contracts.Rate) (*contracts.RateLimitResult, error) { + res, err := a.limiter.Allow(ctx, key, limiter.Rate{ + Limit: rate.Limit, + Period: rate.Period, + }) + if err != nil { + return nil, err + } + return &contracts.RateLimitResult{ + Allowed: res.Allowed, + Remaining: res.Remaining, + ResetAfter: res.ResetAfter, + RetryAfter: res.RetryAfter, + }, nil +} + +func (a *memoryLimiterAdapter) AllowN(ctx context.Context, key string, rate contracts.Rate, n int) (*contracts.RateLimitResult, error) { + res, err := a.limiter.AllowN(ctx, key, limiter.Rate{ + Limit: rate.Limit, + Period: rate.Period, + }, n) + if err != nil { + return nil, err + } + return &contracts.RateLimitResult{ + Allowed: res.Allowed, + Remaining: res.Remaining, + ResetAfter: res.ResetAfter, + RetryAfter: res.RetryAfter, + }, nil +} + +func (a *memoryLimiterAdapter) Reset(ctx context.Context, key string) error { + return a.limiter.Reset(ctx, key) +} diff --git a/backend/plugins/infra/cache_memory/plugin.go b/backend/plugins/infra/cache_memory/plugin.go index 74c815ef..249ef9ad 100644 --- a/backend/plugins/infra/cache_memory/plugin.go +++ b/backend/plugins/infra/cache_memory/plugin.go @@ -79,5 +79,6 @@ func (p *Plugin) Apply(ctx *core.Context) error { } core.Provide[contracts.CacheService](ctx, svc) + core.Provide[contracts.LimiterService](ctx, newMemoryLimiterService()) return nil } diff --git a/backend/plugins/infra/cache_memory/plugin_test.go b/backend/plugins/infra/cache_memory/plugin_test.go index 66c5fdba..71cdc283 100644 --- a/backend/plugins/infra/cache_memory/plugin_test.go +++ b/backend/plugins/infra/cache_memory/plugin_test.go @@ -78,4 +78,14 @@ func TestCacheMemoryPlugin(t *testing.T) { var tempVal string err = cacheSvc.Get(reqCtx, "temp_key", &tempVal) assert.ErrorIs(t, err, contracts.ErrCacheMiss) + + // 6. LimiterService + limiterSvc, err := core.Inject[contracts.LimiterService](ctx) + require.NoError(t, err) + require.NotNil(t, limiterSvc) + + rateRes, err := limiterSvc.Allow(reqCtx, "key_a", contracts.Rate{Limit: 2, Period: time.Minute}) + require.NoError(t, err) + assert.True(t, rateRes.Allowed) + assert.Equal(t, 1, rateRes.Remaining) } diff --git a/backend/plugins/infra/config/source.go b/backend/plugins/infra/config/source.go index fd6c711f..7da2de08 100644 --- a/backend/plugins/infra/config/source.go +++ b/backend/plugins/infra/config/source.go @@ -16,8 +16,14 @@ import ( "github.com/spf13/viper" ) -// DefaultFileName is the configuration file looked up when CONFIG_PATH is unset. -const DefaultFileName = "config.yaml" +// DefaultBaseFileName is the default base configuration file path looked up relative to the workspace. +const DefaultBaseFileName = "manifest/config/config.default.yaml" + +// DefaultOverrideFileName is the default user override configuration file path looked up relative to the workspace. +const DefaultOverrideFileName = "manifest/config/config.yaml" + +// DefaultFileName is kept for backwards compatibility. +const DefaultFileName = DefaultOverrideFileName // EnvOnlyOrigin is reported by Describe when no configuration file was loaded. const EnvOnlyOrigin = "" @@ -33,47 +39,94 @@ type Option func(*Source) func WithPath(path string) Option { return func(s *Source) { s.path = path + s.pinned = true + } +} + +// WithDefaultPath pins the default base configuration file path. +func WithDefaultPath(path string) Option { + return func(s *Source) { + s.defaultPath = path } } // Source implements core.ConfigSource over a configuration file plus the process environment. type Source struct { - v *viper.Viper - path string - found bool + v *viper.Viper + path string + defaultPath string + pinned bool + defaultFound bool + overrideFound bool + found bool } -// NewSource loads the configuration file. A missing file is not an error: the source -// then serves environment values only, matching the behaviour the previous pkg/config -// loader had for deployments that configure everything through the environment. +// NewSource loads configuration files. By default, it loads manifest/config/config.default.yaml +func (s *Source) resolvePaths() { + if s.pinned { + return + } + if s.defaultPath == "" { + s.defaultPath, _ = findExistingUpward(DefaultBaseFileName) + } + if s.path == "" { + s.path = os.Getenv("CONFIG_PATH") + } + if s.path == "" { + s.path, _ = findExistingUpward(DefaultOverrideFileName) + } +} + +func readConfigFile(v *viper.Viper, path string, merge bool) (bool, error) { + if path == "" { + return false, nil + } + + v.SetConfigFile(path) + var err error + if merge { + err = v.MergeInConfig() + } else { + err = v.ReadInConfig() + } + + switch { + case err == nil: + return true, nil + case isNotFound(err): + return false, nil + default: + if _, statErr := os.Stat(path); statErr == nil { //nolint:gosec // path is vetted by bounded upward search or config options + return false, fmt.Errorf("infra/config: read %s: %w", path, err) + } + return false, nil + } +} + +// NewSource loads configuration files. By default, it loads manifest/config/config.default.yaml +// and merges manifest/config/config.yaml (or CONFIG_PATH) on top if present. A missing file is +// not an error: the source then serves environment values only. func NewSource(opts ...Option) (*Source, error) { s := &Source{} for _, opt := range opts { opt(s) } - - if s.path == "" { - s.path = os.Getenv("CONFIG_PATH") - } - if s.path == "" { - s.path = findConfigPath(DefaultFileName) - } + s.resolvePaths() v := viper.New() - v.SetConfigFile(s.path) - err := v.ReadInConfig() - switch { - case err == nil: - s.found = true - case isNotFound(err): - // No file: fall through to environment-only lookups. - default: - if _, statErr := os.Stat(s.path); statErr == nil { //nolint:gosec // s.path comes from CONFIG_PATH or a bounded upward search - return nil, fmt.Errorf("infra/config: read %s: %w", s.path, err) - } + var err error + s.defaultFound, err = readConfigFile(v, s.defaultPath, false) + if err != nil { + return nil, err } + s.overrideFound, err = readConfigFile(v, s.path, s.defaultFound) + if err != nil { + return nil, err + } + + s.found = s.defaultFound || s.overrideFound s.v = v return s, nil } @@ -104,23 +157,25 @@ func (s *Source) Describe() string { if !s.found { return EnvOnlyOrigin } - return s.path + if s.overrideFound { + return s.path + } + return s.defaultPath } -// findConfigPath searches upward from the working directory so tests and binaries run -// from backend/ still find the repository-root configuration file. -func findConfigPath(configPath string) string { - if _, err := os.Stat(configPath); err == nil { - return configPath +// findExistingUpward searches upward from the working directory for a relative file path. +func findExistingUpward(relativeFilePath string) (string, bool) { + if _, err := os.Stat(relativeFilePath); err == nil { + return relativeFilePath, true } dir := "." for range maxSearchDepth { dir += "/.." - path := dir + "/" + configPath + path := dir + "/" + relativeFilePath if _, err := os.Stat(path); err == nil { - return path + return path, true } } - return configPath + return "", false } diff --git a/backend/plugins/infra/config/source_test.go b/backend/plugins/infra/config/source_test.go index 8372e156..39481735 100644 --- a/backend/plugins/infra/config/source_test.go +++ b/backend/plugins/infra/config/source_test.go @@ -91,3 +91,117 @@ func TestSourcePrefersConfigPathEnvironmentVariable(t *testing.T) { require.True(t, ok) assert.Equal(t, ":9100", value) } + +func TestSourceFindsManifestConfigInParentDirectory(t *testing.T) { + tempRoot := t.TempDir() + manifestConfigDir := filepath.Join(tempRoot, "manifest", "config") + require.NoError(t, os.MkdirAll(manifestConfigDir, 0o755)) + require.NoError(t, os.WriteFile( + filepath.Join(manifestConfigDir, "config.yaml"), + []byte("app:\n addr: \":9200\"\n"), + 0o600, + )) + + subDir := filepath.Join(tempRoot, "backend", "cmd") + require.NoError(t, os.MkdirAll(subDir, 0o755)) + + origWd, err := os.Getwd() + require.NoError(t, err) + require.NoError(t, os.Chdir(subDir)) + t.Cleanup(func() { + _ = os.Chdir(origWd) + }) + + t.Setenv("CONFIG_PATH", "") + + src, err := config.NewSource() + require.NoError(t, err) + value, ok := src.Lookup("app.addr") + require.True(t, ok) + assert.Equal(t, ":9200", value) +} + +func TestSourceDefaultConfigMergedWithOverride(t *testing.T) { + tempRoot := t.TempDir() + manifestConfigDir := filepath.Join(tempRoot, "manifest", "config") + require.NoError(t, os.MkdirAll(manifestConfigDir, 0o755)) + + // 1. Write default configuration + require.NoError(t, os.WriteFile( + filepath.Join(manifestConfigDir, "config.default.yaml"), + []byte("app:\n addr: \":8000\"\n node_id: 1\n env: \"development\"\n"), + 0o600, + )) + + // 2. Write override configuration (only overrides addr) + require.NoError(t, os.WriteFile( + filepath.Join(manifestConfigDir, "config.yaml"), + []byte("app:\n addr: \":9500\"\n"), + 0o600, + )) + + subDir := filepath.Join(tempRoot, "backend") + require.NoError(t, os.MkdirAll(subDir, 0o755)) + + origWd, err := os.Getwd() + require.NoError(t, err) + require.NoError(t, os.Chdir(subDir)) + t.Cleanup(func() { + _ = os.Chdir(origWd) + }) + + t.Setenv("CONFIG_PATH", "") + + src, err := config.NewSource() + require.NoError(t, err) + + // Overridden by config.yaml + addr, ok := src.Lookup("app.addr") + require.True(t, ok) + assert.Equal(t, ":9500", addr) + + // Retained from config.default.yaml + nodeID, ok := src.Lookup("app.node_id") + require.True(t, ok) + assert.Equal(t, 1, nodeID) + + envVal, ok := src.Lookup("app.env") + require.True(t, ok) + assert.Equal(t, "development", envVal) +} + +func TestSourceDefaultConfigUsedWhenOverrideAbsent(t *testing.T) { + tempRoot := t.TempDir() + manifestConfigDir := filepath.Join(tempRoot, "manifest", "config") + require.NoError(t, os.MkdirAll(manifestConfigDir, 0o755)) + + // Write only default configuration + require.NoError(t, os.WriteFile( + filepath.Join(manifestConfigDir, "config.default.yaml"), + []byte("app:\n addr: \":8080\"\n env: \"testing\"\n"), + 0o600, + )) + + subDir := filepath.Join(tempRoot, "backend") + require.NoError(t, os.MkdirAll(subDir, 0o755)) + + origWd, err := os.Getwd() + require.NoError(t, err) + require.NoError(t, os.Chdir(subDir)) + t.Cleanup(func() { + _ = os.Chdir(origWd) + }) + + t.Setenv("CONFIG_PATH", "") + + src, err := config.NewSource() + require.NoError(t, err) + + addr, ok := src.Lookup("app.addr") + require.True(t, ok) + assert.Equal(t, ":8080", addr) + + envVal, ok := src.Lookup("app.env") + require.True(t, ok) + assert.Equal(t, "testing", envVal) +} diff --git a/backend/plugins/infra/storage/plugin.go b/backend/plugins/infra/storage/plugin.go index 3cd105af..1fd7fff9 100644 --- a/backend/plugins/infra/storage/plugin.go +++ b/backend/plugins/infra/storage/plugin.go @@ -48,25 +48,11 @@ func (p *Plugin) Name() string { // Apply mounts the storage service into the Context. func (p *Plugin) Apply(ctx *core.Context) error { - // Bind DBService - if db, err := core.Inject[contracts.DBService](ctx); err == nil && db != nil { + core.Bind[contracts.DBService](ctx, func(db contracts.DBService) { objectstore.SetDBService(db) diskcache.SetDBService(db) - } else { - core.When[contracts.DBService](ctx, func(db contracts.DBService) { - objectstore.SetDBService(db) - diskcache.SetDBService(db) - }) - } - - // Bind CacheService - if cache, err := core.Inject[contracts.CacheService](ctx); err == nil && cache != nil { - objectstore.SetCacheService(cache) - } else { - core.When[contracts.CacheService](ctx, func(cache contracts.CacheService) { - objectstore.SetCacheService(cache) - }) - } + }) + core.Bind[contracts.CacheService](ctx, objectstore.SetCacheService) ctx.OnDispose(func() error { objectstore.SetDBService(nil) diff --git a/docker-compose.yaml b/docker-compose.yaml deleted file mode 100644 index 88a7300e..00000000 --- a/docker-compose.yaml +++ /dev/null @@ -1,147 +0,0 @@ -# OpenFlare default compose file (`docker compose` / `docker-compose.yaml`). -# The Wavelet upstream stack is in docker-compose.wavelet.yml and is not used here. -services: - openflare: - build: - context: . - dockerfile: docker/Dockerfile - args: - VERSION: v0.9.9 -# image: ghcr.io/rain-kl/openflare:latest - restart: unless-stopped - env_file: .env - environment: - TZ: ${TZ:-Asia/Shanghai} - OTEL_EXPORTER_OTLP_ENDPOINT: ${OTEL_EXPORTER_OTLP_ENDPOINT:-http://jaeger:4317} - OTEL_EXPORTER_OTLP_INSECURE: ${OTEL_EXPORTER_OTLP_INSECURE:-true} - OTEL_SAMPLING_RATE: ${OTEL_SAMPLING_RATE:-1.0} - ports: - - "3000:3000" - volumes: - - ./uploads:/app/uploads - - ./data/sqlite:/app/data - depends_on: - postgres: - condition: service_healthy - redis: - condition: service_healthy - clickhouse: - condition: service_healthy - jaeger: - condition: service_started - - postgres: - image: postgres:17-alpine - restart: unless-stopped - ports: - - "5432:5432" - environment: - POSTGRES_DB: ${DB_NAME:-openflare} - POSTGRES_USER: ${DB_USERNAME:-openflare} - POSTGRES_PASSWORD: ${DB_PASSWORD:-replace-with-strong-password} - volumes: - - ./data/postgres_data:/var/lib/postgresql/data - healthcheck: - test: ["CMD-SHELL", "pg_isready -U ${DB_USERNAME:-openflare} -d ${DB_NAME:-openflare}"] - interval: 10s - timeout: 5s - retries: 5 - - redis: - image: valkey/valkey:8.0-alpine - restart: unless-stopped - command: ["valkey-server", "--appendonly", "yes"] - ports: - - "${REDIS_PORT:-6379}:6379" - volumes: - - ./data/valkey:/data - healthcheck: - test: ["CMD", "valkey-cli", "ping"] - interval: 10s - timeout: 5s - retries: 5 - start_period: 5s - - jaeger: - image: jaegertracing/jaeger:${JAEGER_VERSION:-2.19.0} - restart: unless-stopped - environment: - TZ: ${TZ:-Asia/Shanghai} - ports: - - "${JAEGER_UI_PORT:-16686}:16686" - - "${JAEGER_OTLP_GRPC_PORT:-4317}:4317" - - "${JAEGER_OTLP_HTTP_PORT:-4318}:4318" - - clickhouse: - image: clickhouse/clickhouse-server:25.3-alpine - restart: unless-stopped - environment: - CLICKHOUSE_DB: ${CLICKHOUSE_NAME:-openflare} - CLICKHOUSE_USER: ${CLICKHOUSE_USERNAME:-default} - CLICKHOUSE_PASSWORD: ${CLICKHOUSE_PASSWORD:-replace-with-clickhouse-password} - CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT: 1 - TZ: ${TZ:-Asia/Shanghai} - ulimits: - nofile: - soft: 262144 - hard: 262144 - ports: - - "8123:8123" - - "9000:9000" - volumes: - - ./data/clickhouse_data:/var/lib/clickhouse - - ./config/clickhouse/performance.xml:/etc/clickhouse-server/config.d/performance.xml:ro - healthcheck: - test: ["CMD", "clickhouse-client", "--user", "${CLICKHOUSE_USERNAME:-default}", "--password", "${CLICKHOUSE_PASSWORD:-replace-with-clickhouse-password}", "--query", "SELECT 1"] - interval: 10s - timeout: 5s - retries: 5 - start_period: 15s - - agent: - build: - context: . - dockerfile: docker/Dockerfile.agent - container_name: openflare-agent - restart: unless-stopped - ports: - - "80:80" - - "443:443" - - "127.0.0.1:18081:18081" - volumes: - - ./data/agent/:/data - environment: - OPENFLARE_SERVER_URL: "http://host.docker.internal:3000" - OPENFLARE_AGENT_TOKEN: "7c7c4c13df0f3a77866bcd8cde492610" - LOG_LEVEL: "debug" - extra_hosts: - - "host.docker.internal:host-gateway" - - relay: - build: - context: . - dockerfile: docker/Dockerfile.relay - container_name: openflare-relay - network_mode: host - restart: unless-stopped - volumes: - - ./data/relay/:/app/data - environment: - OPENFLARE_SERVER_URL: http://host.docker.internal:3000 - OPENFLARE_DISCOVERY_TOKEN: 85464eeb72c49abc430569d6b9c77f78 - LOG_LEVEL: "debug" - extra_hosts: - - "host.docker.internal:host-gateway" - - flared: - build: - context: . - dockerfile: docker/Dockerfile.flared - container_name: openflare-flared - network_mode: "host" - restart: unless-stopped - volumes: - - ./data/flared/:/app/data - environment: - OPENFLARE_SERVER_URL: "http://host.docker.internal:3000" - OPENFLARE_TUNNEL_TOKEN: deb0783ac1e264a9d86440169aca0f09 diff --git a/docker-compose.yml b/docker-compose.yml index c1c322e4..5a10d241 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -1,6 +1,147 @@ -# Isolated Wavelet compose filename. -# OpenFlare's default stack is docker-compose.yaml (`docker compose` prefers -# that name over this one). This stub is kept so `git merge wavelet/main` -# cannot restore a wavelet product stack under the default compose name. -include: - - path: docker-compose.yaml +# OpenFlare default compose file (`docker compose` / `docker-compose.yaml`). +# The Wavelet upstream stack is in docker-compose.wavelet.yml and is not used here. +services: + openflare: + build: + context: . + dockerfile: manifest/docker/Dockerfile + args: + VERSION: v0.9.9 + # image: ghcr.io/rain-kl/openflare:latest + restart: unless-stopped + env_file: .env + environment: + TZ: ${TZ:-Asia/Shanghai} + OTEL_EXPORTER_OTLP_ENDPOINT: ${OTEL_EXPORTER_OTLP_ENDPOINT:-http://jaeger:4317} + OTEL_EXPORTER_OTLP_INSECURE: ${OTEL_EXPORTER_OTLP_INSECURE:-true} + OTEL_SAMPLING_RATE: ${OTEL_SAMPLING_RATE:-1.0} + ports: + - "3000:3000" + volumes: + - ./uploads:/app/uploads + - ./data/sqlite:/app/data + depends_on: + postgres: + condition: service_healthy + redis: + condition: service_healthy + clickhouse: + condition: service_healthy + jaeger: + condition: service_started + + postgres: + image: postgres:17-alpine + restart: unless-stopped + ports: + - "5432:5432" + environment: + POSTGRES_DB: ${DB_NAME:-openflare} + POSTGRES_USER: ${DB_USERNAME:-openflare} + POSTGRES_PASSWORD: ${DB_PASSWORD:-replace-with-strong-password} + volumes: + - ./data/postgres_data:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U ${DB_USERNAME:-openflare} -d ${DB_NAME:-openflare}"] + interval: 10s + timeout: 5s + retries: 5 + + redis: + image: valkey/valkey:8.0-alpine + restart: unless-stopped + command: ["valkey-server", "--appendonly", "yes"] + ports: + - "${REDIS_PORT:-6379}:6379" + volumes: + - ./data/valkey:/data + healthcheck: + test: ["CMD", "valkey-cli", "ping"] + interval: 10s + timeout: 5s + retries: 5 + start_period: 5s + + jaeger: + image: jaegertracing/jaeger:${JAEGER_VERSION:-2.19.0} + restart: unless-stopped + environment: + TZ: ${TZ:-Asia/Shanghai} + ports: + - "${JAEGER_UI_PORT:-16686}:16686" + - "${JAEGER_OTLP_GRPC_PORT:-4317}:4317" + - "${JAEGER_OTLP_HTTP_PORT:-4318}:4318" + + clickhouse: + image: clickhouse/clickhouse-server:25.3-alpine + restart: unless-stopped + environment: + CLICKHOUSE_DB: ${CLICKHOUSE_NAME:-openflare} + CLICKHOUSE_USER: ${CLICKHOUSE_USERNAME:-default} + CLICKHOUSE_PASSWORD: ${CLICKHOUSE_PASSWORD:-replace-with-clickhouse-password} + CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT: 1 + TZ: ${TZ:-Asia/Shanghai} + ulimits: + nofile: + soft: 262144 + hard: 262144 + ports: + - "8123:8123" + - "9000:9000" + volumes: + - ./data/clickhouse_data:/var/lib/clickhouse + - ./config/clickhouse/performance.xml:/etc/clickhouse-server/config.d/performance.xml:ro + healthcheck: + test: ["CMD", "clickhouse-client", "--user", "${CLICKHOUSE_USERNAME:-default}", "--password", "${CLICKHOUSE_PASSWORD:-replace-with-clickhouse-password}", "--query", "SELECT 1"] + interval: 10s + timeout: 5s + retries: 5 + start_period: 15s + + agent: + build: + context: . + dockerfile: manifest/docker/Dockerfile.agent + container_name: openflare-agent + restart: unless-stopped + ports: + - "80:80" + - "443:443" + - "127.0.0.1:18081:18081" + volumes: + - ./data/agent/:/data + environment: + OPENFLARE_SERVER_URL: "http://host.docker.internal:3000" + OPENFLARE_AGENT_TOKEN: "7c7c4c13df0f3a77866bcd8cde492610" + LOG_LEVEL: "debug" + extra_hosts: + - "host.docker.internal:host-gateway" + + relay: + build: + context: . + dockerfile: manifest/docker/Dockerfile.relay + container_name: openflare-relay + network_mode: host + restart: unless-stopped + volumes: + - ./data/relay/:/app/data + environment: + OPENFLARE_SERVER_URL: http://host.docker.internal:3000 + OPENFLARE_DISCOVERY_TOKEN: 85464eeb72c49abc430569d6b9c77f78 + LOG_LEVEL: "debug" + extra_hosts: + - "host.docker.internal:host-gateway" + + flared: + build: + context: . + dockerfile: manifest/docker/Dockerfile.flared + container_name: openflare-flared + network_mode: "host" + restart: unless-stopped + volumes: + - ./data/flared/:/app/data + environment: + OPENFLARE_SERVER_URL: "http://host.docker.internal:3000" + OPENFLARE_TUNNEL_TOKEN: deb0783ac1e264a9d86440169aca0f09 diff --git a/docs/DEPLOYMENT.md b/docs/DEPLOYMENT.md index ff2ee328..235c53a4 100644 --- a/docs/DEPLOYMENT.md +++ b/docs/DEPLOYMENT.md @@ -23,8 +23,8 @@ ## 二、 部署配置准备 -系统在启动前会从当前目录加载 `config.yaml` 配置文件。 -生产环境部署前,请复制 `config.example.yaml` 为 `config.yaml`,并至少确认以下关键参数的配置: +系统在启动前会默认加载 `manifest/config/config.default.yaml`,并自动读取 `manifest/config/config.yaml`(或 `CONFIG_PATH` 环境变量指定的文件)进行覆盖。 +生产环境部署前,可在 `manifest/config/config.yaml` 中按需覆盖以下关键参数: ```yaml app: diff --git a/docs/WAVELET_DEVELOPER_GUIDE.md b/docs/WAVELET_DEVELOPER_GUIDE.md index 29c0b694..af901386 100644 --- a/docs/WAVELET_DEVELOPER_GUIDE.md +++ b/docs/WAVELET_DEVELOPER_GUIDE.md @@ -33,10 +33,9 @@ - [第二部分:整个项目的目录结构划分与包职责定义](#第二部分整个项目的目录结构划分与包职责定义) - [第三部分:框架核心提供给插件调用的公用能力矩阵 (Context Capability Matrix)](#第三部分框架核心提供给插件调用的公用能力矩阵-context-capability-matrix) - [第四部分:Cordis 插件分层开发规范与代码模板 (Plugin Layered Architecture & Code Templates)](#第四部分cordis-插件分层开发规范与代码模板-plugin-layered-architecture--code-templates) - - [1. 分型与选型策略 (模式 1 vs 模式 2)](#1-分型与选型策略-模式-1-vs-模式-2) - - [2. 模式 1:扁平自包含分层规范与完整代码模板](#2-模式-1扁平自包含分层规范与完整代码模板) - - [3. 模式 2:严格子包物理分层规范与完整代码模板](#3-模式-2严格子包物理分层规范与完整代码模板) - - [4. 各层核心职责边界与严格禁止防线 (Guardrails)](#4-各层核心职责边界与严格禁止防线-guardrails) + - [1. 统一标准分层架构与目录规范 (以 custom_example 为基准)](#1-统一标准分层架构与目录规范-以-custom_example-为基准) + - [2. 核心分层代码模板与实现范例](#2-核心分层代码模板与实现范例) + - [3. 各层核心职责边界与严格禁止防线 (Guardrails)](#3-各层核心职责边界与严格禁止防线-guardrails) --- @@ -603,9 +602,9 @@ Wavelet/ │ └── admin/ # 系统管理台与监控面板插件 │ └── downstream/ # 【下游二开项目模板与脚手架】 - ├── custom_plugins/ # 下游自定义业务插件 - ├── config.yaml # 声明启用的插件与配置文件 - └── main.go # 下游项目组合启动入口 + ├── README.md # 下游插件开发指南 + └── plugins/ # 下游自定义业务插件目录 + └── custom_example/ # 官方标准插件开发基准模板(含完整分层结构) ``` ### 各层职责与禁止规则 (Guardrails): @@ -615,11 +614,12 @@ Wavelet/ 2. **`core/contracts/`**: - **职责**:仅定义公开的 Go Interface 和公共 DTO。 - **严禁**:严禁包含任何具体实现逻辑或 SQL 操作。 -3. **`plugins/`**: - - **职责**:所有业务逻辑和驱动实现的归宿。遵循标准分层架构(Layered Architecture / MVC 变体)。 - - **分层模式选型**: - - **模式 1(极简单文件分层,极简微型插件专用)**:单 package 内部仅各保留 1 个对应文件(`plugin.go`, `handlers.go`, `service.go`, `repository.go`, `models.go`, `errs.go`, `migrations/`)。仅适用于单一实体、极小代码量 (<500行) 的微型插件。 - - **模式 2(标准独立子包分层架构,官方推荐标准)**:按职责严格物理分包(`plugin.go`, `handler/`, `service/`, `repository/`, `model/`, `errs/`, `migrations/`)。**子包内文件以纯业务实体命名(如 `user.go`、`config.go`),严禁在根包平铺 `handlers_*`、`service_*`、`repository_*` 等前缀文件**。编译器级强约束 `handler -> service -> repository -> model` 单向依赖。 +3. **`plugins/` 与 `downstream/`**: + - **职责**:所有业务逻辑和驱动实现的归宿。遵循统一标准分层架构(Layered Architecture / MVC 变体)。 + - **统一分层架构与标准开发模板**: + - **唯一基准模板**:以 [`backend/downstream/plugins/custom_example`](file:///Users/ryan/Code/Go/Wavelet/backend/downstream/plugins/custom_example) 为全项目统一基准模板。 + - **物理子包隔离结构**:包含 `plugin.go`, `consts/`, `controller/`, `service/`, `dao/`, `model/` [含 `entity/`, `do/`], `migrations/` [含 `postgres/`, `sqlite/`]。 + - **命名与依赖禁令**:**严禁在根包平铺 `handlers_*`、`service_*`、`dao_*` 等前缀文件**,子包内文件直接以纯业务实体命名(如 `hello.go`、`user.go`)。严格约束 `controller -> service -> dao -> model` 单向依赖。 - **严禁**:插件之间严禁跨包 import 内部私有代码,跨插件调用一律走 `contracts` 接口或 `EventBus`。 --- @@ -652,116 +652,105 @@ Wavelet/ # 第四部分:Cordis 插件分层开发规范与代码模板 (Plugin Layered Architecture & Code Templates) -为了统一规范 Wavelet 所有官方插件与下游业务二开插件的研发质量,每个插件在内部遵循 **标准分层架构(Layered Architecture / MVC 变体)**。 +为了统一规范 Wavelet 所有官方插件与下游业务二开插件的研发质量,全项目插件**统一以 [`backend/downstream/plugins/custom_example`](file:///Users/ryan/Code/Go/Wavelet/backend/downstream/plugins/custom_example) 为基准开发模板**,遵循物理子包隔离的标准分层架构(`controller -> service -> dao -> model`)。 -## 1. 分型与选型策略 (模式 1 vs 模式 2) +## 1. 统一标准分层架构与目录规范 (以 `custom_example` 为基准) -根据业务复杂度和规模采用不同的物理包组织方式: - -``` - ┌────────────────────────┐ - │ 插件分层模式选型策略 │ - └───────────┬────────────┘ - │ - ┌───────────────────────┴───────────────────────┐ - ▼ ▼ -【模式 1:极简单文件分层】 【模式 2:标准独立子包分层】 -适合:极简微型/Demo插件 (<500行) 适合:标准/中大型业务插件 (推荐标准) -结构:单 Package,每层仅对应 1 个同名文件 结构:严格分包 handler/, service/, repository/, model/ -禁令:严禁根目录平铺 handlers_* 等前缀文件 规范:子包内以业务实体命名 (如 user.go, order.go) +```text +backend/downstream/plugins/custom_example/ (或 backend/plugins/domain//) +├── plugin.go # [插件根入口] 实现 core.Plugin,负责装配依赖、路由与扩展点注册 +│ +├── consts/ # package consts:常量定义、配置键名与模块错误标识 +│ └── consts.go +│ +├── controller/ # package controller:HTTP API 接入层 (参数绑定、会话提取、信封响应) +│ └── hello/ # 业务分组子包 +│ └── hello.go # 业务接口 Handler 实现(直接以业务命名,禁止 controller_hello.go) +│ +├── service/ # package service:核心业务逻辑层 (业务用例、事务编排、事件发布) +│ └── hello.go # 业务用例实现(纯 Go 逻辑,禁止依赖 *gin.Context) +│ +├── dao/ # package dao:数据持久化访问层 DAL (GORM CRUD、SQL 转义防注入) +│ └── hello.go # 数据库访问实现(直接以业务命名,禁止 dao_hello.go) +│ +├── model/ # package model:纯数据实体与传输对象 (零 Web/数据库框架依赖) +│ ├── entity/ # 数据库映射实体 (TableName() 必须带专属表前缀) +│ │ └── hello.go +│ └── do/ # 领域对象、请求 Request DTO 与响应 Response DTO +│ └── hello.go +│ +└── migrations/ # Goose SQL 独立迁移嵌入目录 (//go:embed) + ├── postgres/ # PostgreSQL 专属迁移 SQL + │ └── 20260901000001_init_hello.sql + └── sqlite/ # SQLite 专属迁移 SQL + └── 20260901000001_init_hello.sql ``` -| 维度 | 模式 1:极简单文件分层 (Single-File Flat) | 模式 2:标准独立子包分层 (Standard Sub-packages) | -| :--- | :--- | :--- | -| **适用场景** | 极简微型插件、单一实体(仅用于小型工具/示例) | 标准业务插件、包含多实体/多接口(**官方推荐标准**) | -| **代码量规模** | 通常 < 500 行 | 通常 ≥ 500 行(如 `upload`, `auth`, `admin`, `order`) | -| **Go 包形态** | 单一 Go Package,各层级仅各 1 个同名文件 | 按职责严格物理子目录分包,编译级强约束单向依赖 | -| **命名禁令** | **严禁在根目录平铺 `handlers_*`、`service_*` 文件** | **子包内文件直接以业务命名(如 `user.go`),禁止带 `handler_*` 前缀** | +> ⚠️ **严禁规则**: +> - **严禁在根目录平铺文件**:严禁在插件根目录下创建 `handlers_*.go`、`service_*.go`、`dao_*.go` 等前缀文件。 +> - **严禁跨层违规调用**:严格约束 `controller -> service -> dao -> model` 单向依赖。 +> - **严禁跨插件私有导入**:跨插件调用一律走 `contracts` 接口或 `EventBus`。 --- -## 2. 模式 1:极简单文件分层规范与完整代码模板 +## 2. 核心分层代码模板与实现范例 -### 2.1 目录结构 -```text -backend/plugins/domain// -├── plugin.go # [Cordis 接入层] 实现 core.Plugin,负责 Apply 组装与扩展点注册 -├── handlers.go # [Handler 层] 单一文件:Gin API Handler -├── service.go # [Service 层] 单一文件:核心业务用例 -├── repository.go # [Repository 层] 单一文件:GORM / DB 操作 -├── models.go # [Model 层] 单一文件:实体与 DTO -├── errs.go # [Error 层] 单一文件:错误常量 -├── plugin_test.go # 插件级单元与集成测试 -└── migrations/ # Goose SQL 嵌入文件 (//go:embed) - └── 20260828000001_init_.sql -``` - -> ⚠️ **严禁规则**:当单一文件膨胀或需要拆分多个业务实体时,**严禁在根目录创建 `handlers_user.go`, `handlers_admin.go`, `service_user.go` 等前缀文件**,必须立即重构并迁移为 **模式 2(标准独立子包分层架构)**! - -### 2.2 核心代码模板 (模式 1) - -#### (1) `plugin.go` (插件入口与装配) +#### (1) `plugin.go` (插件装配入口) ```go -package order +package hello import ( "embed" - "reflect" "github.com/Rain-kl/Wavelet/core" "github.com/Rain-kl/Wavelet/core/contracts" "github.com/gin-gonic/gin" ) -//go:embed migrations/*.sql -var orderMigrations embed.FS +//go:embed migrations/postgres/*.sql +var pgMigrations embed.FS -const PluginName = "domain.order" +//go:embed migrations/sqlite/*.sql +var sqliteMigrations embed.FS -type Plugin struct { - svc *OrderService -} +type Plugin struct{} func New() *Plugin { return &Plugin{} } func (p *Plugin) Name() string { - return PluginName -} - -func (p *Plugin) Inject() []reflect.Type { - return []reflect.Type{ - reflect.TypeFor[contracts.DBService](), - } + return "custom_example" } func (p *Plugin) Apply(ctx *core.Context) error { // 1. 注册专属数据库迁移 - ctx.Migrations().Register("order", orderMigrations) + ctx.Migrations().Register("custom_example", pgMigrations) - // 2. 初始化持久层与服务层 - repo := newOrderRepository(ctx) - p.svc = newOrderService(ctx, repo) + // 2. 解析依赖并装配各层 + var authSvc contracts.AuthService + if err := core.Using[contracts.AuthService](ctx, func(svc contracts.AuthService) { authSvc = svc }); err != nil { + return err + } - // 3. 注册 HTTP 路由组 - authSvc, _ := core.Inject[contracts.AuthService](ctx) - group := ctx.Router().Group("/api/v1/orders") - if authSvc != nil { - group.Use(authSvc.RequireAuthMiddleware()) - } - { - group.POST("", p.handleCreateOrder) - group.GET("/:id", p.handleGetOrderDetail) - } + // 3. 注册 HTTP 路由组与中间件 + g := ctx.Router().Group("/api/v1/custom", authSvc.RequireAuthMiddleware().(gin.HandlerFunc)) + g.GET("/hello", func(c *gin.Context) { + user, err := authSvc.GetCurrentUser(c.Request.Context()) + if err != nil { + c.JSON(401, gin.H{"error": "unauthorized"}) + return + } + c.JSON(200, gin.H{"message": "Hello " + user.Username}) + }) return nil } ``` -#### (2) `handlers.go` (Controller 层) +#### (2) `controller/hello/hello.go` (HTTP 控制器层) ```go -package order +package hello import ( "net/http" @@ -771,41 +760,39 @@ import ( "github.com/gin-gonic/gin" ) -// @Summary 创建订单 -// @Description 创建新的用户订单 -// @Tags Order -// @Accept json -// @Produce json -// @Param request body CreateOrderRequest true "创建参数" -// @Success 200 {object} response.Envelope{data=OrderDTO} "成功" -// @Failure 400 {object} response.Envelope "参数错误" -// @Router /api/v1/orders [post] -func (p *Plugin) handleCreateOrder(c *gin.Context) { - var req CreateOrderRequest - if err := c.ShouldBindJSON(&req); err != nil { - response.AbortBadRequest(c, errBindParamsFailed) - return - } +type Controller struct { + svc HelloService +} +func NewController(svc HelloService) *Controller { + return &Controller{svc: svc} +} + +// @Summary 获取欢迎信息 +// @Tags Hello +// @Produce json +// @Success 200 {object} response.Envelope{data=do.HelloResponse} "成功" +// @Router /api/v1/custom/hello [get] +func (ctrl *Controller) GetHello(c *gin.Context) { user, ok := oauth.GetCurrentUser(c) if !ok { - response.AbortUnauthorized(c, errUnauthorized) + response.AbortUnauthorized(c, "errUnauthorized") return } - order, err := p.svc.CreateOrder(c.Request.Context(), user.ID, req) + res, err := ctrl.svc.SayHello(c.Request.Context(), user.ID) if err != nil { - response.AbortInternal(c, errCreateOrderFailed) + response.AbortInternal(c, "errInternalServer") return } - c.JSON(http.StatusOK, response.OK(order)) + c.JSON(http.StatusOK, response.OK(res)) } ``` -#### (3) `service.go` (Service 业务逻辑层) +#### (3) `service/hello.go` (业务逻辑层) ```go -package order +package service import ( "context" @@ -813,44 +800,31 @@ import ( "github.com/Rain-kl/Wavelet/core" ) -type OrderService struct { - ctx *core.Context - repo *orderRepository +type HelloService struct { + ctx *core.Context + dao HelloDAO } -func newOrderService(ctx *core.Context, repo *orderRepository) *OrderService { - return &OrderService{ctx: ctx, repo: repo} +func NewHelloService(ctx *core.Context, dao HelloDAO) *HelloService { + return &HelloService{ctx: ctx, dao: dao} } -func (s *OrderService) CreateOrder(ctx context.Context, userID string, req CreateOrderRequest) (*OrderDTO, error) { - order := &OrderModel{ - UserID: userID, - Amount: req.Amount, - Status: "pending", - } - - if err := s.repo.Create(ctx, order); err != nil { +func (s *HelloService) SayHello(ctx context.Context, userID string) (*do.HelloResponse, error) { + record, err := s.dao.GetByUserID(ctx, userID) + if err != nil { return nil, err } // 发射领域事件 - s.ctx.Events().Emit(ctx, "order:created", OrderCreatedEvent{ - OrderID: order.ID, - UserID: order.UserID, - Amount: order.Amount, - }) + s.ctx.Events().Emit(ctx, "custom:hello_visited", map[string]any{"user_id": userID}) - return &OrderDTO{ - ID: order.ID, - Amount: order.Amount, - Status: order.Status, - }, nil + return &do.HelloResponse{Message: "Hello " + record.Name}, nil } ``` -#### (4) `repository.go` (Repository 数据访问层) +#### (4) `dao/hello.go` (数据访问持久化层 DAL) ```go -package order +package dao import ( "context" @@ -861,178 +835,104 @@ import ( "gorm.io/gorm" ) -type orderRepository struct { +type HelloDAO struct { ctx *core.Context } -func newOrderRepository(ctx *core.Context) *orderRepository { - return &orderRepository{ctx: ctx} +func NewHelloDAO(ctx *core.Context) *HelloDAO { + return &HelloDAO{ctx: ctx} } -func (r *orderRepository) getDB(ctx context.Context) *gorm.DB { - if dbSvc, err := core.Inject[contracts.DBService](r.ctx); err == nil && dbSvc != nil { +func (d *HelloDAO) getDB(ctx context.Context) *gorm.DB { + if dbSvc, err := core.Inject[contracts.DBService](d.ctx); err == nil && dbSvc != nil { return dbSvc.GetDB().WithContext(ctx) } return nil } -func (r *orderRepository) Create(ctx context.Context, order *OrderModel) error { - return r.getDB(ctx).Create(order).Error +func (d *HelloDAO) GetByUserID(ctx context.Context, userID string) (*entity.HelloEntity, error) { + var item entity.HelloEntity + err := d.getDB(ctx).Where("user_id = ?", userID).First(&item).Error + return &item, err } -func (r *orderRepository) SearchByKeyword(ctx context.Context, keyword string) ([]OrderModel, error) { - var list []OrderModel +func (d *HelloDAO) SearchByKeyword(ctx context.Context, keyword string) ([]entity.HelloEntity, error) { + var list []entity.HelloEntity // SQL LIKE 防注入与通配符转义规范 safeKeyword := util.EscapeLike(keyword) + "%" - err := r.getDB(ctx).Where("status LIKE ? ESCAPE '\\'", safeKeyword).Find(&list).Error + err := d.getDB(ctx).Where("name LIKE ? ESCAPE '\\'", safeKeyword).Find(&list).Error return list, err } ``` -#### (5) `models.go` 与 `errs.go` +#### (5) `model/entity/` 与 `model/do/` ```go -// models.go -package order +// model/entity/hello.go +package entity import "time" -type OrderModel struct { +type HelloEntity struct { ID string `gorm:"column:id;primaryKey;size:64" json:"id"` UserID string `gorm:"column:user_id;index;size:64;not null" json:"user_id"` - Amount int64 `gorm:"column:amount;not null" json:"amount"` - Status string `gorm:"column:status;size:32;index;not null;default:'pending'" json:"status"` + Name string `gorm:"column:name;size:128;not null" json:"name"` CreatedAt time.Time `gorm:"column:created_at;autoCreateTime" json:"created_at"` UpdatedAt time.Time `gorm:"column:updated_at;autoUpdateTime" json:"updated_at"` } -func (OrderModel) TableName() string { - return "w_orders" -} - -type CreateOrderRequest struct { - Amount int64 `json:"amount" binding:"required,gt=0"` -} - -type OrderDTO struct { - ID string `json:"id"` - Amount int64 `json:"amount"` - Status string `json:"status"` -} - -type OrderCreatedEvent struct { - OrderID string `json:"order_id"` - UserID string `json:"user_id"` - Amount int64 `json:"amount"` +func (HelloEntity) TableName() string { + return "w_custom_hello" } ``` ```go -// errs.go -package order +// model/do/hello.go +package do + +type HelloResponse struct { + Message string `json:"message"` +} + +type CreateHelloRequest struct { + Name string `json:"name" binding:"required,max=128"` +} +``` + +#### (6) `consts/consts.go` (常量与错误定义) +```go +package consts const ( - errBindParamsFailed = "errBindParamsFailed" - errUnauthorized = "errUnauthorized" - errCreateOrderFailed = "errCreateOrderFailed" + PluginName = "custom_example" + + // 模块内部错误码标识 (camelCase) + ErrUserNotFound = "errUserNotFound" + ErrInvalidParams = "errInvalidParams" + ErrOperationFail = "errOperationFail" ) ``` --- -## 3. 模式 2:标准独立子包物理分层规范与完整代码模板 (推荐标准) - -用于标准与中大型业务插件,各层使用独立的 Go package 物理隔离。 - -### 3.1 目录结构与文件命名规约 -```text -backend/plugins/domain/order/ -├── plugin.go # [插件根入口] 实现 core.Plugin,装配各子包并向 Cordis 注册 -│ -├── handler/ # package handler:HTTP API 接入层(或 controller/) -│ ├── router.go # 路由组挂载与中间件绑定 -│ └── order.go # 订单相关 Handler(以业务直接命名,禁止 handlers_order.go) -│ -├── service/ # package service:核心业务逻辑层 -│ ├── service.go # 业务用例接口定义 (Service Interface) -│ └── order.go # 订单业务用例实现(以业务直接命名,禁止 service_order.go) -│ -├── repository/ # package repository:数据访问持久化层 (DAL) -│ ├── repository.go # 仓储通用方法与工厂 -│ └── order.go # 订单仓储持久化实现(以业务直接命名,禁止 repository_order.go) -│ -├── model/ # package model (或 models/):纯领域实体与传输对象(零外部框架依赖) -│ ├── entity.go # 数据库映射实体 (TableName() 必须带 w__ 前缀) -│ ├── dto.go # 请求与响应 DTO -│ └── events.go # 领域事件结构体 -│ -├── errs/ # package errs:错误常量与错误码定义 (或根目录 errs.go) -│ └── errs.go -│ -└── migrations/ # Goose SQL 独立迁移嵌入文件 (//go:embed) - └── 20260828000001_init_order.sql -``` - -### 3.2 模式 2 核心装配代码范例 (`plugin.go`) -```go -package order - -import ( - "embed" - - "github.com/Rain-kl/Wavelet/core" - "github.com/Rain-kl/Wavelet/core/contracts" - "github.com/Rain-kl/Wavelet/plugins/domain/order/handler" - "github.com/Rain-kl/Wavelet/plugins/domain/order/repository" - "github.com/Rain-kl/Wavelet/plugins/domain/order/service" -) - -//go:embed migrations/*.sql -var orderMigrations embed.FS - -type Plugin struct{} - -func (p *Plugin) Name() string { - return "domain.order" -} - -func (p *Plugin) Apply(ctx *core.Context) error { - // 1. 注册迁移 - ctx.Migrations().Register("order", orderMigrations) - - // 2. 构造数据层与服务层 - repo := repository.NewOrderRepository(ctx) - svc := service.NewOrderService(ctx, repo) - - // 3. 构造 Handler 并挂载路由 - h := handler.NewOrderHandler(svc) - authSvc, _ := core.Inject[contracts.AuthService](ctx) - handler.RegisterRoutes(ctx.Router(), h, authSvc) - - return nil -} -``` - ---- - -## 4. 各层核心职责边界与严格禁止防线 (Guardrails) +## 3. 各层核心职责边界与严格禁止防线 (Guardrails) ```text ┌──────────────────────────────────────────────────────────────────┐ -│ Controller / Handler 层 (HTTP 接入) │ +│ Controller 层 (HTTP API 接入) │ │ • 参数绑定 ShouldBindJSON • 用户会话 oauth.GetCurrentUser │ │ • 统一信封 response.OK/Abort* • 严禁 SQL 操作 / 严禁重度业务 │ └─────────────────────────────────┬────────────────────────────────┘ - │ 调用 Service (入参 context.Context) + │ 调用 Service (入参首位 context.Context) ▼ ┌──────────────────────────────────────────────────────────────────┐ │ Service 层 (业务用例 & 领域逻辑) │ │ • 纯 Go 逻辑 (零 Web 依赖) • 事务编排 ctx.DB().Transaction │ │ • 领域事件 ctx.Events().Emit • 严禁 import gin / c.JSON │ └─────────────────────────────────┬────────────────────────────────┘ - │ 调用 Repository 接口 + │ 调用 DAO 接口 ▼ ┌──────────────────────────────────────────────────────────────────┐ -│ Repository 层 (数据持久化 DAL) │ +│ DAO 层 (数据持久化 DAL) │ │ • GORM CRUD 与查询 • EscapeLike 通配符安全转义 │ │ • 严禁反向依赖 Service/Controller • 严禁越权读写其他插件数据表 │ └─────────────────────────────────┬────────────────────────────────┘ @@ -1040,12 +940,12 @@ func (p *Plugin) Apply(ctx *core.Context) error { ▼ ┌──────────────────────────────────────────────────────────────────┐ │ Model 层 (纯实体 & DTO) │ -│ • TableName() 带专属表前缀 • 请求/响应结构体 │ +│ • TableName() 带专属表前缀 • entity/ 与 do/ 明确拆分 │ │ • 零值与 DB 默认值匹配 • 无任何上层包依赖 │ └──────────────────────────────────────────────────────────────────┘ ``` -1. **表单一所有者原则 (Single Owner Principle)**:数据表有且仅由所属插件操作(表名统一前缀 `w__*`),跨插件一律通过公开契约 Interface 或 EventBus 协同。 +1. **表单一所有者原则 (Single Owner Principle)**:数据表有且仅由所属插件声明与维护(表名统一前缀 `w__*`),跨插件一律通过公开契约 Interface 或 EventBus 协同。 2. **LIKE 查询安全防注入**:所有涉及用户输入的模糊查询,必须经过 `util.EscapeLike` 转义通配符并显式声明 `ESCAPE '\\'` 语法。 3. **Goroutine 安全**:并发任务统一使用 `util.Go`,杜绝直接使用裸 `go func()`。 diff --git a/docs/changelog/index.md b/docs/changelog/index.md index 7e0cc19e..cac78333 100644 --- a/docs/changelog/index.md +++ b/docs/changelog/index.md @@ -16,9 +16,11 @@ sidebar: false - Handler 把「记录不存在 → 404、其它错误 → 400」的分支改走上游 `response.AbortNotFoundIfMissing` / `AbortBadRequestOnError`,不再在 OpenFlare 里各写一份。 - 控制面 `server` 插件按限界上下文重排目录:去掉 `openflare/` 与 `router/v1` 嵌套;业务在 `domain/`(site/fleet/pages 等),共享内核在 `kernel/`(model/repository 与适配器),HTTP 装配在 `httpapi`。接口路径与表结构不变。 - `server` 插件把 stamp、of_* SQL 与 ClickHouse 迁入单一 `migrate/` 包,updater 提到 `server/updater/`;删除已停用的 76 条历史迁移。全新安装会写入 OpenFlare 定时任务与产品配置默认值,已 stamp 的升级库不重插。 +- 彻底治理跨组件调用与规约违例:严格遵循 Cordis 插件分层与单一表所有者原则,全面消除业务对上游内部实现的私有 import,统一面向 `backend/core/contracts` 编程;将契约 DTO 持久化解耦并增强通用标准库序列化支持回流 Wavelet 上游,全量架构规约检查与下游测试通过率达 100%。 ### 💄 其他/体验 +- 对齐 Wavelet 架构配置与构建布局(提交 807343c8):配置收敛至 `manifest/config/`(以 `config.default.yaml` 作为基础配置并支持 `config.yaml` 覆盖),Docker 镜像构建文件统一收敛至 `manifest/docker/` 并清理根目录冗余 `docker/` 目录,根目录 docker-compose 配置保持不变并切至新构建路径。 - 内嵌前端拷贝目标改为 `backend/plugins/drivers/driver_http/dist`,与上游 `//go:embed all:dist` 对齐;发布工作流改为读取 `backend/go.mod`。仓库内 `.gitconfig` 提供 `merge.ours` 驱动,合并上游时保留 OpenFlare 自有路径;Wavelet 的 `build-image.yml` 与 `docker-compose.yml` 已隔离,避免 canary 发布成 wavelet 镜像。 - 后端代码整体迁入 `backend/`,与上游 Wavelet 的仓库布局对齐(Cordis 插件化改造第一阶段),模块名保持 `Wavelet` 以保证上游包路径逐字一致。构建、测试、镜像与发布链路已同步调整,`make build-all` / `make dev` / `make swagger` 等本地命令用法不变;HTTP 接口与控制台行为均无变化。 - 引入上游 Cordis 微内核与平台插件到 `backend/{core,pkg,plugins}`(与上游逐字一致,可用 `scripts/sync-upstream.sh` 重复同步),并新增 `backend/openflare/share/` 承载多插件共享资源(控制消息协议、GeoIP、边缘日志)。此阶段仅落位结构与共享层,尚未改变运行时行为。 diff --git a/docs/design/index.md b/docs/design/index.md index 654ffe3d..0197383a 100644 --- a/docs/design/index.md +++ b/docs/design/index.md @@ -83,7 +83,7 @@ OpenFlare 已收敛为**单 monorepo**(Go 模块 `OpenFlare`)。控制面 Se | `pkg/` | 跨组件共享库(协议、渲染、GeoIP 等) | | `scripts/` | Swagger 生成、安装脚本等 | | `docs/` | VitePress 文档站与设计基线 | -| `docker/` | 各组件 Dockerfile | +| `manifest/docker/` | 各组件 Dockerfile | | `uploads/`、`data/` | 运行时上传目录与静态数据(`.gitignore` 忽略) | ### 1. Server 分层(`main.go` + `internal/`) diff --git a/docs/en/design/index.md b/docs/en/design/index.md index 37551a59..36809998 100644 --- a/docs/en/design/index.md +++ b/docs/en/design/index.md @@ -83,7 +83,7 @@ When contributing code, strictly follow this physical layering and directory div | `pkg/` | cross-component shared libs (protocol, rendering, GeoIP, etc.) | | `scripts/` | Swagger generation, install scripts, etc. | | `docs/` | VitePress docs site and design baseline | -| `docker/` | per-component Dockerfiles | +| `manifest/docker/` | per-component Dockerfiles | | `uploads/`, `data/` | runtime upload dir and static data (`.gitignore`d) | ### 1. Server Layering (`main.go` + `internal/`) diff --git a/docs/superpowers/plans/2026-08-27-cordis-plugin-architecture.md b/docs/superpowers/plans/2026-08-27-cordis-plugin-architecture.md deleted file mode 100644 index a1667639..00000000 --- a/docs/superpowers/plans/2026-08-27-cordis-plugin-architecture.md +++ /dev/null @@ -1,332 +0,0 @@ -# Wavelet Cordis 微内核与全插件化改造实施计划 - -> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. - -**Goal:** 将 Wavelet 架构重构为基于 Cordis 理念的微内核与全插件化架构,支持一切能力插件化、自包含数据迁移、多运行切面(API/Worker/Schedule/All)与下游极简二开扩展。 - -**Architecture:** -- **Core (`core/`)**: 纯净微内核,提供 Context 服务总线、泛型 IoC 容器(`Provide/Inject/Using`)、强类型 EventBus、生命周期状态机与 6 大扩展点协议,零外部业务依赖。 -- **Drivers (`plugins/drivers/`)**: Gin HTTP Server、Asynq Worker、Asynq Scheduler 封装为标准运行时驱动插件。 -- **Infra Plugins (`plugins/infra/`)**: 数据库(GORM/DBResolver)、三层缓存(RAM/Redis/PubSub)、日志(Zap/Otel)、对象存储插件化。 -- **Domain Plugins (`plugins/domain/`)**: Auth、User、MessageGateway、RiskControl、Admin 模块拆分为扁平自包含插件,自带独立 Goose 迁移。 - -**Tech Stack:** Go 1.25+, Gin, GORM, Asynq, Redis, Zap, OpenTelemetry, Goose, Viper, Cobra. - -## Global Constraints -- 保持 `core/` 绝对纯净,禁止 import Gin、GORM、Asynq 或具体业务包。 -- 插件之间严禁相互跨包 import 具体实现,跨插件交互一律通过 `core/contracts` 接口或 `ctx.Events()` 事件总线。 -- 严格遵循 Go 单元测试规范,测试临时目录统一使用 `t.TempDir()`,测试覆盖率严格达标。 -- 完成每个 Task 后必须确保代码能通过 `go build ./...` 与 `go test ./...` 检验并及时提交 Git。 - ---- - -### Task 1: 微内核基础契约与泛型 Context 服务总线 (`core/`) - -**Files:** -- Create: `core/types.go` -- Create: `core/manifest.go` -- Create: `core/container.go` -- Create: `core/context.go` -- Test: `core/context_test.go` - -**Interfaces:** -- Produces: `core.Plugin`, `core.Manifest`, `core.Context`, `core.Provide[T]`, `core.Inject[T]`, `core.Using[T]` - -- [ ] **Step 1: 编写 Context 与 IoC 容器的失败测试** - -```go -// core/context_test.go -package core_test - -import ( - "context" - "testing" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - "github.com/Rain-kl/Wavelet/core" -) - -type SampleService interface { - Greet(name string) string -} - -type sampleServiceImpl struct{} - -func (s *sampleServiceImpl) Greet(name string) string { - return "Hello, " + name -} - -func TestContextProvideAndInject(t *testing.T) { - ctx := core.NewContext(context.Background()) - core.Provide[SampleService](ctx, &sampleServiceImpl{}) - - svc, err := core.Inject[SampleService](ctx) - require.NoError(t, err) - assert.Equal(t, "Hello, Wavelet", svc.Greet("Wavelet")) -} - -func TestContextUsing(t *testing.T) { - ctx := core.NewContext(context.Background()) - var called bool - - err := core.Using(ctx, func(s SampleService) { - called = true - assert.Equal(t, "Hello, Cordis", s.Greet("Cordis")) - }) - assert.Error(t, err, "service not ready yet") - assert.False(t, called) - - core.Provide[SampleService](ctx, &sampleServiceImpl{}) - err = core.Using(ctx, func(s SampleService) { - called = true - assert.Equal(t, "Hello, Cordis", s.Greet("Cordis")) - }) - assert.NoError(t, err) - assert.True(t, called) -} -``` - -- [ ] **Step 2: 运行测试验证失败** - -Run: `go test -v ./core` -Expected: FAIL with compilation error (package not found). - -- [ ] **Step 3: 实现 Core 核心接口与泛型容器** - -编写 `core/types.go`、`core/manifest.go`、`core/container.go`、`core/context.go`,提供基于反射与类型推导的安全泛型服务存取、Scope 隔离与 Disposer 回调链。 - -- [ ] **Step 4: 运行测试验证通过** - -Run: `go test -v ./core` -Expected: PASS - -- [ ] **Step 5: 提交 Task 1 代码** - -```bash -git add core/ -git commit -m "feat(core): implement context service hub and generic ioc container" -``` - ---- - -### Task 2: 领域扩展点规范与强类型 EventBus (`core/extpoints/`, `core/events.go`) - -**Files:** -- Create: `core/events.go` -- Create: `core/extpoints/router.go` -- Create: `core/extpoints/migration.go` -- Create: `core/extpoints/task.go` -- Create: `core/extpoints/schedule.go` -- Create: `core/extpoints/setting.go` -- Test: `core/events_test.go` -- Test: `core/extpoints/extpoints_test.go` - -**Interfaces:** -- Consumes: `core.Context` -- Produces: `core.EventBus`, `core.RouterExtension`, `core.MigrationExtension`, `core.TaskExtension`, `core.ScheduleExtension`, `core.SettingExtension` - -- [ ] **Step 1: 编写 EventBus 与扩展点测试用例** - -```go -// core/events_test.go -package core_test - -import ( - "context" - "testing" - "github.com/stretchr/testify/assert" - "github.com/Rain-kl/Wavelet/core" -) - -type UserRegisteredEvent struct { - UserID string -} - -func TestEventBusPublishSubscribe(t *testing.T) { - bus := core.NewEventBus() - var receivedID string - - bus.On("user:registered", func(ctx context.Context, e UserRegisteredEvent) error { - receivedID = e.UserID - return nil - }) - - err := bus.Emit(context.Background(), "user:registered", UserRegisteredEvent{UserID: "u_999"}) - assert.NoError(t, err) - assert.Equal(t, "u_999", receivedID) -} -``` - -- [ ] **Step 2: 运行测试验证失败** - -Run: `go test -v ./core/...` -Expected: FAIL - -- [ ] **Step 3: 实现 EventBus 与 6 大扩展点适配器** - -编写 `core/events.go` 及 `core/extpoints/` 下各个领域的挂载收集器(Router 注册收集、Goose embed.FS 聚合器、Task/Schedule 声明表、Setting 模式注册表)。 - -- [ ] **Step 4: 运行测试验证通过** - -Run: `go test -v ./core/...` -Expected: PASS - -- [ ] **Step 5: 提交 Task 2 代码** - -```bash -git add core/ -git commit -m "feat(core): add typed eventbus and domain extension points" -``` - ---- - -### Task 3: 运行时驱动插件下沉 (`plugins/drivers/`) - -**Files:** -- Create: `plugins/drivers/driver_http/plugin.go` -- Create: `plugins/drivers/driver_asynq_worker/plugin.go` -- Create: `plugins/drivers/driver_asynq_cron/plugin.go` -- Test: `plugins/drivers/drivers_test.go` - -**Interfaces:** -- Consumes: `core.Plugin`, `core.Driver`, `core.Context` -- Produces: `DriverTypeHTTP`, `DriverTypeWorker`, `DriverTypeScheduler` - -- [ ] **Step 1: 编写 Driver 生命周期测试用例** - -测试驱动在接收到 `Start(ctx)` 和 `Stop(ctx)` 信号时的平滑启动与退出状态。 - -- [ ] **Step 2: 编写 Driver 实现** - -将 Gin HTTP Server、Asynq Worker Server、Asynq Scheduler 封装为标准 `core.Driver`,并在 `Apply(ctx)` 时挂载到 Context 驱动树。 - -- [ ] **Step 3: 运行驱动单元测试** - -Run: `go test -v ./plugins/drivers/...` -Expected: PASS - -- [ ] **Step 4: 提交 Task 3 代码** - -```bash -git add plugins/drivers/ -git commit -m "feat(plugins): implement runtime drivers for http, asynq worker, and cron" -``` - ---- - -### Task 4: 基础设施服务插件化 (`plugins/infra/`) - -**Files:** -- Create: `plugins/infra/database/plugin.go` (提供 GORM DBService) -- Create: `plugins/infra/cache/plugin.go` (提供 RAM/Redis 三层缓存) -- Create: `plugins/infra/logger/plugin.go` (提供 Zap/Otel 结构化日志) -- Create: `plugins/infra/storage/plugin.go` (提供统一对象存储) -- Test: `plugins/infra/infra_test.go` - -**Interfaces:** -- Produces: `contracts.DBService`, `contracts.CacheService`, `contracts.LoggerService`, `contracts.StorageService` - -- [ ] **Step 1: 编写基础设施插件注入与提取测试** -- [ ] **Step 2: 实现 4 大基础设施插件并封装现有 pkg 与 infra 底座** -- [ ] **Step 3: 运行基础设施测试验证** - -Run: `go test -v ./plugins/infra/...` -Expected: PASS - -- [ ] **Step 4: 提交 Task 4 代码** - -```bash -git add plugins/infra/ -git commit -m "feat(plugins): package database, cache, logger, and storage as infra plugins" -``` - ---- - -### Task 5: 业务领域插件化重构 (`plugins/domain/`) - -**Files:** -- Create: `plugins/domain/auth/` (认证、Session、Passkey、专属 migrations) -- Create: `plugins/domain/user/` (用户资料、角色权限、专属 migrations) -- Create: `plugins/domain/message_gateway/` (Bot网关、推送通道、Worker消费) -- Create: `plugins/domain/risk_control/` (IP限流、风控中间件) -- Create: `plugins/domain/admin/` (控制台、系统设置) -- Test: `plugins/domain/domain_test.go` - -**Interfaces:** -- Consumes: `contracts.DBService`, `contracts.CacheService`, `contracts.LoggerService` -- Produces: `contracts.AuthService`, `contracts.UserService` - -- [ ] **Step 1: 编写 Auth 与 User 插件业务装配与独立迁移测试** -- [ ] **Step 2: 将各业务模块迁移为扁平自包含插件,嵌入专属 Goose SQL 迁移** -- [ ] **Step 3: 运行业务插件集成测试** - -Run: `go test -v ./plugins/domain/...` -Expected: PASS - -- [ ] **Step 4: 提交 Task 5 代码** - -```bash -git add plugins/domain/ -git commit -m "feat(plugins): migrate auth, user, message_gateway, risk_control, admin to domain plugins" -``` - ---- - -### Task 6: 统一装配入口与运行时切面分发器 (`core/app.go`, `cmd/`) - -**Files:** -- Create: `core/app.go` -- Modify: `internal/cmd/root.go` -- Modify: `internal/cmd/api.go` -- Modify: `internal/cmd/worker.go` -- Modify: `internal/cmd/scheduler.go` -- Modify: `internal/cmd/all.go` -- Test: `core/app_test.go` - -**Interfaces:** -- Consumes: `core.App`, `core.Plugin`, `core.Driver` -- Produces: 统一 CLI 启动与优雅停机流程 - -- [ ] **Step 1: 编写 App 生命周期与 Profile 调度测试** -- [ ] **Step 2: 实现 `core.App` 编排引擎,无缝接入 `wavelet api / worker / schedule / all`** -- [ ] **Step 3: 运行启动与角色切面集成验证** - -Run: `go test -v ./core -run TestAppProfileDispatch` -Expected: PASS - -- [ ] **Step 4: 提交 Task 6 代码** - -```bash -git add core/ internal/cmd/ -git commit -m "feat(core): implement app profile lifecycle dispatcher and wire cli commands" -``` - ---- - -### Task 7: 下游脚手架、自定义示例插件与端到端验证 - -**Files:** -- Create: `downstream/custom_plugins/order/plugin.go` -- Create: `downstream/main.go` -- Create: `downstream/config.yaml` -- Test: `downstream/e2e_test.go` - -- [ ] **Step 1: 编写下游自定义业务插件并在下游 `main.go` 组装启动** -- [ ] **Step 2: 执行全量 E2E 测试,验证数据迁移、HTTP 路由访问、Worker 任务消费与平滑停机** -- [ ] **Step 3: 运行全局质量门禁检查** - -Run: -```bash -make test -make code-check -make format -``` -Expected: 全部 PASS,0 lint 报错。 - -- [ ] **Step 4: 提交 Task 7 代码** - -```bash -git add downstream/ -git commit -m "feat(downstream): add starter scaffold, example custom plugin, and e2e tests" -``` - diff --git a/docs/superpowers/plans/2026-08-28-cordis-architecture-alignment.md b/docs/superpowers/plans/2026-08-28-cordis-architecture-alignment.md deleted file mode 100644 index d3ea6bb8..00000000 --- a/docs/superpowers/plans/2026-08-28-cordis-architecture-alignment.md +++ /dev/null @@ -1,221 +0,0 @@ -# Cordis Architecture Alignment & Refactoring Implementation Plan - -> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. - -**Goal:** Implement Cordis spatiotemporal composability (scoped revertible effects and reactive fiber lifecycle state machine) in `backend/core`, and eliminate cross-plugin direct imports in domain repositories. - -**Architecture:** -1. Build scoped extension proxies on `core.Context` that automatically attach unregister callbacks to `ctx.OnDispose` in LIFO order upon registration. -2. Introduce `core/fiber.go` implementing the Fiber state machine (`PENDING -> LOADING -> ACTIVE -> UNLOADING -> DISPOSED`) with a reactive reconciler in `App`/`Container` ensuring dependency confluence. -3. Clean up defensive boundaries in `backend/plugins/domain/user` by removing direct `database.DB(ctx)` imports in favor of `contracts.DBService`. - -**Tech Stack:** Go 1.24+, GORM, Gin, Asynq, Cordis micro-kernel paradigm. - -## Global Constraints - -- Strictly preserve `backend/pkg/util/` purity (no Gin/GORM imports). -- Zero physically hardcoded temp directories in tests (use `t.TempDir()`). -- All Go error returns and logging must adhere to project standards. -- Follow Conventional Commits (`feat(core): ...`, `refactor(user): ...`). - ---- - -### Task 1: Scoped Revertible Effects for Core Extpoints - -**Files:** -- Create: `backend/core/scoped_extpoints.go` -- Modify: `backend/core/context.go` -- Modify: `backend/core/extpoints/task.go` -- Modify: `backend/core/extpoints/schedule.go` -- Modify: `backend/core/extpoints/setting.go` -- Test: `backend/core/context_test.go` - -**Interfaces:** -- Consumes: `core.Context`, `extpoints.RouterExtension`, `extpoints.TaskExtension`, `extpoints.ScheduleExtension`, `extpoints.SettingExtension`, `core.EventBus` -- Produces: Scoped extension methods on `Context` that automatically register LIFO disposers when routes, tasks, schedules, settings, and events are registered. - -- [ ] **Step 1: Write the failing test for scoped extpoints automatic teardown** - -In `backend/core/context_test.go`, add test cases verifying that registering routes, tasks, schedules, settings, and event listeners on a child context automatically registers unregister callbacks, and calling `childCtx.Dispose()` completely rolls them back: - -```go -func TestContext_ScopedExtpoints_RevertibleEffects(t *testing.T) { - root := NewContext(context.Background()) - child := root.Fork() - - // Register route, task, schedule, setting, event on child - rd := child.Router().GET("/test-route", func() {}) - assert.Equal(t, 1, len(root.Router().Routes())) - - child.Events().On("test:event", func() {}) - assert.Equal(t, 1, root.Events().Listeners("test:event")) - - // Dispose child - err := child.Dispose() - assert.NoError(t, err) - - // All child effects should be revoked - assert.Equal(t, 0, len(root.Router().Routes())) - assert.Equal(t, 0, root.Events().Listeners("test:event")) -} -``` - -- [ ] **Step 2: Run test to verify it fails** - -Run: `go test -v ./backend/core -run TestContext_ScopedExtpoints_RevertibleEffects` -Expected: FAIL (because current `Router().GET()` does not bind unregistration to `child.OnDispose`). - -- [ ] **Step 3: Implement Scoped Extpoints and Context bindings** - -1. In `backend/core/extpoints/task.go`, ensure `Unregister(taskType string) bool` exists. -2. In `backend/core/extpoints/schedule.go`, ensure `Unregister(name string) bool` exists. -3. In `backend/core/extpoints/setting.go`, ensure `Unregister(key string) bool` exists. -4. In `backend/core/scoped_extpoints.go` (or `context.go`), create scoped wrappers for `RouterExtension`, `TaskExtension`, `ScheduleExtension`, `SettingExtension` and `EventBus` that tie registrations to `ctx.OnDispose`. - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `go test -v ./backend/core -run TestContext_ScopedExtpoints_RevertibleEffects` -Expected: PASS - -- [ ] **Step 5: Commit** - -```bash -git add backend/core/ -git commit -m "feat(core): implement scoped revertible effects for context extpoints" -``` - ---- - -### Task 2: Plugin Fiber State Machine and Reactive Coeffects (Confluence) - -**Files:** -- Create: `backend/core/fiber.go` -- Create: `backend/core/fiber_test.go` -- Modify: `backend/core/app.go` -- Modify: `backend/core/types.go` -- Modify: `backend/core/container.go` - -**Interfaces:** -- Consumes: `core.Plugin`, `core.Context`, `core.Container` -- Produces: `core.DependentPlugin`, `core.Fiber`, `core.FiberState`, `App.Reconcile()` - -- [ ] **Step 1: Write the failing test for Fiber state machine and out-of-order registration confluence** - -In `backend/core/fiber_test.go`: - -```go -func TestFiber_ConfluenceAndReactiveActivation(t *testing.T) { - app := NewApp() - - // Plugin B depends on contracts.DBService, but is registered BEFORE DatabasePlugin (Plugin A) - pluginB := &mockDependentPlugin{ - name: "plugin-b", - deps: []reflect.Type{reflect.TypeFor[contracts.DBService]()}, - } - pluginA := &mockDBPlugin{name: "database"} - - app.Use(pluginB, pluginA) - - err := app.Start(context.Background()) - assert.NoError(t, err) - - // Verify both plugins reached FiberActive state and B executed Apply successfully after A provided DBService - assert.True(t, pluginB.applied) - assert.True(t, pluginA.applied) - - _ = app.Stop() -} -``` - -- [ ] **Step 2: Run test to verify it fails** - -Run: `go test -v ./backend/core -run TestFiber_ConfluenceAndReactiveActivation` -Expected: FAIL (because current `app.ApplyPlugins()` applies in static slice order without dependency reconciliation). - -- [ ] **Step 3: Implement Fiber State Machine and Reconciler** - -1. In `backend/core/types.go`, declare: -```go -type DependentPlugin interface { - Plugin - Inject() []reflect.Type -} -``` -2. In `backend/core/fiber.go`, implement `Fiber` with states (`FiberPending`, `FiberLoading`, `FiberActive`, `FiberUnloading`, `FiberDisposed`), child scoped context, and state transition methods. -3. In `backend/core/app.go`, integrate Fibers into `App` and implement iterative dependency reconciliation during `ApplyPlugins` and on dynamic `Provide`. - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `go test -v ./backend/core` -Expected: ALL PASS - -- [ ] **Step 5: Commit** - -```bash -git add backend/core/ -git commit -m "feat(core): implement plugin fiber state machine and reactive dependency reconciler" -``` - ---- - -### Task 3: Domain Plugin Isolation & Boundary Enforcement - -**Files:** -- Modify: `backend/plugins/domain/user/repository.go` -- Modify: `backend/plugins/domain/user/service.go` -- Modify: `backend/plugins/domain/user/handlers.go` -- Modify: `backend/plugins/domain/user/plugin.go` -- Test: `backend/plugins/domain/user/plugin_test.go` - -**Interfaces:** -- Consumes: `contracts.DBService` via `core.Inject` / `ctx.DB()` -- Produces: Decoupled User repository without direct `Wavelet/plugins/infra/database` imports. - -- [ ] **Step 1: Write/update test verifying User repository works with injected DBService** - -In `backend/plugins/domain/user/plugin_test.go`, test user CRUD operations resolving `contracts.DBService` through Context. - -- [ ] **Step 2: Run test to verify current state** - -Run: `go test -v ./backend/plugins/domain/user/...` - -- [ ] **Step 3: Refactor user repository to eliminate direct `plugins/infra/database` imports** - -In `backend/plugins/domain/user/repository.go`: -- Remove `import "Wavelet/plugins/infra/database"`. -- Obtain `*gorm.DB` via `ctx` (e.g. from context using `contracts.DBService` or context value). - -- [ ] **Step 4: Run tests to verify they pass** - -Run: `go test -v ./backend/plugins/domain/user/...` -Expected: PASS - -- [ ] **Step 5: Commit** - -```bash -git add backend/plugins/domain/user/ -git commit -m "refactor(user): decouple repository from direct database infra import" -``` - ---- - -### Task 4: Full Suite Verification & Quality Gate - -**Files:** -- Entire repository - -- [ ] **Step 1: Run all backend tests** - -Run: `cd backend && go test -v ./...` -Expected: ALL PASS - -- [ ] **Step 2: Run code-check and format** - -Run: `make code-check && make format` -Expected: 0 lint errors, clean formatting. - -- [ ] **Step 3: Commit any formatting or lint fixes** - -```bash -git commit -am "chore: format and verify code quality" -``` diff --git a/docs/superpowers/plans/2026-08-28-cordis-architecture-refactor.md b/docs/superpowers/plans/2026-08-28-cordis-architecture-refactor.md deleted file mode 100644 index 489731dd..00000000 --- a/docs/superpowers/plans/2026-08-28-cordis-architecture-refactor.md +++ /dev/null @@ -1,120 +0,0 @@ -# Cordis 架构重构实施计划 (Cordis Architecture Refactor Implementation Plan) - -> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. - -**Goal:** 依据 Cordis 时空可组合性元框架,彻底消除 Wavelet 后端的包级静态单例、`init()` 隐式副作用建连以及跨插件私有实现依赖,实现微内核纯洁化与契约驱动解耦。 - -**Architecture:** -1. 移除 `backend/core/context.go` 中的特权服务快捷方法(`DB()` / `Cache()`)。 -2. 将 `infra/database` 与 `infra/cache` 的连接初始化移至 `Plugin.Apply(ctx)`,并在 `ctx.OnDispose` 中注册 LIFO 逆操作(Close)。 -3. 重构全部 8 个 Domain 业务插件(`auth`、`user`、`admin`、`cap`、`message_gateway`、`risk_control`、`system`、`upload`),彻底斩断对 `infra/database`、`infra/cache` 及其他插件内部包的直接 import,统一面向 `contracts.DBService` / `contracts.CacheService`。 -4. 清除 `admin` 等插件的包级全局变量。 - -**Tech Stack:** Go 1.24+, GORM, Redis (go-redis/v9), Cordis micro-kernel, Goose migration. - -## Global Constraints - -- 严禁任何业务插件跨包 import `Wavelet/plugins/infra/database` 或 `Wavelet/plugins/infra/cache`。 -- 严禁跨插件 import 私有实现包(如 `admin` import `risk_control/logstore`)。 -- 保持 `backend/pkg/util/` 绝对纯净,禁止导入 Web/数据库框架。 -- 重构后必须确保 `go test ./...`、`make code-check` 与 `make format` 全部 0 错误通过。 - ---- - -### Task 1: 微内核纯洁化 (`backend/core/`) - -**Files:** -- Modify: `backend/core/context.go:240-260` -- Test: `backend/core/context_test.go` - -**Interfaces:** -- Consumes: `core.Context`, `core.Inject` -- Produces: 纯净无特权方法的 `core.Context` - -- [ ] **Step 1: 编写/更新 Context 纯洁性测试** -- [ ] **Step 2: 移除 `Context.DB()` 与 `Context.Cache()` 方法** -- [ ] **Step 3: 运行 `go test ./backend/core/...` 验证通过** - ---- - -### Task 2: 基础设施插件生命周期可逆化 (`backend/plugins/infra/`) - -**Files:** -- Modify: `backend/plugins/infra/database/postgres.go` -- Modify: `backend/plugins/infra/database/plugin.go` -- Modify: `backend/plugins/infra/cache/redis.go` -- Modify: `backend/plugins/infra/cache/plugin.go` -- Test: `backend/plugins/infra/infra_test.go` - -**Interfaces:** -- Consumes: `core.Plugin`, `contracts.DBService`, `contracts.CacheService` -- Produces: `contracts.DBService` 与 `contracts.CacheService`(带 `ctx.OnDispose` 逆操作) - -- [ ] **Step 1: 移除 `infra/database` 中的 `func init()` 及全局 `var db`,在 `Plugin.Apply` 中建连并注册 `ctx.OnDispose(sqlDB.Close)`** -- [ ] **Step 2: 移除 `infra/cache` 中的 `func init()` 及全局 `var Redis`,在 `Plugin.Apply` 中建连并注册 `ctx.OnDispose(client.Close)`** -- [ ] **Step 3: 运行 `go test ./backend/plugins/infra/...` 验证通过** - ---- - -### Task 3: 核心 Domain 插件防线重塑(Auth & User 插件) - -**Files:** -- Modify: `backend/plugins/domain/auth/*` -- Modify: `backend/plugins/domain/user/*` -- Test: `backend/plugins/domain/auth/plugin_test.go` -- Test: `backend/plugins/domain/user/plugin_test.go` - -**Interfaces:** -- Consumes: `contracts.DBService`, `contracts.CacheService` -- Produces: `contracts.AuthService`, `contracts.UserService` - -- [ ] **Step 1: 移除 `auth` 插件中对 `Wavelet/plugins/infra/database` 和 `cache` 的 import,改用插件持有的 `contracts.DBService` 与 `contracts.CacheService`** -- [ ] **Step 2: 移除 `user` 插件中对 `Wavelet/plugins/infra/database` 和 `cache` 的 import,改用 `contracts.DBService` 与 `contracts.CacheService`** -- [ ] **Step 3: 运行 `go test ./backend/plugins/domain/auth/... ./backend/plugins/domain/user/...` 验证通过** - ---- - -### Task 4: 业务 Domain 插件防线重塑(Cap, MessageGateway, RiskControl, System, Upload) - -**Files:** -- Modify: `backend/plugins/domain/cap/*` -- Modify: `backend/plugins/domain/message_gateway/*` -- Modify: `backend/plugins/domain/risk_control/*` -- Modify: `backend/plugins/domain/system/*` -- Modify: `backend/plugins/domain/upload/*` -- Test: `backend/plugins/domain/domain_test.go` - -**Interfaces:** -- Consumes: `contracts.DBService`, `contracts.CacheService` - -- [ ] **Step 1: 改造 `cap`、`message_gateway`、`risk_control`、`system`、`upload` 插件,移除所有 `infra/database` 和 `infra/cache` 的直接 import** -- [ ] **Step 2: 统一各插件内部 Repository / Service 的 DB / Cache 获取途径** -- [ ] **Step 3: 运行各插件单测验证通过** - ---- - -### Task 5: Admin 插件解耦与包级全局状态清除 - -**Files:** -- Modify: `backend/plugins/domain/admin/*` -- Test: `backend/plugins/domain/admin/plugin_test.go` - -**Interfaces:** -- Consumes: `contracts.DBService`, `contracts.CacheService`, `contracts.UserService`, `contracts.AuthService`, `ctx.Tasks()` - -- [ ] **Step 1: 移除 `admin` 插件中对 `risk_control/logstore`、`driver_asynq_worker`、`infra/storage/diskcache` 等私有包的直接 import** -- [ ] **Step 2: 清除 `admin/plugin.go` 中的 `globalUserSvc`、`globalAuthSvc`、`globalCoreCtx` 等包级变量** -- [ ] **Step 3: 运行 `go test ./backend/plugins/domain/admin/...` 验证通过** - ---- - -### Task 6: 组装层对齐与全量质量门禁验证 - -**Files:** -- Modify: `backend/cmd/app.go` -- Modify: `backend/cmd/*` - -- [ ] **Step 1: 检查并适配 `cmd/app.go` 及启动指令,确保 Goose 迁移与驱动正确接入新版 `DBService`** -- [ ] **Step 2: 运行全局跨包 import 检查:`grep -r "Wavelet/plugins/infra/database" backend/plugins/domain/` 必须为空** -- [ ] **Step 3: 运行全量单元测试与基准测试:`go test ./...`** -- [ ] **Step 4: 运行质量门禁:`make code-check && make format`** diff --git a/docs/superpowers/plans/2026-08-28-migration-split-plan.md b/docs/superpowers/plans/2026-08-28-migration-split-plan.md deleted file mode 100644 index 48abc22a..00000000 --- a/docs/superpowers/plans/2026-08-28-migration-split-plan.md +++ /dev/null @@ -1,326 +0,0 @@ -# 迁移脚本拆分执行计划 - -## 背景现状 - -| 维度 | 实际状态 | -|---|---| -| 总 SQL 文件 | 26 个全局 (`pkg/migrator/goose/`) + 5 个插件 (`plugins/domain/*/migrations/`) + 1 个 ClickHouse | -| 实际运行的迁移 | **仅 26 个全局文件**(通过 `gooseEngine` → `pkg/migrator.Migrate()`) | -| 插件注册的迁移 | 3 个 (`auth`, `user`, `message_gateway`) — 注册了但被 `gooseEngine` 丢弃 | -| 有迁移文件但未注册的插件 | `admin`(2 个文件,0 个调用) | -| 无迁移文件的插件 | `upload`, `risk_control`, `cap`, `driver_asynq_worker`, `driver_asynq_cron` | -| ClickHouse 迁移 | 1 个文件 (`w_user_access_logs`),通过 `pkg/migrator.MigrateClickHouse()` 单独运行 | - -## 表所有者映射 - -以下列表基于"单一所有者原则",每个表精确映射到一个插件: - -| 表名 | 所有者插件 | 涉及全局迁移 | -|---|---|---| -| `w_users` | `domain/user` | 20260609 (create), 20260614 (seed system user) | -| `w_access_tokens` | `domain/auth` | 20260609 (create), 20260610 (is_admin), 20260611 (rm last_used_at) | -| `w_auth_sources` | `domain/auth` | 20260609 (create) | -| `w_external_accounts` | `domain/auth` | 20260609 (create) | -| `w_system_configs` | `domain/admin` | 20260609→20260611 (rename+seeds×7), 20260613 (TEXT), 20260816 (log_db) | -| `w_templates` | `domain/admin` | 20260609→20260611 (rename) | -| `w_schedules` | `driver_asynq_cron` | 20260610 (create), 20260611 (identity), 20260614 (update cleanup) | -| `w_task_executions` | `driver_asynq_worker` | 20260609→20260611 (rename) | -| `w_uploads` | `domain/upload` | 20260609→20260611 (rename), 20260613 (access_mode), 20260617 (indexes), 20260618 (drop storage_driver) | -| `w_upload_stats` | `domain/upload` | 20260617 (create+backfill) | -| `w_push_events` | `domain/message_gateway` | 20260614 (create), 20260615 (task_type), 20260616 (cleanup) | -| `w_push_histories` | `domain/message_gateway` | 20260614 (create) | -| `w_push_channels` | `domain/message_gateway` | 20260614 (create) | -| `w_message_channels` | `domain/message_gateway` | 20260816 (create) | -| `w_message_bindings` | `domain/message_gateway` | 20260816 (create) | -| `w_message_pairing_codes` | `domain/message_gateway` | 20260816 (create) | -| `w_user_access_logs` | `domain/risk_control` | 20260816 (create) + ClickHouse | - -## 执行步骤(共 8 步) - ---- - -### 步骤 1:创建 Bootstrap 迁移(保留在 `pkg/migrator`) - -**文件**:`pkg/migrator/goose/postgres/00001_bootstrap.sql` - -将以下全局迁移合并为一个 bootstrap 文件: -- **`202606090001_initial_schema.sql`** → 创建 `users`, `auth_sources`, `external_accounts`, `access_tokens`, `system_configs`, `uploads`, `task_executions`, `templates`(全部无前缀旧名) -- **`202606110003_rename_tables_to_w_prefix.sql`** → 全部重命名为 `w_` 前缀 - -**合并后,bootstrap 文件直接创建带 `w_` 前缀的表**,不再需要 rename 步骤: - -```sql --- +goose Up -CREATE TABLE IF NOT EXISTS w_users ( - id BIGINT PRIMARY KEY, - username VARCHAR(64) UNIQUE, - ... -); -CREATE TABLE IF NOT EXISTS w_access_tokens (...); -CREATE TABLE IF NOT EXISTS w_auth_sources (...); -CREATE TABLE IF NOT EXISTS w_external_accounts (...); -CREATE TABLE IF NOT EXISTS w_system_configs ( - key VARCHAR(64) PRIMARY KEY, - value TEXT NOT NULL, - ... -); -CREATE TABLE IF NOT EXISTS w_uploads (...); -CREATE TABLE IF NOT EXISTS w_task_executions (...); -CREATE TABLE IF NOT EXISTS w_templates (...); -CREATE TABLE IF NOT EXISTS w_schedules (...); -``` - -> **为什么保留在 `pkg/migrator`**:这些是平台的"初始化基座"——无论哪些插件启用,这些表都存在。将 bootstrap 放到 `pkg/migrator` 之下回避了循环依赖问题(例如 `w_system_configs` 属于 admin,但 bootstrap 时 admin 插件尚未 apply)。 - ---- - -### 步骤 2:修复 `gooseEngine` 支持插件迁移 - -**文件**:`cmd/app.go` - -```go -type gooseEngine struct{} - -func (e *gooseEngine) Migrate(_ context.Context, entries []core.MigrationEntry) error { - // 1. 先跑 bootstrap(初始化基座) - _ = migrator.Migrate() - - // 2. 再跑每个插件注册的迁移 - for _, entry := range entries { - gormDB := database.DB(context.Background()) - if gormDB == nil { - continue - } - sqlDB, err := gormDB.DB() - if err != nil { - return err - } - - goose.SetBaseFS(entry.FS) - if err := goose.SetDialect(gooseDialect()); err != nil { - return err - } - dir := entry.Dir - if dir == "" { - dir = "migrations" - } - if err := goose.Up(sqlDB, dir); err != nil { - return fmt.Errorf("migrate %s: %w", entry.PluginID, err) - } - } - - // 3. ClickHouse 迁移 - _ = migrator.MigrateClickHouse() - - return nil -} -``` - -依赖项:`gooseDialect()` 从 `pkg/migrator` 导出。 - ---- - -### 步骤 3:按表所有者拆分迁移到各插件 - -| 全局源文件 | 目标插件 | 迁移文件名 | -|---|---|---| -| `202606100002` (access_tokens is_admin) | `domain/auth` | `migrations/00002_add_access_token_is_admin.sql` | -| `202606110001` (drop last_used_at) | `domain/auth` | `migrations/00003_drop_access_token_last_used_at.sql` | -| `202606140003` (system user seed) | `domain/user` | `migrations/00002_seed_system_user.sql` | -| `202606110004` (file_access_whitelist seed) | `domain/admin` | `migrations/00003_seed_file_access_whitelist.sql` | -| `202606110005` (disk_cache configs seed) | `domain/admin` | `migrations/00004_seed_disk_cache_configs.sql` | -| `202606120002` (update_upstream_repo seed) | `domain/admin` | `migrations/00005_seed_upstream_repo_config.sql` | -| `202606130002` (system_configs value TEXT) | `domain/admin` | `migrations/00006_expand_config_value.sql` | -| `202606130003` (storage_config seed) | `domain/admin` | `migrations/00007_seed_storage_config.sql` | -| `202608160002` (log database configs) | `domain/admin` | `migrations/00008_seed_log_db_configs.sql` | -| `202606120001` (login_session_ttl) | `domain/auth` | `migrations/00004_seed_login_session_ttl.sql` | -| `202606130001` (w_uploads access_mode) | `domain/upload` | `migrations/00001_add_access_mode.sql` | -| `202606170001` (upload indexes) | `domain/upload` | `migrations/00002_add_composite_indexes.sql` | -| `202606170002` (upload stats table) | `domain/upload` | `migrations/00003_create_upload_stats.sql` | -| `202606170003` (backfill stats) | `domain/upload` | `migrations/00004_backfill_upload_stats.sql` | -| `202606180001` (drop storage_driver) | `domain/upload` | `migrations/00005_drop_storage_driver.sql` | -| `202606140001` (push tables) | `domain/message_gateway` | `migrations/00002_create_push_tables.sql` | -| `202606140004` (push channels) | `domain/message_gateway` | `migrations/00003_create_push_channels.sql` | -| `202606150001` (push task_type) | `domain/message_gateway` | `migrations/00004_add_push_task_type.sql` | -| `202606160001` (remove push config) | `domain/message_gateway` | `migrations/00005_remove_push_config.sql` | -| `202608160003` (message gateway tables) | `domain/message_gateway` | `migrations/00006_create_message_tables.sql` | -| `202606100001` (schedules) | `driver_asynq_cron` | `migrations/00001_create_schedules.sql` | -| `202606110002` (schedules identity) | `driver_asynq_cron` | `migrations/00002_alter_schedules_identity.sql` | -| `202606140005` (update cleanup schedule) | `driver_asynq_cron` | `migrations/00003_update_cleanup_schedule.sql` | -| `202608160001` (user access logs) | `domain/risk_control/logstore` | `migrations/00001_create_access_logs.sql` | -| `202608160002` (log_database configs) | `domain/admin` | (合并到 admin 步骤 7) | - ---- - -### 步骤 4:补充缺失的 `go:embed` 和 `Register()` 调用 - -**`plugins/domain/admin/plugin.go`**: -```go -//go:embed migrations/*.sql -var adminMigrations embed.FS - -// 在 Apply() 中: -ctx.Migrations().Register("admin", adminMigrations) -``` - -**`plugins/domain/upload/plugin.go`**: -```go -//go:embed migrations/*.sql -var uploadMigrations embed.FS - -// 在 Apply() 中: -ctx.Migrations().Register("upload", uploadMigrations) -``` - -**`plugins/domain/risk_control/plugin.go`**: -```go -// go:embed 由 logstore 子包自行处理(它已有自己的 moved 文件) -// 在 Apply() 中: -ctx.Migrations().Register("risk_control/logstore", logstoreMigrationFS) -``` - -**`plugins/drivers/driver_asynq_cron/plugin.go`**: -```go -//go:embed migrations/*.sql -var cronMigrations embed.FS - -// 在 Apply() 中: -ctx.Migrations().Register("driver_asynq_cron", cronMigrations) -``` - ---- - -### 步骤 5:解决 Admin 插件迁移与 Bootstrap 的冲突 - -当前 `admin/migrations/00001` 执行 `CREATE TABLE IF NOT EXISTS w_system_configs (...)`,但 bootstrap 已在步骤 1 中创建过这张表。需要: -1. **保持 `IF NOT EXISTS`** 保证幂等性 -2. **从 admin migration 中移除 `w_schedules` 和 `w_task_executions` 的 CREATE**(它们在 bootstrap 中创建,属于 driver 插件) -3. **仅保留 admin 自己的表**:`w_system_configs`, `w_templates` -4. Seed 数据使用 `ON CONFLICT DO NOTHING` 避免重复: - -当前 admin 的 seed 包含 29 个系统配置,其中约 14 个与全局迁移重复。整理后的 admin seed 应: - -```sql -INSERT INTO w_system_configs (...) VALUES - ('cap_login_enabled', 'false', ...), - ('cap_auto_solve', 'true', ...), - -- ... (所有 29 个配置) -ON CONFLICT (key) DO NOTHING; -``` - -> 全局迁移中 `202606110004` 到 `202608160002` 的 7 个种子 INSERT 将被迁移到 admin,全部使用 `ON CONFLICT DO NOTHING`。 - ---- - -### 步骤 6:清理已迁移的全局文件 - -拆分完成后,从 `pkg/migrator/goose/postgres/` 中删除以下文件: - -``` -202606100002_access_token_is_admin.sql -202606100001_create_schedules.sql -202606110001_remove_access_token_last_used_at.sql -202606110002_alter_schedules_id_auto_increment.sql -202606110004_add_file_access_whitelist_config.sql -202606110005_add_disk_cache_configs.sql -202606120001_add_login_session_ttl_config.sql -202606120002_add_update_upstream_repository_config.sql -202606130001_add_upload_access_mode.sql -202606130002_expand_system_config_value.sql -202606130003_add_storage_config.sql -202606140001_create_push_tables.sql -202606140003_add_system_user.sql -202606140004_create_push_channels.sql -202606140005_update_system_cleanup_schedule.sql -202606150001_add_task_type_to_push_events.sql -202606160001_remove_push_config.sql -202606170001_add_upload_composite_indexes.sql -202606170002_create_upload_stats_table.sql -202606170003_backfill_upload_stats.sql -202606180001_drop_upload_storage_driver.sql -202608160001_create_user_access_logs.sql -202608160002_log_database_configs.sql -202608160003_create_message_gateway.sql -``` - -**保留在 `pkg/migrator/goose/postgres/` 的仅限**: -``` -00001_bootstrap.sql (合并后的初始化基座) -``` - -**注意**:`202606110003_rename_tables_to_w_prefix.sql` 也被合并进 bootstrap。`202606090001_initial_schema.sql` 也被合并掉。 - ---- - -### 步骤 7:更新 `pkg/migrator` 导出 `gooseDialect()` - -在 `pkg/migrator/migrator.go` 中将 `gooseDialect()` 和 `migrationDir()` 改为导出,供 `cmd/app.go` 的 `gooseEngine.Migrate()` 引用。 - ---- - -### 步骤 8:验证 + 提交 - -```bash -cd /Users/ryan/Code/Go/Wavelet - -# 1. 编译验证 -go build -mod=mod ./... -go vet ./... - -# 2. 架构门禁验证 -make code-check - -# 3. 验证插件迁移注册完整性 -grep -rn 'go:embed.*migrations' plugins/domain/*/plugin.go plugins/drivers/*/plugin.go -grep -rn 'Migrations()\.Register' plugins/domain/*/plugin.go plugins/drivers/*/plugin.go -# → 每个有 migrations/ 目录的插件既要有 go:embed 又要有 Register() - -# 4. 验证 admin 插件迁移完整性 -grep -rn 'w_schedules\|w_task_executions' plugins/domain/admin/migrations/ -# → 不应有(这些属于 driver 插件) - -# 5. 提交 -git add -A && git commit -m "refactor(migration): split global SQL into per-plugin migrations - -- Merge 26 global SQLs into bootstrap + per-plugin migrations -- Fix gooseEngine to iterate plugin-registered MigrationEntry -- Add go:embed + Register() to admin, upload, risk_control, driver_asynq_cron -- Remove 23 migrated SQL files from pkg/migrator/goose/ -- Keep only bootstrap in pkg/migrator/goose/ -- All CREATE TABLE use IF NOT EXISTS, all INSERT use ON CONFLICT DO NOTHING" -``` - ---- - -## 依赖关系图 - -``` -Bootstrap (pkg/migrator) - ├── 创建 w_users, w_access_tokens, w_auth_sources, w_external_accounts - ├── 创建 w_system_configs, w_templates, w_schedules, w_task_executions - ├── 创建 w_uploads, w_upload_stats - └── 创建所有 w_ 前缀表 - │ - ├─ auth/00002 (access_tokens is_admin) - ├─ auth/00003 (drop last_used_at) - ├─ auth/00004 (login_session_ttl seed) - │ - ├─ user/00002 (system user seed) - │ - ├─ admin/00001 (w_system_configs, w_templates) [IF NOT EXISTS] - ├─ admin/00002 (29 config seeds + 2 template seeds) - ├─ admin/00003–00008 (拆分后的种子迁移) - │ - ├─ upload/00001–00005 (access_mode → indexes → stats → backfill → drop) - │ - ├─ message_gateway/00001 (w_message_* tables) - ├─ message_gateway/00002–00006 (push tables → channels → task_type → cleanup) - │ - ├─ driver_asynq_cron/00001–00003 (schedules → identity → cleanup) - │ - ├─ driver_asynq_worker/00001 (task_executions — 如果有追加操作) - │ - └─ risk_control/logstore/00001 (w_user_access_logs) -``` - -所有步骤执行的迁移顺序由 Goose 的文件名前缀控制。Bootstrap 使用 `00001_`,每个插件的迁移从 `00002_` 开始编号(`00001` 留给插件自身表 CREATE,若插件 bootstrap 已创建则从 `00002` 开始)。 \ No newline at end of file diff --git a/docs/superpowers/plans/2026-08-28-zero-redis-pluggable-architecture.md b/docs/superpowers/plans/2026-08-28-zero-redis-pluggable-architecture.md deleted file mode 100644 index 90de7a1e..00000000 --- a/docs/superpowers/plans/2026-08-28-zero-redis-pluggable-architecture.md +++ /dev/null @@ -1,84 +0,0 @@ -# Zero-Redis Pluggable Architecture Implementation Plan - -> **Goal**: Extract Redis into optional plugins and introduce lightweight in-process equivalents (`cache_memory`, `driver_inproc_worker`, `driver_inproc_cron`), enabling zero-Redis monolithic and embedded deployment modes. - -- **Architecture Spec**: [`docs/superpowers/specs/2026-08-28-zero-redis-pluggable-architecture-design.md`](file:///Users/ryan/Code/Go/Wavelet/docs/superpowers/specs/2026-08-28-zero-redis-pluggable-architecture-design.md) -- **Branch**: `main` - ---- - -## Proposed Changes - -### 1. In-Memory Cache Infrastructure Plugin (`backend/plugins/infra/cache_memory`) - -#### [NEW] `backend/plugins/infra/cache_memory/plugin.go` -- Implements `core.Plugin` (`Name() == "cache_memory"`). -- Applies `contracts.CacheService` to the Context via `core.Provide[contracts.CacheService](ctx, memCacheSvc)`. - -#### [NEW] `backend/plugins/infra/cache_memory/cache.go` -- Implements `contracts.CacheService` using `pkg/cache/ram`. -- Dispatches in-process invalidation notifications via `ctx.Events().Emit("cache:invalidate", key)`. - -#### [NEW] `backend/plugins/infra/cache_memory/plugin_test.go` -- Unit tests for Get, Set, Delete, GetOrSet, TTL expiration, and event bus emission. - ---- - -### 2. In-Process Async Worker Driver (`backend/plugins/drivers/driver_inproc_worker`) - -#### [NEW] `backend/plugins/drivers/driver_inproc_worker/plugin.go` -- Implements `core.Plugin` & `core.Driver` (`Type() == core.DriverTypeWorker`). -- Scans and executes registered tasks from `ctx.Tasks().Tasks()`. - -#### [NEW] `backend/plugins/drivers/driver_inproc_worker/executor.go` -- In-memory buffered channel queue and worker goroutine pool managed via `util.Go`. -- Supports execution timeout, retry with backoff, and graceful shutdown. - -#### [NEW] `backend/plugins/drivers/driver_inproc_worker/plugin_test.go` -- Unit tests for in-process task execution, concurrency limit, retry on error, and graceful shutdown. - ---- - -### 3. In-Process Cron Scheduler Driver (`backend/plugins/drivers/driver_inproc_cron`) - -#### [NEW] `backend/plugins/drivers/driver_inproc_cron/plugin.go` -- Implements `core.Plugin` & `core.Driver` (`Type() == core.DriverTypeScheduler`). -- Reads `ctx.Schedules().Schedules()` and schedules jobs using `robfig/cron/v3`. - -#### [NEW] `backend/plugins/drivers/driver_inproc_cron/scheduler.go` -- Handles Cron expression registration, job triggering, and graceful stopping. - -#### [NEW] `backend/plugins/drivers/driver_inproc_cron/plugin_test.go` -- Unit tests verifying cron job scheduling, execution tracking, and stop behavior. - ---- - -### 4. Admin Domain Decoupling from Redis - -#### [MODIFY] `backend/plugins/domain/admin/repository.go` -- Introduce in-memory `RingBuffer` for task output streams when Redis is nil. -- Fallback task log lookups to `RingBuffer` and `w_task_executions` table. - -#### [MODIFY] `backend/plugins/domain/admin/system_config_cache.go` -- Guard Redis PubSub listener so that when Redis is nil, it gracefully falls back to local event bus updates without spawning disconnected subscriber loops. - ---- - -### 5. Application Assembly & Profile Switching - -#### [MODIFY] `backend/cmd/app.go` -- Switch dynamically between Redis plugins (`cache`, `driver_asynq_worker`, `driver_asynq_cron`) and In-Process plugins (`cache_memory`, `driver_inproc_worker`, `driver_inproc_cron`) based on `config.Config.Redis.Enabled`. - -#### [MODIFY] `backend/cmd/app_test.go` -- Add test verifying application bootstrap in both `Redis.Enabled = true` and `Redis.Enabled = false` states. - ---- - -## Verification Plan - -### Automated Tests -1. **In-Memory Cache Tests**: `go test -v ./backend/plugins/infra/cache_memory/...` -2. **In-Process Worker Tests**: `go test -v ./backend/plugins/drivers/driver_inproc_worker/...` -3. **In-Process Cron Tests**: `go test -v ./backend/plugins/drivers/driver_inproc_cron/...` -4. **Full Test Suite**: `cd backend && go test ./...` -5. **Quality Gate**: `make code-check && make format` diff --git a/docs/superpowers/plans/2026-08-29-cordis-config-extension.md b/docs/superpowers/plans/2026-08-29-cordis-config-extension.md deleted file mode 100644 index 35e6a2ca..00000000 --- a/docs/superpowers/plans/2026-08-29-cordis-config-extension.md +++ /dev/null @@ -1,2585 +0,0 @@ -# Cordis 配置扩展点(框架与门禁)Implementation Plan - -> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking. - -**Goal:** 在微内核中落地"插件声明配置字段、内核按声明解析、门禁决定插件激活"的配置扩展点,并证明其解析结果与现有 `backend/pkg/config` 逐 key 等价。 - -**Architecture:** `core/extpoints` 实现纯 stdlib 的配置引擎(声明注册表 + 优先级解析 + 冲突校验 + 脱敏导出),viper 装载隔离在 `plugins/infra/config` 适配器内;`App.Prepare()` 作为解析屏障,`Fiber` 新增 `FiberSkipped` 态承载门禁结果。 - -**Tech Stack:** Go 1.25.7、`github.com/spf13/viper v1.21.0`(仅适配器)、`github.com/stretchr/testify`(测试)、`github.com/google/go-cmp`(对拍)、golangci-lint(gofumpt + cyclop/funlen/mnd/revive/dupl/gosec)。 - -**Spec:** `docs/superpowers/specs/2026-08-29-cordis-config-extension-design.md` - ---- - -## 本计划范围(对应 spec §7.3 的 P1 + P2) - -本计划交付**内核能力**,不改动任何业务插件与 `cmd`:完成后旧的全局单例 `config.Config` 仍然是生产路径的唯一配置来源,应用行为零变化,新增能力由单测与新旧对拍证明。 - -**下一个计划(P3 + P4,本计划完成后另行编写)** 才做 27 个消费文件的迁移、`pkg/idgen` 解耦与 `backend/pkg/config` 删除。分期理由:迁移的声明结构体写法依赖本计划定稿的 tag 与 API 形状,先写会在实现过程中失真。 - -## 文件结构 - -| 文件 | 职责 | 动作 | -| :--- | :--- | :--- | -| `backend/core/extpoints/config.go` | 配置引擎的抽象与声明注册表:`ConfigSource`、`ConfigBinding`、`ConfigEntry`、`ConfigView`、`ConfigExtension`、`ConfigRegistry.Declare` + tag 遍历 + 冲突校验 | Create | -| `backend/core/extpoints/config_value.go` | 值解码:`convertValue` 及其 bool/int/uint/float/string/duration/slice/struct 分支 | Create | -| `backend/core/extpoints/config_resolve.go` | `Resolve` 优先级链、`Bind` 赋值、只读访问器、`Entries` 脱敏导出 | Create | -| `backend/core/extpoints/config_test.go` | 引擎单测(外部测试包 `extpoints_test`,fake source) | Create | -| `backend/core/config.go` | `ConfigGet[T]` 泛型读取入口 | Create | -| `backend/core/config_test.go` | 泛型读取与 `Context.Config()` 接入测试 | Create | -| `backend/core/types.go` | `ConfigExtension`/`ConfigBinding`/`ConfigEntry`/`ConfigSource` 别名 + `ConfigGatedPlugin` 可选接口 | Modify | -| `backend/core/context.go` | `config` 字段、`NewContext` 初始化、`Fork` 共享、`Config()` 访问器 | Modify | -| `backend/core/fiber.go` | `FiberSkipped` 状态与 `Skip()` | Modify | -| `backend/core/app.go` | `WithConfigSource`/`WithConfigDecl`/`Prepare`/`SetShutdownTimeout`、`Use` 收集声明、`reconcileLocked` 门禁求值 | Modify | -| `backend/plugins/infra/config/source.go` | viper + yaml 适配器,实现 `core.ConfigSource`,保留 `CONFIG_PATH` 与向上查找语义 | Create | -| `backend/plugins/infra/config/source_test.go` | 适配器单测(`t.TempDir()` + `t.Setenv`) | Create | -| `backend/pkg/config/config.go` | 抽出可重入 `load(configPath string, testMode bool)`(仅重构,行为不变) | Modify | -| `backend/pkg/config/parity_test.go` | 新旧解析对拍(临时文件,P4 随旧包删除) | Create then Delete in P4 | -| `scripts/check_cordis_architecture.sh` | 微内核禁 viper 检查项 | Modify | - -**约束提示:** 所有新增导出符号必须带符合 `go-documentation` 规范的文档注释(`revive` 会检查);测试统一用 `t.TempDir()`,禁止相对路径创建临时目录。 - ---- - -## Task 1: 配置引擎的抽象与声明注册表 - -**Files:** -- Create: `backend/core/extpoints/config.go` -- Test: `backend/core/extpoints/config_test.go` - -- [ ] **Step 1: 写失败的测试** - -创建 `backend/core/extpoints/config_test.go`: - -```go -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package extpoints_test - -import ( - "Wavelet/core/extpoints" - "testing" - "time" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -// fakeSource is an in-memory extpoints.ConfigSource used by configuration engine tests. -type fakeSource struct { - values map[string]any - env map[string]string -} - -func newFakeSource() *fakeSource { - return &fakeSource{values: map[string]any{}, env: map[string]string{}} -} - -func (f *fakeSource) Lookup(path string) (any, bool) { - v, ok := f.values[path] - return v, ok -} - -func (f *fakeSource) LookupEnv(name string) (string, bool) { - v, ok := f.env[name] - return v, ok -} - -func (f *fakeSource) Describe() string { return "fake" } - -// redisConfig mirrors how a plugin declares the configuration it reads. -type redisConfig struct { - Enabled bool `config:"enabled" env:"REDIS_ENABLED" default:"false" autoEnable:"REDIS_ADDR"` - Addrs []string `config:"addrs" env:"REDIS_ADDR"` - DB int `config:"db" env:"REDIS_DB"` - KeyPrefix string `config:"key_prefix" env:"REDIS_KEY_PREFIX"` - Dial time.Duration `config:"dial_timeout" env:"REDIS_DIAL_TIMEOUT"` - Ignored string `config:"-"` - private string -} - -func TestDeclareRegistersTaggedLeafKeys(t *testing.T) { - r := extpoints.NewConfigRegistry(newFakeSource()) - - require.NoError(t, r.Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: &redisConfig{}})) - - keys := make([]string, 0) - for _, e := range r.Entries() { - keys = append(keys, e.Key) - } - assert.Equal(t, []string{ - "redis.addrs", "redis.db", "redis.dial_timeout", "redis.enabled", "redis.key_prefix", - }, keys) -} - -func TestDeclareRejectsNonStructPointerTarget(t *testing.T) { - r := extpoints.NewConfigRegistry(newFakeSource()) - - assert.ErrorIs(t, r.Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: redisConfig{}}), - extpoints.ErrConfigTarget) - assert.ErrorIs(t, r.Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: (*redisConfig)(nil)}), - extpoints.ErrConfigTarget) -} - -func TestDeclareAllowsIdenticalDuplicateAndRejectsConflictingMetadata(t *testing.T) { - r := extpoints.NewConfigRegistry(newFakeSource()) - binding := extpoints.ConfigBinding{Prefix: "redis", Target: &redisConfig{}} - require.NoError(t, r.Declare("cache", binding)) - require.NoError(t, r.Declare("cache_memory", binding), "identical shared declarations must be allowed") - - type conflictingConfig struct { - Enabled bool `config:"enabled" env:"REDIS_ON" default:"true"` - } - err := r.Declare("driver_http", extpoints.ConfigBinding{Prefix: "redis", Target: &conflictingConfig{}}) - require.ErrorIs(t, err, extpoints.ErrConfigConflict) - assert.Contains(t, err.Error(), "redis.enabled") - assert.Contains(t, err.Error(), "cache") - assert.Contains(t, err.Error(), "driver_http") -} -``` - -- [ ] **Step 2: 运行测试确认失败** - -Run: `cd backend && go test ./core/extpoints/ -run 'TestDeclare' -v` -Expected: 编译失败,报 `undefined: extpoints.NewConfigRegistry`、`undefined: extpoints.ConfigBinding`、`undefined: extpoints.ErrConfigTarget`、`undefined: extpoints.ErrConfigConflict`。 - -- [ ] **Step 3: 写最小实现** - -创建 `backend/core/extpoints/config.go`: - -```go -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package extpoints - -import ( - "errors" - "fmt" - "reflect" - "strings" - "sync" - "time" -) - -// Sentinel errors returned by the configuration extension point. -var ( - // ErrConfigConflict is returned when the same key is declared with disagreeing metadata. - ErrConfigConflict = errors.New("extpoints: conflicting configuration declarations") - - // ErrConfigType is returned when a value cannot be converted to the declared type. - ErrConfigType = errors.New("extpoints: configuration value type mismatch") - - // ErrConfigInvalid is returned when a resolved value violates a declared value range. - // Reserved for source-level value checks; per-plugin value ranges are validated by the - // declaring plugin after Bind (see spec §4.3 C1). - ErrConfigInvalid = errors.New("extpoints: invalid configuration value") - - // ErrConfigUnknownKey is returned when a configuration key was never declared. - ErrConfigUnknownKey = errors.New("extpoints: unknown configuration key") - - // ErrConfigNotResolved is returned when typed reads happen before resolution. - ErrConfigNotResolved = errors.New("extpoints: configuration not resolved; run App.Prepare first") - - // ErrConfigTarget is returned when a binding target is not an addressable struct pointer. - ErrConfigTarget = errors.New("extpoints: configuration binding target must be a non-nil struct pointer") - - // ErrConfigNoSource is returned when resolution is attempted without a registered source. - ErrConfigNoSource = errors.New("extpoints: no configuration source registered") -) - -// Configuration origin labels reported by ConfigView.Origin and ConfigEntry.Origin. -const ( - // OriginEnv marks a value that came from an environment variable. - OriginEnv = "env" - // OriginAutoEnable marks a boolean enabled by the presence of another environment variable. - OriginAutoEnable = "auto-enable" - // OriginFile marks a value that came from the configuration file. - OriginFile = "file" - // OriginDefault marks a value that came from a declaration default. - OriginDefault = "default" -) - -// durationType distinguishes time.Duration from plain int64 during tag walking and decoding. -var durationType = reflect.TypeFor[time.Duration]() - -// ConfigSource abstracts where raw configuration values come from, keeping the -// micro-kernel free of concrete loaders such as viper. -type ConfigSource interface { - // Lookup returns the raw value stored at a dotted path in the configuration file. - Lookup(path string) (any, bool) - // LookupEnv returns the raw value of an environment variable. - LookupEnv(name string) (string, bool) - // Describe returns a human readable identity for the source, used in diagnostics. - Describe() string -} - -// ConfigBinding declares that a plugin reads every `config` tagged field of Target -// under a dotted configuration prefix. -type ConfigBinding struct { - // Prefix is the dotted configuration path, e.g. "redis". An empty prefix means - // each field's `config` tag is already a full path. - Prefix string - // Target must be a non-nil pointer to a struct carrying `config` tags. - Target any -} - -// configField is a single leaf discovered while walking a binding struct's tags. -// key is the fully qualified dotted path used for resolution; path is the raw `config` -// tag value used to locate the Go field again during Bind. -type configField struct { - key string - path string - env string - autoEnable string - def string - secret bool - typ reflect.Type -} - -// configDecl is the registered form of a configField, attributed to its declaring plugin. -type configDecl struct { - key string - pluginID string - env string - autoEnable string - def string - secret bool - typ reflect.Type -} - -// ConfigEntry is a redacted, self-describing view of one effective configuration key. -type ConfigEntry struct { - Key string - PluginID string - Env string - Origin string - Value string -} - -// ConfigView is the read-only surface over effective configuration values. -// Keys are dotted paths such as "redis.enabled". -type ConfigView interface { - Value(key string) (any, bool) - String(key, fallback string) string - Bool(key string, fallback bool) bool - Int(key string, fallback int) int - Duration(key string, fallback time.Duration) time.Duration - Strings(key string) []string - WasSet(envName string) bool - Origin(key string) string -} - -// ConfigExtension is the plugin-facing configuration extension point mounted on the -// root Context and shared by every forked plugin scope. -type ConfigExtension interface { - ConfigView - - // SetSource installs the raw value source after construction, letting the composition - // root build the adapter once the kernel Context already exists. - SetSource(src ConfigSource) - // Declare registers plugin-owned configuration bindings before Apply runs. - Declare(pluginID string, bindings ...ConfigBinding) error - // Bind resolves and assigns the configuration values for a tagged struct. - Bind(prefix string, target any) error - // Resolve computes the effective value of every declared key once. - Resolve() error - // Resolved reports whether Resolve has already run. - Resolved() bool - // Entries returns the redacted effective configuration ordered by key. - Entries() []ConfigEntry -} - -// ConfigRegistry implements ConfigExtension. Declarations are additive; values are -// computed once by Resolve and reused by every later read. -type ConfigRegistry struct { - mu sync.RWMutex - src ConfigSource - decls map[string]*configDecl - order []string - values map[string]any - origins map[string]string - resolved bool -} - -// NewConfigRegistry creates an empty configuration registry. A nil src is allowed so -// that the kernel can construct the registry before the composition root injects one. -func NewConfigRegistry(src ConfigSource) *ConfigRegistry { - return &ConfigRegistry{ - src: src, - decls: make(map[string]*configDecl), - values: make(map[string]any), - origins: make(map[string]string), - } -} - -// SetSource installs the raw value source. It is intended for the composition root, -// which builds the adapter after the kernel Context already exists. -func (r *ConfigRegistry) SetSource(src ConfigSource) { - r.mu.Lock() - defer r.mu.Unlock() - r.src = src -} - -// Declare registers every `config` tagged leaf of each binding's target struct. -// Repeated declarations of the same key are accepted only when their env, default, -// auto-enable and secret metadata agree; disagreement is ErrConfigConflict. -func (r *ConfigRegistry) Declare(pluginID string, bindings ...ConfigBinding) error { - r.mu.Lock() - defer r.mu.Unlock() - - for _, b := range bindings { - if err := r.declareBinding(pluginID, b); err != nil { - return err - } - } - return nil -} - -func (r *ConfigRegistry) declareBinding(pluginID string, b ConfigBinding) error { - target, err := bindingStruct(b.Target, b.Prefix) - if err != nil { - return err - } - - fields, err := walkConfigFields(target.Type(), b.Prefix) - if err != nil { - return err - } - for _, f := range fields { - if err := r.addDecl(pluginID, f); err != nil { - return err - } - } - return nil -} - -// bindingStruct validates that a binding or bind target is a usable struct pointer. -func bindingStruct(target any, prefix string) (reflect.Value, error) { - rv := reflect.ValueOf(target) - if !rv.IsValid() || rv.Kind() != reflect.Pointer || rv.IsNil() || rv.Elem().Kind() != reflect.Struct { - return reflect.Value{}, fmt.Errorf("%w: prefix %q received %T", ErrConfigTarget, prefix, target) - } - return rv.Elem(), nil -} - -// walkConfigFields collects leaf configuration declarations from `config` tagged fields. -// A field without a `config` tag is skipped, except for embedded structs which are -// recursed into so their own tags resolve under the same prefix. -func walkConfigFields(t reflect.Type, prefix string) ([]configField, error) { - var out []configField - - for i := 0; i < t.NumField(); i++ { - sf := t.Field(i) - if sf.PkgPath != "" { - continue - } - - path := sf.Tag.Get("config") - if path == "-" { - continue - } - if path == "" { - if sf.Type.Kind() == reflect.Struct && sf.Type != durationType { - nested, err := walkConfigFields(sf.Type, prefix) - if err != nil { - return nil, err - } - out = append(out, nested...) - } - continue - } - - out = append(out, configField{ - key: joinKey(prefix, path), - path: path, - env: sf.Tag.Get("env"), - autoEnable: sf.Tag.Get("autoEnable"), - def: sf.Tag.Get("default"), - secret: strings.EqualFold(sf.Tag.Get("secret"), "true"), - typ: sf.Type, - }) - } - - return out, nil -} - -func joinKey(prefix, path string) string { - if prefix == "" { - return path - } - return prefix + "." + path -} - -// addDecl records one leaf, enforcing the shared-declaration consistency rule. -func (r *ConfigRegistry) addDecl(pluginID string, f configField) error { - if existing, ok := r.decls[f.key]; ok { - if existing.env != f.env || existing.def != f.def || - existing.autoEnable != f.autoEnable || existing.secret != f.secret { - return fmt.Errorf( - "%w: key %q declared by plugin %q and plugin %q with disagreeing env/default/autoEnable/secret metadata", - ErrConfigConflict, f.key, existing.pluginID, pluginID) - } - return nil - } - - r.decls[f.key] = &configDecl{ - key: f.key, pluginID: pluginID, env: f.env, - autoEnable: f.autoEnable, def: f.def, secret: f.secret, typ: f.typ, - } - r.order = append(r.order, f.key) - return nil -} -``` - -- [ ] **Step 4: 补齐测试所需的占位实现** - -此时 `Entries`、`Resolve`、`Bind`、访问器尚未实现,Step 1 的测试用到 `Entries`。先创建 `backend/core/extpoints/config_resolve.go` 骨架,仅让 `Entries` 返回声明清单(值与来源在 Task 3/4 填充): - -```go -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package extpoints - -import "sort" - -// Entries returns the effective configuration as redacted, key-sorted entries. -func (r *ConfigRegistry) Entries() []ConfigEntry { - r.mu.RLock() - defer r.mu.RUnlock() - - keys := append([]string(nil), r.order...) - sort.Strings(keys) - - out := make([]ConfigEntry, 0, len(keys)) - for _, key := range keys { - d := r.decls[key] - out = append(out, ConfigEntry{ - Key: d.key, PluginID: d.pluginID, Env: d.env, - Origin: r.origins[key], Value: "pending", - }) - } - return out -} -``` - -- [ ] **Step 5: 运行测试确认通过** - -Run: `cd backend && go test ./core/extpoints/ -run 'TestDeclare' -v` -Expected: `--- PASS: TestDeclareRegistersTaggedLeafKeys`、`--- PASS: TestDeclareRejectsNonStructPointerTarget`、`--- PASS: TestDeclareAllowsIdenticalDuplicateAndRejectsConflictingMetadata`,`ok Wavelet/core/extpoints`。 - -- [ ] **Step 6: 格式与静态检查** - -Run: `cd backend && golangci-lint fmt ./core/extpoints/ && golangci-lint run ./core/extpoints/` -Expected: 无告警输出,退出码 0。 - -- [ ] **Step 7: 提交** - -```bash -git add backend/core/extpoints/config.go backend/core/extpoints/config_resolve.go backend/core/extpoints/config_test.go -git commit -m "feat(core): add configuration declaration registry" -``` - ---- - -## Task 2: 值解码(标量、时长、切片、结构体) - -**Files:** -- Create: `backend/core/extpoints/config_value.go` -- Test: `backend/core/extpoints/config_test.go`(追加) - -- [ ] **Step 1: 追加失败的测试** - -在 `backend/core/extpoints/config_test.go` 末尾追加(`fakeSource`、`redisConfig` 复用 Task 1 的定义): - -```go -// queueConfig is a composite element mirroring worker.queues in config.yaml. -type queueConfig struct { - Name string `config:"name"` - Priority int `config:"priority"` -} - -type workerConfig struct { - Concurrency int `config:"concurrency" env:"WORKER_CONCURRENCY"` - Queues []queueConfig `config:"queues"` -} - -type sessionConfig struct { - Secret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"` - Age int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"` -} - -func TestResolveScalarDurationAndSlice(t *testing.T) { - src := newFakeSource() - src.values["redis.db"] = 1 - src.values["redis.dial_timeout"] = "5s" - src.values["redis.addrs"] = []any{"127.0.0.1:6379"} - src.env["REDIS_KEY_PREFIX"] = "refresh:" - - r := extpoints.NewConfigRegistry(src) - require.NoError(t, r.Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: &redisConfig{}})) - require.NoError(t, r.Resolve()) - - var got redisConfig - require.NoError(t, r.Bind("redis", &got)) - assert.Equal(t, redisConfig{ - Addrs: []string{"127.0.0.1:6379"}, DB: 1, KeyPrefix: "refresh:", Dial: 5 * time.Second, - }, got) -} - -func TestResolveFillsSliceFromScalarEnvironmentValue(t *testing.T) { - src := newFakeSource() - src.env["REDIS_ADDR"] = "redis:6379" - - r := extpoints.NewConfigRegistry(src) - require.NoError(t, r.Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: &redisConfig{}})) - require.NoError(t, r.Resolve()) - - assert.Equal(t, []string{"redis:6379"}, r.Strings("redis.addrs")) -} - -func TestResolveCompositeSliceOfStructs(t *testing.T) { - src := newFakeSource() - src.values["worker.concurrency"] = 20 - src.values["worker.queues"] = []any{ - map[string]any{"name": "webhook", "priority": 10}, - map[string]any{"name": "default", "priority": 3}, - } - - r := extpoints.NewConfigRegistry(src) - require.NoError(t, r.Declare("asynq_worker", extpoints.ConfigBinding{Prefix: "worker", Target: &workerConfig{}})) - require.NoError(t, r.Resolve()) - - var got workerConfig - require.NoError(t, r.Bind("worker", &got)) - assert.Equal(t, workerConfig{ - Concurrency: 20, - Queues: []queueConfig{{Name: "webhook", Priority: 10}, {Name: "default", Priority: 3}}, - }, got) -} - -func TestResolveReportsTypeMismatchOnBadEnvironmentValue(t *testing.T) { - src := newFakeSource() - src.env["WORKER_CONCURRENCY"] = "many" - - r := extpoints.NewConfigRegistry(src) - require.NoError(t, r.Declare("asynq_worker", extpoints.ConfigBinding{Prefix: "worker", Target: &workerConfig{}})) - - err := r.Resolve() - require.ErrorIs(t, err, extpoints.ErrConfigType) - assert.Contains(t, err.Error(), "worker.concurrency") - assert.Contains(t, err.Error(), "WORKER_CONCURRENCY") -} -``` - -- [ ] **Step 2: 运行测试确认失败** - -Run: `cd backend && go test ./core/extpoints/ -run 'TestResolve' -v` -Expected: 编译失败,报 `r.Resolve undefined`、`r.Bind undefined`、`r.Strings undefined`。 - -- [ ] **Step 3: 写实现** - -创建 `backend/core/extpoints/config_value.go`: - -```go -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package extpoints - -import ( - "fmt" - "reflect" - "strconv" - "strings" - "time" -) - -// convertValue coerces a raw value coming from the configuration file or an -// environment variable into the declared Go type. -func convertValue(raw any, typ reflect.Type) (any, error) { - if typ == durationType { - return convertDuration(raw) - } - - switch typ.Kind() { - case reflect.Bool: - return convertBool(raw) - case reflect.String: - return convertString(raw) - case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: - return convertInt(raw, typ) - case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64: - return convertUint(raw, typ) - case reflect.Float32, reflect.Float64: - return convertFloat(raw, typ) - case reflect.Slice: - return convertSlice(raw, typ) - case reflect.Struct: - return convertStruct(raw, typ) - default: - return nil, fmt.Errorf("%w: %s is not a supported configuration type", ErrConfigType, typ) - } -} - -func convertBool(raw any) (any, error) { - switch v := raw.(type) { - case bool: - return v, nil - case string: - parsed, err := strconv.ParseBool(strings.TrimSpace(v)) - if err != nil { - return nil, fmt.Errorf("%w: %q is not a boolean", ErrConfigType, v) - } - return parsed, nil - default: - return nil, fmt.Errorf("%w: %v is not a boolean", ErrConfigType, raw) - } -} - -func convertString(raw any) (any, error) { - switch v := raw.(type) { - case string: - return v, nil - case bool, int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64, float32, float64: - return fmt.Sprint(v), nil - default: - return nil, fmt.Errorf("%w: %v is not a string", ErrConfigType, raw) - } -} - -// numericString extracts the textual form of a value so environment overrides, -// which always arrive as strings, share one parsing path with file values. -func numericString(raw any) (string, bool) { - switch v := raw.(type) { - case string: - return strings.TrimSpace(v), true - case bool, int, int8, int16, int32, int64, uint, uint8, uint16, uint32, uint64, float32, float64: - return fmt.Sprint(v), true - default: - return "", false - } -} - -func convertInt(raw any, typ reflect.Type) (any, error) { - text, ok := numericString(raw) - if !ok { - return nil, fmt.Errorf("%w: %v is not an integer", ErrConfigType, raw) - } - parsed, err := strconv.ParseInt(text, 10, typ.Bits()) - if err != nil { - return nil, fmt.Errorf("%w: %q is not a valid %s", ErrConfigType, text, typ) - } - out := reflect.New(typ).Elem() - out.SetInt(parsed) - return out.Interface(), nil -} - -func convertUint(raw any, typ reflect.Type) (any, error) { - text, ok := numericString(raw) - if !ok { - return nil, fmt.Errorf("%w: %v is not an unsigned integer", ErrConfigType, raw) - } - parsed, err := strconv.ParseUint(text, 10, typ.Bits()) - if err != nil { - return nil, fmt.Errorf("%w: %q is not a valid %s", ErrConfigType, text, typ) - } - out := reflect.New(typ).Elem() - out.SetUint(parsed) - return out.Interface(), nil -} - -func convertFloat(raw any, typ reflect.Type) (any, error) { - text, ok := numericString(raw) - if !ok { - return nil, fmt.Errorf("%w: %v is not a float", ErrConfigType, raw) - } - parsed, err := strconv.ParseFloat(text, typ.Bits()) - if err != nil { - return nil, fmt.Errorf("%w: %q is not a valid %s", ErrConfigType, text, typ) - } - out := reflect.New(typ).Elem() - out.SetFloat(parsed) - return out.Interface(), nil -} - -// convertDuration accepts both Go duration strings such as "200ms" and integer -// nanoseconds, mirroring what the previous viper based decoding supported. -func convertDuration(raw any) (any, error) { - switch v := raw.(type) { - case time.Duration: - return v, nil - } - - text, ok := numericString(raw) - if !ok { - return nil, fmt.Errorf("%w: %v is not a duration", ErrConfigType, raw) - } - if parsed, err := time.ParseDuration(text); err == nil { - return parsed, nil - } - nanos, err := strconv.ParseInt(text, 10, 64) - if err != nil { - return nil, fmt.Errorf("%w: %q is not a valid duration", ErrConfigType, text) - } - return time.Duration(nanos), nil -} - -func convertSlice(raw any, typ reflect.Type) (any, error) { - items, ok := sliceItems(raw) - if !ok { - // Scalar-to-single-element promotion keeps REDIS_ADDR populating redis.addrs. - items = []any{raw} - } - - out := reflect.MakeSlice(typ, 0, len(items)) - for _, item := range items { - converted, err := convertValue(item, typ.Elem()) - if err != nil { - return nil, err - } - out = reflect.Append(out, reflect.ValueOf(converted)) - } - return out.Interface(), nil -} - -// sliceItems normalises the several slice shapes a loader may produce. -func sliceItems(raw any) ([]any, bool) { - switch v := raw.(type) { - case []any: - return v, true - case []string: - items := make([]any, len(v)) - for i, s := range v { - items[i] = s - } - return items, true - } - - rv := reflect.ValueOf(raw) - if rv.IsValid() && rv.Kind() == reflect.Slice { - items := make([]any, rv.Len()) - for i := 0; i < rv.Len(); i++ { - items[i] = rv.Index(i).Interface() - } - return items, true - } - return nil, false -} - -func convertStruct(raw any, typ reflect.Type) (any, error) { - table, ok := asStringMap(raw) - if !ok { - return nil, fmt.Errorf("%w: %v is not a mapping, cannot decode into %s", ErrConfigType, raw, typ) - } - - fields, err := walkConfigFields(typ, "") - if err != nil { - return nil, err - } - - out := reflect.New(typ).Elem() - for _, f := range fields { - item, present := table[f.key] - if !present || item == nil { - continue - } - converted, err := convertValue(item, f.typ) - if err != nil { - return nil, fmt.Errorf("%w: %s.%s: %w", ErrConfigType, typ.Name(), f.key, err) - } - out.FieldByName(indexFieldName(typ, f.path)).Set(reflect.ValueOf(converted)) - } - return out.Interface(), nil -} - -// asStringMap normalises the two map shapes produced by YAML decoders. -func asStringMap(raw any) (map[string]any, bool) { - switch v := raw.(type) { - case map[string]any: - return v, true - case map[any]any: - out := make(map[string]any, len(v)) - for key, val := range v { - name, ok := key.(string) - if !ok { - return nil, false - } - out[name] = val - } - return out, true - default: - return nil, false - } -} - -// indexFieldName maps a declared config path back to the Go struct field carrying it. -func indexFieldName(t reflect.Type, key string) string { - for i := 0; i < t.NumField(); i++ { - if t.Field(i).Tag.Get("config") == key { - return t.Field(i).Name - } - } - return "" -} -``` - -- [ ] **Step 4: 写解析与绑定实现** - -在 `backend/core/extpoints/config_resolve.go` **末尾追加**下列实现(保留 Task 1 写入的文件头、`Entries` 占位与 `sort` import;本步骤新增用到 `errors`、`fmt`、`reflect`,不要引入 `sync`/`time`): - -```go -// Resolve computes the effective value of every declared key. Priority is, in order: -// an explicit environment override, an auto-enable trigger, the configuration file, -// then the declared default. Resolution is idempotent; later declarations resolve lazily. -func (r *ConfigRegistry) Resolve() error { - r.mu.Lock() - defer r.mu.Unlock() - - if r.src == nil { - return ErrConfigNoSource - } - - var errs []error - for _, key := range r.order { - if _, done := r.values[key]; done { - continue - } - if err := r.resolveLocked(key); err != nil { - errs = append(errs, err) - } - } - r.resolved = true - - return errors.Join(errs...) -} - -// Resolved reports whether Resolve has already run. -func (r *ConfigRegistry) Resolved() bool { - r.mu.RLock() - defer r.mu.RUnlock() - return r.resolved -} - -// resolveLocked computes one key. The caller must hold r.mu. -func (r *ConfigRegistry) resolveLocked(key string) error { - d, ok := r.decls[key] - if !ok { - return fmt.Errorf("%w: %s", ErrConfigUnknownKey, key) - } - - if d.env != "" { - if raw, found := r.src.LookupEnv(d.env); found { - value, err := convertValue(raw, d.typ) - if err != nil { - return fmt.Errorf("%w: key %q from environment %s: %w", ErrConfigType, key, d.env, err) - } - r.values[key], r.origins[key] = value, OriginEnv - return nil - } - } - - if d.autoEnable != "" && d.typ.Kind() == reflect.Bool { - if _, found := r.src.LookupEnv(d.autoEnable); found { - r.values[key], r.origins[key] = true, OriginAutoEnable - return nil - } - } - - if raw, found := r.src.Lookup(key); found { - value, err := convertValue(raw, d.typ) - if err != nil { - return fmt.Errorf("%w: key %q from %s: %w", ErrConfigType, key, r.src.Describe(), err) - } - r.values[key], r.origins[key] = value, OriginFile - return nil - } - - if d.def != "" { - value, err := convertValue(d.def, d.typ) - if err != nil { - return fmt.Errorf("%w: default %q for key %q: %w", ErrConfigType, d.def, key, err) - } - r.values[key], r.origins[key] = value, OriginDefault - return nil - } - - r.values[key] = reflect.New(d.typ).Elem().Interface() - r.origins[key] = "" - return nil -} - -// Bind resolves the tagged fields of target and assigns them in place. Prefixes that -// were never declared self-register, so only gates need DeclareConfig. -func (r *ConfigRegistry) Bind(prefix string, target any) error { - r.mu.Lock() - defer r.mu.Unlock() - - if r.src == nil { - return ErrConfigNoSource - } - if !r.resolved { - return fmt.Errorf("%w: Bind(%q, %T) ran before App.Prepare", ErrConfigNotResolved, prefix, target) - } - - elem, err := bindingStruct(target, prefix) - if err != nil { - return err - } - fields, err := walkConfigFields(elem.Type(), prefix) - if err != nil { - return err - } - - for _, f := range fields { - if _, declared := r.decls[f.key]; !declared { - if err := r.addDecl("bind:"+prefix, f); err != nil { - return err - } - } - if _, done := r.values[f.key]; !done { - if err := r.resolveLocked(f.key); err != nil { - return err - } - } - } - - for _, f := range fields { - value := r.values[f.key] - field := elem.FieldByName(indexFieldName(elem.Type(), f.path)) - if !field.IsValid() || !field.CanSet() { - return fmt.Errorf("%w: field for key %q is not settable", ErrConfigTarget, f.key) - } - rv := reflect.ValueOf(value) - if !rv.Type().AssignableTo(field.Type()) { - return fmt.Errorf("%w: key %q resolves to %s, field expects %s", - ErrConfigType, f.key, rv.Type(), field.Type()) - } - field.Set(rv) - } - return nil -} -``` - -注意:`ErrConfigUnknownKey` 等全部哨兵错误已在 Task 1 的 `config.go` 错误块中定义,本步骤不要重复声明。 - -- [ ] **Step 5: 运行测试确认通过** - -Run: `cd backend && go test ./core/extpoints/ -run 'TestResolve|TestDeclare' -v` -Expected: 全部 `--- PASS`,`ok Wavelet/core/extpoints`。若报 `Entries redeclared`,说明 Step 4 误把整文件替换而非追加。 - -- [ ] **Step 6: 格式与静态检查** - -Run: `cd backend && golangci-lint fmt ./core/extpoints/ && golangci-lint run ./core/extpoints/` -Expected: 无告警。(`convertValue` 保持 8 个分支,若 `cyclop` 仍报复杂度过高,把 `Slice`/`Struct` 两分支拆成独立函数,不要放宽 lint 配置。) - -- [ ] **Step 7: 提交** - -```bash -git add backend/core/extpoints/ -git commit -m "feat(core): resolve declared configuration with env and file precedence" -``` - ---- - -## Task 3: 只读访问器、泛型读取与脱敏导出 - -**Files:** -- Modify: `backend/core/extpoints/config_resolve.go`(追加访问器) -- Modify: `backend/core/extpoints/config.go`(`Entries` 用到的 secret 判断) -- Create: `backend/core/config.go` -- Test: `backend/core/extpoints/config_test.go`(追加)、`backend/core/config_test.go` - -- [ ] **Step 1: 追加失败的测试** - -在 `backend/core/extpoints/config_test.go` 末尾追加: - -```go -func TestViewAccessorsAndOrigins(t *testing.T) { - src := newFakeSource() - src.values["redis.db"] = 1 - src.env["REDIS_ADDR"] = "redis:6379" - src.env["REDIS_ENABLED"] = "false" - - r := extpoints.NewConfigRegistry(src) - require.NoError(t, r.Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: &redisConfig{}})) - require.NoError(t, r.Declare("auth", extpoints.ConfigBinding{Prefix: "app", Target: &sessionConfig{}})) - require.NoError(t, r.Resolve()) - - assert.Equal(t, extpoints.OriginEnv, r.Origin("redis.addrs")) - assert.Equal(t, "redis:6379", r.Strings("redis.addrs")[0]) - assert.False(t, r.Bool("redis.enabled", true)) - assert.Equal(t, 1, r.Int("redis.db", 0)) - assert.Equal(t, "86400", r.String("redis.missing", "86400")) - assert.True(t, r.WasSet("REDIS_ADDR")) - assert.False(t, r.WasSet("REDIS_NOPE")) -} - -func TestAutoEnableBeatsFileValueButLosesToExplicitEnv(t *testing.T) { - src := newFakeSource() - src.env["REDIS_ADDR"] = "redis:6379" - src.values["redis.enabled"] = false - - r := extpoints.NewConfigRegistry(src) - require.NoError(t, r.Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: &redisConfig{}})) - require.NoError(t, r.Resolve()) - assert.True(t, r.Bool("redis.enabled", false), "REDIS_ADDR presence implies enabled") - assert.Equal(t, extpoints.OriginAutoEnable, r.Origin("redis.enabled")) - - explicit := newFakeSource() - explicit.env["REDIS_ADDR"] = "redis:6379" - explicit.env["REDIS_ENABLED"] = "false" - - r2 := extpoints.NewConfigRegistry(explicit) - require.NoError(t, r2.Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: &redisConfig{}})) - require.NoError(t, r2.Resolve()) - assert.False(t, r2.Bool("redis.enabled", true), "explicit REDIS_ENABLED must win over auto-enable") - assert.Equal(t, extpoints.OriginEnv, r2.Origin("redis.enabled")) -} - -func TestEntriesRedactSecretsAndReportDefaults(t *testing.T) { - r := extpoints.NewConfigRegistry(newFakeSource()) - require.NoError(t, r.Declare("auth", extpoints.ConfigBinding{Prefix: "app", Target: &sessionConfig{}})) - require.NoError(t, r.Resolve()) - - entries := map[string]extpoints.ConfigEntry{} - for _, e := range r.Entries() { - entries[e.Key] = e - } - - assert.Equal(t, extpoints.RedactedValue, entries["app.session_secret"].Value) - assert.Equal(t, extpoints.OriginDefault, entries["app.session_age"].Origin) - assert.Equal(t, "86400", entries["app.session_age"].Value) -} - -func TestBindRejectsReadsBeforeSourceIsRegistered(t *testing.T) { - r := extpoints.NewConfigRegistry(nil) - require.NoError(t, r.Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: &redisConfig{}})) - - assert.ErrorIs(t, r.Resolve(), extpoints.ErrConfigNoSource) - - var cfg redisConfig - assert.ErrorIs(t, r.Bind("redis", &cfg), extpoints.ErrConfigNoSource) -} -``` - -新增脱敏常量到 `backend/core/extpoints/config.go`(Task 1 未定义它,因为此处才首次使用): - -```go -// RedactedValue replaces the printed value of keys declared with secret:"true". -const RedactedValue = "******" -``` - -- [ ] **Step 2: 运行测试确认失败** - -Run: `cd backend && go test ./core/extpoints/ -run 'TestView|TestAutoEnable|TestEntries|TestBindRejects' -v` -Expected: 编译失败,报 `r.Bool undefined`、`extpoints.RedactedValue undefined` 等;`Entries` 已有但返回 `"pending"` 占位,故 `TestEntriesRedactSecretsAndReportDefaults` 亦失败。 - -- [ ] **Step 3: 实现访问器** - -在 `backend/core/extpoints/config_resolve.go` 末尾追加: - -```go -// Value returns the resolved value for key, lazily resolving it when a source is -// available. Unresolvable and missing keys report false rather than an error so -// gates and diagnostics can keep using fallback accessors. -func (r *ConfigRegistry) Value(key string) (any, bool) { - r.mu.Lock() - defer r.mu.Unlock() - - if r.src == nil { - value, ok := r.values[key] - return value, ok - } - if _, done := r.values[key]; !done { - if _, declared := r.decls[key]; !declared { - return nil, false - } - if err := r.resolveLocked(key); err != nil { - return nil, false - } - } - value, ok := r.values[key] - return value, ok -} - -// String returns the string value of key or fallback when absent or mismatched. -func (r *ConfigRegistry) String(key, fallback string) string { - if value, ok := r.Value(key); ok { - if converted, err := convertString(value); err == nil { - return converted.(string) - } - } - return fallback -} - -// Bool returns the boolean value of key or fallback when absent or mismatched. -func (r *ConfigRegistry) Bool(key string, fallback bool) bool { - if value, ok := r.Value(key); ok { - if converted, err := convertBool(value); err == nil { - return converted.(bool) - } - } - return fallback -} - -// Int returns the int value of key or fallback when absent or mismatched. -func (r *ConfigRegistry) Int(key string, fallback int) int { - if value, ok := r.Value(key); ok { - if converted, err := convertInt(value, reflect.TypeFor[int]()); err == nil { - return int(converted.(int)) - } - } - return fallback -} - -// Duration returns the time.Duration value of key or fallback when absent or mismatched. -func (r *ConfigRegistry) Duration(key string, fallback time.Duration) time.Duration { - if value, ok := r.Value(key); ok { - if converted, err := convertDuration(value); err == nil { - return converted.(time.Duration) - } - } - return fallback -} - -// Strings returns the []string value of key, or nil when absent. -func (r *ConfigRegistry) Strings(key string) []string { - value, ok := r.Value(key) - if !ok { - return nil - } - converted, err := convertSlice(value, reflect.TypeFor[[]string]()) - if err != nil { - return nil - } - list, _ := converted.([]string) - return list -} - -// WasSet reports whether an environment variable is present, regardless of its value. -func (r *ConfigRegistry) WasSet(envName string) bool { - r.mu.RLock() - defer r.mu.RUnlock() - if r.src == nil { - return false - } - _, found := r.src.LookupEnv(envName) - return found -} - -// Origin reports where a key's effective value came from; "" means the zero value. -func (r *ConfigRegistry) Origin(key string) string { - r.mu.RLock() - defer r.mu.RUnlock() - return r.origins[key] -} -``` - -把 `Entries` 的占位实现替换为真实值与脱敏: - -```go -// Entries returns the effective configuration as redacted, key-sorted entries. -func (r *ConfigRegistry) Entries() []ConfigEntry { - r.mu.Lock() - defer r.mu.Unlock() - - keys := append([]string(nil), r.order...) - sort.Strings(keys) - - out := make([]ConfigEntry, 0, len(keys)) - for _, key := range keys { - d := r.decls[key] - if _, done := r.values[key]; !done && r.src != nil { - _ = r.resolveLocked(key) - } - out = append(out, ConfigEntry{ - Key: d.key, - PluginID: d.pluginID, - Env: d.env, - Origin: r.origins[key], - Value: formatEntryValue(r.values[key], d.secret), - }) - } - return out -} - -// formatEntryValue renders one effective value for diagnostics, masking secrets. -func formatEntryValue(value any, secret bool) string { - if secret { - return RedactedValue - } - if value == nil { - return "" - } - return fmt.Sprint(value) -} -``` - -- [ ] **Step 4: 实现泛型读取入口** - -创建 `backend/core/config.go`: - -```go -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package core - -import ( - "fmt" - - "Wavelet/core/extpoints" -) - -// ConfigGet reads one resolved configuration value with its declared type. It is the -// generic counterpart of the fallback accessors on ConfigView, and returns -// ErrConfigNotResolved when the key has neither been declared nor resolved. -func ConfigGet[T any](view extpoints.ConfigView, key string) (T, error) { - var zero T - if view == nil { - return zero, extpoints.ErrConfigNotResolved - } - - raw, ok := view.Value(key) - if !ok { - return zero, fmt.Errorf("%w: %s", extpoints.ErrConfigUnknownKey, key) - } - - value, ok := raw.(T) - if !ok { - return zero, fmt.Errorf("%w: key %q holds %T, want %T", extpoints.ErrConfigType, key, raw, zero) - } - return value, nil -} -``` - -创建 `backend/core/config_test.go`: - -```go -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package core_test - -import ( - "Wavelet/core" - "Wavelet/core/extpoints" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -type otelConfig struct { - SamplingRate float64 `config:"sampling_rate" env:"OTEL_SAMPLING_RATE"` -} - -func TestConfigGetReturnsDeclaredType(t *testing.T) { - r := extpoints.NewConfigRegistry(nil) - require.NoError(t, r.Declare("host", extpoints.ConfigBinding{Prefix: "otel", Target: &otelConfig{}})) - - rate, err := core.ConfigGet[float64](r, "otel.sampling_rate") - require.ErrorIs(t, err, extpoints.ErrConfigUnknownKey) - assert.Equal(t, 0.0, rate) -} -``` - -- [ ] **Step 5: 运行测试确认通过** - -Run: `cd backend && go test ./core/... -run 'TestView|TestAutoEnable|TestEntries|TestBindRejects|TestConfigGet' -v` -Expected: 全部 `--- PASS`。 - -- [ ] **Step 6: 格式与静态检查** - -Run: `cd backend && golangci-lint fmt ./core/... && golangci-lint run ./core/...` -Expected: 无告警。若 `dupl` 因 `convertInt`/`convertUint` 结构相似报警,为其中之一加注释说明类型不同不可合并,或拆出公共 reflect 设置函数;不得关闭 `dupl`。 - -- [ ] **Step 7: 提交** - -```bash -git add backend/core/ -git commit -m "feat(core): add read-only config view accessors and generic getter" -``` - ---- - -## Task 4: 把配置注册表挂到 Context 并在 types.go 导出别名 - -**Files:** -- Modify: `backend/core/context.go`(`config` 字段、`NewContext`、`Fork`、`Config()`) -- Modify: `backend/core/types.go`(别名) -- Test: `backend/core/config_test.go`(追加)、`backend/core/context_test.go`(追加断言) - -- [ ] **Step 1: 追加失败的测试** - -在 `backend/core/config_test.go` 末尾追加: - -```go -func TestContextConfigIsSharedAcrossForks(t *testing.T) { - ctx := core.NewContext(nil) - child := ctx.Fork() - - require.NoError(t, child.Config().Declare("cache", extpoints.ConfigBinding{Prefix: "redis", Target: &otelConfig{}})) - - rate, err := core.ConfigGet[float64](ctx.Config(), "redis.sampling_rate") - require.ErrorIs(t, err, extpoints.ErrConfigUnknownKey) - assert.Zero(t, rate) - assert.False(t, ctx.Config().Resolved()) -} -``` - -在 `backend/core/context_test.go` 中已有的 Context 构造测试里追加一行断言(沿用该文件现有测试函数与变量名): - -```go - require.NotNil(t, ctx.Config(), "every Context must expose the configuration extension") -``` - -- [ ] **Step 2: 运行测试确认失败** - -Run: `cd backend && go test ./core/ -run 'TestContextConfigIsSharedAcrossForks' -v` -Expected: 编译失败,报 `ctx.Config undefined`、`core.ConfigBinding undefined`(测试里改用 `extpoints.ConfigBinding` 后该项消除)。 - -- [ ] **Step 3: 写实现** - -`backend/core/context.go` 的 `Context` 结构体字段中,在 `settings` 之后加一行: - -```go - settings extpoints.SettingExtension - config extpoints.ConfigExtension -``` - -`NewContext` 的返回字面量中,在 `settings: extpoints.NewSettingRegistry(),` 之后加: - -```go - config: extpoints.NewConfigRegistry(nil), -``` - -`ForkWithContext` 的 child 字面量中,在 `settings: c.settings,` 之后加: - -```go - config: c.config, -``` - -在 `Setting()` 别名方法之后加访问器: - -```go -// Config returns the process-level configuration extension point. The registry is -// shared by every fork because configuration declarations are global facts, and it -// intentionally carries no per-scope disposers: values are resolved once before Apply. -func (c *Context) Config() extpoints.ConfigExtension { - return c.config -} -``` - -`backend/core/types.go` 末尾追加别名与新接口: - -```go -// ConfigExtension re-exports extpoints.ConfigExtension. -type ConfigExtension = extpoints.ConfigExtension - -// ConfigSource re-exports extpoints.ConfigSource. -type ConfigSource = extpoints.ConfigSource - -// ConfigBinding re-exports extpoints.ConfigBinding. -type ConfigBinding = extpoints.ConfigBinding - -// ConfigView re-exports extpoints.ConfigView. -type ConfigView = extpoints.ConfigView - -// ConfigEntry re-exports extpoints.ConfigEntry. -type ConfigEntry = extpoints.ConfigEntry - -// ConfigGatedPlugin is an optional interface for plugins whose activation depends on -// configuration. The kernel evaluates the gate before any Apply runs, so keys read by -// ConfigEnabled must be published through DeclareConfig. -type ConfigGatedPlugin interface { - Plugin - - // DeclareConfig publishes the configuration bindings consumed by ConfigEnabled. - DeclareConfig() []extpoints.ConfigBinding - - // ConfigEnabled reports whether this plugin should activate for the resolved values. - ConfigEnabled(view extpoints.ConfigView) bool -} -``` - -- [ ] **Step 4: 运行测试确认通过** - -Run: `cd backend && go test ./core/... -v -run 'TestContextConfig|TestFiber|TestApp|TestConfigGet'` -Expected: 新增用例 `--- PASS`,既有用例无回归。 - -- [ ] **Step 5: 提交** - -```bash -git add backend/core/context.go backend/core/types.go backend/core/config_test.go backend/core/context_test.go -git commit -m "feat(core): mount the configuration extension point on the kernel Context" -``` - ---- - -## Task 5: Fiber 门禁跳过态 - -**Files:** -- Modify: `backend/core/fiber.go` -- Test: `backend/core/fiber_test.go`(追加) - -- [ ] **Step 1: 追加失败的测试** - -在 `backend/core/fiber_test.go` 末尾追加: - -```go -// gatedPlugin is a minimal plugin used to exercise configuration gating. -type gatedPlugin struct { - name string - enabled bool - applied bool -} - -func (g *gatedPlugin) Name() string { return g.name } -func (g *gatedPlugin) Apply(ctx *core.Context) error { - g.applied = true - return nil -} -func (g *gatedPlugin) DeclareConfig() []extpoints.ConfigBinding { - return []extpoints.ConfigBinding{{Prefix: "gate", Target: &gateConfig{}}} -} -func (g *gatedPlugin) ConfigEnabled(view extpoints.ConfigView) bool { - return view.Bool("gate.enabled", false) == g.enabled -} - -type gateConfig struct { - Enabled bool `config:"enabled" env:"GATE_ENABLED"` -} - -func TestFiberSkipMovesToSkippedStateAndDisposesScope(t *testing.T) { - root := core.NewContext(nil) - plugin := &gatedPlugin{name: "cache", enabled: true} - f := core.NewFiber(root, plugin) - require.Equal(t, core.FiberPending, f.State()) - - require.NoError(t, f.Skip()) - - assert.Equal(t, core.FiberSkipped, f.State()) - assert.True(t, f.Skipped()) - assert.False(t, plugin.applied, "a skipped plugin must never reach Apply") - assert.NoError(t, f.Unload(), "unloading a skipped fiber is a no-op") -} - -func TestFiberSkipIsIdempotentForActiveFibers(t *testing.T) { - root := core.NewContext(nil) - f := core.NewFiber(root, &gatedPlugin{name: "cache", enabled: true}) - require.NoError(t, f.Load()) - - require.NoError(t, f.Skip()) - assert.Equal(t, core.FiberActive, f.State(), "Skip only applies to pending fibers") -} -``` - -- [ ] **Step 2: 运行测试确认失败** - -Run: `cd backend && go test ./core/ -run 'TestFiberSkip' -v` -Expected: 编译失败,报 `undefined: core.FiberSkipped`、`f.Skip undefined`、`f.Skipped undefined`。 - -- [ ] **Step 3: 写实现** - -`backend/core/fiber.go` 的 `FiberState` 常量块中,在 `FiberDisposed` 之后追加一个状态: - -```go - // FiberSkipped indicates the plugin never activated because its configuration gate - // evaluated to false, so an alternative provider took over. - FiberSkipped FiberState = "SKIPPED" -``` - -`Load` 之后追加 `Skip`: - -```go -// Skip transitions a pending plugin to FiberSkipped and releases its scoped Context. -// Active plugins are left untouched, making the call safe to replay during reconcile. -func (f *Fiber) Skip() error { - f.mu.Lock() - if f.state != FiberPending { - f.mu.Unlock() - return nil - } - f.state = FiberSkipped - f.mu.Unlock() - - return f.ctx.Dispose() -} - -// Skipped reports whether the plugin was excluded by its configuration gate. -func (f *Fiber) Skipped() bool { - return f.State() == FiberSkipped -} -``` - -> **不要改 `DependenciesSatisfied`**:依赖能否满足完全由 IoC 容器解析决定。被跳过的插件从未执行 `Apply`,也就没有 `core.Provide`,其消费者自然解析不到服务并在 `Reconcile` 里报"waiting for"。用 Fiber 状态做短路是错误语义。 - -- [ ] **Step 4: 运行测试确认通过** - -Run: `cd backend && go test ./core/ -run 'TestFiber' -v` -Expected: 新增两个用例 `--- PASS`,既有 `TestFiber_ConfluenceAndReactiveActivation`、`TestFiber_UnsatisfiedDependencyReturnsError` 无回归。 - -- [ ] **Step 5: 提交** - -```bash -git add backend/core/fiber.go backend/core/fiber_test.go -git commit -m "feat(core): add skipped fiber state for configuration gates" -``` - ---- - -## Task 6: App 装配选项、解析屏障与门禁求值 - -**Files:** -- Modify: `backend/core/app.go` -- Test: `backend/core/app_test.go`(追加) - -- [ ] **Step 1: 追加失败的测试** - -在 `backend/core/app_test.go` 末尾追加(复用 Task 5 的 `gatedPlugin`/`gateConfig`;`mapSource` 为本测试自备的内存源): - -```go -// mapSource implements core.ConfigSource over static maps. -type mapSource struct { - values map[string]any - env map[string]string -} - -func (m *mapSource) Lookup(path string) (any, bool) { - v, ok := m.values[path] - return v, ok -} -func (m *mapSource) LookupEnv(name string) (string, bool) { - v, ok := m.env[name] - return v, ok -} -func (m *mapSource) Describe() string { return "map" } - -func newGateSource(enabled bool) *mapSource { - src := &mapSource{values: map[string]any{"gate.enabled": enabled}, env: map[string]string{}} - return src -} - -func TestAppPrepareResolvesAndGatesPlugins(t *testing.T) { - redisLike := &gatedPlugin{name: "cache", enabled: true} - redisAlt := &gatedPlugin{name: "cache_memory", enabled: false} - - app := core.NewApp( - core.WithProfile(core.ProfileAPI), - core.WithConfigSource(newGateSource(true)), - ) - app.Use(redisLike, redisAlt) - require.NoError(t, app.Prepare()) - - cacheFiber, ok := app.Fiber("cache") - require.True(t, ok) - require.Equal(t, core.FiberPending, cacheFiber.State(), "Prepare only builds the resolution barrier") - assert.True(t, app.Context().Config().Resolved()) - - require.NoError(t, app.Reconcile()) - - require.Equal(t, core.FiberActive, cacheFiber.State()) - - memoryFiber, ok := app.Fiber("cache_memory") - require.True(t, ok) - assert.Equal(t, core.FiberSkipped, memoryFiber.State()) - assert.False(t, redisAlt.applied) -} - -func TestAppGatesPluginsMountedAfterPrepare(t *testing.T) { - app := core.NewApp(core.WithConfigSource(newGateSource(true))) - require.NoError(t, app.Prepare()) - - late := &gatedPlugin{name: "cache_memory", enabled: false} - app.Use(late) - require.NoError(t, app.Reconcile()) - - fiber, ok := app.Fiber("cache_memory") - require.True(t, ok) - assert.Equal(t, core.FiberSkipped, fiber.State(), - "plugins added after Prepare must still be gated") -} - -func TestAppApplyPluginsGatesImplicitly(t *testing.T) { - redisLike := &gatedPlugin{name: "cache", enabled: true} - app := core.NewApp(core.WithConfigSource(newGateSource(false))) - app.Use(redisLike) - - require.NoError(t, app.ApplyPlugins()) - - fiber, ok := app.Fiber("cache") - require.True(t, ok) - assert.Equal(t, core.FiberSkipped, fiber.State(), "ApplyPlugins must resolve and gate implicitly") -} - -func TestAppPrepareReportsConfigurationErrors(t *testing.T) { - src := &mapSource{values: map[string]any{"gate.enabled": "yes"}, env: map[string]string{}} - app := core.NewApp(core.WithConfigSource(src)) - app.Use(&gatedPlugin{name: "cache", enabled: true}) - - err := app.Prepare() - require.Error(t, err) - assert.Contains(t, err.Error(), "gate.enabled") -} - -func TestAppSetShutdownTimeoutOverridesDefault(t *testing.T) { - app := core.NewApp() - app.SetShutdownTimeout(0) - assert.NotZero(t, app.ShutdownTimeout(), "zero durations must not shrink the kernel fallback") - - app.SetShutdownTimeout(45 * time.Second) - assert.Equal(t, 45*time.Second, app.ShutdownTimeout()) -} -``` - -- [ ] **Step 2: 运行测试确认失败** - -Run: `cd backend && go test ./core/ -run 'TestAppPrepare|TestAppStartRunsPrepare|TestAppSetShutdown' -v` -Expected: 编译失败,报 `core.WithConfigSource undefined`、`app.Prepare undefined`、`app.ShutdownTimeout undefined`。 - -- [ ] **Step 3: 写实现** - -`backend/core/app.go` 中,`App` 结构体追加两个字段: - -```go - migrationEngine MigrationEngine - shutdownTimeout time.Duration - configSource ConfigSource - prepared bool -``` - -`AppOption` 区追加选项(放在 `WithShutdownTimeout` 之后): - -```go -// WithConfigSource installs the raw configuration source adapter, typically built by -// an infrastructure package outside the kernel, before any plugin is applied. -func WithConfigSource(src ConfigSource) AppOption { - return func(a *App) { - if src == nil { - return - } - a.configSource = src - a.ctx.Config().SetSource(src) - } -} - -// WithConfigDecl lets the composition root declare the configuration it reads itself, -// so host-level values participate in conflict validation and redacted reporting. -func WithConfigDecl(pluginID string, bindings ...ConfigBinding) AppOption { - return func(a *App) { - if len(bindings) == 0 { - return - } - if err := a.ctx.Config().Declare(pluginID, bindings...); err != nil { - a.applyErr = err - } - } -} -``` - -`App` 结构体再加 `applyErr error` 字段,并在 `NewApp` 末尾返回前保持原逻辑(`applyErr` 由 `Prepare` 首次上报)。 - -追加 `Prepare`、`ShutdownTimeout`、`SetShutdownTimeout`: - -```go -// Prepare resolves declared configuration and evaluates plugin gates. It is idempotent -// and runs implicitly from ApplyPlugins, so callers that need resolved values earlier -// (for example to size a shutdown budget) can invoke it explicitly. -func (a *App) Prepare() error { - a.mu.Lock() - defer a.mu.Unlock() - - if err := a.applyErr; err != nil { - return err - } - return a.prepareLocked() -} - -func (a *App) prepareLocked() error { - if a.prepared { - return nil - } - - if err := a.ctx.Config().Resolve(); err != nil { - return err - } - a.prepared = true - - return nil -} - -// ShutdownTimeout returns the graceful shutdown budget for the application. -func (a *App) ShutdownTimeout() time.Duration { - a.mu.RLock() - defer a.mu.RUnlock() - return a.shutdownTimeout -} - -// SetShutdownTimeout replaces the graceful shutdown budget, ignoring non-positive values. -func (a *App) SetShutdownTimeout(timeout time.Duration) *App { - a.mu.Lock() - defer a.mu.Unlock() - if timeout > 0 { - a.shutdownTimeout = timeout - } - return a -} -``` - -> **门禁为什么在 `reconcileLocked` 内求值而不是 `Prepare` 里一次性遍历**:`App.Use` 可以在 `Prepare` 之后继续挂载插件(下游定制与动态装配)。只在 `Prepare` 求值会留下一批永不判定的门禁;放在调和循环里则任何时刻新挂载的插件都会被正确判定,且 `Fiber.Skip` 自带"仅 Pending 可跳过"守卫,重复遍历安全。 - -`Use` 中,为每个成功登记的插件收集声明(放在 `a.pluginMap[name] = p` 之前): - -```go - if gated, ok := p.(ConfigGatedPlugin); ok { - if err := a.ctx.Config().Declare(name, gated.DeclareConfig()...); err != nil { - if a.applyErr == nil { - a.applyErr = err - } - } - } -``` - -`ApplyPlugins` 与 `reconcileLocked` 接入屏障(`ApplyPlugins` 已持锁,改调用 `prepareLocked`): - -```go -func (a *App) ApplyPlugins() error { - a.mu.Lock() - if a.applied { - a.mu.Unlock() - return nil - } - a.applied = true - - if err := a.applyErr; err != nil { - a.mu.Unlock() - return err - } - if err := a.prepareLocked(); err != nil { - a.mu.Unlock() - return err - } - a.mu.Unlock() - - return a.Reconcile() -} -``` - -把 `reconcileLocked` 的内层循环替换为带门禁判定的版本(其余保持不变): - -```go -func (a *App) reconcileLocked() error { - if err := a.prepareLocked(); err != nil { - return err - } - - view := a.ctx.Config() - - for { - progress := false - for _, f := range a.fibers { - if f.State() != FiberPending { - continue - } - - if gated, ok := f.plugin.(ConfigGatedPlugin); ok { - if !view.Resolved() { - continue - } - if !gated.ConfigEnabled(view) { - if err := f.Skip(); err != nil { - return fmt.Errorf("core: skip gated plugin %q: %w", f.Name(), err) - } - continue - } - } - - if f.DependenciesSatisfied(a.ctx) { - if err := f.Load(); err != nil { - return fmt.Errorf("core: load fiber %q failed: %w", f.Name(), err) - } - progress = true - } - } - if !progress { - break - } - } - - // ...existing unsatisfied-dependency reporting unchanged -} -``` - -`unsatisfied` 收集循环无需改动:它只统计 `FiberPending`,被门禁排除的插件已是 `FiberSkipped`。 - -- [ ] **Step 4: 运行测试确认通过** - -Run: `cd backend && go test ./core/ -v` -Expected: 全部 `--- PASS`,包括既有 `TestApp*`、`TestFiber*`、`TestContext*`。 - -- [ ] **Step 5: 全量回归(应用行为必须不变)** - -Run: `cd backend && go test ./... && go build -o /dev/null ./...` -Expected: 全绿。此时业务插件与 `cmd` 仍走旧的全局单例,因此运行行为与迁移前完全一致——这是本计划的关键安全属性。 - -- [ ] **Step 6: 格式与静态检查** - -Run: `cd backend && golangci-lint fmt ./core/... && golangci-lint run ./core/...` -Expected: 无告警。 - -- [ ] **Step 7: 提交** - -```bash -git add backend/core/app.go backend/core/app_test.go -git commit -m "feat(core): add config resolution barrier and plugin gating to App" -``` - ---- - -## Task 7: viper 配置源适配器(plugins/infra/config) - -**Files:** -- Create: `backend/plugins/infra/config/source.go` -- Create: `backend/plugins/infra/config/source_test.go` - -- [ ] **Step 1: 写失败的测试** - -创建 `backend/plugins/infra/config/source_test.go`: - -```go -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -package config_test - -import ( - "os" - "path/filepath" - "testing" - - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - "Wavelet/plugins/infra/config" -) - -const sampleYAML = "" + - "app:\n addr: \":8000\"\n node_id: 1\n" + - "database:\n enabled: false\n port: 5432\n slow_threshold: 200ms\n" + - "redis:\n addrs:\n - \"127.0.0.1:6379\"\n" - -func writeConfig(t *testing.T) string { - t.Helper() - - dir := t.TempDir() - path := filepath.Join(dir, "config.yaml") - require.NoError(t, os.WriteFile(path, []byte(sampleYAML), 0o600)) - return path -} - -func TestSourceLooksUpNestedPaths(t *testing.T) { - src, err := config.NewSource(config.WithPath(writeConfig(t))) - require.NoError(t, err) - - value, ok := src.Lookup("database.port") - require.True(t, ok) - assert.Equal(t, 5432, value) - - _, ok = src.Lookup("database.missing") - assert.False(t, ok) -} - -func TestSourceTreatsUnsetFileAsEnvOnly(t *testing.T) { - missing := filepath.Join(t.TempDir(), "absent.yaml") - - src, err := config.NewSource(config.WithPath(missing)) - require.NoError(t, err, "a missing configuration file must fall back to environment values") - - _, ok := src.Lookup("app.addr") - assert.False(t, ok) - assert.Equal(t, config.EnvOnlyOrigin, src.Describe()) -} - -func TestSourceRejectsMalformedFile(t *testing.T) { - dir := t.TempDir() - path := filepath.Join(dir, "config.yaml") - require.NoError(t, os.WriteFile(path, []byte("app: [unclosed\n"), 0o600)) - - _, err := config.NewSource(config.WithPath(path)) - require.Error(t, err) -} - -func TestSourceLookupEnvReadsProcessEnvironment(t *testing.T) { - t.Setenv("WAVELET_SOURCE_PROBE", "present") - - src, err := config.NewSource(config.WithPath(writeConfig(t))) - require.NoError(t, err) - - value, ok := src.LookupEnv("WAVELET_SOURCE_PROBE") - require.True(t, ok) - assert.Equal(t, "present", value) - - _, ok = src.LookupEnv("WAVELET_SOURCE_ABSENT") - assert.False(t, ok) -} -``` - -- [ ] **Step 2: 运行测试确认失败** - -Run: `cd backend && go test ./plugins/infra/config/ -v` -Expected: 编译失败,报 `package Wavelet/plugins/infra/config is not in std` / `undefined: config.NewSource`。 - -- [ ] **Step 3: 写实现** - -创建 `backend/plugins/infra/config/source.go`: - -```go -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Package config adapts viper to the kernel configuration source contract. It is a -// runtime adapter rather than a core.Plugin: it owns no routes, services or tasks, and -// therefore never appears in app.Use. Keeping it out of core preserves the micro-kernel -// rule against importing concrete runtime dependencies. -package config - -import ( - "errors" - "fmt" - "os" - - "github.com/spf13/viper" -) - -// DefaultFileName is the configuration file looked up when CONFIG_PATH is unset. -const DefaultFileName = "config.yaml" - -// EnvOnlyOrigin is reported by Describe when no configuration file was loaded. -const EnvOnlyOrigin = "" - -// maxSearchDepth bounds the upward directory walk so a misconfigured working set -// cannot make the loader scan the whole filesystem. -const maxSearchDepth = 5 - -// Option configures a Source. -type Option func(*Source) - -// WithPath pins the configuration file, bypassing CONFIG_PATH and the upward search. -func WithPath(path string) Option { - return func(s *Source) { - s.path = path - } -} - -// Source implements core.ConfigSource over a configuration file plus the process environment. -type Source struct { - v *viper.Viper - path string - found bool -} - -// NewSource loads the configuration file. A missing file is not an error: the source -// then serves environment values only, matching the previous pkg/config behaviour. -func NewSource(opts ...Option) (*Source, error) { - s := &Source{} - for _, opt := range opts { - opt(s) - } - - if s.path == "" { - s.path = os.Getenv("CONFIG_PATH") - } - if s.path == "" { - s.path = findConfigPath(DefaultFileName) - } - - v := viper.New() - v.SetConfigFile(s.path) - - err := v.ReadInConfig() - switch { - case err == nil: - s.found = true - case isNotFound(err): - // fall through to environment-only lookups - default: - if _, statErr := os.Stat(s.path); statErr == nil { //nolint:gosec // s.path comes from CONFIG_PATH or a bounded upward search - return nil, fmt.Errorf("infra/config: read %s: %w", s.path, err) - } - } - - s.v = v - return s, nil -} - -func isNotFound(err error) bool { - var notFound viper.ConfigFileNotFoundError - return errors.As(err, ¬Found) || errors.Is(err, os.ErrNotExist) -} - -// Lookup returns the raw value stored at a dotted path, or false when the file was not -// loaded or the path is absent. -func (s *Source) Lookup(path string) (any, bool) { - if !s.found || !s.v.IsSet(path) { - return nil, false - } - return s.v.Get(path), true -} - -// LookupEnv reads a process environment variable. -func (s *Source) LookupEnv(name string) (string, bool) { - return os.LookupEnv(name) -} - -// Describe returns the loaded file path, or EnvOnlyOrigin when running on environment values. -func (s *Source) Describe() string { - if !s.found { - return EnvOnlyOrigin - } - return s.path -} - -// findConfigPath searches upward from the working directory so tests and binaries run -// from backend/ still find the repository-root configuration file. -func findConfigPath(configPath string) string { - if _, err := os.Stat(configPath); err == nil { - return configPath - } - - dir := "." - for range maxSearchDepth { - dir += "/.." - path := dir + "/" + configPath - if _, err := os.Stat(path); err == nil { - return path - } - } - return configPath -} -``` - -- [ ] **Step 4: 运行测试确认通过** - -Run: `cd backend && go test ./plugins/infra/config/ -v` -Expected: 四个用例全部 `--- PASS`。 - -- [ ] **Step 5: 格式与静态检查** - -Run: `cd backend && golangci-lint fmt ./plugins/infra/config/ && golangci-lint run ./plugins/infra/config/` -Expected: 无告警。 - -- [ ] **Step 6: 提交** - -```bash -git add backend/plugins/infra/config/ -git commit -m "feat(infra): add viper backed configuration source adapter" -``` - ---- - -## Task 8: 旧配置包可重入重构 + 新旧对拍 - -**Files:** -- Modify: `backend/pkg/config/config.go:57-108` -- Create: `backend/pkg/config/parity_test.go`(临时文件,P4 随旧包一并删除) - -- [ ] **Step 1: 重构旧加载器为可重入函数** - -把 `backend/pkg/config/config.go` 的 `init()` 拆成 `load` + `init`。原实现使用包级 `viper` 全局并在 `init` 里内联全部步骤,对拍需要能反复调用且不受 `isTest()` 干扰: - -```go -// load reads configuration from configPath, applies defaults and environment overrides, -// and optionally disables external services for in-test runs. -func load(configPath string, testMode bool) *configModel { - v := viper.New() - v.SetConfigFile(configPath) - - if err := v.ReadInConfig(); err != nil { - var notFound viper.ConfigFileNotFoundError - if !errors.As(err, ¬Found) { - if _, statErr := os.Stat(configPath); statErr == nil { //nolint:gosec // configPath is loaded from CONFIG_PATH environment variable - log.Fatalf("[Config] read config failed: %v\n", err) - } - } - log.Println("[Config] no config file found, using environment variables only") - v.SetConfigType("yaml") - if err := v.ReadConfig(strings.NewReader("")); err != nil { - log.Fatalf("[Config] failed to init empty config: %v\n", err) - } - } - - var c configModel - if err := v.Unmarshal(&c); err != nil { - log.Fatalf("[Config] parse config failed: %v\n", err) - } - - applyDefaults(&c) - applyEnvOverrides(&c) - applyDefaults(&c) - - if testMode { - c.Database.Enabled = false - c.Database.SQLitePath = ":memory:" - c.Redis.Enabled = false - c.ClickHouse.Enabled = false - } - - return &c -} - -func init() { - configPath := os.Getenv("CONFIG_PATH") - if configPath == "" { - configPath = findConfigPath("config.yaml") - } - - Config = load(configPath, isTest()) - - printConfig(Config) -} -``` - -同步调整:`import` 增加 `"errors"`;删除原 `init` 里的 `viper.SetConfigFile`/`viper.AutomaticEnv`/`viper.ReadInConfig` 等包级调用与 `if _, ok := err.(viper.ConfigFileNotFoundError); !ok` 断言(改用上面的 `errors.As`)。 - -> **行为等价说明(评审时核对)**:`viper.AutomaticEnv()` 只影响按键读取,`Unmarshal` 走的是 `AllKeys`,因此去掉 `AutomaticEnv` 不改变解析结果;环境变量覆盖仍由 `applyEnvOverrides` 负责。 - -- [ ] **Step 2: 运行旧测试确认无回归** - -Run: `cd backend && go test ./pkg/config/ ./cmd/ -v -run 'TestApplyEnvOverrides|Test' 2>&1 | tail -30` -Expected: `pkg/config` 的 `TestApplyEnvOverridesRedisMaintNotifications` 通过;`cmd` 包既有用例结果与改动前一致(先运行一次改动前的 `go test ./cmd/` 记录基线)。 - -- [ ] **Step 3: 写对拍测试** - -创建 `backend/pkg/config/parity_test.go`: - -```go -// Copyright 2026 Arctel.net -// SPDX-License-Identifier: Apache-2.0 - -// Temporary migration harness: proves the new kernel configuration engine resolves -// every key identically to pkg/config before the legacy singleton is deleted in P4. -// Delete this file together with backend/pkg/config. -package config - -import ( - "os" - "path/filepath" - "reflect" - "testing" - "time" - - "github.com/google/go-cmp/cmp" - "github.com/spf13/viper" - - "Wavelet/core/extpoints" -) - -// yamlSource is a test-local core.ConfigSource over the repository config file. -// It deliberately does not import plugins/infra/config: backend/pkg must not depend on -// upper layers even in tests, and the adapter has its own coverage in its package tests. -type yamlSource struct { - v *viper.Viper -} - -func newYAMLSource(t *testing.T, path string) *yamlSource { - t.Helper() - - v := viper.New() - v.SetConfigFile(path) - require.NoError(t, v.ReadInConfig()) - return &yamlSource{v: v} -} - -func (s *yamlSource) Lookup(path string) (any, bool) { - if !s.v.IsSet(path) { - return nil, false - } - return s.v.Get(path), true -} - -func (s *yamlSource) LookupEnv(name string) (string, bool) { return os.LookupEnv(name) } - -func (s *yamlSource) Describe() string { return s.v.ConfigFileUsed() } - -// engineAppConfig mirrors appConfig with engine tags. -type engineAppConfig struct { - AppName string `config:"app_name" env:"APP_NAME"` - Env string `config:"env" env:"APP_ENV"` - Addr string `config:"addr" env:"APP_ADDR"` - NodeID int64 `config:"node_id" env:"APP_NODE_ID"` - APIPrefix string `config:"api_prefix" env:"APP_API_PREFIX"` - GracefulShutdownTimeout int `config:"graceful_shutdown_timeout" env:"APP_GRACEFUL_SHUTDOWN_TIMEOUT"` - SessionCookieName string `config:"session_cookie_name" env:"APP_SESSION_COOKIE_NAME"` - SessionSecret string `config:"session_secret" env:"APP_SESSION_SECRET" secret:"true"` - SessionDomain string `config:"session_domain" env:"APP_SESSION_DOMAIN"` - SessionAge int `config:"session_age" env:"APP_SESSION_AGE" default:"86400"` - SessionHTTPOnly bool `config:"session_http_only" env:"APP_SESSION_HTTP_ONLY"` - SessionSecure bool `config:"session_secure" env:"APP_SESSION_SECURE"` -} - -type engineDatabaseConfig struct { - Enabled bool `config:"enabled" env:"DB_ENABLED" default:"false" autoEnable:"DB_HOST"` - SQLitePath string `config:"sqlite_path" env:"SQLITE_PATH"` - Host string `config:"host" env:"DB_HOST"` - Port int `config:"port" env:"DB_PORT"` - Username string `config:"username" env:"DB_USERNAME"` - Password string `config:"password" env:"DB_PASSWORD" secret:"true"` - Database string `config:"database" env:"DB_NAME"` - MaxIdleConn int `config:"max_idle_conn" env:"DB_MAX_IDLE_CONN"` - MaxOpenConn int `config:"max_open_conn" env:"DB_MAX_OPEN_CONN"` - ConnMaxLifetime int `config:"conn_max_lifetime" env:"DB_CONN_MAX_LIFETIME"` - ConnMaxIdleTime int `config:"conn_max_idle_time" env:"DB_CONN_MAX_IDLE_TIME"` - LogLevel string `config:"log_level" env:"DB_LOG_LEVEL"` - SSLMode string `config:"ssl_mode" env:"DB_SSL_MODE"` - TimeZone string `config:"time_zone" env:"DB_TIMEZONE"` - ApplicationName string `config:"application_name" env:"DB_APPLICATION_NAME"` - SearchPath string `config:"search_path" env:"DB_SEARCH_PATH"` - PreferSimpleProtocol bool `config:"prefer_simple_protocol" env:"DB_PREFER_SIMPLE_PROTOCOL"` - StatementCacheCapacity int `config:"statement_cache_capacity" env:"DB_STATEMENT_CACHE_CAPACITY"` - DefaultQueryExecMode string `config:"default_query_exec_mode" env:"DB_DEFAULT_QUERY_EXEC_MODE"` - Replicas []engineReplicaConfig `config:"replicas"` - SlowThreshold time.Duration `config:"slow_threshold" env:"DB_SLOW_THRESHOLD"` -} - -type engineRedisConfig struct { - Enabled bool `config:"enabled" env:"REDIS_ENABLED" default:"false" autoEnable:"REDIS_ADDR"` - Addrs []string `config:"addrs" env:"REDIS_ADDR"` - Username string `config:"username" env:"REDIS_USERNAME"` - Password string `config:"password" env:"REDIS_PASSWORD" secret:"true"` - DB int `config:"db" env:"REDIS_DB"` - ClusterMode bool `config:"cluster_mode" env:"REDIS_CLUSTER_MODE"` - MasterName string `config:"master_name" env:"REDIS_MASTER_NAME"` - KeyPrefix string `config:"key_prefix" env:"REDIS_KEY_PREFIX"` - PoolSize int `config:"pool_size" env:"REDIS_POOL_SIZE"` - MinIdleConn int `config:"min_idle_conn" env:"REDIS_MIN_IDLE_CONN"` - DialTimeout int `config:"dial_timeout" env:"REDIS_DIAL_TIMEOUT"` - ReadTimeout int `config:"read_timeout" env:"REDIS_READ_TIMEOUT"` - WriteTimeout int `config:"write_timeout" env:"REDIS_WRITE_TIMEOUT"` - MaxRetries int `config:"max_retries" env:"REDIS_MAX_RETRIES"` - PoolTimeout int `config:"pool_timeout" env:"REDIS_POOL_TIMEOUT"` - ConnMaxIdleTime int `config:"conn_max_idle_time" env:"REDIS_CONN_MAX_IDLE_TIME"` - MaintNotifications bool `config:"maint_notifications" env:"REDIS_MAINT_NOTIFICATIONS"` -} - -type engineClickHouseConfig struct { - Enabled bool `config:"enabled" env:"CLICKHOUSE_ENABLED" default:"false" autoEnable:"CLICKHOUSE_HOST"` - Hosts []string `config:"hosts" env:"CLICKHOUSE_HOST"` - Username string `config:"username" env:"CLICKHOUSE_USERNAME"` - Password string `config:"password" env:"CLICKHOUSE_PASSWORD" secret:"true"` - Database string `config:"database" env:"CLICKHOUSE_NAME"` - MaxIdleConn int `config:"max_idle_conn" env:"CLICKHOUSE_MAX_IDLE_CONN"` - MaxOpenConn int `config:"max_open_conn" env:"CLICKHOUSE_MAX_OPEN_CONN"` - ConnMaxLifetime int `config:"conn_max_lifetime" env:"CLICKHOUSE_CONN_MAX_LIFETIME"` - DialTimeout int `config:"dial_timeout" env:"CLICKHOUSE_DIAL_TIMEOUT"` - BlockBufferSize uint8 `config:"block_buffer_size" env:"CLICKHOUSE_BLOCK_BUFFER_SIZE"` -} - -type engineLogConfig struct { - Level string `config:"level" env:"LOG_LEVEL"` - Format string `config:"format" env:"LOG_FORMAT"` - Output string `config:"output" env:"LOG_OUTPUT"` - FilePath string `config:"file_path" env:"LOG_FILE_PATH"` - MaxSize int `config:"max_size" env:"LOG_MAX_SIZE"` - MaxAge int `config:"max_age" env:"LOG_MAX_AGE"` - MaxBackups int `config:"max_backups" env:"LOG_MAX_BACKUPS"` - Compress bool `config:"compress" env:"LOG_COMPRESS"` -} - -type engineOtelConfig struct { - SamplingRate float64 `config:"sampling_rate" env:"OTEL_SAMPLING_RATE"` - TracerName string `config:"tracer_name" env:"OTEL_TRACER_NAME" default:"github.com/Rain-kl/Wavelet"` -} - -// engineReplicaConfig mirrors databaseReplicaConfig, a composite element of database.replicas. -type engineReplicaConfig struct { - Host string `config:"host"` - Port int `config:"port"` - Username string `config:"username"` - Password string `config:"password"` -} - -// engineQueueConfig and engineWorkerConfig mirror the worker section, whose defaults -// legitimately move to driver_asynq_worker in P3; the repository file declares them -// explicitly, so parity is unaffected. -type engineQueueConfig struct { - Name string `config:"name"` - Priority int `config:"priority"` -} - -type engineWorkerConfig struct { - Concurrency int `config:"concurrency" env:"WORKER_CONCURRENCY"` - StrictPriority bool `config:"strict_priority" env:"WORKER_STRICT_PRIORITY"` - Queues []engineQueueConfig `config:"queues"` -} - -func repositoryConfig(t *testing.T) string { - t.Helper() - - path := filepath.Join("..", "..", "config.yaml") - if _, statErr := os.Stat(path); statErr != nil { - t.Skip("repository root config.yaml is unavailable") - } - return path -} - -// flatten exports a struct into dotted leaf paths rendered as text. Both sides of the -// parity assertion use distinct Go types for the same shape, so values are compared -// textually instead of handing cmp a cross-type diff. -func flatten(prefix string, v reflect.Value, out map[string]string) { - t := v.Type() - - for i := 0; i < t.NumField(); i++ { - field := t.Field(i) - if field.PkgPath != "" { - continue - } - - fv := v.Field(i) - path := prefix + "." + field.Name - if fv.Kind() == reflect.Struct && fv.Type() != durationType { - flatten(path, fv, out) - continue - } - out[path] = fmt.Sprint(fv.Interface()) - } -} - -// durationType mirrors the engine's own notion of a scalar duration field. -var durationType = reflect.TypeFor[time.Duration]() - -func TestEngineParityWithLegacyLoader(t *testing.T) { - path := repositoryConfig(t) - - scenarios := []struct { - name string - env map[string]string - }{ - {name: "file only", env: nil}, - { - name: "implicit enable from hosts", - env: map[string]string{ - "DB_HOST": "postgres", "REDIS_ADDR": "redis:6379", "CLICKHOUSE_HOST": "ch:9000", - }, - }, - { - name: "explicit flags win over implicit enable", - env: map[string]string{ - "DB_HOST": "postgres", "DB_ENABLED": "false", - "REDIS_ADDR": "redis:6379", "REDIS_ENABLED": "false", - "CLICKHOUSE_HOST": "ch:9000", "CLICKHOUSE_ENABLED": "false", - }, - }, - { - name: "scalar overrides and duration parsing", - env: map[string]string{ - "LOG_LEVEL": "debug", "APP_ADDR": ":9999", "DB_SLOW_THRESHOLD": "1s", - "REDIS_MAINT_NOTIFICATIONS": "true", "OTEL_SAMPLING_RATE": "0.5", - }, - }, - } - - for _, scenario := range scenarios { - t.Run(scenario.name, func(t *testing.T) { - for name, value := range scenario.env { - t.Setenv(name, value) - } - - legacy := load(path, false) - - src := newYAMLSource(t, path) - - engine := extpoints.NewConfigRegistry(src) - require.NoError(t, engine.Declare("parity", - extpoints.ConfigBinding{Prefix: "app", Target: &engineAppConfig{}}, - extpoints.ConfigBinding{Prefix: "database", Target: &engineDatabaseConfig{}}, - extpoints.ConfigBinding{Prefix: "redis", Target: &engineRedisConfig{}}, - extpoints.ConfigBinding{Prefix: "clickhouse", Target: &engineClickHouseConfig{}}, - extpoints.ConfigBinding{Prefix: "log", Target: &engineLogConfig{}}, - extpoints.ConfigBinding{Prefix: "otel", Target: &engineOtelConfig{}}, - extpoints.ConfigBinding{Prefix: "worker", Target: &engineWorkerConfig{}}, - )) - require.NoError(t, engine.Resolve()) - - var app engineAppConfig - var database engineDatabaseConfig - var redis engineRedisConfig - var clickhouse engineClickHouseConfig - var log engineLogConfig - var otel engineOtelConfig - var worker engineWorkerConfig - for _, binding := range []struct { - prefix string - target any - }{ - {"app", &app}, {"database", &database}, {"redis", &redis}, - {"clickhouse", &clickhouse}, {"log", &log}, {"otel", &otel}, {"worker", &worker}, - } { - require.NoError(t, engine.Bind(binding.prefix, binding.target)) - } - - // The legacy loader keeps two code-level defaults outside its tags; the engine - // expresses them as declared defaults, so normalise before diffing (spec C1). - if legacy.App.SessionAge <= 0 { - legacy.App.SessionAge = 86400 - } - if legacy.Otel.TracerName == "" { - legacy.Otel.TracerName = "github.com/Rain-kl/Wavelet" - } - - legacyFlat := map[string]string{} - flatten("app", reflect.ValueOf(legacy.App), legacyFlat) - flatten("database", reflect.ValueOf(legacy.Database), legacyFlat) - flatten("redis", reflect.ValueOf(legacy.Redis), legacyFlat) - flatten("clickhouse", reflect.ValueOf(legacy.ClickHouse), legacyFlat) - flatten("log", reflect.ValueOf(legacy.Log), legacyFlat) - flatten("otel", reflect.ValueOf(legacy.Otel), legacyFlat) - flatten("worker", reflect.ValueOf(legacy.Worker), legacyFlat) - - engineFlat := map[string]string{} - flatten("app", reflect.ValueOf(app), engineFlat) - flatten("database", reflect.ValueOf(database), engineFlat) - flatten("redis", reflect.ValueOf(redis), engineFlat) - flatten("clickhouse", reflect.ValueOf(clickhouse), engineFlat) - flatten("log", reflect.ValueOf(log), engineFlat) - flatten("otel", reflect.ValueOf(otel), engineFlat) - flatten("worker", reflect.ValueOf(worker), engineFlat) - - assert.Empty(t, cmp.Diff(legacyFlat, engineFlat), "engine resolution drifted from legacy loader") - }) - } -} -``` - -测试文件的 import 块必须包含:`fmt`、`os`、`path/filepath`、`reflect`、`testing`、`time`、`github.com/google/go-cmp/cmp`、`github.com/spf13/viper`、`github.com/stretchr/testify/assert`、`github.com/stretchr/testify/require`、`Wavelet/core/extpoints`。 - -- [ ] **Step 4: 运行对拍确认等价** - -Run: `cd backend && go test ./pkg/config/ -run TestEngineParityWithLegacyLoader -v` -Expected: 四个场景全部 `--- PASS`,无任何 `drifted` 断言输出。若出现 drift,逐项核对是否为 spec §4.3 登记的 C1–C5 有意差异:是则在该场景补注释说明差异来源,否则按 drift 修复引擎。 - -- [ ] **Step 5: 确认既有测试与构建仍全绿** - -Run: `cd backend && go test ./... && go build -o /dev/null ./...` -Expected: 全绿。 - -- [ ] **Step 6: 提交** - -```bash -git add backend/pkg/config/ -git commit -m "refactor(config): make legacy loader reentrant and add engine parity test" -``` - ---- - -## Task 9: 架构门禁脚本禁止 core 引入 viper - -**Files:** -- Modify: `scripts/check_cordis_architecture.sh:57-65` - -- [ ] **Step 1: 扩展检查项** - -把 `CORE_FRAMEWORK_IMPORTS` 的匹配式加入 viper、mapstructure 与 gin 的兄弟框架,使配置装载依赖无法渗进内核: - -```bash -CORE_FRAMEWORK_IMPORTS=$(rg -n '"github.com/gin-gonic/gin"|"gorm.io/gorm"|"github.com/hibiken/asynq"|"github.com/robfig/cron|"github.com/spf13/viper"|"github.com/mitchellh/mapstructure"' \ - "${BACKEND_DIR}/core/" --glob '*.go' -g '!*contracts*' -g '!*_test.go' || true) -``` - -同步更新失败提示文案,列出新增的两个包: - -```bash - log_fail "backend/core/ 严禁导入具体 Web/ORM/Worker/Config 运行时框架 (gin, gorm, asynq, cron, viper, mapstructure):" -``` - -- [ ] **Step 2: 运行脚本确认通过** - -Run: `./scripts/check_cordis_architecture.sh` -Expected: `✓ backend/core/ 无重型框架依赖`,末行 `✓ 所有 Cordis 架构合规性检查全部通过 (0 Violations)!`,退出码 0。 - -- [ ] **Step 3: 反向验证检查有效性** - -Run: `printf 'package core\n\nimport _ "github.com/spf13/viper"\n' > backend/core/zz_probe_tmp.go && ./scripts/check_cordis_architecture.sh; STATUS=$?; rm backend/core/zz_probe_tmp.go; exit $STATUS` -Expected: 报 `✗ [FAIL] backend/core/ 严禁导入具体 Web/ORM/Worker/Config 运行时框架`,退出码非 0。(临时文件必须删除,用后确认 `git status --short` 干净。) - -- [ ] **Step 4: 提交** - -```bash -git add scripts/check_cordis_architecture.sh -git commit -m "chore(arch): forbid viper and mapstructure inside the micro-kernel" -``` - ---- - -## Task 10: 收尾验证与文档同步 - -**Files:** -- Modify: `AGENTS.md`(Cordis 分层清单补配置扩展点条目) -- Modify: `docs/WAVELET_WHITE_PAPER.md`(§5.1 顶级分包说明) - -- [ ] **Step 1: 写内核侧用法文档** - -在 `AGENTS.md` 的“严格遵循事项 (Guardrails)”中,`扩展点自包含注册` 列表内 `动态配置` 一条之后补充: - -```markdown - - **静态配置声明**:插件在 `Apply` 中通过 `ctx.Config().Bind("", &cfg)` 读取自己声明的配置(字段用 `config` / `env` / `default` / `autoEnable` / `secret` tag 声明);需要在 `Apply` 之前被门禁求值的键,必须在 `DeclareConfig()` 中提前声明。**严禁**新增全局配置单例或在 `backend/pkg/` 读取配置。 -``` - -- [ ] **Step 2: 修正白皮书漂移** - -在 `docs/WAVELET_WHITE_PAPER.md` §5.1 的顶级分包说明后追加一条: - -```markdown -- **配置所有权下沉**:`backend/pkg/config` 全局单例已废除。`core/extpoints` 只提供配置声明与解析引擎(不 import viper),`plugins/infra/config` 承担文件与环境装载,读哪些字段由各插件自行声明;组合根不再跨插件判断配置选实现,改由 `ConfigGatedPlugin` 门禁 + `FiberSkipped` 决定激活方。 -``` - -- [ ] **Step 3: 全量门禁** - -Run: `make code-check` -Expected: 架构脚本 0 Violations;`golangci-lint run` 无输出;前端 tsc 与 eslint 无新增错误(本计划未触碰前端,若报既有错误需确认为基线问题)。 - -- [ ] **Step 4: 格式化** - -Run: `make format` -Expected: gofumpt 无 diff 或自动格式化;`git status --short` 中出现的文件需一并纳入下一步。 - -- [ ] **Step 5: 实跑确认应用行为未变** - -Run: `cd backend && go run main.go all 2>&1 | head -40` -Expected: banner 正常输出、迁移日志正常、`[Config] loaded configuration` 出现,进程正常启动后 Ctrl-C 优雅退出。**这是 P1/P2 的验收线:新能力已就位,生产路径仍走旧单例,行为必须与改动前一致。** - -- [ ] **Step 6: 提交** - -```bash -git add AGENTS.md docs/WAVELET_WHITE_PAPER.md -git commit -m "docs(config): record the configuration extension point ownership rules" -``` - ---- - -## 完成标准(本计划) - -1. `cd backend && go test ./...` 与 `go build ./...` 全绿;`make code-check` 与 `./scripts/check_cordis_architecture.sh` 零违规。 -2. 引擎单测覆盖:优先级四档、`autoEnable` 与显式 env 的相对优先、标量 env→切片、`time.Duration`、结构体切片、冲突校验、脱敏导出、无 source 时的错误路径。 -3. `TestEngineParityWithLegacyLoader` 四个场景零 drift,证明除 spec §4.3 的 C1–C5 外解析结果与旧实现逐 key 等价。 -4. 同时挂载门禁谓词相反的两个插件时,恰好一个 `FiberActive`、一个 `FiberSkipped`,且被跳过者 `Apply` 从未执行。 -5. `core/` 不出现 viper/mapstructure import,且架构脚本能主动拦截该违规。 -6. 生产启动路径行为未变(仍由 `config.Config` 供值),业务插件与 `cmd` 零改动。 - ---- - -## 与 spec 的偏差(实施时须回写 spec) - -计划编写阶段的自审发现三处与本 spec 已批准版本不一致的实现,均属实现期发现的正确性问题。在 Task 10 落库时一并回写 `docs/superpowers/specs/2026-08-29-cordis-config-extension-design.md`: - -| # | spec 原述 | 本计划实现 | 理由 | -| :--- | :--- | :--- | :--- | -| S1 | §4.3 C1 与 §6 把 `app.session_age<=0` 的判定列为内核解析错误(`ErrConfigInvalid`) | 引擎不做值域校验;`ErrConfigInvalid` 保留为源级校验占位,值域由声明者在 `Bind` 之后自行校验(auth 校验 `SessionAge > 0` 并使 `Apply` 失败) | 引擎被设计成不认识任何业务 key 的语义,让它知道"session_age 必须为正"会破坏该不变式并把业务规则塞进内核 | -| S2 | §3.2 `ConfigView` 含 `Source(key) string` | 更名 `Origin(key)`,并新增 `Value(key) (any, bool)`;`ConfigExtension` 新增 `SetSource(src)` 与 `Resolved()` | `Source` 与类型名 `ConfigSource` 在同文件内易混淆;`Value` 是 `core.ConfigGet[T]` 的支撑(Go 方法不能带类型参数);`SetSource` 进接口以避免 `WithConfigSource` 里的运行时类型断言 | -| S3 | §3.5 只给出 `WithShutdownTimeout` 与 `app.Prepare()` | 额外新增 `App.ShutdownTimeout()` 读取器与 `SetShutdownTimeout(d) *App`(链式,与既有 `WithProfile` 风格一致) | 组合根需要在 `Prepare()` 之后把已解析的预算写回内核,原先只有构造期选项,无法表达该顺序 | - -回写时同步修正 §7.3 的分期编号:本计划覆盖 P1 + P2,P3 + P4 由后续计划承接(`pkg/idgen` 解耦、27 个文件迁移、删除 `backend/pkg/config`)。 diff --git a/docs/superpowers/specs/2026-08-27-cordis-downstream-developer-guide.md b/docs/superpowers/specs/2026-08-27-cordis-downstream-developer-guide.md deleted file mode 100644 index 50af64b9..00000000 --- a/docs/superpowers/specs/2026-08-27-cordis-downstream-developer-guide.md +++ /dev/null @@ -1,594 +0,0 @@ -# Wavelet Cordis 插件化架构实战开发指南与标准规范 - -- **文档类型**: 下游开发者手册 / 架构实战指南 (Cookbook & Architecture Reference) -- **目标受众**: 官方插件开发者、下游业务二开工程师、架构师 -- **版本**: v1.0.0 (2026-08-27) - ---- - -# 目录 -- [第一部分:下游项目实战开发指南与 22 个高频开发场景解答](#第一部分下游项目实战开发指南与-22-个高频开发场景解答) - - [场景 1:插件必须要实现哪些方法与契约?](#场景-1插件必须要实现哪些方法与契约) - - [场景 2:插件间如何进行单向服务调用?](#场景-2插件间如何进行单向服务调用) - - [场景 3:插件间存在双向/循环调用时如何解决(杜绝 import cycle)?](#场景-3插件间存在双向循环调用时如何解决杜绝-import-cycle) - - [场景 4:如何开发并注册一个 HTTP API 接口?如何添加路由中间件?](#场景-4如何开发并注册一个-http-api-接口如何添加路由中间件) - - [场景 5:如何获取当前登录用户信息?](#场景-5如何获取当前登录用户信息) - - [场景 6:如何开发并注册一个 Asynq 异步 Worker 任务?](#场景-6如何开发并注册一个-asynq-异步-worker-任务) - - [场景 7:如何开发并注册一个 Cron 定时任务?](#场景-7如何开发并注册一个-cron-定时任务) - - [场景 8:数据库表结构如何声明?ORM 模型规范是什么?](#场景-8数据库表结构如何声明orm-模型规范是什么) - - [场景 9:数据库如何做独立迁移?Goose SQL 怎么组织?](#场景-9数据库如何做独立迁移goose-sql-怎么组织) - - [场景 10:如果有多个业务插件需要读写同一张表怎么办?](#场景-10如果有多个业务插件需要读写同一张表怎么办) - - [场景 11:如果跨插件操作多张表,如何确保事务一致性?](#场景-11如果跨插件操作多张表如何确保事务一致性) - - [场景 12:如何发布和订阅领域事件 (EventBus)?](#场景-12如何发布和订阅领域事件-eventbus) - - [场景 13:如何向系统注册插件自定义配置(config.yaml 与管理台热加载设置)?](#场景-13如何向系统注册插件自定义配置configyaml-与管理台热加载设置) - - [场景 14:如何使用多层缓存(RAM L1 + Redis L2 + PubSub 同步)?](#场景-14如何使用多层缓存ram-l1--redis-l2--pubsub-同步) - - [场景 15:如何使用分布式锁 (DistLock) 防止并发超卖与重复消费?](#场景-15如何使用分布式锁-distlock-防止并发超卖与重复消费) - - [场景 16:如何向管理后台动态注册监控数据与管理控制台?](#场景-16如何向管理后台动态注册监控数据与管理控制台) - - [场景 17:插件如何实现健康检查探针与就绪检查 (Health Check)?](#场景-17插件如何实现健康检查探针与就绪检查-health-check) - - [场景 18:插件如何扩展其他插件的能力(如新增一种 OAuth 登录提供商 / 新增消息推送渠道)?](#场景-18插件如何扩展其他插件的能力如新增一种-oauth-登录提供商--新增消息推送渠道) - - [场景 19:插件如何编写单元测试与集成测试(Mock 上下文与依赖打桩)?](#场景-19插件如何编写单元测试与集成测试mock-上下文与依赖打桩) - - [场景 20:以不同角色(api / worker / schedule / all)启动时,插件代码如何适配?](#场景-20以不同角色api--worker--schedule--all启动时插件代码如何适配) - - [场景 21:当某个插件流量暴增需要独立拆分为微服务时,如何零成本平滑改造?](#场景-21当某个插件流量暴增需要独立拆分为微服务时如何零成本平滑改造) - - [场景 22:插件如何安全处理文件上传与大文件摄取 (upload.Ingest)?](#场景-22插件如何安全处理文件上传与大文件摄取-uploadingest) -- [第二部分:整个项目的目录结构划分与包职责定义](#第二部分整个项目的目录结构划分与包职责定义) -- [第三部分:框架核心提供给插件调用的公用能力矩阵 (Context Capability Matrix)](#第三部分框架核心提供给插件调用的公用能力矩阵-context-capability-matrix) - ---- - -# 第一部分:下游项目实战开发指南与 22 个高频开发场景解答 - -### 场景 1:插件必须要实现哪些方法与契约? -每个插件必须实现 `core.Plugin` 接口,仅需提供两个核心方法:`Name()` 与 `Apply(ctx *core.Context)`。 - -```go -package myplugin - -import "github.com/Rain-kl/Wavelet/core" - -type Plugin struct{} - -// 1. Name: 返回全局唯一的插件标识符(建议遵循命名空间规范,如 "biz.order") -func (p *Plugin) Name() string { - return "biz.order" -} - -// 2. Apply: 核心装载入口,所有的路由注册、任务注册、服务提供与依赖消费均在此完成 -func (p *Plugin) Apply(ctx *core.Context) error { - // 在此编写装载逻辑 - return nil -} -``` - ---- - -### 场景 2:插件间如何进行单向服务调用? -**规则**:插件之间**禁止直接相互 import 具体实现包**。调用方仅面向 `core/contracts` 中的纯 Interface 编程,运行时通过 Context 解析。 - -```go -// 1. 插件 A (提供者 plugins/user) 将服务注入 Context -func (p *UserPlugin) Apply(ctx *core.Context) error { - userSvc := NewUserServiceImpl(ctx.DB()) - ctx.Provide[contracts.UserService](userSvc) - return nil -} - -// 2. 插件 B (消费者 plugins/order) 声明依赖并调用 -func (p *OrderPlugin) Apply(ctx *core.Context) error { - return ctx.Using(func(userSvc contracts.UserService) { - // userSvc 已由容器自动注入就绪 - v1 := ctx.Router().Group("/api/v1/orders") - v1.POST("", func(c *gin.Context) { - userInfo, err := userSvc.GetUserProfile(c.Request.Context(), "user_123") - // 处理订单逻辑... - }) - }) -} -``` - ---- - -### 场景 3:插件间存在双向/循环调用时如何解决(杜绝 import cycle)? -**问题场景**:`auth` 登录成功后需要查 `user` 资料;`user` 重置密码后需要调 `auth` 吊销 session。若两个 package 互相 import,Go 编译器会报 `import cycle not allowed`。 - -**Cordis 解法**: -1. 接口均定义在 `core/contracts`,双方只依赖 `core/contracts`。 -2. 运行时采用 **延迟注入 (Lazy Resolution / Inject)** 或 **事件解耦 (EventBus)**: - -```go -// plugins/auth/service.go -func (s *AuthServiceImpl) OnLoginSuccess(c context.Context, uid string) { - // 延迟注入 UserService,不发生 package 级循环导入 - userSvc, err := core.Inject[contracts.UserService](s.ctx) - if err == nil { - userSvc.UpdateLastLoginTime(c, uid) - } -} -``` -*更加推荐的方式是发射领域事件*(见场景 12),由 `user` 插件自愿监听,彻底消除相互调用的硬依赖。 - ---- - -### 场景 4:如何开发并注册一个 HTTP API 接口?如何添加路由中间件? -插件通过 `ctx.Router()` 声明路由。微内核支持标准 Gin 路由组与中间件挂载: - -```go -func (p *OrderPlugin) Apply(ctx *core.Context) error { - // 获取全局或 auth 插件提供的中间件 - authSvc, _ := core.Inject[contracts.AuthService](ctx) - - // 创建带版本前缀和鉴权中间件的路由组 - group := ctx.Router().Group("/api/v1/orders", authSvc.RequireAuthMiddleware()) - - // 注册 Handler - group.GET("", p.handleListOrders) - group.POST("", p.handleCreateOrder) - group.GET("/:id", p.handleGetOrderDetail) - - return nil -} -``` - ---- - -### 场景 5:如何获取当前登录用户信息? -`auth` 插件会在上下文中注入当前用户 Session。业务 Handler 可直接调用统一 Helper: - -```go -func (p *OrderPlugin) handleCreateOrder(c *gin.Context) { - // 1. 从当前 Gin 请求上下文中提取认证用户信息 - currentUser, ok := oauth.GetCurrentUser(c) - if !ok { - response.AbortUnauthorized(c, errs.ErrUnauthorized) - return - } - - log.Printf("当前下单用户 ID: %s, 权限角色: %s", currentUser.ID, currentUser.Role) - // 2. 正常业务处理... -} -``` - ---- - -### 场景 6:如何开发并注册一个 Asynq 异步 Worker 任务? -```go -func (p *OrderPlugin) Apply(ctx *core.Context) error { - // 1. 注册 Asynq 任务类型与消费处理器 - ctx.Task().Register("order:cancel_timeout", p.handleTimeoutCancelTask) - return nil -} - -// 2. 任务执行函数 -func (p *OrderPlugin) handleTimeoutCancelTask(ctx context.Context, t *asynq.Task) error { - var payload OrderTimeoutPayload - if err := json.Unmarshal(t.Payload(), &payload); err != nil { - return err - } - // 执行超时关单业务逻辑... - return nil -} - -// 3. 业务中异步投递任务 -func (p *OrderPlugin) EnqueueTimeoutCheck(ctx context.Context, orderID string) { - p.ctx.TaskClient().EnqueueContext(ctx, asynq.NewTask("order:cancel_timeout", payloadBytes), asynq.ProcessIn(15*time.Minute)) -} -``` - ---- - -### 场景 7:如何开发并注册一个 Cron 定时任务? -```go -func (p *ReportPlugin) Apply(ctx *core.Context) error { - // 每天凌晨 2 点执行日报汇总任务 - ctx.Schedule().RegisterCron("0 2 * * *", "report:daily_summary", DailyReportPayload{Type: "all"}) - return nil -} -``` - ---- - -### 场景 8:数据库表结构如何声明?ORM 模型规范是什么? -**规范**: -1. 表名必须带有插件专有前缀(如 `w_order_`、`w_auth_`),避免跨插件表名冲突。 -2. 零值与数据库默认值严格对齐;禁止物理外键,显式建索引。 -3. 必须通过 GORM 结构体清晰声明 `gorm:"..."` 标签与 `json:"..."`。 - -```go -package models - -import "time" - -type Order struct { - ID string `gorm:"column:id;primaryKey;size:64" json:"id"` - UserID string `gorm:"column:user_id;index;size:64;not null" json:"user_id"` - Amount int64 `gorm:"column:amount;not null" json:"amount"` - Status string `gorm:"column:status;size:32;index;not null;default:'pending'" json:"status"` - CreatedAt time.Time `gorm:"column:created_at;autoCreateTime" json:"created_at"` - UpdatedAt time.Time `gorm:"column:updated_at;autoUpdateTime" json:"updated_at"` - DeletedAt *time.Time `gorm:"column:deleted_at;index" json:"-"` -} - -func (Order) TableName() string { - return "w_orders" -} -``` - ---- - -### 场景 9:数据库如何做独立迁移?Goose SQL 怎么组织? -**彻底告别集中大迁移目录**。每个插件在内部目录建立 `migrations/`,并通过 `//go:embed` 打包注入: - -```go -// plugins/order/plugin.go -package order - -import ( - "embed" - "github.com/Rain-kl/Wavelet/core" -) - -//go:embed migrations/*.sql -var orderMigrations embed.FS - -func (p *Plugin) Apply(ctx *core.Context) error { - // 注册本插件的专属迁移(系统启动时自动按版本号执行) - ctx.Migrations().Register("order", orderMigrations) - return nil -} -``` - -#### SQL 迁移脚本规范 (`plugins/order/migrations/00001_initial.sql`): - -每个插件只需维护一个 `00001_initial.sql`,包含该插件的全部建表语句与种子数据。 - -```sql --- +goose Up --- +goose StatementBegin -CREATE TABLE IF NOT EXISTS w_orders ( - id VARCHAR(64) PRIMARY KEY, - user_id VARCHAR(64) NOT NULL, - amount BIGINT NOT NULL, - status VARCHAR(32) NOT NULL DEFAULT 'pending', - created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, - updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP -); -CREATE INDEX IF NOT EXISTS idx_w_orders_user_id ON w_orders(user_id); - --- 种子数据(使用 ON CONFLICT DO NOTHING 保证幂等) -INSERT INTO w_orders (id, user_id, amount, status, created_at, updated_at) -VALUES ('init_001', 'system', 0, 'completed', CURRENT_TIMESTAMP, CURRENT_TIMESTAMP) -ON CONFLICT (id) DO NOTHING; --- +goose StatementEnd - --- +goose Down --- +goose StatementBegin -DROP TABLE IF EXISTS w_orders; --- +goose StatementEnd -``` - -#### 版本管理机制 - -所有插件共享一张 `w_schema_versions` 表,以 `plugin_id` 区分: - -``` -w_schema_versions (plugin_id, version_id, applied_at) -``` - -启动时,引擎遍历每个插件: -1. 查询 `w_schema_versions WHERE plugin_id = 'order'` 获取当前最大版本号 -2. 扫描插件 `migrations/` 目录下的 `.sql` 文件 -3. 如果存在未应用的版本号 → 执行迁移 -4. 如果全部已应用 → 跳过 - -```sql --- 查看全局迁移状态 -SELECT * FROM w_schema_versions ORDER BY plugin_id, version_id; -``` - ---- - -### 场景 10:如果有多个业务插件需要读写同一张表怎么办? -**黄金准则**:**表有且仅有一个所有者插件 (Single Owner Principle)**。 -* 严禁插件 B 直接通过 SQL 修改插件 A 拥有的核心表(如订单插件直接修改用户表)。 -* **合法模式 1(服务调用)**:插件 A 提供 `UserService.DeductBalance(uid, amount)`,插件 B 调用该接口。 -* **合法模式 2(只读视图 / 共享查询 DTO)**:如果仅仅是高频联合查询(报表),插件 A 暴露只读查询接口,或通过数据库只读从库直接投影。 - ---- - -### 场景 11:如果跨插件操作多张表,如何确保事务一致性? -在插件化和微服务就绪体系下,**跨插件的强分布式事务是反模式**。 - -1. **同插件内多表操作**:直接使用本地数据库事务: - ```go - err := ctx.DB().Transaction(func(tx *gorm.DB) error { - if err := tx.Create(&order).Error; err != nil { return err } - if err := tx.Create(&orderItem).Error; err != nil { return err } - return nil - }) - ``` -2. **跨插件操作(如创建订单 + 扣减库存 + 发送通知)**: - * 采用 **最终一致性 (Eventual Consistency / Saga 模式)**。 - * 本地事务成功后,发射 `OrderCreatedEvent` 到 EventBus; - * 库存插件监听到事件后扣减库存,若失败则发布补偿事件触发订单取消。 - ---- - -### 场景 12:如何发布和订阅领域事件 (EventBus)? -```go -// 1. 定义强类型事件结构 -type OrderPaidEvent struct { - OrderID string `json:"order_id"` - UserID string `json:"user_id"` - PayAmount int64 `json:"pay_amount"` -} - -// 2. 插件 A 发布事件 -ctx.Events().Emit("order:paid", OrderPaidEvent{OrderID: "ord_1", UserID: "u_1", PayAmount: 9900}) - -// 3. 插件 B 订阅事件 -ctx.Events().On("order:paid", func(c context.Context, e OrderPaidEvent) error { - log.Printf("收到支付成功事件,开始为用户 %s 发放权益", e.UserID) - return nil -}) -``` - ---- - -### 场景 13:如何向系统注册插件自定义配置(config.yaml 与管理台热加载设置)? -```go -type OrderConfig struct { - MaxItemsPerOrder int `yaml:"max_items" json:"max_items"` - AutoCancelMins int `yaml:"auto_cancel_mins" json:"auto_cancel_mins"` -} - -func (p *OrderPlugin) Apply(ctx *core.Context) error { - var cfg OrderConfig - // 1. 自动从 config.yaml 中的 plugins.order 节点绑定配置 - ctx.Config().Bind("plugins.order", &cfg) - - // 2. 注册为管理台可动态修改的系统参数 - ctx.Settings().Register(core.SettingSchema{ - Key: "order.auto_cancel_mins", - Default: 15, - Description: "未支付订单自动取消时间 (分钟)", - }) - return nil -} -``` - ---- - -### 场景 14:如何使用多层缓存(RAM L1 + Redis L2 + PubSub 同步)? -框架提供三层穿透缓存能力,防止缓存击穿与雪崩: - -```go -func (s *OrderService) GetOrderWithCache(ctx context.Context, orderID string) (*Order, error) { - var order Order - err := s.ctx.Cache().GetOrSet(ctx, "order:"+orderID, &order, 10*time.Minute, func() (any, error) { - // Cache Miss 回源查 DB - var dbOrder Order - if err := s.db.WithContext(ctx).First(&dbOrder, "id = ?", orderID).Error; err != nil { - return nil, err - } - return &dbOrder, nil - }) - return &order, err -} - -// 当订单更新时,广播失效所有节点的 L1 内存缓存与 L2 Redis 缓存 -func (s *OrderService) InvalidateCache(ctx context.Context, orderID string) { - s.ctx.Cache().Delete(ctx, "order:"+orderID) -} -``` - ---- - -### 场景 15:如何使用分布式锁 (DistLock) 防止并发超卖与重复消费? -```go -func (s *OrderService) ProcessPayment(ctx context.Context, orderID string) error { - // 获取分布式锁,租期 5 秒 - unlock, err := s.ctx.DistLock().Lock(ctx, "lock:order:pay:"+orderID, 5*time.Second) - if err != nil { - return fmt.Errorf("当前订单正在处理中,请勿重复提交") - } - defer unlock() // 确保释放 - - // 执行扣款操作... - return nil -} -``` - ---- - -### 场景 16:如何向管理后台动态注册监控数据与管理控制台? -插件可以向管理后台扩展点注入自己的仪表盘指标和诊断探针: - -```go -func (p *OrderPlugin) Apply(ctx *core.Context) error { - ctx.Admin().RegisterMetric("order_count_today", func(c context.Context) any { - var count int64 - ctx.DB().Model(&models.Order{}).Where("created_at >= ?", todayStart()).Count(&count) - return count - }) - return nil -} -``` - ---- - -### 场景 17:插件如何实现健康检查探针与就绪检查 (Health Check)? -```go -func (p *PaymentPlugin) Apply(ctx *core.Context) error { - ctx.Health().RegisterProbe("payment_gateway", func(ctx context.Context) error { - // 测试第三方支付网关网络连通性 - return pingPaymentGateway(ctx) - }) - return nil -} -``` - ---- - -### 场景 18:插件如何扩展其他插件的能力(如新增一种 OAuth 登录提供商 / 新增消息推送渠道)? -采用 **注册表扩展点模式 (Registry Pattern)**: - -```go -// 1. 下游编写微信登录插件 plugins/oauth_wechat -func (p *WeChatOAuthPlugin) Apply(ctx *core.Context) error { - return ctx.Using(func(authRegistry contracts.AuthRegistry) { - // 向核心 auth 插件注入微信 OAuth 实现 - authRegistry.RegisterOAuthProvider("wechat", &WeChatProvider{...}) - }) -} -``` - ---- - -### 场景 19:插件如何编写单元测试与集成测试(Mock 上下文与依赖打桩)? -微内核提供轻量测试脚手架 `coretest`: - -```go -func TestOrderCreate(t *testing.T) { - // 1. 创建内存测试专用 Context - ctx := coretest.NewMockContext(t) - - // 2. Mock 依赖的 UserService - mockUserSvc := &MockUserService{ReturnUser: &contracts.UserDTO{ID: "u_1", Balance: 1000}} - ctx.Provide[contracts.UserService](mockUserSvc) - - // 3. 装载插件 - plugin := &OrderPlugin{} - require.NoError(t, plugin.Apply(ctx)) - - // 4. 发起 HTTP 接口测试 - w := ctx.PerformRequest("POST", "/api/v1/orders", `{"item_id":"item_1"}`) - assert.Equal(t, 200, w.Code) -} -``` - ---- - -### 场景 20:以不同角色(api / worker / schedule / all)启动时,插件代码如何适配? -**开发者无需做任何特殊处理**! -插件只需在一个 `Apply` 方法中把自己的路由、任务、调度全部注册进 `Context`。微内核调度器会根据运行命令自动按需激活对应的运行时驱动,不匹配的能力保持休眠。 - ---- - -### 场景 21:当某个插件流量暴增需要独立拆分为微服务时,如何零成本平滑改造? -```go -// 1. 之前单体模式:在 main.go 中加载本地实现 -app.Use(&auth.Plugin{}) // 进程内直接运行 - -// 2. 拆分为微服务后:只需将 main.go 替换为 gRPC 客户端代理插件! -app.Use(&auth_grpc_client.Plugin{RemoteAddr: "auth-service.prod:9000"}) - -// 3. 所有依赖 auth 的业务插件(如 order, user)业务代码 0 处修改! -``` - ---- - -### 场景 22:插件如何安全处理文件上传与大文件摄取 (upload.Ingest)? -**严格规则**:禁止插件自行直接写入对象存储底层 Bucket 或直连底层文件系统。统一走平台摄取服务: - -```go -func (p *OrderPlugin) handleUploadInvoice(c *gin.Context) { - fileHeader, _ := c.FormFile("file") - - // 使用平台统一摄取引擎(自动计算哈希、防重传、生成签名 URL 与入库追踪) - ingestResult, err := upload.IngestFormFile(c.Request.Context(), fileHeader, upload.IngestPolicy{ - AllowedTypes: []string{"image/png", "application/pdf"}, - MaxSizeBytes: 10 * 1024 * 1024, - }) - if err != nil { - response.AbortBadRequest(c, errs.ErrUploadFailed) - return - } - - c.JSON(200, response.OK(gin.H{"file_url": ingestResult.URL})) -} -``` - ---- - -# 第二部分:整个项目的目录结构划分与包职责定义 - -```text -Wavelet/ -├── cmd/ # CLI 命令分发与装配入口 -│ ├── root.go # Cobra 根命令 -│ ├── server.go # 综合启动器(支持 api/worker/schedule/all profile) -│ └── migrate.go # 数据库独立迁移命令行工具 -│ -├── core/ # 【微内核引擎 (Zero Business Logic)】 -│ ├── context.go # Context 上下文总线与 Fork 树 -│ ├── container.go # 基于泛型的 IoC 服务注册与解析器 -│ ├── events.go # 强类型领域事件总线 (EventBus) -│ ├── lifecycle.go # 启动/停止生命周期编排状态机 -│ ├── contracts/ # 【跨插件标准服务契约 (纯 Interface)】 -│ │ ├── auth.go # AuthService 契约 -│ │ ├── user.go # UserService 契约 -│ │ ├── cache.go # CacheService 契约 -│ │ └── database.go # DBService 契约 -│ └── extpoints/ # 扩展点定义 (Router, Task, Migration, Setting) -│ -├── plugins/ # 【官方标准插件库 (完全高内聚闭包)】 -│ ├── drivers/ # 运行时驱动插件 -│ │ ├── driver_http/ # Gin Web HTTP 驱动 -│ │ ├── driver_asynq_worker/ # Asynq Worker 并发消费驱动 -│ │ └── driver_asynq_cron/ # Asynq Cron 调度器驱动 -│ │ -│ ├── infra/ # 基础设施服务插件 -│ │ ├── database/ # GORM 多数据源与读写分离插件 -│ │ ├── cache/ # RAM + Redis + PubSub 缓存插件 -│ │ ├── logger/ # Zap + Otel 分布式链路追踪日志插件 -│ │ └── storage/ # S3 / OSS / Local 对象存储插件 -│ │ -│ └── domain/ # 业务领域能力插件 -│ ├── auth/ # OAuth / Session / Passkey 认证插件 -│ ├── user/ # 用户资料 / 权限 / 角色插件 -│ ├── message_gateway/ # Bot 网关 / 渠道推送插件 -│ ├── risk_control/ # 访问控制 / IP 限流 / 安全风控插件 -│ └── admin/ # 系统管理台与监控面板插件 -│ -└── downstream/ # 【下游二开项目模板与脚手架】 - ├── custom_plugins/ # 下游自定义业务插件 - ├── config.yaml # 声明启用的插件与配置文件 - └── main.go # 下游项目组合启动入口 -``` - -### 各层职责与禁止规则 (Guardrails): -1. **`core/`**: - - **职责**:纯抽象,提供 IoC、Context、EventBus 和 Lifecycle。 - - **严禁**:严禁 import 任何具体业务包,严禁 import `gin`、`gorm`、`asynq`。 -2. **`core/contracts/`**: - - **职责**:仅定义公开的 Go Interface 和公共 DTO。 - - **严禁**:严禁包含任何具体实现逻辑或 SQL 操作。 -3. **`plugins/`**: - - **职责**:所有业务逻辑和驱动实现的归宿。遵循标准分层架构(Layered Architecture / MVC 变体)。 - - **分层模式选型**: - - **模式 1(极简单文件分层,微型插件专用)**:单 package 极简结构(仅单文件 `plugin.go`, `handlers.go`, `service.go`, `repository.go`, `models.go`, `errs.go`, `migrations/`)。适用于单一实体、极小代码量 (<500行) 的微型插件。 - - **模式 2(标准独立子包分层架构,官方推荐标准)**:按职责严格物理分包(`plugin.go`, `handler/`, `service/`, `repository/`, `model/`, `errs/`, `migrations/`)。**子包内文件以纯业务实体命名(如 `user.go`、`config.go`),严禁在根包平铺 `handlers_*`、`service_*`、`repository_*` 等前缀文件**。编译器级强约束 `handler -> service -> repository -> model` 单向依赖。 - - **严禁**:插件之间严禁跨包 import 内部私有代码,跨插件调用一律走 `contracts` 接口或 `EventBus`。 - ---- - -# 第三部分:框架核心提供给插件调用的公用能力矩阵 (Context Capability Matrix) - -每个插件在 `Apply(ctx *core.Context)` 时,都可以无缝调用微内核暴露的以下标准能力: - -| 扩展点方法 | 返回类型 | 功能说明 | 适用场景 | -| :--- | :--- | :--- | :--- | -| `ctx.Router()` | `RouterExtension` | 声明 HTTP 路由、前缀分组与挂载中间件 | 暴露 API 接口、Web 控制台 | -| `ctx.Task()` | `TaskExtension` | 注册 Asynq 异步任务消费处理器 | 耗时后台任务、异步消息发送 | -| `ctx.Schedule()` | `ScheduleExtension`| 注册 Cron 定时调度任务 | 定时报表统计、周期性清理 | -| `ctx.Migrations()` | `MigrationExtension`| 注册插件专属的 Goose SQL 迁移嵌入系统 | 自建数据表、版本升级 | -| `ctx.Events()` | `EventBus` | 强类型领域事件的发布与订阅 (Emit / On) | 跨插件完全解耦通知与状态同步 | -| `ctx.Settings()` | `SettingExtension` | 声明动态可配置项(支持热更新) | 业务参数配置、管理台可调节参数 | -| `ctx.DB()` | `*gorm.DB` | 获取全局受事务与 Trace 保护的 GORM 数据源 | 数据持久化 CRUD | -| `ctx.Cache()` | `CacheService` | 三层穿透缓存(RAM L1 + Redis L2 + PubSub 广播)| 高频读数据性能加速 | -| `ctx.DistLock()` | `DistLockService` | 基于 Redis 的工业级分布式锁 | 防并发超卖、防重复执行 | -| `ctx.Logger()` | `Logger` | 携带链路 TraceID 的结构化日志记录器 | 业务日志打印与审计 | -| `ctx.Storage()` | `StorageService` | 统一对象存储读写引擎 | 文件摄取、图片持久化 | -| `core.Provide[T]`| `void` | 向全局 IoC 容器注册本插件提供的强类型服务 | 暴露自身能力给其他插件消费 | -| `core.Inject[T]` | `(T, error)` | 从全局 IoC 容器中按类型获取服务实例 | 消费其他插件暴露的服务 | -| `core.Using[T]` | `error` | 响应式声明依赖,当服务就绪时执行回调 | 声明前置依赖关系 | - diff --git a/docs/superpowers/specs/2026-08-27-cordis-plugin-architecture-design.md b/docs/superpowers/specs/2026-08-27-cordis-plugin-architecture-design.md deleted file mode 100644 index ab6b30cb..00000000 --- a/docs/superpowers/specs/2026-08-27-cordis-plugin-architecture-design.md +++ /dev/null @@ -1,279 +0,0 @@ -# Wavelet Cordis 微内核与全插件化架构设计规范 - -- **创建日期**: 2026-08-27 -- **状态**: Approved Design -- **架构代号**: Cordis-Wavelet (Next 5-Year Foundation) - ---- - -## 1. 背景与目标 - -### 1.1 现状与痛点 -Wavelet 当前采用中心化显式装配架构(`internal/platform/bootstrap` 与 `internal/router`),业务逻辑集中在 `internal/apps/` 下。 -随着业务功能的快速拓展,现有架构暴露出以下瓶颈: -1. **模块高耦合**:新增功能需要横跨多个中心化目录(`apps/`、`router/`、`bootstrap/`、`migrator/`、`task/handlers/`)进行插桩,难以做到随插随用与物理隔离。 -2. **下游扩展困难**:二次开发项目无法在不修改核心源码的前提下灵活扩展或替换业务模块。 -3. **缺乏清晰的运行切面**:API、Worker、Scheduler 启动模式依赖手动条件判断,维护成本高。 - -### 1.2 改造核心目标 -1. **微内核 (Micro-Kernel)**:内核仅提供上下文总线(Context)、依赖注入(IoC)、生命周期状态机与扩展点协议,内核本身零具体业务依赖。 -2. **一切皆插件 (All-in-Plugins)**:数据库、缓存、日志、HTTP 服务、任务处理、认证鉴权、消息网关及业务能力全部以插件形式挂载在 Context 上。 -3. **下游一等公民支持**:下游项目通过声明式 `app.Use(&MyPlugin{})` 引入官方或自定义插件,编译为单一高性能二进制文件。 -4. **面向未来 5 年的分布式与微服务就绪 (Monolith-First, Microservice-Ready)**:基于强类型接口契约,单体模式下零开销内存调用,高并发下支持透明替换为 gRPC/RPC 客户端插件完成微服务拆分。 - ---- - -## 2. 核心架构模型 (Core Architecture) - -``` -+-----------------------------------------------------------------------------------+ -| 下游业务项目 (Downstream Application) | -| main.go: app.Use(&logger.Plugin{}).Use(&auth.Plugin{})... | -+-----------------------------------------------------------------------------------+ - │ - ▼ -+-----------------------------------------------------------------------------------+ -| Wavelet Core (微内核上下文总线) | -| - Context (服务树与扩展点总线) - Lifecycle Manager (生命周期编排) | -| - Service Hub (泛型 IoC 容器) - EventBus (强类型领域事件总线) | -+-----------------------------------------------------------------------------------+ - │ │ - ▼ 注册与驱动 ▼ 挂载能力 -+------------------------------------+ +-------------------------------------------+ -| 运行时驱动插件 (Driver Plugins) | | 业务领域插件 (Domain Plugins) | -| - driver-http (Gin Web 引擎) | | - plugin-auth (认证/Session/OAuth) | -| - driver-worker (Asynq 消费池) | | - plugin-user (用户资料/角色权限) | -| - driver-cron (Asynq 定时调度器) | | - plugin-msg-gateway (消息通道与推送) | -| - driver-database (GORM 数据源) | | - plugin-risk-control (访问风控与限流) | -| - driver-cache (RAM/Redis 缓存) | | - [下游自定义插件] (业务私有插件) | -+------------------------------------+ +-------------------------------------------+ -``` - ---- - -## 3. 微内核协议契约与设计规范 - -### 3.1 插件契约 (`core.Plugin`) -所有官方插件与下游自定义插件均实现统一的 `Plugin` 接口: - -```go -package core - -import "context" - -// Plugin 插件统一契约 -type Plugin interface { - // Name 插件唯一标识(如 "auth", "database", "message_gateway") - Name() string - // Apply 核心装载入口:通过 Context 提供服务、注册路由、声明任务与监听事件 - Apply(ctx *Context) error -} -``` - -### 3.2 运行时驱动契约 (`core.Driver`) -HTTP 服务、Worker 消费池、Cron 调度器不硬编码在内核中,而是作为标准 `Driver` 挂载: - -```go -package core - -type DriverType string - -const ( - DriverTypeHTTP DriverType = "http" - DriverTypeWorker DriverType = "worker" - DriverTypeScheduler DriverType = "schedule" -) - -// Driver 是具备事件循环或监听端口的运行时引擎 -type Driver interface { - Type() DriverType - Start(ctx context.Context) error - Stop(ctx context.Context) error -} -``` - -### 3.3 Context 统一服务总线与泛型注入 -```go -package core - -// Provide 向 Context 注册强类型服务实现 -func Provide[T any](ctx *Context, service T) - -// Inject 从 Context 获取已注册的服务 -func Inject[T any](ctx *Context) (T, error) - -// Using 声明式依赖注入(当且仅当依赖的服务全部就绪时激活回调) -func Using[T1 any](ctx *Context, fn func(s1 T1)) error -func Using2[T1, T2 any](ctx *Context, fn func(s1 T1, s2 T2)) error -``` - ---- - -## 4. 领域扩展点规范 (Domain Extension Points) - -微内核提供 6 大标准扩展点,供插件高内聚地声明自己的资源: - -### 4.1 HTTP 路由扩展 (`ctx.Router()`) -```go -type RouterExtension interface { - Group(relativePath string, handlers ...gin.HandlerFunc) *gin.RouterGroup - Use(middleware ...gin.HandlerFunc) -} -``` - -### 4.2 数据迁移扩展 (`ctx.Migrations()`) -每个插件通过 Go 内置 `embed.FS` 打包专属的 Goose SQL 文件,彻底消除单体大迁移目录的合并冲突: - -```go -type MigrationExtension interface { - // Register 注册插件专属的 SQL 迁移文件系统 - Register(pluginID string, fsys fs.FS, dir ...string) -} -``` - -**版本隔离机制**:所有插件共享一张 `w_schema_versions` 表,以 `plugin_id` 列区分。运行时引擎(`gooseEngine`)实现 `goosedb.Store` 接口,对该表执行 `plugin_id` 限定的 CRUD 操作,确保各插件的版本互不干扰。 - -```sql -w_schema_versions ( - plugin_id VARCHAR(64) NOT NULL, - version_id BIGINT NOT NULL, - applied_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP, - PRIMARY KEY (plugin_id, version_id) -) -``` - -**启动流程**: -1. `ApplyPlugins()` 阶段:各插件调用 `ctx.Migrations().Register("order", embedFS)` 收集迁移 -2. `RunMigrations()` 阶段:引擎遍历所有 `entries`,为每个插件创建 `goose.NewProvider(dialect, sqlDB, entry.FS, goose.WithStore(store))` -3. `provider.Up()` 查询 `w_schema_versions WHERE plugin_id = 'order'` 决定版本,执行增量迁移 - -### 4.3 异步任务与定时调度扩展 (`ctx.Task()` & `ctx.Schedule()`) -```go -type TaskExtension interface { - Register(taskType string, handler asynq.HandlerFunc) -} - -type ScheduleExtension interface { - RegisterCron(spec string, taskType string, payload any) -} -``` - -### 4.4 领域事件总线 (`ctx.Events()`) -用于跨插件完全解耦通信,单机模式走内存通道,集群模式无缝升级为 Redis Stream / NATS: -```go -type EventBus interface { - On(topic string, handler any) - Emit(topic string, payload any) error -} -``` - -### 4.5 动态系统设置扩展 (`ctx.Settings()`) -```go -type SettingExtension interface { - RegisterSchema(pluginID string, schema any) -} -``` - ---- - -## 5. 插件形态与目录布局规范 - -插件结构遵循 **“扁平 (Flat)、自包含 (Self-Contained)、就近组织 (Colocated)”** 原则,杜绝不必要的 DDD 样板代码。 - -### 5.1 官方插件目录结构 -```text -plugins/ -├── database/ # 数据库驱动插件 -│ ├── plugin.go # 注册 DBService 与连接池 -│ └── service.go -├── auth/ # 认证插件 -│ ├── plugin.go # 插件装载入口:ctx.Provide[AuthService] + 路由挂载 -│ ├── service.go # AuthService 接口实现 (登录/Token/Session) -│ ├── handlers.go # HTTP Controller -│ ├── models.go # GORM 实体定义 -│ └── migrations/ # 专属 Goose SQL 迁移 -│ └── 001_auth_init.sql -├── message_gateway/ # 消息网关插件 -│ ├── plugin.go # 路由挂载 + Worker 任务注册 -│ ├── channels.go # Telegram / QQ / Webhook 各渠道实现 -│ └── models.go -└── [下游自定义插件]/ # 下游业务方自研插件 - ├── plugin.go - └── models.go -``` - ---- - -## 6. 插件间引用关系与协同规范 - -为杜绝 Go 语言的 `import cycle not allowed` 错误并保持插件的独立可替换性,插件间交互严格遵循以下 3 大模式: - -1. **服务槽位与延迟绑定(用于跨插件直接调用)**: - * 双方互不 import 对方包,仅面向 `core/contracts` 暴露的 Interface 编程。 - * 运行时通过 `core.Inject[contracts.UserService](ctx)` 获取服务。 -2. **事件总线广播(用于通知与状态联动)**: - * 登录成功、密码修改、订单创建等事件统一通过 `ctx.Events().Emit()` 广播,下游自愿监听。 -3. **注册表扩展点模式(用于功能插件扩充主插件能力)**: - * 主插件向 Context 提供注册表(如 `OAuthProviderRegistry`),扩充插件在 `Apply` 中向注册表添加自己的 Provider 实现。 - ---- - -## 7. 运行切面与启动路径 (Runtime Profiles) - -CLI 命令仅作为**切面激活器 (Target Selector)**,业务插件无需感知当前的运行角色: - -``` -[CLI: wavelet api / worker / schedule / all] - ↓ -1. App Bootstrap: 加载所有已配置插件并构建 Context - ↓ -2. Apply Phase: 执行所有 plugin.Apply(ctx),收集路由、任务、调度与迁移 - └─ 各插件调用 ctx.Migrations().Register("auth", authMigrations) 等 - ↓ -3. Migration Engine: 遍历所有 entries,逐插件创建 Goose Provider 执行迁移 - └─ 每个插件使用独立的 sharedStore(pluginID),共享同一张 w_schema_versions 表 - └─ provider.Up() 检查 w_schema_versions WHERE plugin_id = 'auth' - └─ 未执行过 → 执行 00001_initial.sql → INSERT 版本记录 - └─ 已执行过 → 跳过 - ↓ -4. Profile Dispatch: - - "api": 激活 DriverTypeHTTP 驱动 (Gin.ListenAndServe) - - "worker": 激活 DriverTypeWorker 驱动 (Asynq.Run) - - "schedule": 激活 DriverTypeScheduler 驱动 (Asynq.Scheduler) - - "all": 激活所有 Driver 实例 (单体一键融合启动) - ↓ -5. Graceful Shutdown: 监听系统信号,逆序安全停机 -``` - ---- - -## 8. 面向未来 5 年的分布式与服务拆分演进 - -```mermaid -graph LR - subgraph Monolith ["阶段 1:单体插件化 (进程内零开销)"] - UserP["plugin-user"] -->|Go Interface 内存调用| AuthP["plugin-auth"] - end - - subgraph Distributed ["阶段 2:高并发微服务拆分 (透明代理替换)"] - UserP2["plugin-user"] -->|相同的 Go Interface| AuthClient["plugin-auth-client (gRPC 代理)"] - AuthClient -.->|gRPC / HTTP/2| RemoteAuth["独立 Auth 微服务集群"] - end -``` - -1. **接口不变性 (Contract Stability)**:所有跨模块调用走 Interface,微服务化拆分时只需引入 RPC 客户端插件替换原插件,调用方业务代码 **0 修改**。 -2. **分布式事件驱动**:进程内 EventBus 通过简单配置可无缝切换为 Redis Stream / NATS / Kafka。 -3. **独立数据分片**:每个插件表名自带命名空间(如 `w_auth_*`),且有独立 Migration,天然支持物理分库分表。 - ---- - -## 9. 渐进式改造实施路线图 - -1. **Phase 1: 微内核基础设施搭建 (`core/` & `core/contracts/`)** - - 实现 Context、泛型 IoC 容器、生命周期状态机与 6 大扩展点协议。 -2. **Phase 2: 运行时驱动插件下沉 (`plugins/driver_*`)** - - 将现有 Gin、Asynq Worker、Asynq Scheduler、GORM、Redis 封装为标准 Driver 插件。 -3. **Phase 3: 官方领域模块插件化拆分 (`plugins/domain_*`)** - - 依次将 `auth`、`user`、`message_gateway`、`risk_control`、`admin` 迁移为标准插件。 -4. **Phase 4: 下游工程脚手架与验证** - - 提供下游开发模板,编写示例自定义插件,端到端验证 API/Worker/Schedule 运行切面与测试覆盖。 diff --git a/docs/superpowers/specs/2026-08-28-cordis-architecture-alignment-design.md b/docs/superpowers/specs/2026-08-28-cordis-architecture-alignment-design.md deleted file mode 100644 index b772e045..00000000 --- a/docs/superpowers/specs/2026-08-28-cordis-architecture-alignment-design.md +++ /dev/null @@ -1,82 +0,0 @@ -# Cordis Architecture Alignment & Refactoring Design - -**Date**: 2026-08-28 -**Topic**: Cordis Meta-framework Alignment (Spatiotemporal Composability, Revertible Effects, Reactive Coeffects & Boundary Isolation) -**Status**: Approved - ---- - -## 1. Background & Objectives - -Wavelet adopts the **Cordis** micro-kernel paradigm (originating from Koishi and DeepSeek Harness) to achieve runtime composability and zero-side-effect lifecycle management. -According to the formal metatheory of Cordis (*A Programming Paradigm for Spatiotemporal Composability*), the runtime must satisfy two orthogonal requirements: -1. **Temporal Composability (时间可组合性)**: Every context mutation/registration must track an inverse operation (Revertible Effects) and automatically roll back in LIFO order upon unloading/disposing. -2. **Spatial Composability (空间可组合性)**: Components declare required coeffects/dependencies (`inject`); when dependencies become available or unavailable, the system reactively activates or deactivates components (Fiber state machine), guaranteeing **Confluence (合流)** regardless of registration order. -3. **Context as the Sole Surface & Defensive Isolation**: Eliminate cross-plugin private imports and global static singletons (`database.DB()`, global configs), strictly enforcing single-owner boundaries and `contracts` programming. - ---- - -## 2. Architecture & Detailed Design - -### 2.1 Revertible Effects & Scoped Extpoints (时间可组合性) - -- **`Context` Scoped Lifetime**: - Each plugin instance is mounted with a dedicated child context `pluginCtx := rootCtx.Fork()`. -- **Automatic Disposer Registration for Extpoints**: - When registrations occur through `pluginCtx`, inverse operations are automatically pushed to `pluginCtx`'s Disposer stack: - - **`Router`**: Registering a route returns a definition with an ID; `pluginCtx` records a disposer that calls `router.UnregisterByID(id)`. - - **`Events`**: `ctx.Events().On(...)` returns a `Disposer`; when called on a scoped context (or via `ctx.On(...)`), it binds to `pluginCtx.OnDispose`. - - **`Tasks`**: Registering an async task binds `tasks.Unregister(taskType)` to `pluginCtx.OnDispose`. - - **`Schedules`**: Registering a cron schedule binds `schedules.Unregister(cronName)` to `pluginCtx.OnDispose`. - - **`Settings`**: Registering setting schemas binds schema deregistration to `pluginCtx.OnDispose`. - - **`Container (Provide)`**: Providing a service type `T` binds `container.remove(T)` to `pluginCtx.OnDispose`. -- **LIFO Teardown Guarantee**: - Calling `pluginCtx.Dispose()` runs all registered disposers in reverse order (LIFO), cleanly revoking routes, event listeners, tasks, schedules, and service bindings without residual side effects. - ---- - -### 2.2 Reactive Coeffects & Fiber Lifecycle (空间可组合性) - -- **Dependency Declaration (`DependentPlugin`)**: - Plugins can optionally implement: - ```go - type DependentPlugin interface { - Plugin - Inject() []reflect.Type - } - ``` -- **Plugin Fiber State Machine**: - ``` - PENDING ──(All dependencies provided)──> LOADING ──(Apply succeeds)──> ACTIVE - ▲ │ - └─────────────(Dependency removed / Plugin unloaded)──────────────────────┘ - ``` - - **States**: `FiberPending`, `FiberLoading`, `FiberActive`, `FiberUnloading`, `FiberDisposed`. - - **Reconciler**: When `core.Provide[T]` registers a service or `core.App.Use` registers a plugin, the reconciler checks all pending fibers. Fibers with satisfied dependencies transition `Pending -> Loading -> Active`. - - **Confluence**: Plugin registration order (`app.Use(A, B)` vs `app.Use(B, A)`) produces the exact same final active state once all dependencies are satisfied. - ---- - -### 2.3 Boundary Defense & Single Owner Enforcement (架构防线) - -- **Eliminate Direct Global Invocations**: - - Refactor `plugins/domain/user/repository.go` and other domain repositories to avoid direct `import "Wavelet/plugins/infra/database"` and direct calls to `database.DB(ctx)`. - - Inject `contracts.DBService` via repository struct or retrieve via `ctx.DB()`. -- **Strict Package Separation**: - - `backend/core/`: Micro-kernel, context, container, fiber, events, scoped extpoints. - - `backend/core/contracts/`: Public interfaces and shared DTOs/events. - - `backend/plugins/infra/`: Infrastructure implementations providing contracts services. - - `backend/plugins/drivers/`: Runtime drivers (HTTP, Asynq Worker, Cron). - - `backend/plugins/domain/`: Domain business logic and single-owner tables. - - `backend/pkg/`: Stateless utilities and algorithm libraries. - ---- - -## 3. Verification Plan - -1. **Unit Tests for Core**: - - `core/fiber_test.go`: Test fiber state transitions, out-of-order registration confluence, and dynamic unloading. - - `core/context_test.go` & `core/extpoints/`: Test automatic scoped disposer tracking for routes, tasks, schedules, and event listeners. -2. **Refactoring Verification for Domain Plugins**: - - Run `go test ./backend/...` across all domain and infra packages. - - Run `make code-check` and verify zero lint regressions. diff --git a/docs/superpowers/specs/2026-08-28-cordis-architecture-refactor-design.md b/docs/superpowers/specs/2026-08-28-cordis-architecture-refactor-design.md deleted file mode 100644 index 7ac5e106..00000000 --- a/docs/superpowers/specs/2026-08-28-cordis-architecture-refactor-design.md +++ /dev/null @@ -1,84 +0,0 @@ -# Cordis 架构重构设计规格书 (Cordis Architecture Refactor Design) - -**日期**: 2026-08-28 -**目标**: 依据 Cordis 时空可组合性元框架(Spatiotemporal Composability)哲学,重构 Wavelet 后端包结构、包职责与插件边界,消除全局静态单例与跨插件私有实现依赖,实现真正的可逆副作用与契约化隔离。 - ---- - -## 1. 背景与核心设计原则 - -Cordis 是一个面向时空可组合性的元框架,核心在于: -1. **时间可组合性 (Temporal Composability / Revertible Effects)**:组件挂载到上下文时产生的任何副作用(数据库连接、Redis 客户端、路由、事件监听、定时任务)必须具备明确的逆操作,在卸载时按 LIFO(后进先出)干净撤销。 -2. **空间可组合性 (Spatial Composability / Reactive Coeffects)**:组件通过 `Inject` 声明依赖;无特权微内核,所有基础设施与业务均以平等插件形态存在;组件之间严格面向抽象服务契约(Contracts)编程,严禁跨包引用私有实现。 -3. **合流定理 (Confluence)**:任何插件的装载/卸载顺序,静止状态等同于从零静态装配,杜绝全局隐藏状态与启动顺序隐式假设。 - ---- - -## 2. 详细重构方案 - -### 2.1 微内核纯洁化 (`backend/core/`) - -#### 改造点: -1. **移除特权辅助方法**: - - 从 `backend/core/context.go` 中移除 `func (c *Context) DB() contracts.DBService` 与 `func (c *Context) Cache() contracts.CacheService`。 - - 所有服务消费方统一面向 `core.Inject[T](ctx)`、`core.MustInject[T](ctx)` 或 `core.Using[T](ctx, ...)`。 -2. **保持依赖注入纯粹性**: - - 内核仅保留:`Context`、`Container`、`Fiber`、`EventBus`、生命周期管理以及通用的扩展点挂载。 - ---- - -### 2.2 基础设施插件生命周期可逆化 (`backend/plugins/infra/`) - -#### 1. 数据库插件 (`plugins/infra/database`) -- **移除隐式副作用**: - - 删除 `postgres.go` 与 `sqlite.go` 中的 `func init() { ... }` 静态建连。 - - 删除包级导出的静态全局变量 `var db *gorm.DB` 以及全局 `DB(ctx)` / `SetDB()`。 -- **生命周期受控与可逆释放**: - - 在 `Plugin.Apply(ctx *core.Context)` 时根据配置建立数据库连接(GORM + underlying `*sql.DB`)。 - - 创建 `contracts.DBService` 实例并通过 `core.Provide[contracts.DBService](ctx, svc)` 注册。 - - 注册 `ctx.OnDispose` 逆操作,在插件卸载时调用 `sqlDB.Close()`。 - -#### 2. 缓存插件 (`plugins/infra/cache`) -- **移除隐式副作用**: - - 删除 `redis.go` 中的 `func init() { ... }` 静态建连。 - - 删除包级导出的全局变量 `var Redis redis.UniversalClient`。 -- **生命周期受控与可逆释放**: - - 在 `Plugin.Apply(ctx *core.Context)` 时初始化 Redis 客户端并构造 `contracts.CacheService`。 - - 通过 `core.Provide[contracts.CacheService](ctx, svc)` 注册。 - - 注册 `ctx.OnDispose` 逆操作,在插件卸载时调用 `client.Close()`。 - ---- - -### 2.3 业务领域插件防线隔离与依赖重构 (`backend/plugins/domain/`) - -#### 1. 消除跨插件私有 Import -- 遍历并重构以下 8 个 Domain 插件: - - `auth` - - `user` - - `admin` - - `cap` - - `message_gateway` - - `risk_control` - - `system` - - `upload` -- **规则**: - - 严禁任何 domain 插件 `import "Wavelet/plugins/infra/database"` 或 `import "Wavelet/plugins/infra/cache"`。 - - 严禁任何 domain 插件直接 import 另一个 domain 插件的具体实现包(如 `admin` 严禁 import `risk_control/logstore` 或 `storage/diskcache`)。 - - 各插件内部的 Repository / Service 统一通过 `core.Inject[contracts.DBService](ctx)` 或插件内部 scoped context 获取数据库连接。 - -#### 2. `admin` 插件解耦与全局变量清除 -- 移除 `admin/plugin.go` 中的包级变量(`globalUserSvc`, `globalAuthSvc`, `globalCoreCtx`)。 -- 将 `admin` 的日志查询、任务触发、缓存清理等管理接口改造为通过 `contracts` 或 `ctx.Tasks()` 访问,消除对 `risk_control`、`driver_asynq_worker` 等的私有依赖。 - ---- - -## 3. 验证与门禁标准 - -1. **编译与依赖检查**: - - 运行 `grep -r "Wavelet/plugins/infra/database" backend/plugins/domain/` 结果为空。 - - 运行 `grep -r "Wavelet/plugins/infra/cache" backend/plugins/domain/` 结果为空。 -2. **自动化测试**: - - 所有既有单元测试与集成测试(`go test ./...`)无回归,全部 PASS。 -3. **代码质量门禁**: - - `make code-check` 静态检查 0 告警通过。 - - `make format` 格式化通过。 diff --git a/docs/superpowers/specs/2026-08-28-cordis-plugin-layered-architecture-spec.md b/docs/superpowers/specs/2026-08-28-cordis-plugin-layered-architecture-spec.md deleted file mode 100644 index 5fd6dd9b..00000000 --- a/docs/superpowers/specs/2026-08-28-cordis-plugin-layered-architecture-spec.md +++ /dev/null @@ -1,140 +0,0 @@ -# Cordis 架构插件标准分层设计规范 (Plugin Layered Architecture Spec) - -- **文档状态**: 已敲定 (Approved) -- **版本**: v1.1.0 (2026-08-28) -- **适用范围**: Wavelet 官方插件 (`backend/plugins/`)、下游定制插件 (`downstream/custom_plugins/`) - ---- - -## 1. 架构总览与分型原则 (Architecture & Selection Strategy) - -在 Wavelet 的 Cordis 微内核架构中,系统通过 **微内核 (`core/`) + 服务契约 (`core/contracts/`) + 自包含插件 (`plugins/`)** 实现高度解耦与单向依赖。 -为了规范插件内部代码组织,插件遵循 **标准分层架构(Layered Architecture / MVC 变体)**,并根据业务复杂度提供两套标准物理包结构: - -| 模式 | 适用场景 | 复杂度特征 | 物理结构形式 | 命名规范核心禁令 | -| :--- | :--- | :--- | :--- | :--- | -| **模式 1:极简单文件自包含**
(Single-File Flat) | 极简微型插件 | 仅有 1 个单一实体、代码量 < 500 行(如极简工具、Demo) | 单 Package,每个层级仅对应 1 个同名文件 (`handlers.go`, `service.go`, `models.go`, `repository.go`) | **严禁在根目录平铺 `handlers_*`、`service_*` 等前缀文件** | -| **模式 2:独立子包分层架构**
(Strict Sub-packages) | 标准/中大型业务插件(**官方推荐标准**) | 包含多实体/多接口、代码量 ≥ 500 行(如 `upload`, `auth`, `admin`, `order` 等) | 严格按层独立子包 (`handler/`, `service/`, `repository/`, `model/`, `errs/`) | **子包内文件直接以业务命名(如 `user.go`, `config.go`),禁止带 `handler_*` / `service_*` 前缀** | - ---- - -## 2. 模式 1:极简单文件自包含规范 (Single-File Flat Package) - -仅适用于极简小型插件(整个插件代码极少且各层只有一个文件)。 - -### 2.1 目录结构 -```text -backend/plugins/domain// -├── plugin.go # [Cordis 接入层] 实现 core.Plugin,负责 Apply 组装与扩展点注册 -├── handlers.go # [Handler 层] 单一文件:Gin API Handler -├── service.go # [Service 层] 单一文件:核心业务用例 -├── repository.go # [Repository 层] 单一文件:GORM / DB 操作 -├── models.go # [Model 层] 单一文件:实体与 DTO -├── errs.go # [Error 层] 单一文件:错误常量 -├── plugin_test.go # 插件测试 -└── migrations/ # Goose SQL 嵌入文件 - └── 20260828000001_init_.sql -``` - -> ⚠️ **严禁规则**:当单一文件膨胀或需要拆分多个业务实体时,**严禁在根目录创建 `handlers_user.go`, `handlers_admin.go`, `service_user.go` 等前缀文件**,必须立即重构并迁移为 **模式 2(独立子包分层架构)**! - ---- - -## 3. 模式 2:标准独立子包分层架构 (Standard Sub-package Architecture - 推荐规范) - -适用于绝大多数业务插件。各层使用独立的 Go package 物理隔离,**在子包内以纯业务实体命名文件**。 - -### 3.1 目录结构与文件命名规约 -```text -backend/plugins/domain// -├── plugin.go # [插件根入口] 实现 core.Plugin,装配各子包并向 Cordis 注册 -│ -├── handler/ # package handler:HTTP API 接入层(或 controller/) -│ ├── router.go # 路由组挂载与中间件绑定 -│ ├── auth.go # 认证相关 Handler(直接命名为 auth.go,禁止 handlers_auth.go) -│ ├── user.go # 用户相关 Handler(直接命名为 user.go,禁止 handlers_user.go) -│ ├── config.go # 配置相关 Handler(直接命名为 config.go,禁止 handlers_config.go) -│ └── logs.go # 日志相关 Handler(直接命名为 logs.go,禁止 handlers_logs.go) -│ -├── service/ # package service:核心领域业务逻辑层 -│ ├── service.go # 顶层 Service 组合与构造工厂 -│ ├── auth.go # 认证业务逻辑(直接命名为 auth.go,禁止 service_auth.go) -│ ├── user.go # 用户业务逻辑(直接命名为 user.go,禁止 service_user.go) -│ ├── config.go # 配置业务逻辑(直接命名为 config.go,禁止 service_config.go) -│ └── logs.go # 日志业务逻辑(直接命名为 logs.go,禁止 service_logs.go) -│ -├── repository/ # package repository:数据访问持久化层 (DAL) -│ ├── repository.go # 仓储通用方法与工厂 -│ ├── user.go # 用户仓储实现(直接命名为 user.go,禁止 repository_user.go) -│ ├── config.go # 配置仓储实现(直接命名为 config.go,禁止 repository_config.go) -│ └── log.go # 日志仓储实现(直接命名为 log.go,禁止 repository_log.go) -│ -├── model/ # package model (或 models/):纯领域实体与传输对象 -│ ├── entity.go # 数据库映射实体 (TableName() 必须带 w__ 前缀) -│ ├── dto.go # 请求入参与响应出参 DTO -│ └── events.go # 插件内部/广播事件结构体定义 -│ -├── errs/ # package errs:错误常量与错误码定义 (或根目录 errs.go) -│ └── errs.go -│ -└── migrations/ # Goose SQL 独立迁移嵌入文件 (//go:embed) - └── 20260828000001_init_.sql -``` - -### 3.2 依赖方向约束 (Strict Dependency Flow) -```mermaid -graph TD - Plugin[plugin.go 入口] --> Handler[handler/ 接入层] - Plugin --> Service[service/ 业务层] - Plugin --> Repository[repository/ 仓储层] - Handler --> Service - Handler --> Model[model/ 实体与DTO] - Handler --> Errs[errs/ 错误常量] - Service --> Repository - Service --> Model - Service --> Errs - Repository --> Model -``` -* **单向依赖铁律**: - 1. `handler/` 依赖 `service/`、`model/`、`errs/`; - 2. `service/` 依赖 `repository/`、`model/`、`errs/`,**严禁 import gin**; - 3. `repository/` 依赖 `model/` 和数据库底层,**严禁反向依赖 service 或 handler**; - 4. `model/` 纯粹由 Go 结构体组成,**严禁依赖上层 handler/service/repository**。 - ---- - -## 4. 各层职责边界与编码守则 (Layer Responsibilities & Guardrails) - -### 4.1 Handler 层 (`handler/`) -1. **参数绑定**:使用 `c.ShouldBindJSON` 或 `c.ShouldBindQuery`。 -2. **上下文提取**:从 `*gin.Context` 提取登录态(如 `oauth.GetCurrentUser(c)`)。 -3. **调用下游**:调用 Service 方法,禁止直接调用 Repository 或编写 SQL。 -4. **统一信封响应**: - - 成功:`c.JSON(http.StatusOK, response.OK(data))` 或 `response.OKNil()`。 - - 失败:使用 `backend/pkg/response` 的 `Abort*` 系列函数(如 `AbortBadRequest`、`AbortUnauthorized`、`AbortNotFound`、`AbortInternal`)。 -5. **Swagger 注释**:每个导出 Handler 必须编写完整的 OpenAPI/Swagger 注解。 - -### 4.2 Service 层 (`service/`) -1. **纯 Go 逻辑**:第一参数必须为 `ctx context.Context`,返回 `(result, error)`。 -2. **禁止依赖 Web 框架**:严禁 import `github.com/gin-gonic/gin`,严禁接收 `*gin.Context`,严禁调用 `c.JSON`/`Abort*`。 -3. **事务编排**:涉及插件内多表原子操作时,通过 `ctx.DB().Transaction(...)` 编排。 -4. **事件驱动解耦**:跨插件业务通知与状态联动统一通过 `ctx.Events().Emit(...)` 广播领域事件,杜绝直接跨插件调用私有方法。 - -### 4.3 Repository 层 (`repository/`) -1. **GORM / SQL 操作**:统一接收 `context.Context`,通过 `db.WithContext(ctx)` 操作数据。 -2. **SQL LIKE 防注入**:所有含用户输入的模糊查询必须调用 `backend/pkg/util.EscapeLike` 并显式声明 `ESCAPE '\\'`。 -3. **表单一所有者原则**:仅操作本插件所属表(前缀 `w__*`),严禁越权 DML/DDL 其他插件所有表。 - -### 4.4 Model 层 (`model/` 或 `models/`) -1. **GORM 映射**:显式实现 `TableName() string` 返回带前缀表名。 -2. **零值对齐**:Go 结构体字段零值必须与数据库默认值匹配。 -3. **无物理外键**:禁止物理外键约束,显式建立单列/复合索引。 - -### 4.5 Plugin 入口 (`plugin.go`) -1. 实现 `core.Plugin` 接口(`Name() string` 与 `Apply(ctx *core.Context) error`)。 -2. 在 `Apply` 中完成: - - 依赖注入与解析(`core.Provide` / `core.Inject` / `ctx.Using`) - - 路由与中间件声明(`ctx.Router().Group(...)`) - - 异步与定时任务注册(`ctx.Task().Register` / `ctx.Schedule().RegisterCron`) - - 配置与设置声明(`ctx.Settings().Register` / `ctx.Config().Bind`) - - 数据库迁移注册(`ctx.Migrations().Register`) diff --git a/docs/superpowers/specs/2026-08-28-zero-redis-pluggable-architecture-design.md b/docs/superpowers/specs/2026-08-28-zero-redis-pluggable-architecture-design.md deleted file mode 100644 index b2fdfd7d..00000000 --- a/docs/superpowers/specs/2026-08-28-zero-redis-pluggable-architecture-design.md +++ /dev/null @@ -1,97 +0,0 @@ -# Zero-Redis Pluggable Architecture Design - -**Date**: 2026-08-28 -**Topic**: Decoupling Redis via Cordis Pluggable Infrastructure and In-Process Drivers (Zero-Redis Monolith Mode) -**Status**: Approved - ---- - -## 1. Background & Objectives - -Currently, the Wavelet platform has direct or indirect couplings with Redis across four areas: -1. **Cache Layer (`infra/cache`)**: Hardcoded initialization of Redis client and L2 cache lookup. -2. **Background Worker & Cron Drivers (`drivers/driver_asynq_*`)**: Asynq requires Redis as message queue and timer broker. -3. **Cross-Node Invalidation (Pub/Sub)**: Invalidation messages directly interact with Redis channels. -4. **Task Execution Log Stream**: Real-time worker output writes directly to Redis pipelines in `admin/repository.go`. - -**Goal**: -In accordance with Cordis's "Everything is a Plugin" and "Single Owner Principle", extract Redis into dedicated optional plugins and provide lightweight in-process equivalents (`cache_memory`, `driver_inproc_worker`, `driver_inproc_cron`) so that standalone monolith deployments, embedded scenarios, and local development can run with zero external Redis dependency. - ---- - -## 2. Architecture & Detailed Design - -### 2.1 Cache Infrastructure Split (`backend/plugins/infra/`) - -`contracts.CacheService` remains the sole contract for caching. Two alternative plugins implement this contract: - -1. **`plugins/infra/cache_memory` (Default for Monolith / Zero-Redis)**: - - Encapsulates `pkg/cache/ram` for fast in-process TTL caching. - - Cache invalidations emit `cache:invalidate` events via `ctx.Events()` locally. - - Provides `core.Provide[contracts.CacheService](ctx, memCacheSvc)`. -2. **`plugins/infra/cache_redis` (Distributed Cluster Mode)**: - - Provides full multi-tier caching: L1 Local RAM + L2 Remote Redis + Redis Pub/Sub invalidation. - - Implements `core.DependentPlugin` (declares dependencies on database / configuration). - - Provides `core.Provide[contracts.CacheService](ctx, redisCacheSvc)`. - ---- - -### 2.2 In-Process Worker & Scheduler Drivers (`backend/plugins/drivers/`) - -Domain plugins register tasks and schedules only against `ctx.Tasks()` and `ctx.Schedules()` extension points, completely oblivious to the underlying runner. - -1. **`plugins/drivers/driver_inproc_worker` (In-Process Worker Driver)**: - - Implements `core.Driver` with `Type() == core.DriverTypeWorker`. - - Maintains an in-memory buffered channel queue and worker goroutine pool (managed via `util.Go` with panic recovery). - - Supports task concurrency limits, exponential backoff retries, and context execution timeouts. -2. **`plugins/drivers/driver_inproc_cron` (In-Process Cron Scheduler Driver)**: - - Implements `core.Driver` with `Type() == core.DriverTypeScheduler`. - - Uses `robfig/cron/v3` to poll and trigger entries in `ctx.Schedules().Schedules()`. -3. **`plugins/drivers/driver_asynq_*` (Distributed Cluster Drivers)**: - - Retains Asynq worker and cron drivers for Redis-backed distributed workloads. - ---- - -### 2.3 Task Execution Log Stream & Event Bus Decoupling - -1. **Task Stream Logs**: - - Provide an in-memory `RingBuffer` (e.g. recent 500 lines per execution). - - When Redis is disabled, logs stream into the `RingBuffer` and flush to `w_task_executions` upon completion. -2. **System Config & Invalidation Broadcast**: - - In single-node mode, `ctx.Events()` in-process event bus handles all notifications immediately. - - In multi-node mode, `cache_redis` bridges events across instances via Redis Pub/Sub. - ---- - -### 2.4 Application Assembly (`backend/cmd/app.go`) - -In `cmd/app.go`, the application declaratively selects the plugin suite based on configuration: - -```go -if config.Config.Redis.Enabled { - app.Use( - cache_redis.New(), - driver_asynq_worker.New(), - driver_asynq_cron.New(), - ) -} else { - app.Use( - cache_memory.New(), - driver_inproc_worker.New(), - driver_inproc_cron.New(), - ) -} -``` - ---- - -## 3. Verification Plan - -1. **Unit Tests**: - - `plugins/infra/cache_memory/plugin_test.go`: Test in-memory cache operations, TTL expiry, and `contracts.CacheService` compliance. - - `plugins/drivers/driver_inproc_worker/plugin_test.go`: Test in-process task dispatch, concurrency, retry, and cancellation. - - `plugins/drivers/driver_inproc_cron/plugin_test.go`: Test in-process cron schedule execution. -2. **Integration Verification**: - - Verify that running the application with `config.Database.Enabled = false` and `config.Redis.Enabled = false` boots cleanly into `all`, `api`, `worker`, and `scheduler` profiles with zero connection errors. -3. **Quality Gate**: - - Run `go test ./...`, `make code-check`, and `make format`. diff --git a/docs/superpowers/specs/2026-08-29-cordis-config-extension-design.md b/docs/superpowers/specs/2026-08-29-cordis-config-extension-design.md deleted file mode 100644 index c2965834..00000000 --- a/docs/superpowers/specs/2026-08-29-cordis-config-extension-design.md +++ /dev/null @@ -1,364 +0,0 @@ -# Cordis 配置扩展点设计 (Config Extension Point) - -- **文档状态**: 已敲定 (Approved) -- **版本**: v1.0.0 (2026-08-29) -- **适用范围**: `backend/core/`(微内核)、`backend/plugins/`(自包含插件)、`backend/cmd/`(组合根)、`backend/pkg/`(无状态基础库) - ---- - -## 0. 背景与动机 - -`backend/pkg/config` 同时承担了三件事:viper 装载 `config.yaml`、环境变量覆盖、以及以全局单例 `config.Config` 暴露全量配置模型。它与架构文档对 `backend/pkg/` 的定位("Stateless utilities and algorithm libraries")冲突,并且带来两个结构性问题: - -1. **配置所有权倒挂**:任何包都能读到全量配置,因此 `cmd` 直接替 `cache` 插件判断 Redis 是否启用、`risk_control` 直接判断 `clickhouse.enabled`。配置的"读者"与"所有者"没有关系约束。 -2. **隐式全局状态**:`init()` 内完成文件搜索、解析与 `log.Fatalf`,并以 `isTest()` 猜测执行上下文来禁用数据库/Redis/ClickHouse;测试通过改写全局单例驱动生产代码路径。 - -本设计把"配置的读取框架"下沉为内核扩展点,把"读哪些字段"的所有权交给各插件自己声明,并一次性迁移全部 27 个消费文件(109 处引用),彻底删除全局单例。 - -`AGENTS.md` 与 `new-setting` skill 中早已写明插件应通过 `ctx.Config().Bind(...)` 绑定静态配置,但该 API 在代码中从未存在——本设计同时修正这一文档漂移。 - ---- - -## 1. 决策记录 - -| # | 决策 | 理由与取舍 | -| :--- | :--- | :--- | -| D1 | **预声明阶段 + 配置门禁** | 内核在 `Apply` 之前收集声明并求值门禁,使组合根不再跨插件读配置选实现。代价是给 Fiber 增加"被门禁跳过"语义。 | -| D2 | **混合读取形态:结构体 `Bind` + 泛型 `Get`** | `redis`/`database` 等 14+ 字段结构体整体消费,逐 key 声明不可读;`app.session_secret` 等单字段不值得为它绑一个结构体。 | -| D3 | **共享声明 + 内核冲突校验** | 配置是进程级只读事实,不存在数据表那种写竞争,因此允许读者各自声明同一 key;由内核强制"重复声明必须一致"兜底。放弃严格单所有权(需为若干配置值另造契约接口,且 `driver_http` 需 session store 连接参数是真实底层依赖)。 | -| D4 | **一次性全量迁移** | 不留双轨,架构一次到位;接受较大的 diff。 | -| D5 | **显式测试缝** | 删除 `isTest()` 魔法,测试通过 `core.WithConfigValues(...)` 注入。放弃"测试环境自动禁用中间件"的安全网,换取语义透明与可并行。 | -| D6 | **内核持抽象,viper 归 infra 适配器** | 微内核防线规定 `core/` 严禁 import 具体运行时依赖。`core` 只依赖 `ConfigSource` 接口,viper/yaml 装载放 `plugins/infra/config`。放弃"全放 core/config"(污染内核纯净性)与"完全插件化 + `contracts.ConfigService`"(门禁求值在 core,而 config 插件 `Apply` 尚未运行,存在鸡生蛋时序问题)。 | -| D7 | **顺带解耦 `pkg/idgen`** | 其 `init()` 读全局配置,导致 `pkg` 反向依赖配置单例。 | - ---- - -## 2. 分层与物理结构 - -```text -backend/core/extpoints/config.go # 配置引擎(仅 stdlib + reflect): - # ConfigSource 接口、声明注册、解析、冲突校验、脱敏 dump -backend/core/config.go # 泛型读取入口与 App 装配选项(Go 方法不支持类型参数) -backend/core/fiber.go # 新增 FiberSkipped 状态与门禁求值 -backend/core/types.go # 新增 ConfigExtension / ConfigBinding / ConfigView 别名 -backend/plugins/infra/config/ # viper + yaml 适配器,实现 core.ConfigSource(非 core.Plugin) -backend/cmd/ # 组合根:host 声明集 + app.Prepare() -backend/pkg/idgen/ # 移除 config 依赖,改为显式 Init(nodeID) -删除 backend/pkg/config/ # 全局单例 config.Config 一并消失 -``` - -职责边界: - -| 单元 | 做什么 | 不做什么 | -| :--- | :--- | :--- | -| `core/extpoints` 配置引擎 | 维护 key 注册表、按优先级解析、类型转换、冲突校验、脱敏输出 | 不知道任何具体 key 的名字,不读文件,不 import viper | -| `plugins/infra/config` | 定位 `config.yaml`(`CONFIG_PATH` → 向上查找)、解析成 raw map、代理 env 查询 | 不含 schema、不含业务字段语义 | -| 各插件 | 声明自己读哪些字段(tag 结构体)、声明门禁谓词 | 不读未声明的 key、不访问他插件的声明类型 | -| `cmd` | 声明 host 级 key、注入 `ConfigSource`、按已解析值初始化 logger/trace/banner | 不做 `if redis.enabled { ... }` 这类跨插件判断 | - -`plugins/infra/config` 不实现 `core.Plugin`,不出现在 `app.Use()` 列表里:它只向内核提供一个 `ConfigSource` 实例,没有服务、路由或任务可注册。它归 `plugins/infra/` 而非 `pkg/`,是因为它封装了具体运行时依赖(viper、文件系统)并持有装载状态,不符合 `pkg/` 的无状态定位。 - -`pkg/idgen` 解耦后,`backend/pkg/` 恢复"不依赖配置源"的无状态定位。 - ---- - -## 3. 核心类型与 API - -### 3.1 声明形态:带 tag 的结构体 - -唯一的批量作者形态是结构体 tag,一个字段同时表达 yaml 路径、env 覆盖名、默认值与敏感标记: - -```go -// plugins/infra/cache/redis_config.go —— redis 配置由 redis 插件自己声明 -type redisConfig struct { - Enabled bool `config:"enabled" env:"REDIS_ENABLED" default:"false" autoEnable:"REDIS_ADDR"` - Addrs []string `config:"addrs" env:"REDIS_ADDR"` - Username string `config:"username" env:"REDIS_USERNAME"` - Password string `config:"password" env:"REDIS_PASSWORD" secret:"true"` - DB int `config:"db" env:"REDIS_DB"` - ClusterMode bool `config:"cluster_mode" env:"REDIS_CLUSTER_MODE"` - MasterName string `config:"master_name" env:"REDIS_MASTER_NAME"` - KeyPrefix string `config:"key_prefix" env:"REDIS_KEY_PREFIX"` - MaintNotifications bool `config:"maint_notifications" env:"REDIS_MAINT_NOTIFICATIONS" default:"false"` - // ...pool/timeout 字段略 -} -``` - -支持的 tag:`config`(yaml 相对路径,必填)、`env`(覆盖用环境变量名)、`default`(字符串形式,缺省时视为未设置)、`autoEnable`(该 env 一旦存在即把本布尔字段置 true)、`secret`(dump 时脱敏)。 - -### 3.2 内核接口 - -```go -// ConfigSource 抽象了"原始值从哪来",由 infra 适配器实现,使内核不绑定 viper。 -type ConfigSource interface { - Lookup(path string) (any, bool) // config.yaml 中的点分路径 - LookupEnv(name string) (string, bool) - Describe() string // 用于日志,如 "config.yaml" 或 "" -} - -// ConfigBinding 把一个结构体绑定到某个 yaml 前缀上,是插件的声明单元。 -type ConfigBinding struct { - Prefix string // "redis";空串表示字段 key 即完整路径 - Target any // 指向带 tag 的结构体的指针 -} - -// ConfigView 是只读的已解析视图,供门禁与零散取值使用。 -type ConfigView interface { - String(key, fallback string) string - Bool(key string, fallback bool) bool - Int(key string, fallback int) int - Duration(key string, fallback time.Duration) time.Duration - Strings(key string) []string - WasSet(envName string) bool - Source(key string) string // "env" | "yaml" | "default",用于诊断 -} - -// ConfigExtension 是挂载在 Context 上的扩展点,根 Context 与所有 Fork 共享。 -type ConfigExtension interface { - ConfigView - Declare(pluginID string, bindings ...ConfigBinding) error - Bind(prefix string, target any) error - Entries() []ConfigEntry // 有效配置的脱敏视图 -} -``` - -`core` 侧导出别名与泛型入口(沿用仓库既有 `core.Provide[T]` / `core.Inject[T]` 风格): - -```go -func ConfigGet[T any](v extpoints.ConfigView, key string) (T, error) -``` - -### 3.3 插件侧用法 - -```go -// 批量绑定(Apply 内) -var cfg redisConfig -if err := ctx.Config().Bind("redis", &cfg); err != nil { - return err -} - -// 单字段读取:带 fallback 的访问器(门禁使用) -secret := ctx.Config().String("app.session_secret", "") - -// 单字段读取:需要区分"未设置"与"设置为零值"时用泛型入口 -rate, err := core.ConfigGet[float64](ctx.Config(), "otel.sampling_rate") -``` - -`ConfigEntry` 是 `Entries()` 返回的诊断单元,只含元数据与脱敏后的值: - -```go -type ConfigEntry struct { - Key string // "redis.password" - PluginID string // 首次声明者,用于冲突报错点名 - Env string - Source string // "env" | "yaml" | "default" - Value string // secret key 输出 "******" -} -``` - -`Bind` 的双重语义:若该 prefix 尚未声明,则按 `Target` 的 tag 自登记;若已声明,则是纯读取。登记时提供的 `env`/`default`/`secret` 元数据一律参与冲突校验(依 D3),因此自登记不会绕过校验。**只有需要早于 `Apply` 求值的插件才必须显式 `DeclareConfig()`。** - -### 3.4 门禁接口与 Fiber 跳过态 - -```go -// ConfigGatedPlugin 是可选接口:让内核在 Apply 之前决定插件是否激活。 -type ConfigGatedPlugin interface { - Plugin - DeclareConfig() []extpoints.ConfigBinding // 门禁所需 key 必须提前声明 - ConfigEnabled(v extpoints.ConfigView) bool -} -``` - -`FiberState` 新增 `FiberSkipped`。`App.reconcileLocked()` 在 `Load()` 前求值门禁:门禁为 false 的 Fiber 置 `FiberSkipped`,不计入依赖 satisfied 判定,也不参与 driver 启动。`App.Stop()` 对 skipped 与 active 一视同仁地按 LIFO 卸载其 scoped Context。 - -受门禁的插件对(现状仅三对,均为 Redis 存在与否的互斥实现): - -| 启用 | 跳过 | 门禁谓词 | -| :--- | :--- | :--- | -| `infra/cache` | `infra/cache_memory` | `redis.enabled` | -| `drivers/driver_asynq_worker` | `drivers/driver_inproc_worker` | `redis.enabled` | -| `drivers/driver_asynq_cron` | `drivers/driver_inproc_cron` | `redis.enabled` | - -### 3.5 组合根 - -```go -src := config.NewSource() // plugins/infra/config:仅定位与 raw 解析 -app := core.NewApp( - core.WithProfile(profile), - core.WithConfigSource(src), - core.WithConfigDecl(hostBinding...), // app.* / log.* / otel.* -) -app.Use( - infradb.New(), logger.New(), storage.New(), - cache.New(), cache_memory.New(), // 不再 if/else,门禁决定 - driver_asynq_worker.New(), driver_inproc_worker.New(), - driver_asynq_cron.New(), driver_inproc_cron.New(), - admin.New(), user.New(), auth.New(), /* ... */ - driver_http.New(), // addr 由插件自己声明读取 -) -if err := app.Prepare(); err != nil { return err } // 解析屏障 + 门禁求值 - -timeout, _ := app.Context().Config().Duration("app.graceful_shutdown_timeout", 30) -app.SetShutdownTimeout(timeout) -``` - -`WithShutdownTimeout(d)` 保留为显式覆盖入口(测试与非标准装配使用),生产路径改为 `Prepare()` 之后由已解析视图经 `SetShutdownTimeout` 设定。`driver_http.New(WithAddr(...))` 选项删除,addr 归 `driver_http` 在 `Apply` 内声明读取。 - ---- - -## 4. 解析语义与启动时序 - -### 4.1 单 key 优先级链 - -```text -1. 显式 env 命中 env:"DB_ENABLED" → 最高优先级 -2. autoEnable env 命中 autoEnable:"DB_HOST" → true(被 1 覆盖) -3. config.yaml 命中 config:"enabled" -4. default tag 兜底 -``` - -需要保留的既有特殊语义: - -- **标量 env 填充切片字段**:`REDIS_ADDR=redis:6379` → `redis.addrs = ["redis:6379"]`;`CLICKHOUSE_HOST` 同理。 -- **隐式启用**:`DB_HOST` → `database.enabled=true`、`REDIS_ADDR` → `redis.enabled=true`、`CLICKHOUSE_HOST` → `clickhouse.enabled=true`;显式 `*_ENABLED` 始终优先于隐式推导。 -- **同一 env 的双重角色**:`REDIS_ADDR` 既是 `redis.addrs` 的值来源,又是 `redis.enabled` 的 `autoEnable` 触发器。引擎按 key 独立解析、允许一个 env 名服务多个 key,实现时不可把它建模成"env → 单一 key"的一对一映射。 -- **时长字段**:`slow_threshold: 200ms` 解析为 `time.Duration`。 -- **文件定位**:`CONFIG_PATH` 优先;否则从工作目录向上最多 5 层查找 `config.yaml`(该文件位于仓库根,`backend/` 为其子目录)。 - -### 4.2 时序 - -```text -config.NewSource() # 读 yaml → raw map;零 schema 知识 - ↓ -core.NewApp(WithConfigSource) # 记录 host 声明 - ↓ -app.Use(...) # 遇 DeclareConfig() 立即登记 binding(叶子 key + env + default + secret) - ↓ -app.Prepare() # ① 冲突校验 ② 逐 key 解析 ③ 脱敏 dump ④ 门禁求值 → FiberSkipped - ↓ -app.Run() → Reconcile/Apply # 插件内 Bind/Get 读取已解析值 -``` - -`App.Start()` 在未显式调用 `Prepare()` 时幂等补做,防止遗漏。冲突校验规则:同一 key 的多份声明必须 `env` 名、`default`、`secret` 三项一致,否则 `Prepare()` 返回错误并点名两个声明者。 - -### 4.3 有意的行为变更 - -| # | 变更 | 现状 | 变更后 | -| :--- | :--- | :--- | :--- | -| C1 | `default` 生效条件 | `applyDefaults` 对零值二次回落(`session_age<=0` → 86400) | 仅当 env 与 yaml 均缺失时生效;`app.session_age<=0` 在 `Prepare()` 判为配置错误(fail fast 优于静默改写) | -| C2 | 测试上下文 | `isTest()` 自动禁用 DB/Redis/ClickHouse 并把 sqlite 指向 `:memory:` | 删除该魔法;测试用 `core.WithConfigValues(...)` 显式声明。未声明 `database.enabled` 时按 default `false` 落 sqlite 后备,其路径沿用 `postgres.go` 既有的 `./data/wavelet.db` 回落——需要内存库的用例必须显式注入 `database.sqlite_path = ":memory:"` | -| C3 | 配置 dump | `printConfig` 明文打印全量结构体,含 `DB_PASSWORD`、`APP_SESSION_SECRET` | 按 `secret:"true"` 脱敏后输出,并标注每个 key 的来源(env/yaml/default) | -| C4 | 队列默认值 | 硬编码在 `pkg/config` 的 `applyEnvOverrides` | 移入唯一消费者 `driver_asynq_worker` 的声明(`webhook`/`whitelist_only`/`default` 三级优先级不变) | -| C5 | 非法 env 值 | `envInt/envBool/envFloat64` 在 `strconv` 失败时静默丢弃 env 值、回落 yaml/default | `Prepare()` 返回 `ErrConfigType` 并点名 key 与非法值 | - -除此之外,解析结果与现状逐 key 等价(由 §7.1 第 4 条的对拍测试证明)。 - ---- - -## 5. 声明归属映射 - -| 声明方 | key 前缀 | 消费者(含跨插件读) | -| :--- | :--- | :--- | -| `cmd` host 声明集 | `app.{env,app_name,addr,node_id,graceful_shutdown_timeout}`、`log.*`、`otel.*` | `cmd/root.go`、`cmd/banner.go`、`core.App` | -| `plugins/infra/cache` | `redis.*`(含 `enabled` 门禁、`autoEnable: REDIS_ADDR`) | `infra/cache`、`driver_http`(session store)、`driver_asynq_worker`、`driver_asynq_cron` | -| `plugins/infra/database` | `database.*`、`clickhouse.*` | `infra/database`、`admin`、`risk_control` | -| `plugins/domain/auth` | `app.session_*`(cookie/secret/age/domain/secure/http_only) | `auth`、`cap`、`message_gateway`、`driver_http` | -| `plugins/drivers/driver_asynq_worker` | `worker.*`(并发、strict_priority、queues 默认值) | 自身 | -| 其余 | 按需就近声明 | — | - -跨插件读同一 key(如 `cap` 读 auth 声明的 `app.session_secret`)依 D3 走共享声明:`cap` 也声明该 key,三份元数据必须与 auth 一致,否则启动失败。 - -**Key 命名约定**:既有 infra key 保持顶层(`redis.*`、`database.*`),以兼容线上 `config.yaml`;新增插件的私有配置归 `plugins..*` 命名空间,与 `new-setting` skill 的描述对齐。 - -### 5.1 `pkg/idgen` 解耦 - -- 删除 `init()` 中对 `config.Config.App.NodeID` 的读取。 -- 新增 `idgen.Init(nodeID int64) error`,由 host 在 `Prepare()` 之后显式调用(值来自 host 声明的 `app.node_id`)。 -- 未初始化时 `NextUint64ID()` panic 并点名"未调用 idgen.Init",而非静默使用 nodeID=0 生成可能与集群冲突的 ID。 -- 11 个调用点的 `idgen.NextUint64ID()` 签名保持不变;依赖 ID 生成的测试需显式 `idgen.Init`。这是本次迁移唯一会触及既有测试文件之处。 - -### 5.2 明确不在范围内 - -本设计只改变配置的**来源与所有权**,不动这些既有全局变量:`cache.Redis`、`driver_asynq_worker.RedisOpt`/`AsynqClient`、`infra/database.db`。它们各自的收敛属于独立议题。 - ---- - -## 6. 错误处理 - -- `Prepare()` 以 `errors.Join` 聚合全部配置错误,哨兵错误:`ErrConfigConflict`(重复声明不一致)、`ErrConfigType`(env 值无法转为目标类型)、`ErrConfigInvalid`(值域校验失败,如 `session_age<=0`)、`ErrConfigNotResolved`(`Prepare()` 之前调用 `Bind`/`Get`,错误信息点名正确调用顺序)。 -- 所有配置错误经 `error` 返回,由 `cmd` 决定终止方式;`core` 与 `extpoints` 内不再有 `log.Fatalf`。 -- `config.yaml` 缺失不是错误(沿用"仅用 env"路径,记一条 info 日志);`CONFIG_PATH` 显式指定但读不到或解析失败 → 返回 error。 -- 门禁 `ConfigEnabled(v ConfigView) bool` 只读已解析值、用带 fallback 的访问器,配置错误已在 `Prepare()` 阶段暴露,因此门禁不引入新的错误源。 -- `Declare` 与 `Bind` 校验 `Target` 必须是非 nil 结构体指针,否则返回 error(不 panic)。 - ---- - -## 7. 测试与验收 - -### 7.1 测试分层 - -1. **引擎单测**(`core/extpoints`):内存 fake `ConfigSource`,表驱动覆盖优先级四档、标量 env→切片、`autoEnable` 与显式 env 的优先关系、冲突校验、脱敏 dump、`time.Duration` 与嵌套结构体 tag 解析、`Prepare()` 前访问的错误路径。 -2. **门禁单测**(`core`):互斥插件对恰好激活一个、被跳过插件不计入依赖 satisfied、`FiberSkipped` 参与 `Stop` 的 LIFO 卸载。 -3. **适配器单测**(`plugins/infra/config`):`t.TempDir()` 写 yaml + `t.Setenv`,禁止相对路径。 -4. **新旧对拍**:迁移期间保留一份临时对拍测试,用仓库现网 `config.yaml` 与 `.env` 逐 key 比较旧 `pkg/config` 与新引擎的输出,证明除 C1–C5 外完全等价;验证通过后随旧包一并删除。 -5. **迁移后插件测试**:改用 `core.WithConfigValues(...)` 显式注入;依赖 ID 生成的测试显式 `idgen.Init`。 - -### 7.2 验收标准 - -1. `backend/pkg/config` 不存在,`grep -rn "pkg/config\|config\.Config" backend/` 零命中。 -2. `core/` 无 viper import;`backend/pkg/` 内不出现任何配置源 import。 -3. `cmd/app.go` 中不存在跨插件配置判断,驱动选型完全由门禁产生。 -4. `.env`、`config.yaml`、docker-compose **零改动**即可启动,行为等价(除已登记的 C1–C5)。 -5. 同时挂载 `cache` 与 `cache_memory` 而仅激活其一——"预声明 + 门禁"的端到端可验证证据;两条路径(Redis 启用 → asynq;禁用 → inproc)各实跑一次。 -6. `make code-check`、`make format`、`go test ./backend/...` 全绿;`go run main.go all` 实跑通过,覆盖 banner、迁移与门禁。 -7. `AGENTS.md`、`new-setting` skill 与白皮书中 `ctx.Config().Bind(...)` 的签名与 key 命名约定更新为已实现的真实 API。 - -### 7.3 实施顺序建议 - -每阶段独立可验证,供实施计划拆分参考: - -| 阶段 | 内容 | 验证 | -| :--- | :--- | :--- | -| P1 | 配置引擎(`core/extpoints/config.go`)+ `plugins/infra/config` 适配器 + 新旧对拍测试 | `go test ./backend/core/...`;对拍输出等价性报告 | -| P2 | 门禁与 `FiberSkipped`、`App.Prepare()` 解析屏障 | `core` 门禁单测;现有测试全绿(此时旧单例仍在,未迁移) | -| P3 | 按 infra → drivers → domain → cmd 顺序迁移 27 个文件;`idgen.Init` 解耦 | 每层迁移后 `go build ./...` + 该层测试;最后实跑两条门禁路径 | -| P4 | 删除 `backend/pkg/config` 与对拍测试;更新 `AGENTS.md`/skill/白皮书 API | §7.2 全部验收项逐条复核 | - ---- - -## 附录 A:迁移清单 - -删除:`backend/pkg/config/{config.go,model.go,config_test.go}` - -新增:`backend/core/extpoints/config.go`、`backend/core/config.go`、`backend/plugins/infra/config/*`、各插件内 `_config.go` 声明文件 - -需改写的 27 个文件: - -| 分组 | 文件 | -| :--- | :--- | -| 组合根 | `cmd/app.go`、`cmd/root.go`、`cmd/banner.go`、`cmd/app_test.go`、`cmd/banner_test.go`、`cmd/redis_plug_test.go` | -| 基础库 | `pkg/idgen/snowflake.go`(连带 `pkg/idgen/snowflake_test.go`) | -| infra | `plugins/infra/cache/redis.go`、`plugins/infra/database/postgres.go`、`plugins/infra/database/clickhouse.go` | -| drivers | `plugins/drivers/driver_http/engine.go`、`plugins/drivers/driver_http/middlewares.go`、`plugins/drivers/driver_asynq_worker/utils.go`、`plugins/drivers/driver_asynq_worker/utils_test.go`、`plugins/drivers/driver_asynq_cron/plugin.go` | -| domain/admin | `plugins/domain/admin/handler/db.go`、`plugins/domain/admin/repository/db.go`、`plugins/domain/admin/service/db.go`、`plugins/domain/admin/service/status.go`、`plugins/domain/admin/service/log_switch.go` | -| domain/其他 | `plugins/domain/auth/session.go`、`plugins/domain/cap/service.go`、`plugins/domain/message_gateway/service/service.go`、`plugins/domain/system/plugin.go`、`plugins/domain/risk_control/middleware.go`、`plugins/domain/risk_control/middleware_test.go`、`plugins/domain/risk_control/logstore/provider.go` | - -> 注:`risk_control/middleware.go`、`cap/service.go`、`message_gateway/service/service.go` 等处以 `config.Config != nil` 做存在性判断的分支,在注入式配置模型下不再可能,迁移时一并消除。 - ---- - -## 8. 落地回写(P1 + P2 已实施) - -实施结果与本设计原述的差异,均已按下列口径落地: - -| # | 设计原述 | 落地结果 | 缘由 | -| :--- | :--- | :--- | :--- | -| R1 | §4.3 C1、§6 把 `app.session_age<=0` 列为内核解析错误 | 引擎不做值域校验,`ErrConfigInvalid` 保留但未在内核使用;值域由声明者在 `Bind` 之后校验(P3 由 auth 承担) | 引擎被设计成不认识任何业务 key 的语义,把业务规则塞进内核会破坏该不变式 | -| R2 | §3.2 `ConfigView.Source(key)` | 更名 `Origin(key)`;新增 `Value(key) (any, bool)`;`ConfigExtension` 增加 `SetSource`、`Resolved` | `Source` 与类型名 `ConfigSource` 同文件易混淆;`Value` 支撑 `core.ConfigGet[T]`(Go 方法不能带类型参数);`SetSource` 进接口以免运行时类型断言 | -| R3 | §3.5 仅有 `WithShutdownTimeout` | 新增 `App.ShutdownTimeout()` 与 `SetShutdownTimeout(d) *App` | 组合根需在 `Prepare()` 之后把已解析预算写回内核,构造期选项无法表达该顺序 | -| R4 | §4.2 时序图把门禁求值画在 `Prepare()` 内 | `Prepare()` 只建立解析屏障,门禁在 `reconcileLocked` 每轮调和中求值 | `App.Use` 可在 `Prepare()` 之后继续挂载插件;只在 `Prepare` 求值会留下一批永不判定的门禁 | -| R5 | 未涉及 | `App` 未注入 `ConfigSource` 时配置能力视为未启用,解析屏障直接放行;实现了 `ConfigGatedPlugin` 却无配置源的插件 fail fast 点名原因 | 内核存在大量不使用配置的装配路径(既有测试与嵌入式用法),不能强制要求配置源;但门禁无数据可依时必须报错,而非静默全激活 | -| R6 | §4.1 隐含"每个 key 都有 env 覆盖" | env 覆盖面完全由声明决定。旧装载器只对部分 key 提供 env(`slow_threshold`、`conn_max_lifetime` 等从未有 env 覆盖),对拍镜像必须精确复刻该覆盖面 | 否则对拍出现假漂移;放宽某 key 的 env 覆盖是 P3 的声明选择,不构成引擎行为变更 | -| R7 | §4.1 "向上最多 5 层查找 `config.yaml`" | 该向上查找会**越出 git worktree 边界**:从 `backend/pkg/config` 出发第 5 层可命中父级检出的 `config.yaml` | 属既有行为、非本次引入,但在 worktree 中开发会静默使用另一份检出的配置。对拍测试已改为以入库的 `config.example.yaml` 所在目录为锚;`config.yaml` 本身被 gitignore,干净克隆中不存在 | - -分期口径:本设计 §7.3 的 P1 + P2 已实施完成;P3(27 个消费文件迁移、`pkg/idgen` 解耦)与 P4(删除 `backend/pkg/config`、移除对拍夹具)由后续计划承接。 diff --git a/docs/swagger.json b/docs/swagger.json index 3bc39020..dafc71ac 100644 --- a/docs/swagger.json +++ b/docs/swagger.json @@ -1016,7 +1016,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.ChannelDTO" + "$ref": "#/definitions/do.ChannelDTO" } } } @@ -1050,7 +1050,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.CreateChannelRequest" + "$ref": "#/definitions/do.CreateChannelRequest" } } ], @@ -1066,7 +1066,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.ChannelDTO" + "$ref": "#/definitions/do.ChannelDTO" } } } @@ -1111,7 +1111,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.Definition" + "$ref": "#/definitions/do.Definition" } } } @@ -1192,7 +1192,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.UpdateChannelRequest" + "$ref": "#/definitions/do.UpdateChannelRequest" } } ], @@ -1208,7 +1208,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.ChannelDTO" + "$ref": "#/definitions/do.ChannelDTO" } } } @@ -1305,7 +1305,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.PushChannel" + "$ref": "#/definitions/entity.PushChannel" } } } @@ -1339,7 +1339,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.CreatePushChannelRequest" + "$ref": "#/definitions/do.CreatePushChannelRequest" } } ], @@ -1355,7 +1355,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.PushChannel" + "$ref": "#/definitions/entity.PushChannel" } } } @@ -1415,7 +1415,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.TestPushChannelRequest" + "$ref": "#/definitions/do.TestPushChannelRequest" } } ], @@ -1462,7 +1462,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.UpdatePushChannelRequest" + "$ref": "#/definitions/do.UpdatePushChannelRequest" } } ], @@ -1478,7 +1478,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.PushChannel" + "$ref": "#/definitions/entity.PushChannel" } } } @@ -1550,7 +1550,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.PushEvent" + "$ref": "#/definitions/entity.PushEvent" } } } @@ -1584,7 +1584,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.CreatePushEventRequest" + "$ref": "#/definitions/do.CreatePushEventRequest" } } ], @@ -1600,7 +1600,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.PushEvent" + "$ref": "#/definitions/entity.PushEvent" } } } @@ -1667,7 +1667,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.UpdatePushEventRequest" + "$ref": "#/definitions/do.UpdatePushEventRequest" } } ], @@ -1859,7 +1859,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.TestPushRequest" + "$ref": "#/definitions/do.TestPushRequest" } } ], @@ -2668,7 +2668,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Schedule" + "$ref": "#/definitions/model.Schedule" } } } @@ -2730,7 +2730,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Schedule" + "$ref": "#/definitions/model.Schedule" } } } @@ -2812,7 +2812,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Schedule" + "$ref": "#/definitions/model.Schedule" } } } @@ -3001,7 +3001,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Template" + "$ref": "#/definitions/model.Template" } } } @@ -3139,7 +3139,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Template" + "$ref": "#/definitions/model.Template" } } } @@ -3219,7 +3219,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/Wavelet_plugins_domain_admin_model.Template" + "$ref": "#/definitions/model.Template" } } } @@ -4921,7 +4921,7 @@ "name": "request", "in": "body", "schema": { - "$ref": "#/definitions/cap.challengeRequest" + "$ref": "#/definitions/dto.ChallengeRequest" } } ], @@ -4937,7 +4937,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/cap.ChallengeResponse" + "$ref": "#/definitions/dto.ChallengeResponse" } } } @@ -4970,7 +4970,7 @@ "name": "request", "in": "body", "schema": { - "$ref": "#/definitions/cap.challengeRequest" + "$ref": "#/definitions/dto.ChallengeRequest" } } ], @@ -4986,7 +4986,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/cap.ChallengeResponse" + "$ref": "#/definitions/dto.ChallengeResponse" } } } @@ -5022,7 +5022,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/cap.redeemRequest" + "$ref": "#/definitions/dto.RedeemRequest" } } ], @@ -5038,7 +5038,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/cap.RedeemResponse" + "$ref": "#/definitions/dto.RedeemResponse" } } } @@ -5062,7 +5062,7 @@ }, "/api/v1/config/public": { "get": { - "description": "返回系统配置表中 visibility 为 1 的配置键值集合", + "description": "返回系统配置表中 visibility 为 1 的扁平键值集合(如 cap_login_enabled)", "consumes": [ "application/json" ], @@ -13215,7 +13215,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.BindingDTO" + "$ref": "#/definitions/do.BindingDTO" } } } @@ -13255,7 +13255,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/model.BindRequest" + "$ref": "#/definitions/do.BindRequest" } } ], @@ -13271,7 +13271,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/model.BindingDTO" + "$ref": "#/definitions/do.BindingDTO" } } } @@ -13368,7 +13368,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/model.PublicChannelDTO" + "$ref": "#/definitions/do.PublicChannelDTO" } } } @@ -13405,7 +13405,7 @@ "in": "body", "required": true, "schema": { - "$ref": "#/definitions/auth.CallbackRequest" + "$ref": "#/definitions/dto.CallbackRequest" } } ], @@ -13421,7 +13421,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.OAuthCallbackResult" + "$ref": "#/definitions/dto.OAuthCallbackResult" } } } @@ -13575,7 +13575,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.OAuthAuthorizeResponse" + "$ref": "#/definitions/dto.OAuthAuthorizeResponse" } } } @@ -13664,7 +13664,7 @@ "data": { "type": "array", "items": { - "$ref": "#/definitions/auth.AuthSourceView" + "$ref": "#/definitions/dto.AuthSourceView" } } } @@ -13702,7 +13702,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.BasicUserInfo" + "$ref": "#/definitions/dto.BasicUserInfo" } } } @@ -13755,7 +13755,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.OAuthAuthorizeResponse" + "$ref": "#/definitions/dto.OAuthAuthorizeResponse" } } } @@ -14390,7 +14390,7 @@ "type": "object", "properties": { "data": { - "$ref": "#/definitions/auth.BasicUserInfo" + "$ref": "#/definitions/dto.BasicUserInfo" } } } @@ -14647,7 +14647,7 @@ }, "/api/v1/user/login": { "post": { - "description": "使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。", + "description": "使用用户名和密码登录系统,验证通过后建立 Session 并返回用户信息。", "consumes": [ "application/json" ], @@ -14657,7 +14657,7 @@ "tags": [ "user" ], - "summary": "用户密码登录", + "summary": "用户登录", "parameters": [ { "description": "登录请求参数", @@ -14682,6 +14682,12 @@ "$ref": "#/definitions/response.Any" } }, + "429": { + "description": "登录尝试过于频繁", + "schema": { + "$ref": "#/definitions/response.Any" + } + }, "500": { "description": "服务内部错误", "schema": { @@ -14817,6 +14823,12 @@ "$ref": "#/definitions/response.Any" } }, + "429": { + "description": "注册尝试过于频繁", + "schema": { + "$ref": "#/definitions/response.Any" + } + }, "500": { "description": "服务内部错误", "schema": { @@ -14870,6 +14882,17 @@ "user" ], "summary": "发送邮箱验证码", + "parameters": [ + { + "description": "目标邮箱", + "name": "request", + "in": "body", + "required": true, + "schema": { + "$ref": "#/definitions/user.sendEmailCodeRequest" + } + } + ], "responses": { "200": { "description": "发送成功", @@ -14882,6 +14905,12 @@ "schema": { "$ref": "#/definitions/response.Any" } + }, + "500": { + "description": "发送失败", + "schema": { + "$ref": "#/definitions/response.Any" + } } } } @@ -15204,36 +15233,6 @@ } } }, - "Wavelet_plugins_domain_admin_model.Schedule": { - "type": "object", - "properties": { - "created_at": { - "type": "string" - }, - "cron": { - "type": "string" - }, - "id": { - "type": "string", - "example": "0" - }, - "is_active": { - "type": "boolean" - }, - "name": { - "type": "string" - }, - "payload": { - "type": "string" - }, - "task_type": { - "type": "string" - }, - "updated_at": { - "type": "string" - } - } - }, "Wavelet_plugins_domain_admin_model.SystemConfig": { "type": "object", "properties": { @@ -15320,41 +15319,6 @@ } } }, - "Wavelet_plugins_domain_admin_model.Template": { - "type": "object", - "properties": { - "content": { - "type": "string" - }, - "created_at": { - "type": "string" - }, - "description": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "is_system": { - "type": "boolean" - }, - "key": { - "type": "string" - }, - "name": { - "type": "string" - }, - "subject": { - "type": "string" - }, - "type": { - "type": "string" - }, - "updated_at": { - "type": "string" - } - } - }, "agent.ActiveConfigMeta": { "type": "object", "properties": { @@ -15662,179 +15626,6 @@ } } }, - "auth.AuthSourceView": { - "type": "object", - "properties": { - "client_secret_configured": { - "type": "boolean" - }, - "display_name": { - "type": "string" - }, - "icon_url": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "is_active": { - "type": "boolean" - }, - "name": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, - "auth.BasicUserInfo": { - "type": "object", - "properties": { - "avatar_url": { - "type": "string" - }, - "bio": { - "type": "string" - }, - "email": { - "type": "string" - }, - "gender": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "is_admin": { - "type": "boolean" - }, - "location": { - "type": "string" - }, - "need_change_password": { - "type": "boolean" - }, - "nickname": { - "type": "string" - }, - "phone": { - "type": "string" - }, - "username": { - "type": "string" - }, - "website": { - "type": "string" - } - } - }, - "auth.CallbackRequest": { - "type": "object", - "required": [ - "code", - "state" - ], - "properties": { - "code": { - "type": "string" - }, - "state": { - "type": "string" - } - } - }, - "auth.OAuthAuthorizeResponse": { - "type": "object", - "properties": { - "authorize_url": { - "type": "string" - } - } - }, - "auth.OAuthCallbackResult": { - "type": "object", - "properties": { - "status": { - "type": "string" - }, - "user": { - "$ref": "#/definitions/auth.BasicUserInfo" - } - } - }, - "cap.ChallengeResponse": { - "type": "object", - "properties": { - "challenge": { - "type": "object", - "properties": { - "c": { - "type": "integer" - }, - "d": { - "type": "integer" - }, - "s": { - "type": "integer" - } - } - }, - "expires": { - "description": "ms timestamp", - "type": "integer" - }, - "token": { - "type": "string" - } - } - }, - "cap.RedeemResponse": { - "type": "object", - "properties": { - "error": { - "type": "string" - }, - "expires": { - "type": "integer" - }, - "success": { - "type": "boolean" - }, - "token": { - "type": "string" - } - } - }, - "cap.challengeRequest": { - "type": "object", - "properties": { - "scope": { - "type": "string" - } - } - }, - "cap.redeemRequest": { - "type": "object", - "required": [ - "solutions", - "token" - ], - "properties": { - "scope": { - "type": "string" - }, - "solutions": { - "type": "array", - "items": { - "type": "integer" - } - }, - "token": { - "type": "string" - } - } - }, "cloudflare.AvailableDomain": { "type": "object", "properties": { @@ -16484,6 +16275,573 @@ } } }, + "do.BindRequest": { + "type": "object", + "properties": { + "channel_id": { + "type": "string" + }, + "code": { + "type": "string" + } + } + }, + "do.BindingDTO": { + "type": "object", + "properties": { + "channel_id": { + "type": "string", + "example": "0" + }, + "channel_name": { + "type": "string" + }, + "channel_type": { + "type": "string" + }, + "created_at": { + "type": "string" + }, + "id": { + "type": "string", + "example": "0" + }, + "platform_user_id": { + "type": "string" + }, + "user_id": { + "type": "string", + "example": "0" + } + } + }, + "do.ChannelDTO": { + "type": "object", + "properties": { + "credentials": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "extra": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "id": { + "type": "string", + "example": "0" + }, + "name": { + "type": "string" + }, + "owner_id": { + "type": "string", + "example": "0" + }, + "owner_scope": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "do.CreateChannelRequest": { + "type": "object", + "properties": { + "credentials": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "extra": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "name": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "do.CreatePushChannelRequest": { + "type": "object", + "required": [ + "name", + "type" + ], + "properties": { + "description": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "name": { + "type": "string" + }, + "other": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "do.CreatePushEventRequest": { + "type": "object", + "properties": { + "channels": { + "type": "array", + "items": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "event_key": { + "type": "string" + }, + "targets": { + "type": "array", + "items": { + "type": "string" + } + }, + "task_type": { + "type": "string" + }, + "template": { + "type": "string" + } + } + }, + "do.Definition": { + "type": "object", + "properties": { + "fields": { + "type": "array", + "items": { + "$ref": "#/definitions/do.Field" + } + }, + "type": { + "type": "string" + } + } + }, + "do.Field": { + "type": "object", + "properties": { + "key": { + "type": "string" + }, + "required": { + "type": "boolean" + }, + "type": { + "type": "string" + } + } + }, + "do.PublicChannelDTO": { + "type": "object", + "properties": { + "id": { + "type": "string", + "example": "0" + }, + "name": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "do.TestPushChannelRequest": { + "type": "object", + "properties": { + "name": { + "type": "string" + }, + "other": { + "type": "string" + }, + "target": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "do.TestPushRequest": { + "type": "object", + "required": [ + "config" + ], + "properties": { + "config": { + "$ref": "#/definitions/push.Config" + }, + "target": { + "type": "string" + } + } + }, + "do.UpdateChannelRequest": { + "type": "object", + "properties": { + "credentials": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "extra": { + "type": "object", + "additionalProperties": { + "type": "string" + } + }, + "name": { + "type": "string" + } + } + }, + "do.UpdatePushChannelRequest": { + "type": "object", + "required": [ + "type" + ], + "properties": { + "description": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "other": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "do.UpdatePushEventRequest": { + "type": "object", + "required": [ + "template" + ], + "properties": { + "channels": { + "type": "array", + "items": { + "type": "string" + } + }, + "enabled": { + "type": "boolean" + }, + "targets": { + "type": "array", + "items": { + "type": "string" + } + }, + "template": { + "type": "string" + } + } + }, + "dto.AuthSourceView": { + "type": "object", + "properties": { + "client_secret_configured": { + "type": "boolean" + }, + "display_name": { + "type": "string" + }, + "icon_url": { + "type": "string" + }, + "id": { + "type": "integer" + }, + "is_active": { + "type": "boolean" + }, + "name": { + "type": "string" + }, + "type": { + "type": "string" + } + } + }, + "dto.BasicUserInfo": { + "type": "object", + "properties": { + "avatar_url": { + "type": "string" + }, + "bio": { + "type": "string" + }, + "email": { + "type": "string" + }, + "gender": { + "type": "string" + }, + "id": { + "type": "string", + "example": "0" + }, + "is_admin": { + "type": "boolean" + }, + "location": { + "type": "string" + }, + "need_change_password": { + "type": "boolean" + }, + "nickname": { + "type": "string" + }, + "phone": { + "type": "string" + }, + "username": { + "type": "string" + }, + "website": { + "type": "string" + } + } + }, + "dto.CallbackRequest": { + "type": "object", + "required": [ + "code", + "state" + ], + "properties": { + "code": { + "type": "string" + }, + "state": { + "type": "string" + } + } + }, + "dto.ChallengeRequest": { + "type": "object", + "properties": { + "scope": { + "type": "string" + } + } + }, + "dto.ChallengeResponse": { + "type": "object", + "properties": { + "challenge": { + "type": "object", + "properties": { + "c": { + "type": "integer" + }, + "d": { + "type": "integer" + }, + "s": { + "type": "integer" + } + } + }, + "expires": { + "description": "ms timestamp", + "type": "integer" + }, + "token": { + "type": "string" + } + } + }, + "dto.OAuthAuthorizeResponse": { + "type": "object", + "properties": { + "authorize_url": { + "type": "string" + } + } + }, + "dto.OAuthCallbackResult": { + "type": "object", + "properties": { + "status": { + "type": "string" + }, + "user": { + "$ref": "#/definitions/dto.BasicUserInfo" + } + } + }, + "dto.RedeemRequest": { + "type": "object", + "required": [ + "solutions", + "token" + ], + "properties": { + "scope": { + "type": "string" + }, + "solutions": { + "type": "array", + "items": { + "type": "integer" + } + }, + "token": { + "type": "string" + } + } + }, + "dto.RedeemResponse": { + "type": "object", + "properties": { + "error": { + "type": "string" + }, + "expires": { + "type": "integer" + }, + "success": { + "type": "boolean" + }, + "token": { + "type": "string" + } + } + }, + "entity.PushChannel": { + "type": "object", + "properties": { + "created_at": { + "type": "string" + }, + "description": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "id": { + "type": "integer" + }, + "name": { + "type": "string" + }, + "other": { + "type": "string" + }, + "token": { + "type": "string" + }, + "type": { + "type": "string" + }, + "updated_at": { + "type": "string" + }, + "url": { + "type": "string" + } + } + }, + "entity.PushEvent": { + "type": "object", + "properties": { + "channels": { + "type": "array", + "items": { + "type": "string" + } + }, + "created_at": { + "type": "string" + }, + "enabled": { + "type": "boolean" + }, + "event_key": { + "type": "string" + }, + "id": { + "type": "integer" + }, + "name": { + "type": "string" + }, + "targets": { + "type": "array", + "items": { + "type": "string" + } + }, + "task_type": { + "type": "string" + }, + "template": { + "type": "string" + }, + "updated_at": { + "type": "string" + } + } + }, "flared.ApplyLogPayload": { "type": "object", "properties": { @@ -16797,46 +17155,6 @@ } } }, - "model.BindRequest": { - "type": "object", - "properties": { - "channel_id": { - "type": "string" - }, - "code": { - "type": "string" - } - } - }, - "model.BindingDTO": { - "type": "object", - "properties": { - "channel_id": { - "type": "string", - "example": "0" - }, - "channel_name": { - "type": "string" - }, - "channel_type": { - "type": "string" - }, - "created_at": { - "type": "string" - }, - "id": { - "type": "string", - "example": "0" - }, - "platform_user_id": { - "type": "string" - }, - "user_id": { - "type": "string", - "example": "0" - } - } - }, "model.BrowserItem": { "type": "object", "properties": { @@ -16848,43 +17166,6 @@ } } }, - "model.ChannelDTO": { - "type": "object", - "properties": { - "credentials": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "extra": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "id": { - "type": "string", - "example": "0" - }, - "name": { - "type": "string" - }, - "owner_id": { - "type": "string", - "example": "0" - }, - "owner_scope": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, "model.ConfigVersion": { "type": "object", "properties": { @@ -16943,91 +17224,6 @@ } } }, - "model.CreateChannelRequest": { - "type": "object", - "properties": { - "credentials": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "extra": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "name": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, - "model.CreatePushChannelRequest": { - "type": "object", - "required": [ - "name", - "type" - ], - "properties": { - "description": { - "type": "string" - }, - "enabled": { - "type": "boolean" - }, - "name": { - "type": "string" - }, - "other": { - "type": "string" - }, - "token": { - "type": "string" - }, - "type": { - "type": "string" - }, - "url": { - "type": "string" - } - } - }, - "model.CreatePushEventRequest": { - "type": "object", - "properties": { - "channels": { - "type": "array", - "items": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "event_key": { - "type": "string" - }, - "targets": { - "type": "array", - "items": { - "type": "string" - } - }, - "task_type": { - "type": "string" - }, - "template": { - "type": "string" - } - } - }, "model.CreateScheduleRequest": { "type": "object", "required": [ @@ -17214,20 +17410,6 @@ } } }, - "model.Definition": { - "type": "object", - "properties": { - "fields": { - "type": "array", - "items": { - "$ref": "#/definitions/model.Field" - } - }, - "type": { - "type": "string" - } - } - }, "model.DispatchTaskRequest": { "type": "object", "required": [ @@ -17290,20 +17472,6 @@ } } }, - "model.Field": { - "type": "object", - "properties": { - "key": { - "type": "string" - }, - "required": { - "type": "boolean" - }, - "type": { - "type": "string" - } - } - }, "model.ListUsersResponse": { "type": "object", "properties": { @@ -17626,92 +17794,31 @@ } } }, - "model.PublicChannelDTO": { + "model.Schedule": { "type": "object", "properties": { + "created_at": { + "type": "string" + }, + "cron": { + "type": "string" + }, "id": { "type": "string", "example": "0" }, - "name": { - "type": "string" - }, - "type": { - "type": "string" - } - } - }, - "model.PushChannel": { - "type": "object", - "properties": { - "created_at": { - "type": "string" - }, - "description": { - "type": "string" - }, - "enabled": { + "is_active": { "type": "boolean" }, - "id": { - "type": "integer" - }, "name": { "type": "string" }, - "other": { + "payload": { "type": "string" }, - "token": { - "type": "string" - }, - "type": { - "type": "string" - }, - "updated_at": { - "type": "string" - }, - "url": { - "type": "string" - } - } - }, - "model.PushEvent": { - "type": "object", - "properties": { - "channels": { - "type": "array", - "items": { - "type": "string" - } - }, - "created_at": { - "type": "string" - }, - "enabled": { - "type": "boolean" - }, - "event_key": { - "type": "string" - }, - "id": { - "type": "integer" - }, - "name": { - "type": "string" - }, - "targets": { - "type": "array", - "items": { - "type": "string" - } - }, "task_type": { "type": "string" }, - "template": { - "type": "string" - }, "updated_at": { "type": "string" } @@ -17886,39 +17993,37 @@ "TaskExecutionStatusFailed" ] }, - "model.TestPushChannelRequest": { + "model.Template": { "type": "object", "properties": { + "content": { + "type": "string" + }, + "created_at": { + "type": "string" + }, + "description": { + "type": "string" + }, + "id": { + "type": "integer" + }, + "is_system": { + "type": "boolean" + }, + "key": { + "type": "string" + }, "name": { "type": "string" }, - "other": { - "type": "string" - }, - "target": { - "type": "string" - }, - "token": { + "subject": { "type": "string" }, "type": { "type": "string" }, - "url": { - "type": "string" - } - } - }, - "model.TestPushRequest": { - "type": "object", - "required": [ - "config" - ], - "properties": { - "config": { - "$ref": "#/definitions/push.Config" - }, - "target": { + "updated_at": { "type": "string" } } @@ -18016,81 +18121,6 @@ } } }, - "model.UpdateChannelRequest": { - "type": "object", - "properties": { - "credentials": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "extra": { - "type": "object", - "additionalProperties": { - "type": "string" - } - }, - "name": { - "type": "string" - } - } - }, - "model.UpdatePushChannelRequest": { - "type": "object", - "required": [ - "type" - ], - "properties": { - "description": { - "type": "string" - }, - "enabled": { - "type": "boolean" - }, - "other": { - "type": "string" - }, - "token": { - "type": "string" - }, - "type": { - "type": "string" - }, - "url": { - "type": "string" - } - } - }, - "model.UpdatePushEventRequest": { - "type": "object", - "required": [ - "template" - ], - "properties": { - "channels": { - "type": "array", - "items": { - "type": "string" - } - }, - "enabled": { - "type": "boolean" - }, - "targets": { - "type": "array", - "items": { - "type": "string" - } - }, - "template": { - "type": "string" - } - } - }, "model.UpdateScheduleRequest": { "type": "object", "required": [ @@ -20517,6 +20547,10 @@ "description": "AppID 或 SMTP 用户名", "type": "string" }, + "other": { + "description": "附加配置 (如 ChatID / UserKey / 扩展 JSON)", + "type": "string" + }, "secret": { "description": "签名密钥或 SMTP 密码/Token", "type": "string" @@ -20885,6 +20919,17 @@ } } }, + "user.sendEmailCodeRequest": { + "type": "object", + "required": [ + "email" + ], + "properties": { + "email": { + "type": "string" + } + } + }, "user.updateProfileRequest": { "type": "object", "properties": { diff --git a/docs/swagger.yaml b/docs/swagger.yaml index 18668c08..0c20ed10 100644 --- a/docs/swagger.yaml +++ b/docs/swagger.yaml @@ -156,26 +156,6 @@ definitions: type: type: string type: object - Wavelet_plugins_domain_admin_model.Schedule: - properties: - created_at: - type: string - cron: - type: string - id: - example: "0" - type: string - is_active: - type: boolean - name: - type: string - payload: - type: string - task_type: - type: string - updated_at: - type: string - type: object Wavelet_plugins_domain_admin_model.SystemConfig: properties: created_at: @@ -233,29 +213,6 @@ definitions: updated_at: type: string type: object - Wavelet_plugins_domain_admin_model.Template: - properties: - content: - type: string - created_at: - type: string - description: - type: string - id: - type: integer - is_system: - type: boolean - key: - type: string - name: - type: string - subject: - type: string - type: - type: string - updated_at: - type: string - type: object agent.ActiveConfigMeta: properties: checksum: @@ -456,119 +413,6 @@ definitions: totalPage: type: integer type: object - auth.AuthSourceView: - properties: - client_secret_configured: - type: boolean - display_name: - type: string - icon_url: - type: string - id: - type: integer - is_active: - type: boolean - name: - type: string - type: - type: string - type: object - auth.BasicUserInfo: - properties: - avatar_url: - type: string - bio: - type: string - email: - type: string - gender: - type: string - id: - type: integer - is_admin: - type: boolean - location: - type: string - need_change_password: - type: boolean - nickname: - type: string - phone: - type: string - username: - type: string - website: - type: string - type: object - auth.CallbackRequest: - properties: - code: - type: string - state: - type: string - required: - - code - - state - type: object - auth.OAuthAuthorizeResponse: - properties: - authorize_url: - type: string - type: object - auth.OAuthCallbackResult: - properties: - status: - type: string - user: - $ref: '#/definitions/auth.BasicUserInfo' - type: object - cap.ChallengeResponse: - properties: - challenge: - properties: - c: - type: integer - d: - type: integer - s: - type: integer - type: object - expires: - description: ms timestamp - type: integer - token: - type: string - type: object - cap.RedeemResponse: - properties: - error: - type: string - expires: - type: integer - success: - type: boolean - token: - type: string - type: object - cap.challengeRequest: - properties: - scope: - type: string - type: object - cap.redeemRequest: - properties: - scope: - type: string - solutions: - items: - type: integer - type: array - token: - type: string - required: - - solutions - - token - type: object cloudflare.AvailableDomain: properties: domain: @@ -996,6 +840,379 @@ definitions: ttl_minutes: type: integer type: object + do.BindRequest: + properties: + channel_id: + type: string + code: + type: string + type: object + do.BindingDTO: + properties: + channel_id: + example: "0" + type: string + channel_name: + type: string + channel_type: + type: string + created_at: + type: string + id: + example: "0" + type: string + platform_user_id: + type: string + user_id: + example: "0" + type: string + type: object + do.ChannelDTO: + properties: + credentials: + additionalProperties: + type: string + type: object + enabled: + type: boolean + extra: + additionalProperties: + type: string + type: object + id: + example: "0" + type: string + name: + type: string + owner_id: + example: "0" + type: string + owner_scope: + type: string + type: + type: string + type: object + do.CreateChannelRequest: + properties: + credentials: + additionalProperties: + type: string + type: object + enabled: + type: boolean + extra: + additionalProperties: + type: string + type: object + name: + type: string + type: + type: string + type: object + do.CreatePushChannelRequest: + properties: + description: + type: string + enabled: + type: boolean + name: + type: string + other: + type: string + token: + type: string + type: + type: string + url: + type: string + required: + - name + - type + type: object + do.CreatePushEventRequest: + properties: + channels: + items: + type: string + type: array + enabled: + type: boolean + event_key: + type: string + targets: + items: + type: string + type: array + task_type: + type: string + template: + type: string + type: object + do.Definition: + properties: + fields: + items: + $ref: '#/definitions/do.Field' + type: array + type: + type: string + type: object + do.Field: + properties: + key: + type: string + required: + type: boolean + type: + type: string + type: object + do.PublicChannelDTO: + properties: + id: + example: "0" + type: string + name: + type: string + type: + type: string + type: object + do.TestPushChannelRequest: + properties: + name: + type: string + other: + type: string + target: + type: string + token: + type: string + type: + type: string + url: + type: string + type: object + do.TestPushRequest: + properties: + config: + $ref: '#/definitions/push.Config' + target: + type: string + required: + - config + type: object + do.UpdateChannelRequest: + properties: + credentials: + additionalProperties: + type: string + type: object + enabled: + type: boolean + extra: + additionalProperties: + type: string + type: object + name: + type: string + type: object + do.UpdatePushChannelRequest: + properties: + description: + type: string + enabled: + type: boolean + other: + type: string + token: + type: string + type: + type: string + url: + type: string + required: + - type + type: object + do.UpdatePushEventRequest: + properties: + channels: + items: + type: string + type: array + enabled: + type: boolean + targets: + items: + type: string + type: array + template: + type: string + required: + - template + type: object + dto.AuthSourceView: + properties: + client_secret_configured: + type: boolean + display_name: + type: string + icon_url: + type: string + id: + type: integer + is_active: + type: boolean + name: + type: string + type: + type: string + type: object + dto.BasicUserInfo: + properties: + avatar_url: + type: string + bio: + type: string + email: + type: string + gender: + type: string + id: + example: "0" + type: string + is_admin: + type: boolean + location: + type: string + need_change_password: + type: boolean + nickname: + type: string + phone: + type: string + username: + type: string + website: + type: string + type: object + dto.CallbackRequest: + properties: + code: + type: string + state: + type: string + required: + - code + - state + type: object + dto.ChallengeRequest: + properties: + scope: + type: string + type: object + dto.ChallengeResponse: + properties: + challenge: + properties: + c: + type: integer + d: + type: integer + s: + type: integer + type: object + expires: + description: ms timestamp + type: integer + token: + type: string + type: object + dto.OAuthAuthorizeResponse: + properties: + authorize_url: + type: string + type: object + dto.OAuthCallbackResult: + properties: + status: + type: string + user: + $ref: '#/definitions/dto.BasicUserInfo' + type: object + dto.RedeemRequest: + properties: + scope: + type: string + solutions: + items: + type: integer + type: array + token: + type: string + required: + - solutions + - token + type: object + dto.RedeemResponse: + properties: + error: + type: string + expires: + type: integer + success: + type: boolean + token: + type: string + type: object + entity.PushChannel: + properties: + created_at: + type: string + description: + type: string + enabled: + type: boolean + id: + type: integer + name: + type: string + other: + type: string + token: + type: string + type: + type: string + updated_at: + type: string + url: + type: string + type: object + entity.PushEvent: + properties: + channels: + items: + type: string + type: array + created_at: + type: string + enabled: + type: boolean + event_key: + type: string + id: + type: integer + name: + type: string + targets: + items: + type: string + type: array + task_type: + type: string + template: + type: string + updated_at: + type: string + type: object flared.ApplyLogPayload: properties: checksum: @@ -1202,33 +1419,6 @@ definitions: url: type: string type: object - model.BindRequest: - properties: - channel_id: - type: string - code: - type: string - type: object - model.BindingDTO: - properties: - channel_id: - example: "0" - type: string - channel_name: - type: string - channel_type: - type: string - created_at: - type: string - id: - example: "0" - type: string - platform_user_id: - type: string - user_id: - example: "0" - type: string - type: object model.BrowserItem: properties: browser: @@ -1236,31 +1426,6 @@ definitions: count: type: integer type: object - model.ChannelDTO: - properties: - credentials: - additionalProperties: - type: string - type: object - enabled: - type: boolean - extra: - additionalProperties: - type: string - type: object - id: - example: "0" - type: string - name: - type: string - owner_id: - example: "0" - type: string - owner_scope: - type: string - type: - type: string - type: object model.ConfigVersion: properties: checksum: @@ -1299,62 +1464,6 @@ definitions: version: type: string type: object - model.CreateChannelRequest: - properties: - credentials: - additionalProperties: - type: string - type: object - enabled: - type: boolean - extra: - additionalProperties: - type: string - type: object - name: - type: string - type: - type: string - type: object - model.CreatePushChannelRequest: - properties: - description: - type: string - enabled: - type: boolean - name: - type: string - other: - type: string - token: - type: string - type: - type: string - url: - type: string - required: - - name - - type - type: object - model.CreatePushEventRequest: - properties: - channels: - items: - type: string - type: array - enabled: - type: boolean - event_key: - type: string - targets: - items: - type: string - type: array - task_type: - type: string - template: - type: string - type: object model.CreateScheduleRequest: properties: cron: @@ -1485,15 +1594,6 @@ definitions: version: type: string type: object - model.Definition: - properties: - fields: - items: - $ref: '#/definitions/model.Field' - type: array - type: - type: string - type: object model.DispatchTaskRequest: properties: end_time: @@ -1535,15 +1635,6 @@ definitions: description: '"select" 或 "exec"' type: string type: object - model.Field: - properties: - key: - type: string - required: - type: boolean - type: - type: string - type: object model.ListUsersResponse: properties: total: @@ -1756,63 +1847,23 @@ definitions: value: type: string type: object - model.PublicChannelDTO: + model.Schedule: properties: + created_at: + type: string + cron: + type: string id: example: "0" type: string - name: - type: string - type: - type: string - type: object - model.PushChannel: - properties: - created_at: - type: string - description: - type: string - enabled: + is_active: type: boolean - id: - type: integer name: type: string - other: + payload: type: string - token: - type: string - type: - type: string - updated_at: - type: string - url: - type: string - type: object - model.PushEvent: - properties: - channels: - items: - type: string - type: array - created_at: - type: string - enabled: - type: boolean - event_key: - type: string - id: - type: integer - name: - type: string - targets: - items: - type: string - type: array task_type: type: string - template: - type: string updated_at: type: string type: object @@ -1930,30 +1981,29 @@ definitions: - TaskExecutionStatusRunning - TaskExecutionStatusSucceeded - TaskExecutionStatusFailed - model.TestPushChannelRequest: + model.Template: properties: + content: + type: string + created_at: + type: string + description: + type: string + id: + type: integer + is_system: + type: boolean + key: + type: string name: type: string - other: - type: string - target: - type: string - token: + subject: type: string type: type: string - url: + updated_at: type: string type: object - model.TestPushRequest: - properties: - config: - $ref: '#/definitions/push.Config' - target: - type: string - required: - - config - type: object model.TestSMTPRequest: properties: smtp_host: @@ -2018,55 +2068,6 @@ definitions: - max_size_mb - ttl_minutes type: object - model.UpdateChannelRequest: - properties: - credentials: - additionalProperties: - type: string - type: object - enabled: - type: boolean - extra: - additionalProperties: - type: string - type: object - name: - type: string - type: object - model.UpdatePushChannelRequest: - properties: - description: - type: string - enabled: - type: boolean - other: - type: string - token: - type: string - type: - type: string - url: - type: string - required: - - type - type: object - model.UpdatePushEventRequest: - properties: - channels: - items: - type: string - type: array - enabled: - type: boolean - targets: - items: - type: string - type: array - template: - type: string - required: - - template - type: object model.UpdateScheduleRequest: properties: cron: @@ -3669,6 +3670,9 @@ definitions: key: description: AppID 或 SMTP 用户名 type: string + other: + description: 附加配置 (如 ChatID / UserKey / 扩展 JSON) + type: string secret: description: 签名密钥或 SMTP 密码/Token type: string @@ -3917,6 +3921,13 @@ definitions: - password - username type: object + user.sendEmailCodeRequest: + properties: + email: + type: string + required: + - email + type: object user.updateProfileRequest: properties: avatar_url: @@ -4884,7 +4895,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.ChannelDTO' + $ref: '#/definitions/do.ChannelDTO' type: array type: object security: @@ -4902,7 +4913,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.CreateChannelRequest' + $ref: '#/definitions/do.CreateChannelRequest' produces: - application/json responses: @@ -4913,7 +4924,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.ChannelDTO' + $ref: '#/definitions/do.ChannelDTO' type: object "400": description: Bad Request @@ -4964,7 +4975,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.UpdateChannelRequest' + $ref: '#/definitions/do.UpdateChannelRequest' produces: - application/json responses: @@ -4975,7 +4986,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.ChannelDTO' + $ref: '#/definitions/do.ChannelDTO' type: object "400": description: Bad Request @@ -5033,7 +5044,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.Definition' + $ref: '#/definitions/do.Definition' type: array type: object security: @@ -5055,7 +5066,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.PushChannel' + $ref: '#/definitions/entity.PushChannel' type: array type: object security: @@ -5073,7 +5084,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.CreatePushChannelRequest' + $ref: '#/definitions/do.CreatePushChannelRequest' produces: - application/json responses: @@ -5084,7 +5095,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.PushChannel' + $ref: '#/definitions/entity.PushChannel' type: object security: - SessionCookie: [] @@ -5129,7 +5140,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.UpdatePushChannelRequest' + $ref: '#/definitions/do.UpdatePushChannelRequest' produces: - application/json responses: @@ -5140,7 +5151,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.PushChannel' + $ref: '#/definitions/entity.PushChannel' type: object security: - SessionCookie: [] @@ -5173,7 +5184,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.TestPushChannelRequest' + $ref: '#/definitions/do.TestPushChannelRequest' produces: - application/json responses: @@ -5200,7 +5211,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.PushEvent' + $ref: '#/definitions/entity.PushEvent' type: array type: object security: @@ -5218,7 +5229,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.CreatePushEventRequest' + $ref: '#/definitions/do.CreatePushEventRequest' produces: - application/json responses: @@ -5229,7 +5240,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.PushEvent' + $ref: '#/definitions/entity.PushEvent' type: object security: - SessionCookie: [] @@ -5277,7 +5288,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.UpdatePushEventRequest' + $ref: '#/definitions/do.UpdatePushEventRequest' produces: - application/json responses: @@ -5379,7 +5390,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.TestPushRequest' + $ref: '#/definitions/do.TestPushRequest' produces: - application/json responses: @@ -5862,7 +5873,7 @@ paths: - properties: data: items: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Schedule' + $ref: '#/definitions/model.Schedule' type: array type: object "401": @@ -5899,7 +5910,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Schedule' + $ref: '#/definitions/model.Schedule' type: object "400": description: Cron 表达式无效、异步任务类型不存在或参数错误 @@ -5990,7 +6001,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Schedule' + $ref: '#/definitions/model.Schedule' type: object "400": description: Cron 表达式无效、参数错误 @@ -6061,7 +6072,7 @@ paths: - properties: data: items: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Template' + $ref: '#/definitions/model.Template' type: array type: object "401": @@ -6189,7 +6200,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Template' + $ref: '#/definitions/model.Template' type: object "401": description: 未登录 @@ -6238,7 +6249,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/Wavelet_plugins_domain_admin_model.Template' + $ref: '#/definitions/model.Template' type: object "400": description: 参数错误 @@ -7221,7 +7232,7 @@ paths: in: body name: request schema: - $ref: '#/definitions/cap.challengeRequest' + $ref: '#/definitions/dto.ChallengeRequest' produces: - application/json responses: @@ -7232,7 +7243,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/cap.ChallengeResponse' + $ref: '#/definitions/dto.ChallengeResponse' type: object "500": description: 内部服务错误 @@ -7250,7 +7261,7 @@ paths: in: body name: request schema: - $ref: '#/definitions/cap.challengeRequest' + $ref: '#/definitions/dto.ChallengeRequest' produces: - application/json responses: @@ -7261,7 +7272,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/cap.ChallengeResponse' + $ref: '#/definitions/dto.ChallengeResponse' type: object "500": description: 内部服务错误 @@ -7281,7 +7292,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/cap.redeemRequest' + $ref: '#/definitions/dto.RedeemRequest' produces: - application/json responses: @@ -7292,7 +7303,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/cap.RedeemResponse' + $ref: '#/definitions/dto.RedeemResponse' type: object "400": description: 参数错误或核销失败 @@ -7309,7 +7320,7 @@ paths: get: consumes: - application/json - description: 返回系统配置表中 visibility 为 1 的配置键值集合 + description: 返回系统配置表中 visibility 为 1 的扁平键值集合(如 cap_login_enabled) produces: - application/json responses: @@ -12204,7 +12215,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.BindingDTO' + $ref: '#/definitions/do.BindingDTO' type: array type: object "401": @@ -12227,7 +12238,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/model.BindRequest' + $ref: '#/definitions/do.BindRequest' produces: - application/json responses: @@ -12238,7 +12249,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/model.BindingDTO' + $ref: '#/definitions/do.BindingDTO' type: object "400": description: Bad Request @@ -12296,7 +12307,7 @@ paths: - properties: data: items: - $ref: '#/definitions/model.PublicChannelDTO' + $ref: '#/definitions/do.PublicChannelDTO' type: array type: object "401": @@ -12331,7 +12342,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.OAuthAuthorizeResponse' + $ref: '#/definitions/dto.OAuthAuthorizeResponse' type: object "400": description: 认证源不存在或未启用 @@ -12355,7 +12366,7 @@ paths: name: request required: true schema: - $ref: '#/definitions/auth.CallbackRequest' + $ref: '#/definitions/dto.CallbackRequest' produces: - application/json responses: @@ -12366,7 +12377,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.OAuthCallbackResult' + $ref: '#/definitions/dto.OAuthCallbackResult' type: object "400": description: state 无效、参数错误或认证源错误 @@ -12459,7 +12470,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.OAuthAuthorizeResponse' + $ref: '#/definitions/dto.OAuthAuthorizeResponse' type: object "400": description: 认证源不存在或未配置 @@ -12510,7 +12521,7 @@ paths: - properties: data: items: - $ref: '#/definitions/auth.AuthSourceView' + $ref: '#/definitions/dto.AuthSourceView' type: array type: object summary: 获取可用登录源 @@ -12529,7 +12540,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.BasicUserInfo' + $ref: '#/definitions/dto.BasicUserInfo' type: object "401": description: 未登录 @@ -12905,7 +12916,7 @@ paths: - $ref: '#/definitions/response.Any' - properties: data: - $ref: '#/definitions/auth.BasicUserInfo' + $ref: '#/definitions/dto.BasicUserInfo' type: object "401": description: 未登录 @@ -13063,7 +13074,7 @@ paths: post: consumes: - application/json - description: 使用用户名和密码登录,登录成功后建立 Session。若管理员已关闭密码登录功能则返回错误。 + description: 使用用户名和密码登录系统,验证通过后建立 Session 并返回用户信息。 parameters: - description: 登录请求参数 in: body @@ -13082,11 +13093,15 @@ paths: description: 用户名或密码错误 schema: $ref: '#/definitions/response.Any' + "429": + description: 登录尝试过于频繁 + schema: + $ref: '#/definitions/response.Any' "500": description: 服务内部错误 schema: $ref: '#/definitions/response.Any' - summary: 用户密码登录 + summary: 用户登录 tags: - user /api/v1/user/logout: @@ -13166,6 +13181,10 @@ paths: description: 参数错误、用户名已存在或注册已关闭 schema: $ref: '#/definitions/response.Any' + "429": + description: 注册尝试过于频繁 + schema: + $ref: '#/definitions/response.Any' "500": description: 服务内部错误 schema: @@ -13197,6 +13216,13 @@ paths: consumes: - application/json description: 向指定邮箱发送验证码(用于注册场景) + parameters: + - description: 目标邮箱 + in: body + name: request + required: true + schema: + $ref: '#/definitions/user.sendEmailCodeRequest' produces: - application/json responses: @@ -13208,6 +13234,10 @@ paths: description: 参数错误 schema: $ref: '#/definitions/response.Any' + "500": + description: 发送失败 + schema: + $ref: '#/definitions/response.Any' summary: 发送邮箱验证码 tags: - user diff --git a/frontend/app/(main)/admin/logs/components/app-logs.tsx b/frontend/app/(main)/admin/logs/components/app-logs.tsx index 4185e3f7..a4fcb5f5 100644 --- a/frontend/app/(main)/admin/logs/components/app-logs.tsx +++ b/frontend/app/(main)/admin/logs/components/app-logs.tsx @@ -255,10 +255,16 @@ export function AppLogs() { // ---- Initialize ------------------------------------------------------ useEffect(() => { - loadHistory(0).then(() => connectWs()); + let cancelled = false; + loadHistory(0).then(() => { + if (cancelled) return; + connectWs(); + }); return () => { - wsRef.current?.close(); + cancelled = true; + const ws = wsRef.current; wsRef.current = null; + ws?.close(); }; // eslint-disable-next-line react-hooks/exhaustive-deps }, []); diff --git a/frontend/app/(main)/admin/push/components/events-tab.tsx b/frontend/app/(main)/admin/push/components/events-tab.tsx index d0acce60..4b9aad4c 100644 --- a/frontend/app/(main)/admin/push/components/events-tab.tsx +++ b/frontend/app/(main)/admin/push/components/events-tab.tsx @@ -712,8 +712,8 @@ export function EventsTab() { {newEventType === 'task' - ? t('taskTemplateVars') - : t('eventTemplateVars')} + ? t.raw('taskTemplateVars') + : t.raw('eventTemplateVars')}