diff --git a/cmd/agent/main.go b/cmd/agent/main.go
index 49303854..73bc2a2f 100644
--- a/cmd/agent/main.go
+++ b/cmd/agent/main.go
@@ -66,6 +66,7 @@ func main() {
"lua_dir", cfg.LuaDir,
"runtime_config_dir", cfg.RuntimeConfigDir,
"mmdb_path", cfg.MMDBPath,
+ "city_mmdb_path", cfg.CityMMDBPath,
)
client := httpclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration())
@@ -81,6 +82,8 @@ func main() {
LuaDir: cfg.LuaDir,
NginxLuaDir: cfg.OpenrestyLuaDir,
RuntimeConfigDir: cfg.RuntimeConfigDir,
+ MMDBPath: cfg.MMDBPath,
+ CityMMDBPath: cfg.CityMMDBPath,
PagesDir: cfg.PagesDir,
OpenrestyObservabilityListen: nginx.ObservabilityListenAddress(cfg.OpenrestyObservabilityPort),
OpenrestyObservabilityPort: cfg.OpenrestyObservabilityPort,
@@ -122,10 +125,9 @@ func main() {
}
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
- geoIPUpdater := &geoipupdate.Updater{
- MMDBPath: cfg.MMDBPath,
- DownloadURL: cfg.MMDBDownloadURL,
- UpdateInterval: cfg.MMDBUpdateInterval.Duration(),
+ geoIPUpdater := newGeoIPUpdater(cfg)
+ if err = geoIPUpdater.EnsureInitialDatabases(ctx); err != nil {
+ slog.Warn("failed to prepare GeoIP databases before agent startup", "error", err)
}
go geoIPUpdater.Run(ctx)
slog.Info("agent process started")
@@ -138,3 +140,13 @@ func main() {
stop()
slog.Info("agent process stopped")
}
+
+func newGeoIPUpdater(cfg *config.Config) *geoipupdate.Updater {
+ return &geoipupdate.Updater{
+ MMDBPath: cfg.MMDBPath,
+ DownloadURL: cfg.MMDBDownloadURL,
+ CityMMDBPath: cfg.CityMMDBPath,
+ CityDownloadURL: cfg.CityMMDBDownloadURL,
+ UpdateInterval: cfg.MMDBUpdateInterval.Duration(),
+ }
+}
diff --git a/cmd/agent/main_test.go b/cmd/agent/main_test.go
new file mode 100644
index 00000000..2a99adf9
--- /dev/null
+++ b/cmd/agent/main_test.go
@@ -0,0 +1,24 @@
+package main
+
+import (
+ "testing"
+ "time"
+
+ "github.com/Rain-kl/Wavelet/internal/apps/agent/config"
+)
+
+func TestNewGeoIPUpdaterWiresCountryAndCity(t *testing.T) {
+ cfg := &config.Config{
+ MMDBPath: "/data/GeoLite2-Country.mmdb",
+ MMDBDownloadURL: "https://geo.example/GeoLite2-Country.mmdb",
+ CityMMDBPath: "/data/GeoLite2-City.mmdb",
+ CityMMDBDownloadURL: "https://geo.example/GeoLite2-City.mmdb",
+ MMDBUpdateInterval: config.MillisecondDuration(time.Hour),
+ }
+ updater := newGeoIPUpdater(cfg)
+ if updater.MMDBPath != cfg.MMDBPath || updater.DownloadURL != cfg.MMDBDownloadURL ||
+ updater.CityMMDBPath != cfg.CityMMDBPath || updater.CityDownloadURL != cfg.CityMMDBDownloadURL ||
+ updater.UpdateInterval != time.Hour {
+ t.Fatalf("GeoIP updater wiring incomplete: %#v", updater)
+ }
+}
diff --git a/docs/changelog/index.md b/docs/changelog/index.md
index 9c07cc11..28d1b78e 100644
--- a/docs/changelog/index.md
+++ b/docs/changelog/index.md
@@ -21,8 +21,20 @@ sidebar: false
## [unreleased]
+### 新增
+
+- WAF 规则支持可视化 DAG 编排、版本冲突保护和有序路由绑定,发布时编译为 OpenResty 纯内存运行图。
+- WAF IP 组支持 checksum 驱动的 Worker 内存热刷新,并补充 City MMDB 地区匹配数据源。
+
+### 变更
+
+- 移除 WAF 规则旧固定黑白名单、地域名单与 PoW 数据库字段;升级后需在发布前重新编排规则。
+
### 修复
+- WAF 规则编辑器新增启用/停用控制,并移除已废弃的规则组侧站点绑定入口;站点规则顺序统一在反代路由详情中管理。
+- 修复 Agent 心跳无法从活动 WAF 运行图发现 IP 组引用的问题,并对发布与同步的完整 IP 组快照增加 20 MiB 聚合容量保护。
+- 修复 Agent 增量同步长期保留已取消引用的 WAF IP 组、最终阻塞合法热更新的问题;配置同步现按活动引用集合权威收敛,实时广播仅更新本地现存组。
- 修复 Pages 部署文件清单前端请求路径与后端路由不一致导致的 404 错误。
## [v3.2.0] - 2026-07-12
diff --git a/docs/design/waf-design.md b/docs/design/waf-design.md
index 9658b98b..973f39ab 100644
--- a/docs/design/waf-design.md
+++ b/docs/design/waf-design.md
@@ -1,132 +1,15 @@
-# WAF 设计文档
+# WAF 设计
-> WAF 规则正在从固定的白名单、黑名单、PoW 判定链演进为可视化 DAG。新的图模型、执行顺序、发布加载和迁移边界以 [WAF 可编排规则设计](./waf-orchestration-design.md) 为准;本文保留 IP 组、GeoIP 与现有运行时背景说明。
+OpenFlare WAF 的现行规则模型是可视化 DAG。节点语义、图约束、多规则顺序、发布编译与迁移边界统一以 [WAF 可编排规则设计](./waf-orchestration-design.md) 为准。
-你会学到:OpenFlare 边缘 Web 应用防火墙(WAF)的核心架构、动态 IP 组异步差分同步模型、OpenResty Lua 高性能缓存方案以及完整的请求过滤与判定逻辑。
+## 系统边界
----
+Server 保存带坐标和修订号的编辑图,发布时再次校验并编译为紧凑运行图;Agent 原子写入快照并 reload OpenResty;请求热路径只遍历 Worker 内存中的不可变图。
-## 需求分析
+IP 组独立于规则拓扑更新。手动、订阅和自动 IP 组由控制面维护,Agent 先原子替换 JSON、最后更新 checksum。协调 Worker 每 5 秒检查 checksum,仅变化时读取完整快照并分发给其它 Worker;失败时保留上一份有效数据。完整运行时快照上限为 20 MiB,Server 发布/同步与 Agent 落盘使用同一序列化校验;OpenResty 使用独立的 64 MiB 共享字典和非淘汰写入,容量不足时拒绝新版本而不破坏已提交快照。
-在互联网公开环境中,Web 应用程序面临着各种各样的安全威胁(如扫描器踩点、刷接口、针对特定地域的恶意网络爬虫、勒索攻击及 CC 攻击等)。如果直接把恶意请求放行给源站(Origin Server),会导致:
-1. **源站负载飙升**:高频的数据库查询与 CPU 运算极易耗尽服务器资源。
-2. **敏感接口被刷**:登录、注册、短信验证码接口容易被恶意滥用导致财产损失。
-3. **数据泄露风险**:恶意的通用漏洞探测行为无法被提前拦截。
+地域节点使用 Country 与 City MMDB。数据库不可用时地域匹配返回 `false` 并限频告警,不允许因数据损坏意外放行其它执行错误。
-因此,OpenFlare 需要在最前端的数据面(OpenResty)构建一套 **高性能、可弹性伸缩的 WAF 过滤引擎**。该引擎能够在最接近用户的边缘层以毫秒级的极低开销对恶意请求进行深度过滤,减轻源站压力,并提供防 CC(PoW 挑战)、IP 黑白名单与地域级别拦截等核心安全防护能力。
+## 安全顺序
----
-
-## 核心功能
-
-OpenFlare WAF 包含以下核心防护维度:
-
-* **IP 级拦截(IP 黑白名单)**:支持单 IP、CIDR 网段过滤,支持将上万 IP 聚合为 IP 组进行高效比对。
-* **地域黑白名单(GeoIP 限制)**:集成 MaxMind 数据库,支持针对国家(Country)和省份/地区(Region)执行精准准入控制。
-* **自定义拦截响应**:支持针对不同的过滤规则自定义阻断状态码(如 403, 418)以及个性化的 HTML 拦截页面。
-* **人机挑战(PoW CC 防护)**:支持无感人机挑战,通过计算 Hash 碰撞防止自动化脚本和僵尸网络(Botnet)对接口进行并发冲击。
-
----
-
-## IP 组设计与动态异步同步
-
-IP 组是 WAF 进行高效黑白名单管控的核心容器。OpenFlare 将 IP 组根据更新频率与产生渠道分为三类:
-
-### 1. IP 组类型
-* **手动 IP 组(Manual)**:由管理员在控制面板上手动输入 IP 或 CIDR 列表。主要用于静态的信任 IP 或长期的封禁。
-* **订阅 IP 组(Subscription)**:配置远程文本(按行分隔)或标准的 JSON 订阅地址。Server 侧的定时任务会周期性抓取远程订阅源并自动解析导入。主要用于集成开源的威胁情报库、云厂商的 IP 范围等。
-* **自动 IP 组(Automatic)**:**最具弹性的动态防护通道**。控制面的定时扫描任务会读取所有节点的访问日志,按照设定的 Expr 规则(例如:“5分钟内请求 `/api/login` 接口触发 401 超过 50 次”)进行聚合分析,一旦匹配,自动将该恶意源 IP 写入封禁组,并指定封禁时长。
-
-### 2. 异步差分同步设计 (不触发 Nginx Reload)
-在传统的 Nginx WAF 设计中,IP 黑名单的更新通常需要重写配置并 reload。如果恶意 IP 封禁以秒级或分钟级高频触发,频繁 reload 会导致 Nginx 频繁新建 Worker 进程并销毁老进程,导致性能骤降。
-
-OpenFlare 采用 **动态 IP 组异步差分同步设计**:
-
-```text
-WAF IP 成员更新 (手动/订阅/自动自动触发)
- |
- v
-Server 更新数据库并计算该 IP 组的全新 MD5 Checksum
- |
- +----------------------------------------+
- | (WebSocket 实时广播) | (心跳兜底比对)
- v v
-Server 立即向所有 Agent 推送变更组的完整成员 Agent 心跳上报本地所有 IP 组的 Checksum 映射表
- | |
- | v
- | Server 发现 Checksum 不一致,下发变更的 IP 组成员
- v |
-Agent 接收成员数据,将其以 JSON 形式写入本地磁盘路径:waf_ip_groups.json
- |
- v (Lua 内存感知)
-OpenResty Lua 引擎通过 MD5 校验和秒级感知文件变化并热更新内存,无需 reload 进程
-```
-
-通过这一架构,上万个高频变动的动态黑名单 IP 的落地和生效,**全程无需 reload 任何 Nginx 进程**,极大地保护了网关的高并发性能。
-
----
-
-## 规则组与网站绑定
-
-* **WAF 规则组(Rule Group)**:WAF 过滤政策的最小逻辑集合。一条规则组内可以包含 IP 黑白名单、IP 组引用、地域限制及防 CC 挑战配置。
-* **全局规则组(Global)**:当规则组被标记为 `is_global = true` 时,该规则组对节点上托管的**所有网站路由**默认生效。
-* **网站绑定绑定(Site Binding)**:网站路由(Proxy Route)可以绑定一个或多个非全局规则组。判定时,会执行 `全局规则组 + 绑定规则组` 的并集逻辑。
-
----
-
-## 实现方案与高性能缓存
-
-WAF 在 OpenResty 的 `access_by_lua` 阶段被触发,核心由 Lua 文件与本地落地的 JSON 配置构成。
-
-### 1. 物理结构
-* `waf_config.json`:包含所有规则组的元数据、国家地域限制、以及网站(Site)与规则组的关联映射。
-* `waf_ip_groups.json`:包含所有同步下来的 IP 组与对应的 IP 列表。
-* `waf/runtime.lua`:WAF 规则比对的实际运行时引擎。
-* `waf/check.lua`:接入层入口,负责包引入与 check() 触发。
-
-### 2. 共享内存字典 (ngx.shared) 高性能缓存设计
-在每次 Web 请求进来时都读取磁盘上的 JSON 文件并进行解码,会导致磁盘 I/O 成为严重的性能瓶颈。
-
-OpenFlare 利用 **OpenResty 共享内存字典 (ngx.shared.openflare_waf_config)** 设计了二级缓存机制:
-
-1. **零文件 I/O 路径**:
- 在 Lua 中,每次执行 `check()` 时,首先利用 `ngx.md5` 瞬间计算本地磁盘 JSON 文件的 MD5 哈希(这一操作几乎为零耗时,因为文件已被操作系统 Page Cache 缓存)。
-2. **哈希比对与热加载**:
- 比对共享内存中存储的缓存哈希键(`_config_hash`)。
- * **若哈希未发生变化**:直接从共享内存字典中读取已解码、存在内存中的 Lua Table 配置,整个校验过程完全基于**共享内存操作**,耗时在 **微秒级** 级别。
- * **若哈希不一致**:说明 Agent 刚刚落地了新的 WAF 规则或 IP 组,Lua 自动读取磁盘文件并使用 `cjson.decode` 解码,解码后的数据及全新的 MD5 写入共享内存,供后续 Worker 进程无缝读取。
-
----
-
-## 应用流程与判定判定控制逻辑
-
-当一个 HTTP/HTTPS 请求到达 OpenResty 后,WAF 会在 `access` 阶段按下图所示的漏斗判决链进行逐步匹配拦截:
-
-### 1. WAF 判定流程图
-
-```mermaid
-flowchart TD
- A[请求进入 access 阶段] --> B[获取当前请求的 Site Name]
- B --> C[在共享内存中加载与此 Site 绑定的所有活跃规则组]
- C --> D{匹配到 IP 白名单 / 白名单 IP 组?}
- D -- 是 (匹配成功) --> E[放行请求 - ALLOW]
- D -- 否 --> F{匹配到国家/地区地域白名单?}
- F -- 是 (匹配成功) --> E
- F -- 否 --> G{匹配到 IP 黑名单 / 黑名单 IP 组?}
- G -- 是 (匹配成功) --> H[阻断请求 - BLOCK]
- G -- 否 --> I{匹配到国家/地区地域黑名单?}
- I -- 是 (匹配成功) --> H
- I -- 否 --> J{是否启用了防 CC PoW 验证?}
- J -- 是 --> K[转交防 CC 模块处理]
- J -- 否 --> L[无安全风险,正常放行]
-
- H --> M[退出并返回规则组配置的自定义状态码与拦截响应体]
-```
-
-### 2. 判决步骤细则
-1. **白名单前置**:
- 为了防止误杀以及保障核心回源流量(如搜索引擎蜘蛛、CDN 回源 IP、办公区出口)的顺畅,WAF **优先匹配 IP 白名单与地域白名单**。一旦白名单匹配成功,直接绕过后续的所有黑名单检测和 CC 挑战,立刻放行。如果请求未命中白名单,则继续向下进行黑名单检测及其他后续判定。
-2. **黑名单强力阻断**:
- 如果在白名单判定中未被捕获,请求将进入黑名单漏斗。一旦请求源 IP 命中 IP 黑名单、命中引用的黑名单 IP 组、或是处于被禁止的国家/地区范围内,Lua 引擎立即将 `ngx.ctx.openflare_waf_blocked` 标记设为 `true`。
-3. **输出响应**:
- 命中黑名单后,Lua 提取匹配到规则组的 `block_status_code`(默认返回 418 / 403)和 `block_response_body`(拦截页面 HTML),通过 `ngx.say()` 输出响应体并执行 `ngx.exit(status)` 平滑退出请求,防止请求继续向后透传。
+启用的全局规则固定前置;路由规则按绑定 sequence 执行。阻止节点立即终止,通过节点仅结束当前规则,全部规则通过后才进入回源链路。未知节点、缺失出口或步数超限一律阻止请求。
diff --git a/docs/design/waf-orchestration-design.md b/docs/design/waf-orchestration-design.md
index 6cc622fd..9bb979ae 100644
--- a/docs/design/waf-orchestration-design.md
+++ b/docs/design/waf-orchestration-design.md
@@ -70,11 +70,11 @@ IP 组采用协调 Worker、共享快照和 Worker 本地对象的两级缓存
1. 请求始终读取当前 Worker 内存中的 IP 组对象,不访问文件或共享字典中的 JSON。
2. 每 5 秒只有一个取得共享锁的 Worker 读取轻量 checksum 文件。
3. checksum 未变化时立即结束,不读取完整 `waf_ip_groups.json`。
-4. checksum 变化时,协调 Worker 读取并验证一次完整 JSON,再把原始快照与新版本写入 `ngx.shared`。
+4. checksum 变化时,协调 Worker 读取并验证一次完整 JSON,再把原始快照按 checksum 写入独立的 64 MiB `ngx.shared.openflare_waf_ip_groups`,最后更新提交指针。
5. 其他 Worker 发现共享版本变化后,从共享内存取得快照、解析并原子替换各自的本地对象,不重复读取磁盘。
6. 刷新失败时继续使用上一份有效对象,限频记录错误,并在下一周期重试。
-Agent 必须先原子替换 IP 组 JSON,最后原子更新 checksum,使 Worker 永远不会把半写入文件识别为新版本。
+Agent 必须先原子替换 IP 组 JSON,最后原子更新 checksum,使 Worker 永远不会把半写入文件识别为新版本。Server 发布/同步和 Agent 落盘共同执行 20 MiB 聚合快照上限;共享字典使用不会强制淘汰旧键的安全写入,失败时保留当前与上一代不可变快照。
## API 与编辑器
diff --git a/docs/docs.go b/docs/docs.go
index b3ecc49f..1a25dfa5 100644
--- a/docs/docs.go
+++ b/docs/docs.go
@@ -10249,17 +10249,16 @@ const docTemplate = `{
"SessionCookie": []
}
],
- "description": "返回全部 WAF 规则组,需要管理员权限",
"produces": [
"application/json"
],
"tags": [
"openflare-waf"
],
- "summary": "列出 WAF 规则组",
+ "summary": "列出 WAF 规则",
"responses": {
"200": {
- "description": "规则组列表",
+ "description": "规则列表",
"schema": {
"allOf": [
{
@@ -10271,7 +10270,7 @@ const docTemplate = `{
"data": {
"type": "array",
"items": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
@@ -10279,12 +10278,6 @@ const docTemplate = `{
]
}
},
- "400": {
- "description": "参数错误",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- },
"401": {
"description": "未登录",
"schema": {
@@ -10311,7 +10304,6 @@ const docTemplate = `{
"SessionCookie": []
}
],
- "description": "创建新的 WAF 规则组,需要管理员权限",
"consumes": [
"application/json"
],
@@ -10321,21 +10313,21 @@ const docTemplate = `{
"tags": [
"openflare-waf"
],
- "summary": "创建 WAF 规则组",
+ "summary": "创建 WAF 规则",
"parameters": [
{
- "description": "规则组参数",
+ "description": "规则名称",
"name": "request",
"in": "body",
"required": true,
"schema": {
- "$ref": "#/definitions/waf.RuleGroupInput"
+ "$ref": "#/definitions/waf.CreateRuleInput"
}
}
],
"responses": {
"200": {
- "description": "创建成功的规则组",
+ "description": "创建成功",
"schema": {
"allOf": [
{
@@ -10345,7 +10337,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
@@ -10386,18 +10378,17 @@ const docTemplate = `{
"SessionCookie": []
}
],
- "description": "按 ID 返回 WAF 规则组详情,需要管理员权限",
"produces": [
"application/json"
],
"tags": [
"openflare-waf"
],
- "summary": "获取 WAF 规则组详情",
+ "summary": "获取 WAF 规则详情",
"parameters": [
{
"type": "integer",
- "description": "规则组 ID",
+ "description": "规则 ID",
"name": "id",
"in": "path",
"required": true
@@ -10405,7 +10396,7 @@ const docTemplate = `{
],
"responses": {
"200": {
- "description": "规则组详情",
+ "description": "规则详情",
"schema": {
"allOf": [
{
@@ -10415,7 +10406,7 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
@@ -10435,7 +10426,7 @@ const docTemplate = `{
}
},
"404": {
- "description": "记录不存在",
+ "description": "无权限或不存在",
"schema": {
"$ref": "#/definitions/response.Any"
}
@@ -10456,18 +10447,17 @@ const docTemplate = `{
"SessionCookie": []
}
],
- "description": "按 ID 删除 WAF 规则组,需要管理员权限",
"produces": [
"application/json"
],
"tags": [
"openflare-waf"
],
- "summary": "删除 WAF 规则组",
+ "summary": "删除 WAF 规则",
"parameters": [
{
"type": "integer",
- "description": "规则组 ID",
+ "description": "规则 ID",
"name": "id",
"in": "path",
"required": true
@@ -10493,7 +10483,7 @@ const docTemplate = `{
}
},
"404": {
- "description": "记录不存在",
+ "description": "无权限或不存在",
"schema": {
"$ref": "#/definitions/response.Any"
}
@@ -10507,14 +10497,13 @@ const docTemplate = `{
}
}
},
- "/api/v1/d/waf/rule-groups/{id}/sites": {
+ "/api/v1/d/waf/rule-groups/{id}/graph": {
"post": {
"security": [
{
"SessionCookie": []
}
],
- "description": "替换 WAF 规则组关联的代理站点列表,需要管理员权限",
"consumes": [
"application/json"
],
@@ -10524,28 +10513,28 @@ const docTemplate = `{
"tags": [
"openflare-waf"
],
- "summary": "替换规则组站点绑定",
+ "summary": "保存 WAF 规则图",
"parameters": [
{
"type": "integer",
- "description": "规则组 ID",
+ "description": "规则 ID",
"name": "id",
"in": "path",
"required": true
},
{
- "description": "站点 ID 列表",
+ "description": "规则图和修订号",
"name": "request",
"in": "body",
"required": true,
"schema": {
- "$ref": "#/definitions/waf.IDsRequest"
+ "$ref": "#/definitions/waf.SaveRuleGraphInput"
}
}
],
"responses": {
"200": {
- "description": "更新后的规则组",
+ "description": "保存成功",
"schema": {
"allOf": [
{
@@ -10555,7 +10544,94 @@ const docTemplate = `{
"type": "object",
"properties": {
"data": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
+ }
+ }
+ }
+ ]
+ }
+ },
+ "400": {
+ "description": "参数或规则图错误",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ },
+ "401": {
+ "description": "未登录",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ },
+ "404": {
+ "description": "无权限或不存在",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ },
+ "409": {
+ "description": "修订冲突",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ },
+ "500": {
+ "description": "内部错误",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ }
+ }
+ }
+ },
+ "/api/v1/d/waf/rule-groups/{id}/meta": {
+ "post": {
+ "security": [
+ {
+ "SessionCookie": []
+ }
+ ],
+ "consumes": [
+ "application/json"
+ ],
+ "produces": [
+ "application/json"
+ ],
+ "tags": [
+ "openflare-waf"
+ ],
+ "summary": "更新 WAF 规则元数据",
+ "parameters": [
+ {
+ "type": "integer",
+ "description": "规则 ID",
+ "name": "id",
+ "in": "path",
+ "required": true
+ },
+ {
+ "description": "规则元数据",
+ "name": "request",
+ "in": "body",
+ "required": true,
+ "schema": {
+ "$ref": "#/definitions/waf.UpdateRuleMetaInput"
+ }
+ }
+ ],
+ "responses": {
+ "200": {
+ "description": "更新成功",
+ "schema": {
+ "allOf": [
+ {
+ "$ref": "#/definitions/response.Any"
+ },
+ {
+ "type": "object",
+ "properties": {
+ "data": {
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
@@ -10575,89 +10651,7 @@ const docTemplate = `{
}
},
"404": {
- "description": "记录不存在",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- },
- "500": {
- "description": "内部错误",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- }
- }
- }
- },
- "/api/v1/d/waf/rule-groups/{id}/update": {
- "post": {
- "security": [
- {
- "SessionCookie": []
- }
- ],
- "description": "按 ID 更新 WAF 规则组,需要管理员权限",
- "consumes": [
- "application/json"
- ],
- "produces": [
- "application/json"
- ],
- "tags": [
- "openflare-waf"
- ],
- "summary": "更新 WAF 规则组",
- "parameters": [
- {
- "type": "integer",
- "description": "规则组 ID",
- "name": "id",
- "in": "path",
- "required": true
- },
- {
- "description": "规则组参数",
- "name": "request",
- "in": "body",
- "required": true,
- "schema": {
- "$ref": "#/definitions/waf.RuleGroupInput"
- }
- }
- ],
- "responses": {
- "200": {
- "description": "更新后的规则组",
- "schema": {
- "allOf": [
- {
- "$ref": "#/definitions/response.Any"
- },
- {
- "type": "object",
- "properties": {
- "data": {
- "$ref": "#/definitions/waf.RuleGroupView"
- }
- }
- }
- ]
- }
- },
- "400": {
- "description": "参数错误",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- },
- "401": {
- "description": "未登录",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- },
- "404": {
- "description": "记录不存在",
+ "description": "无权限或不存在",
"schema": {
"$ref": "#/definitions/response.Any"
}
@@ -18531,6 +18525,14 @@ const docTemplate = `{
}
}
},
+ "waf.CreateRuleInput": {
+ "type": "object",
+ "properties": {
+ "name": {
+ "type": "string"
+ }
+ }
+ },
"waf.IDsRequest": {
"type": "object",
"properties": {
@@ -18716,139 +18718,97 @@ const docTemplate = `{
}
}
},
- "waf.PoWConfig": {
+ "waf.RuleEdge": {
"type": "object",
"properties": {
- "algorithm": {
+ "id": {
"type": "string"
},
- "blacklist": {
- "$ref": "#/definitions/waf.PoWListConfig"
+ "source": {
+ "type": "string"
},
- "challenge_ttl": {
- "type": "integer"
+ "source_handle": {
+ "type": "string"
},
- "difficulty": {
- "type": "integer"
- },
- "session_ttl": {
- "type": "integer"
- },
- "whitelist": {
- "$ref": "#/definitions/waf.PoWListConfig"
+ "target": {
+ "type": "string"
}
}
},
- "waf.PoWListConfig": {
+ "waf.RuleGraph": {
"type": "object",
"properties": {
- "ip_cidrs": {
+ "edges": {
"type": "array",
"items": {
- "type": "string"
+ "$ref": "#/definitions/waf.RuleEdge"
}
},
- "ips": {
+ "nodes": {
"type": "array",
"items": {
- "type": "string"
+ "$ref": "#/definitions/waf.RuleNode"
}
},
- "path_regexes": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "paths": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "user_agents": {
- "type": "array",
- "items": {
- "type": "string"
- }
+ "schema_version": {
+ "type": "integer"
}
}
},
- "waf.RuleGroupInput": {
+ "waf.RuleNode": {
"type": "object",
"properties": {
- "block_response_body": {
- "type": "string"
- },
- "block_status_code": {
- "type": "integer"
- },
- "country_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "country_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "enabled": {
- "type": "boolean"
- },
- "ip_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "ip_blacklist_group_ids": {
+ "config": {
"type": "array",
"items": {
"type": "integer"
}
},
- "ip_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "ip_whitelist_group_ids": {
- "type": "array",
- "items": {
- "type": "integer"
- }
- },
- "name": {
+ "id": {
"type": "string"
},
- "pow_config": {
- "type": "array",
- "items": {
- "type": "integer"
- }
+ "label": {
+ "type": "string"
},
- "pow_enabled": {
- "type": "boolean"
+ "position": {
+ "$ref": "#/definitions/waf.RulePosition"
},
- "region_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "region_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
+ "type": {
+ "$ref": "#/definitions/waf.RuleNodeType"
}
}
},
- "waf.RuleGroupView": {
+ "waf.RuleNodeType": {
+ "type": "string",
+ "enum": [
+ "start",
+ "allow",
+ "block",
+ "ip_match",
+ "geo_match",
+ "pow"
+ ],
+ "x-enum-varnames": [
+ "RuleNodeStart",
+ "RuleNodeAllow",
+ "RuleNodeBlock",
+ "RuleNodeIPMatch",
+ "RuleNodeGeoMatch",
+ "RuleNodePoW"
+ ]
+ },
+ "waf.RulePosition": {
+ "type": "object",
+ "properties": {
+ "x": {
+ "type": "number"
+ },
+ "y": {
+ "type": "number"
+ }
+ }
+ },
+ "waf.RuleView": {
"type": "object",
"properties": {
"applied_site_count": {
@@ -18860,86 +18820,43 @@ const docTemplate = `{
"type": "integer"
}
},
- "block_response_body": {
- "type": "string"
- },
- "block_status_code": {
- "type": "integer"
- },
- "country_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "country_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
"created_at": {
"type": "string"
},
"enabled": {
"type": "boolean"
},
+ "graph": {
+ "$ref": "#/definitions/waf.RuleGraph"
+ },
"id": {
"type": "integer"
},
- "ip_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "ip_blacklist_group_ids": {
- "type": "array",
- "items": {
- "type": "integer"
- }
- },
- "ip_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "ip_whitelist_group_ids": {
- "type": "array",
- "items": {
- "type": "integer"
- }
- },
"is_global": {
"type": "boolean"
},
"name": {
"type": "string"
},
- "pow_config": {
- "$ref": "#/definitions/waf.PoWConfig"
- },
- "pow_enabled": {
- "type": "boolean"
- },
- "region_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "region_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
+ "revision": {
+ "type": "integer"
},
"updated_at": {
"type": "string"
}
}
},
+ "waf.SaveRuleGraphInput": {
+ "type": "object",
+ "properties": {
+ "graph": {
+ "$ref": "#/definitions/waf.RuleGraph"
+ },
+ "revision": {
+ "type": "integer"
+ }
+ }
+ },
"waf.SiteRuleGroupsView": {
"type": "object",
"properties": {
@@ -18952,11 +18869,11 @@ const docTemplate = `{
"applied_rule_groups": {
"type": "array",
"items": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
},
"global_rule_group": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
},
"route_id": {
"type": "integer"
@@ -18964,11 +18881,22 @@ const docTemplate = `{
"rule_groups": {
"type": "array",
"items": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
},
+ "waf.UpdateRuleMetaInput": {
+ "type": "object",
+ "properties": {
+ "enabled": {
+ "type": "boolean"
+ },
+ "name": {
+ "type": "string"
+ }
+ }
+ },
"zone.DomainInput": {
"type": "object",
"properties": {
diff --git a/docs/guide/waf-usage.md b/docs/guide/waf-usage.md
index 0196900d..c9c31e66 100644
--- a/docs/guide/waf-usage.md
+++ b/docs/guide/waf-usage.md
@@ -1,176 +1,25 @@
# WAF 安全防护使用
-你会学到:OpenFlare 边缘 Web 应用防火墙 (WAF) 的工作原理、防护维度,如何管理与引用三类 IP 组(手动、订阅与基于 Expr 的自动 IP 组),配置防 CC 挑战(PoW 人机验证)与地域级拦截,以及如何在不 reload 进程的情况下实现 IP 组成员的秒级热更新。
+OpenFlare WAF 使用可视化有向无环图编排规则。新建规则时只填写名称,系统创建默认的“开始 → 通过”图并进入编辑器。
----
+## 节点与连线
-## 核心概念
+- **开始**:每条规则唯一,沿 `next` 进入图。
+- **通过**:结束当前规则;若路由仍有后续规则则继续执行。
+- **阻止**:立即按配置的状态码和 HTML 响应终止请求。
+- **IP 匹配**:配置 IP、CIDR 或 IP 组,分别连接 `true`、`false`。
+- **地域匹配**:按国家或地区代码分支;City MMDB 不可用时按未匹配处理。
+- **PoW**:未完成挑战时接管请求,验证通过后沿 `next` 继续。
-在配置安全策略前,你需要理解 WAF 的几个核心组成部分:
+服务端会拒绝循环、悬空出口、不可达节点、重复端口连接和无效配置。保存时携带页面加载得到的 `revision`;发生 409 冲突时应重新加载,避免覆盖他人修改。
-| 概念 | 说明 | 作用范围与生效方式 |
-| --- | --- | --- |
-| **WAF 规则组 (Rule Group)** | 安全规则的逻辑集合。包括:IP 黑白名单(直接录入或引用 IP 组)、国家/地区地域限制、防 CC 挑战(PoW)以及自定义拦截响应。 | 支持全局生效或绑定到单个/多个网站。**修改规则组定义必须发布并激活配置版本**。 |
-| **IP 组 (IP Group)** | 存放单个 IP 或 CIDR 网段的列表容器。分为**手动**、**订阅**与**自动**三类。WAF 规则组可通过 ID 引用 IP 组。 | 属于动态资源。**IP 组成员的增减支持 WebSocket 秒级无缝热同步,无需 reload 进程**。 |
-| **人机挑战 (CC PoW)** | 基于 Proof of Work (工作量证明) 的人机验证挑战。通过让浏览器计算特定难度的哈希碰撞,静默阻断恶意刷接口的自动化脚本与 Bot,保障正常用户体验。 | 位于规则组内的配置 Tab。**修改 PoW 参数必须发布并激活配置版本**。 |
+## 绑定与生效
----
+启用的全局规则固定最先执行;路由绑定的自定义规则严格按列表顺序执行。调整顺序后需要发布配置版本,规则拓扑才会随 OpenResty reload 生效。
-## 推荐配置顺序
+IP 组成员是动态资源。Agent 每 5 秒检查 checksum,变化后在 Worker 间更新内存快照,无需重新发布规则或 reload。手动、订阅与自动 IP 组均可被 IP 匹配节点引用。单次完整 IP 组运行时快照最多 20 MiB;超过上限时发布或同步会返回错误,并继续使用上一份有效快照。
-配置网站的安全防护时,推荐按这个顺序进行:
+> [!IMPORTANT]
+> 从旧固定黑白名单/地域/PoW 表单升级时,规则图会重置为“开始 → 通过”,旧策略字段不会迁移。请在发布新版本前逐条重新编排并验证规则。
-1. 进入左侧菜单 **「安全性」->「IP 组」**,创建所需的 **手动 IP 组** (如开发者白名单) 或 **自动 IP 组** (如根据 404 扫描自动封禁的 IP)。
-2. 创建或编辑 **WAF 规则组**(菜单路径 **「安全性」->「WAF」**):
- * 绑定需要引用或阻断的 IP 组。
- * 配置国家或省份的地域黑白名单限制。
- * (可选) 在 `PoW` 标签页配置人机挑战参数。
- * 在 `拦截返回` 标签页设定自定义状态码(如 403, 418)和 HTML 拦截页。
-3. 将规则组关联到对应的 **路由规则**(在 **「规则管理」** 页面编辑对应规则,并在「WAF」选项卡中勾选关联规则组)。
-4. 发布并激活配置版本,使边缘节点 (Agent) 开始应用 WAF 规则过滤流量。
-
----
-
-## 详细步骤指南
-
-### 第一步:管理与配置 IP 组
-
-IP 组是进行大批量 IP 过滤的基石。OpenFlare 提供了极富弹性的三类 IP 组:
-
-#### 1. 手动 IP 组 (Manual)
-* **用途**:静态维护一些确定受信任或确定需长期拦截的 IP/网段。
-* **配置**:点击「创建 IP 组」-> 类型选择「手动」-> 按行直接填入 IP 或 CIDR 格式(例如 `192.168.1.100` 或 `10.0.0.0/24`)。
-
-#### 2. 订阅 IP 组 (Subscription)
-* **用途**:接入第三方开源威胁情报库、云厂商公布的官方网段(如 Cloudflare, GitHub Action IP 列表),或团队内部统一维护的动态 IP 源。
-* **配置参数**:
- * **订阅 URL**:必须是合法的 `http` 或 `https` 链接。
- * **订阅格式**:支持 `Text` 与 `JSON` 两种数据格式:
- * **Text 格式**:纯文本格式。按行分隔读取 IP/CIDR,会自动过滤掉以 `#` 开头的注释行和空白行。
- * **JSON 格式**:当订阅源是一个结构化的 JSON 响应时,需要编写 **映射规则 (Mapping Rule)** 从 JSON 数据中提取 IP 列表。
- * **映射规则**:使用类似 JSONPath 的轻量点语法定位 IP 数组,支持以 `[]` 展开数组。例如:
- * 若 JSON 结构为 `{"data": {"ips": ["1.1.1.1", "2.2.2.2"]}}`,则映射规则填写 `$.data.ips[]`(或 `data.ips[]`)。
- * 若 JSON 根节点本身即为字符串数组(如 `["1.1.1.1", "2.2.2.2"]`),映射规则留空或填写 `$` 即可。
- * **同步间隔 (分钟)**:该订阅组自动同步的周期,默认为 `1440` 分钟(24小时),允许范围为 `5` 至 `43200` 分钟。
-* **安全限额与同步频率**:
- * 为防止恶意或超大订阅源造成系统负担,单次抓取上限限制为 **2 MiB**,网络拉取超时为 15 秒。
- * Server 默认每 5 分钟在后台扫描一次到期的订阅 IP 组并拉取同步。
-
-> [!TIP]
-> 关于 WAF 的动态 IP 组异步差分同步模型(WebSocket 实时热同步、不触发 Nginx Reload 机制)以及高性能 Lua 缓存方案等底层设计细节,请参阅 [WAF 设计](../design/waf-design.md)。
-
-#### 3. 自动 IP 组 (Automatic)
-* **用途**:**最具杀伤力的防扫描、防爆破自动通道**。
-* **配置**:类型选择「自动」-> 编写 Expr 日志聚合逻辑。你可以直接引用系统内置的预设:
- * **单 IP 404 高频扫描**:`request_count > 100 && StatusRatio(404) >= 0.8` (单个 IP 最近一小时请求超 100 次且 404 响应占比超 80%)。
- * **单 IP 直连访问异常**:`ip_host_count > 50 && ip_host_ratio > 0.5` (绕过域名直接通过 IP 地址进行高频请求)。
-* **测试与立即执行**:保存前可点击 **「测试规则」** 按钮预览当前日志窗口被命中的 IP。保存后可点击 **「立即执行」** 直接聚合日志并生成封禁名单。
-
-> [!TIP]
-> 自动 IP 组的详细语法和可用指标请参阅 [WAF 自动 IP 组规则语法](./waf-ip-group-expr.md)。
-
----
-
-### 第二步:创建与配置 WAF 规则组
-
-1. 导航至左侧菜单 **「安全性」->「WAF」**,点击 **「创建规则组」**。
-2. 填写规则组名称(如 `production-api-shield`),选择是否为「全局规则组」。
-3. 进入规则组详情,在下方几个配置 Tab 中依次设置:
-
-#### 1. 黑白名单配置 (Allow / Block Lists)
-* **直录 IP**:可直接在框内按行填入临时需要白名单放行或黑名单阻断的单个 IP 或网段。
-* **IP 组引用**:点击「绑定 IP 组」,选择你在第一步中配置好的手动、自动或订阅 IP 组。白名单引用会直接放行,黑名单引用则直接阻断。
-
-#### 2. 地域限制 (GeoIP)
-* **说明**:OpenFlare 集成了 GeoIP 地理位置解析。
-* **配置**:可开启地域限制开关,模式可选择「仅允许」或「禁止」。
-* * 例如,若你的服务只服务于国内,可以将模式设为「仅允许」,并在国家列表中勾选 `中国`。
-* * 支持细化到具体省份/地区(Region),一键拦截特定地理区域的恶意流量。
-
-#### 3. 人机挑战配置 (PoW CC 防护)
-* **说明**:开启防 CC 的人机挑战。当请求触发防CC机制时,浏览器会渲染一个静默挑战页面,并在几百毫秒内完成数学计算(哈希碰撞)。通过后会被写入 Cookie,后续访问直接放行。此过程对真实用户几乎无感,但能完美拦截不支持 JS/不具备计算能力的爆破脚本与 CC 僵尸工具。
-* **核心参数**:
- * **开启状态**:启用/禁用。
- * **哈希难度**:控制碰撞难度(建议设定为 `4` 或 `5`)。
- * **Cookie 有效期**:挑战通过后,在多长时间内免验证(例如 `3600` 秒)。
- * **自定义挑战 HTML**:可定制挑战中的 Loading 页面风格,让其融入你的业务设计。
-
-#### 4. 拦截返回 (Block Response)
-* **说明**:设定 WAF 规则拦截恶意请求时的返回行为。
-* **配置**:
- * **拦截状态码**:可自定义拦截响应的 HTTP 状态码,例如标准的 `403`,或带有趣味性质的 `418 (I'm a teapot)`。
- * **拦截响应体**:可在此输入自定义的 HTML 内容,展示给被拦截的攻击者(如:“WAF 拦截:你的请求已被记录”)。
-
----
-
-### 第三步:将规则组关联到路由规则
-
-规则组配置完成后,并不会自动生效,你需要将其与具体的路由规则绑定。
-
-* **关联配置步骤**:进入 **「规则管理」** 页面,点击进入对应反代或静态托管规则的详情,切换到 **「WAF」** 选项卡,勾选并绑定刚才创建的 WAF 规则组。
-
-> [!NOTE]
-> 如果规则组被标记为 **「全局规则组 (is_global)」**,它将自动应用到网关上托管的**所有网站**,无需手动执行绑定。
-
----
-
-### 第四步:发布并生效配置
-
-1. 如果你修改了 **规则组定义**、**GeoIP 范围**、**PoW 防CC难度** 或 **网站的绑定关系**:
- * 你需要点击管理端右上角的 **「配置预览」** -> **「发布并激活」**。
- * Agent 拉取并校验新版本后,将重写本地 OpenResty 核心配置文件(`waf_config.json` 等)并平滑重载进程使策略生效。
-2. 如果你只是更新了 **IP 组的成员名单**(如:在手动 IP 组中删减了一个 IP,或者自动 IP 组定时聚合出了一批新的封禁 IP):
- * **不需要做任何发布操作!**
- * Server 会在数据库更新后立即计算 IP 组全新的 Checksum 摘要。
- * 控制面会通过 **WebSocket 长连接实时向所有在线的 Agent 广播** 变更的 IP 组成员,Agent 接收后会增量覆写到本地的运行时磁盘文件 `waf_ip_groups.json`。
- * OpenResty Lua 引擎在处理新请求时,会在微秒级计算文件哈希,若发现 Checksum 变更则实时重载入内存字典(`ngx.shared`),**整个过程全程不需要 reload 任何 Nginx 服务,对线上高并发业务毫无影响**。
- * 即使 WebSocket 连接意外中断,Agent 也会在每周期心跳中上报本地 Checksum,由 Server 差分补齐下发,确保万无一失。
-
----
-
-## WAF 判定逻辑 (过滤漏斗)
-
-当一个外部请求到达 OpenResty 数据面时,WAF 运行时引擎会以微秒级的极速开销进行如下判决流检测。只要判定出明确结果,即不再向下执行:
-
-```text
- 请求进入 access 阶段
- │
- ▼
- 获取当前请求绑定的所有规则组 (全局规则组 + 自定义规则组)
- │
- ▼
- 1. 匹配 IP 白名单 / 白名单 IP 组? ──────(是)─────► [ 放行 (ALLOW) ]
- │ (否)
- ▼
- 2. 匹配国家 / 省份地域白名单? ────────(是)─────► [ 放行 (ALLOW) ]
- │ (否)
- ▼
- 3. 匹配 IP 黑名单 / 黑名单 IP 组? ──────(是)─────► [ 拦截 (BLOCK) ] ──► 返回自定义状态码与HTML拦截页
- │ (否)
- ▼
- 4. 匹配国家 / 省份地域黑名单? ────────(是)─────► [ 拦截 (BLOCK) ] ──► 返回自定义状态码与HTML拦截页
- │ (否)
- ▼
- 5. 该站点是否启用了 PoW CC 防护?
- ├───(是)───► [ 校验 PoW Cookie ] ──(验证通过)──► [ 放行 (ALLOW) ]
- │ │
- │ (未通过)
- │ ▼
- │ [ 渲染 PoW 挑战页 ] ──(计算正确)──► 写入 Cookie 并放行
- ▼
- 6. 未触发任何策略,属于正常业务流量 ───────────────► [ 放行 (ALLOW) ]
-```
-
----
-
-## 最佳实践与调优建议
-
-* **白名单放行语义**:一旦某个生效规则组配置了 IP 白名单、白名单 IP 组或地域白名单,且请求命中了其中至少一条白名单规则,该请求将被直接放行,并优先绕过后续的黑名单与 PoW 检查;未命中的请求则会继续进行黑名单等后续防护校验。
-* **白名单前置与保护**:在部署高强度黑名单或地域屏蔽前,建议首先创建一个「受信任 IP 组」,放入你团队的办公室出口 IP、本地开发 IP 以及可能访问你的第三方回调源站 IP(如微信、支付宝支付回调地址),并在规则组的**白名单**中优先引入。这可以有效防止误杀,确保信任的 IP 即使命中黑名单或 CC 限制也能无阻碍访问。
-* **合理微调 PoW 难度**:人机 CC 挑战的哈希碰撞计算(`challenge_difficulty`)是一把双刃剑。
- * 难度值 `3`:几乎瞬间完成计算,防 CC 强度低。
- * 难度值 `4`:普通手机/低端浏览器在 100~300ms 内完成计算,防护性能良好。
- * 难度值 `5`:需要 500ms~2s,防护性强,但低配端可能会感觉稍显卡顿。
- * 难度值 `6` 及以上:计算量呈指数级上升,可能导致移动端用户浏览器 CPU 持续打满卡死。**因此强烈建议在生产环境选用 `4` 或 `5`**。
-* **善用“测试规则”**:对于自动 IP 组,在点击保存之前务必点击 **「测试规则」**。通过分析当前窗口内被命中的 IP 列表,确认你的 Expr 表达式阈值(如请求数、404占比等)配置是否过宽或过紧,防止由于阈值配置不合理导致大面积误封正常用户。
-* **分离静态与动态黑名单**:不要将需要长期封禁的静态恶意 IP 填入自动封禁组(因为自动聚合的名单随时会被新的执行窗口覆盖)。应该将确定的恶意 IP 录入到一个专门的「手动封禁 IP 组」中,并让规则组同时引用该手动组与自动组。
+架构、图校验和失败回滚细节见 [WAF 可编排规则设计](../design/waf-orchestration-design.md)。
diff --git a/docs/plan/20260713-waf-orchestration.md b/docs/plan/20260713-waf-orchestration.md
index b3c3c4d5..6ea1976c 100644
--- a/docs/plan/20260713-waf-orchestration.md
+++ b/docs/plan/20260713-waf-orchestration.md
@@ -8,6 +8,12 @@
**Tech Stack:** Go 1.25、Gin、GORM、goose、PostgreSQL/SQLite、OpenResty Lua、Next.js 16 App Router、React 19、TypeScript、`@xyflow/react`、TanStack Query、shadcn/ui、Vitest。
+## 实现状态(2026-07-13)
+
+Tasks 1–11 已实现,包含三段数据库迁移、图模型与编译器、规则 API、发布快照、OpenResty 内存执行器、IP 组协调刷新、React Flow 编辑器、有序绑定、GeoLite2 City/Country 支持以及中文文档与 Swagger 更新。
+
+当前工作区已完成 `go test ./...`、前端全量 Vitest(54 项)、`make swagger`、`make code-check` 与 `git diff --check` 验证。Next.js 生产构建在本机持续停留于 Turbopack 的 `Creating an optimized production build ...`,未返回编译错误或成功状态,故不计为通过。
+
## Global Constraints
- 每张图恰好一个 `start` 和一个 `allow`;`block` 可多个;图必须无环、无悬空、无不可达节点,所有路径必须抵达 `allow` 或 `block`。
@@ -374,7 +380,7 @@ Commit: `feat(agent): execute waf graphs from worker memory`
**Interfaces:**
- Produces: `waf_ip_groups.json.checksum`;Lua `ip_groups.current()` 返回 Worker 本地对象;协调刷新间隔固定 5 秒。
-- Consumes: 现有 Agent IP 组同步 payload 与 `ngx.shared.openflare_waf_config`。
+- Consumes: 现有 Agent IP 组同步 payload 与独立的 `ngx.shared.openflare_waf_ip_groups`(64 MiB);完整运行时快照上限为 20 MiB。
- [ ] **Step 1: 写失败测试**
diff --git a/docs/reference/configuration.md b/docs/reference/configuration.md
index 0b4df4c4..af66a549 100644
--- a/docs/reference/configuration.md
+++ b/docs/reference/configuration.md
@@ -286,6 +286,8 @@ Server 的所有核心基础配置定义在 `config.yaml` 中,且均支持环
| `OPENFLARE_MMDB_PATH` | WAF GeoIP mmdb 路径,可覆盖 `agent.json` | 空 |
| `OPENFLARE_MMDB_UPDATE_INTERVAL` | WAF GeoIP mmdb 更新间隔,可覆盖 `agent.json` | 空 |
| `OPENFLARE_MMDB_DOWNLOAD_URL` | WAF GeoIP mmdb 下载地址,可覆盖 `agent.json` | 空 |
+| `OPENFLARE_CITY_MMDB_PATH` | WAF 地区匹配 City MMDB 路径,可覆盖 `agent.json` | 空 |
+| `OPENFLARE_CITY_MMDB_DOWNLOAD_URL` | WAF City MMDB 下载地址,可覆盖 `agent.json` | 空 |
---
@@ -315,8 +317,10 @@ Server 的所有核心基础配置定义在 `config.yaml` 中,且均支持环
| `runtime_config_dir` | Agent 运行时配置写入目录,如 `pow_config.json` | 否 | `data_dir/etc/openflare` |
| `pages_dir` | Pages 静态部署包解压与当前部署目录 | 否 | `data_dir/var/lib/openflare/pages` |
| `mmdb_path` | WAF GeoIP mmdb 文件路径 | 否 | `data_dir/etc/openflare/GeoLite2-Country.mmdb` |
+| `city_mmdb_path` | WAF 地区匹配 City MMDB 文件路径 | 否 | `data_dir/etc/openflare/GeoLite2-City.mmdb` |
| `mmdb_update_interval` | WAF GeoIP mmdb 更新间隔 | 否 | `86400000` 毫秒 (24h) |
| `mmdb_download_url` | WAF GeoIP mmdb 下载地址 | 否 | 内置 GeoLite2 Country 下载地址 |
+| `city_mmdb_download_url` | WAF City MMDB 下载地址 | 否 | 内置 GeoLite2 City 下载地址 |
| `observability_buffer_path` | 观测补报缓冲文件路径 | 否 | `data_dir/var/lib/openflare/observability-buffer.json` |
| `observability_replay_minutes` | 自动补传最近观测窗口分钟数 | 否 | `15` |
| `state_path` | Agent 本地状态文件路径 | 否 | `data_dir/var/lib/openflare/agent-state.json` |
diff --git a/docs/swagger.json b/docs/swagger.json
index 5bd7ede5..e35868c6 100644
--- a/docs/swagger.json
+++ b/docs/swagger.json
@@ -10242,17 +10242,16 @@
"SessionCookie": []
}
],
- "description": "返回全部 WAF 规则组,需要管理员权限",
"produces": [
"application/json"
],
"tags": [
"openflare-waf"
],
- "summary": "列出 WAF 规则组",
+ "summary": "列出 WAF 规则",
"responses": {
"200": {
- "description": "规则组列表",
+ "description": "规则列表",
"schema": {
"allOf": [
{
@@ -10264,7 +10263,7 @@
"data": {
"type": "array",
"items": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
@@ -10272,12 +10271,6 @@
]
}
},
- "400": {
- "description": "参数错误",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- },
"401": {
"description": "未登录",
"schema": {
@@ -10304,7 +10297,6 @@
"SessionCookie": []
}
],
- "description": "创建新的 WAF 规则组,需要管理员权限",
"consumes": [
"application/json"
],
@@ -10314,21 +10306,21 @@
"tags": [
"openflare-waf"
],
- "summary": "创建 WAF 规则组",
+ "summary": "创建 WAF 规则",
"parameters": [
{
- "description": "规则组参数",
+ "description": "规则名称",
"name": "request",
"in": "body",
"required": true,
"schema": {
- "$ref": "#/definitions/waf.RuleGroupInput"
+ "$ref": "#/definitions/waf.CreateRuleInput"
}
}
],
"responses": {
"200": {
- "description": "创建成功的规则组",
+ "description": "创建成功",
"schema": {
"allOf": [
{
@@ -10338,7 +10330,7 @@
"type": "object",
"properties": {
"data": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
@@ -10379,18 +10371,17 @@
"SessionCookie": []
}
],
- "description": "按 ID 返回 WAF 规则组详情,需要管理员权限",
"produces": [
"application/json"
],
"tags": [
"openflare-waf"
],
- "summary": "获取 WAF 规则组详情",
+ "summary": "获取 WAF 规则详情",
"parameters": [
{
"type": "integer",
- "description": "规则组 ID",
+ "description": "规则 ID",
"name": "id",
"in": "path",
"required": true
@@ -10398,7 +10389,7 @@
],
"responses": {
"200": {
- "description": "规则组详情",
+ "description": "规则详情",
"schema": {
"allOf": [
{
@@ -10408,7 +10399,7 @@
"type": "object",
"properties": {
"data": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
@@ -10428,7 +10419,7 @@
}
},
"404": {
- "description": "记录不存在",
+ "description": "无权限或不存在",
"schema": {
"$ref": "#/definitions/response.Any"
}
@@ -10449,18 +10440,17 @@
"SessionCookie": []
}
],
- "description": "按 ID 删除 WAF 规则组,需要管理员权限",
"produces": [
"application/json"
],
"tags": [
"openflare-waf"
],
- "summary": "删除 WAF 规则组",
+ "summary": "删除 WAF 规则",
"parameters": [
{
"type": "integer",
- "description": "规则组 ID",
+ "description": "规则 ID",
"name": "id",
"in": "path",
"required": true
@@ -10486,7 +10476,7 @@
}
},
"404": {
- "description": "记录不存在",
+ "description": "无权限或不存在",
"schema": {
"$ref": "#/definitions/response.Any"
}
@@ -10500,14 +10490,13 @@
}
}
},
- "/api/v1/d/waf/rule-groups/{id}/sites": {
+ "/api/v1/d/waf/rule-groups/{id}/graph": {
"post": {
"security": [
{
"SessionCookie": []
}
],
- "description": "替换 WAF 规则组关联的代理站点列表,需要管理员权限",
"consumes": [
"application/json"
],
@@ -10517,28 +10506,28 @@
"tags": [
"openflare-waf"
],
- "summary": "替换规则组站点绑定",
+ "summary": "保存 WAF 规则图",
"parameters": [
{
"type": "integer",
- "description": "规则组 ID",
+ "description": "规则 ID",
"name": "id",
"in": "path",
"required": true
},
{
- "description": "站点 ID 列表",
+ "description": "规则图和修订号",
"name": "request",
"in": "body",
"required": true,
"schema": {
- "$ref": "#/definitions/waf.IDsRequest"
+ "$ref": "#/definitions/waf.SaveRuleGraphInput"
}
}
],
"responses": {
"200": {
- "description": "更新后的规则组",
+ "description": "保存成功",
"schema": {
"allOf": [
{
@@ -10548,7 +10537,94 @@
"type": "object",
"properties": {
"data": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
+ }
+ }
+ }
+ ]
+ }
+ },
+ "400": {
+ "description": "参数或规则图错误",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ },
+ "401": {
+ "description": "未登录",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ },
+ "404": {
+ "description": "无权限或不存在",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ },
+ "409": {
+ "description": "修订冲突",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ },
+ "500": {
+ "description": "内部错误",
+ "schema": {
+ "$ref": "#/definitions/response.Any"
+ }
+ }
+ }
+ }
+ },
+ "/api/v1/d/waf/rule-groups/{id}/meta": {
+ "post": {
+ "security": [
+ {
+ "SessionCookie": []
+ }
+ ],
+ "consumes": [
+ "application/json"
+ ],
+ "produces": [
+ "application/json"
+ ],
+ "tags": [
+ "openflare-waf"
+ ],
+ "summary": "更新 WAF 规则元数据",
+ "parameters": [
+ {
+ "type": "integer",
+ "description": "规则 ID",
+ "name": "id",
+ "in": "path",
+ "required": true
+ },
+ {
+ "description": "规则元数据",
+ "name": "request",
+ "in": "body",
+ "required": true,
+ "schema": {
+ "$ref": "#/definitions/waf.UpdateRuleMetaInput"
+ }
+ }
+ ],
+ "responses": {
+ "200": {
+ "description": "更新成功",
+ "schema": {
+ "allOf": [
+ {
+ "$ref": "#/definitions/response.Any"
+ },
+ {
+ "type": "object",
+ "properties": {
+ "data": {
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
@@ -10568,89 +10644,7 @@
}
},
"404": {
- "description": "记录不存在",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- },
- "500": {
- "description": "内部错误",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- }
- }
- }
- },
- "/api/v1/d/waf/rule-groups/{id}/update": {
- "post": {
- "security": [
- {
- "SessionCookie": []
- }
- ],
- "description": "按 ID 更新 WAF 规则组,需要管理员权限",
- "consumes": [
- "application/json"
- ],
- "produces": [
- "application/json"
- ],
- "tags": [
- "openflare-waf"
- ],
- "summary": "更新 WAF 规则组",
- "parameters": [
- {
- "type": "integer",
- "description": "规则组 ID",
- "name": "id",
- "in": "path",
- "required": true
- },
- {
- "description": "规则组参数",
- "name": "request",
- "in": "body",
- "required": true,
- "schema": {
- "$ref": "#/definitions/waf.RuleGroupInput"
- }
- }
- ],
- "responses": {
- "200": {
- "description": "更新后的规则组",
- "schema": {
- "allOf": [
- {
- "$ref": "#/definitions/response.Any"
- },
- {
- "type": "object",
- "properties": {
- "data": {
- "$ref": "#/definitions/waf.RuleGroupView"
- }
- }
- }
- ]
- }
- },
- "400": {
- "description": "参数错误",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- },
- "401": {
- "description": "未登录",
- "schema": {
- "$ref": "#/definitions/response.Any"
- }
- },
- "404": {
- "description": "记录不存在",
+ "description": "无权限或不存在",
"schema": {
"$ref": "#/definitions/response.Any"
}
@@ -18524,6 +18518,14 @@
}
}
},
+ "waf.CreateRuleInput": {
+ "type": "object",
+ "properties": {
+ "name": {
+ "type": "string"
+ }
+ }
+ },
"waf.IDsRequest": {
"type": "object",
"properties": {
@@ -18709,139 +18711,97 @@
}
}
},
- "waf.PoWConfig": {
+ "waf.RuleEdge": {
"type": "object",
"properties": {
- "algorithm": {
+ "id": {
"type": "string"
},
- "blacklist": {
- "$ref": "#/definitions/waf.PoWListConfig"
+ "source": {
+ "type": "string"
},
- "challenge_ttl": {
- "type": "integer"
+ "source_handle": {
+ "type": "string"
},
- "difficulty": {
- "type": "integer"
- },
- "session_ttl": {
- "type": "integer"
- },
- "whitelist": {
- "$ref": "#/definitions/waf.PoWListConfig"
+ "target": {
+ "type": "string"
}
}
},
- "waf.PoWListConfig": {
+ "waf.RuleGraph": {
"type": "object",
"properties": {
- "ip_cidrs": {
+ "edges": {
"type": "array",
"items": {
- "type": "string"
+ "$ref": "#/definitions/waf.RuleEdge"
}
},
- "ips": {
+ "nodes": {
"type": "array",
"items": {
- "type": "string"
+ "$ref": "#/definitions/waf.RuleNode"
}
},
- "path_regexes": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "paths": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "user_agents": {
- "type": "array",
- "items": {
- "type": "string"
- }
+ "schema_version": {
+ "type": "integer"
}
}
},
- "waf.RuleGroupInput": {
+ "waf.RuleNode": {
"type": "object",
"properties": {
- "block_response_body": {
- "type": "string"
- },
- "block_status_code": {
- "type": "integer"
- },
- "country_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "country_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "enabled": {
- "type": "boolean"
- },
- "ip_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "ip_blacklist_group_ids": {
+ "config": {
"type": "array",
"items": {
"type": "integer"
}
},
- "ip_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "ip_whitelist_group_ids": {
- "type": "array",
- "items": {
- "type": "integer"
- }
- },
- "name": {
+ "id": {
"type": "string"
},
- "pow_config": {
- "type": "array",
- "items": {
- "type": "integer"
- }
+ "label": {
+ "type": "string"
},
- "pow_enabled": {
- "type": "boolean"
+ "position": {
+ "$ref": "#/definitions/waf.RulePosition"
},
- "region_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "region_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
+ "type": {
+ "$ref": "#/definitions/waf.RuleNodeType"
}
}
},
- "waf.RuleGroupView": {
+ "waf.RuleNodeType": {
+ "type": "string",
+ "enum": [
+ "start",
+ "allow",
+ "block",
+ "ip_match",
+ "geo_match",
+ "pow"
+ ],
+ "x-enum-varnames": [
+ "RuleNodeStart",
+ "RuleNodeAllow",
+ "RuleNodeBlock",
+ "RuleNodeIPMatch",
+ "RuleNodeGeoMatch",
+ "RuleNodePoW"
+ ]
+ },
+ "waf.RulePosition": {
+ "type": "object",
+ "properties": {
+ "x": {
+ "type": "number"
+ },
+ "y": {
+ "type": "number"
+ }
+ }
+ },
+ "waf.RuleView": {
"type": "object",
"properties": {
"applied_site_count": {
@@ -18853,86 +18813,43 @@
"type": "integer"
}
},
- "block_response_body": {
- "type": "string"
- },
- "block_status_code": {
- "type": "integer"
- },
- "country_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "country_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
"created_at": {
"type": "string"
},
"enabled": {
"type": "boolean"
},
+ "graph": {
+ "$ref": "#/definitions/waf.RuleGraph"
+ },
"id": {
"type": "integer"
},
- "ip_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "ip_blacklist_group_ids": {
- "type": "array",
- "items": {
- "type": "integer"
- }
- },
- "ip_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "ip_whitelist_group_ids": {
- "type": "array",
- "items": {
- "type": "integer"
- }
- },
"is_global": {
"type": "boolean"
},
"name": {
"type": "string"
},
- "pow_config": {
- "$ref": "#/definitions/waf.PoWConfig"
- },
- "pow_enabled": {
- "type": "boolean"
- },
- "region_blacklist": {
- "type": "array",
- "items": {
- "type": "string"
- }
- },
- "region_whitelist": {
- "type": "array",
- "items": {
- "type": "string"
- }
+ "revision": {
+ "type": "integer"
},
"updated_at": {
"type": "string"
}
}
},
+ "waf.SaveRuleGraphInput": {
+ "type": "object",
+ "properties": {
+ "graph": {
+ "$ref": "#/definitions/waf.RuleGraph"
+ },
+ "revision": {
+ "type": "integer"
+ }
+ }
+ },
"waf.SiteRuleGroupsView": {
"type": "object",
"properties": {
@@ -18945,11 +18862,11 @@
"applied_rule_groups": {
"type": "array",
"items": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
},
"global_rule_group": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
},
"route_id": {
"type": "integer"
@@ -18957,11 +18874,22 @@
"rule_groups": {
"type": "array",
"items": {
- "$ref": "#/definitions/waf.RuleGroupView"
+ "$ref": "#/definitions/waf.RuleView"
}
}
}
},
+ "waf.UpdateRuleMetaInput": {
+ "type": "object",
+ "properties": {
+ "enabled": {
+ "type": "boolean"
+ },
+ "name": {
+ "type": "string"
+ }
+ }
+ },
"zone.DomainInput": {
"type": "object",
"properties": {
diff --git a/docs/swagger.yaml b/docs/swagger.yaml
index 6aed4726..97f3f440 100644
--- a/docs/swagger.yaml
+++ b/docs/swagger.yaml
@@ -3641,6 +3641,11 @@ definitions:
website:
type: string
type: object
+ waf.CreateRuleInput:
+ properties:
+ name:
+ type: string
+ type: object
waf.IDsRequest:
properties:
ids:
@@ -3762,94 +3767,69 @@ definitions:
updated_at:
type: string
type: object
- waf.PoWConfig:
+ waf.RuleEdge:
properties:
- algorithm:
+ id:
type: string
- blacklist:
- $ref: '#/definitions/waf.PoWListConfig'
- challenge_ttl:
- type: integer
- difficulty:
- type: integer
- session_ttl:
- type: integer
- whitelist:
- $ref: '#/definitions/waf.PoWListConfig'
- type: object
- waf.PoWListConfig:
- properties:
- ip_cidrs:
- items:
- type: string
- type: array
- ips:
- items:
- type: string
- type: array
- path_regexes:
- items:
- type: string
- type: array
- paths:
- items:
- type: string
- type: array
- user_agents:
- items:
- type: string
- type: array
- type: object
- waf.RuleGroupInput:
- properties:
- block_response_body:
+ source:
type: string
- block_status_code:
+ source_handle:
+ type: string
+ target:
+ type: string
+ type: object
+ waf.RuleGraph:
+ properties:
+ edges:
+ items:
+ $ref: '#/definitions/waf.RuleEdge'
+ type: array
+ nodes:
+ items:
+ $ref: '#/definitions/waf.RuleNode'
+ type: array
+ schema_version:
type: integer
- country_blacklist:
- items:
- type: string
- type: array
- country_whitelist:
- items:
- type: string
- type: array
- enabled:
- type: boolean
- ip_blacklist:
- items:
- type: string
- type: array
- ip_blacklist_group_ids:
+ type: object
+ waf.RuleNode:
+ properties:
+ config:
items:
type: integer
type: array
- ip_whitelist:
- items:
- type: string
- type: array
- ip_whitelist_group_ids:
- items:
- type: integer
- type: array
- name:
+ id:
type: string
- pow_config:
- items:
- type: integer
- type: array
- pow_enabled:
- type: boolean
- region_blacklist:
- items:
- type: string
- type: array
- region_whitelist:
- items:
- type: string
- type: array
+ label:
+ type: string
+ position:
+ $ref: '#/definitions/waf.RulePosition'
+ type:
+ $ref: '#/definitions/waf.RuleNodeType'
type: object
- waf.RuleGroupView:
+ waf.RuleNodeType:
+ enum:
+ - start
+ - allow
+ - block
+ - ip_match
+ - geo_match
+ - pow
+ type: string
+ x-enum-varnames:
+ - RuleNodeStart
+ - RuleNodeAllow
+ - RuleNodeBlock
+ - RuleNodeIPMatch
+ - RuleNodeGeoMatch
+ - RuleNodePoW
+ waf.RulePosition:
+ properties:
+ x:
+ type: number
+ "y":
+ type: number
+ type: object
+ waf.RuleView:
properties:
applied_site_count:
type: integer
@@ -3857,59 +3837,30 @@ definitions:
items:
type: integer
type: array
- block_response_body:
- type: string
- block_status_code:
- type: integer
- country_blacklist:
- items:
- type: string
- type: array
- country_whitelist:
- items:
- type: string
- type: array
created_at:
type: string
enabled:
type: boolean
+ graph:
+ $ref: '#/definitions/waf.RuleGraph'
id:
type: integer
- ip_blacklist:
- items:
- type: string
- type: array
- ip_blacklist_group_ids:
- items:
- type: integer
- type: array
- ip_whitelist:
- items:
- type: string
- type: array
- ip_whitelist_group_ids:
- items:
- type: integer
- type: array
is_global:
type: boolean
name:
type: string
- pow_config:
- $ref: '#/definitions/waf.PoWConfig'
- pow_enabled:
- type: boolean
- region_blacklist:
- items:
- type: string
- type: array
- region_whitelist:
- items:
- type: string
- type: array
+ revision:
+ type: integer
updated_at:
type: string
type: object
+ waf.SaveRuleGraphInput:
+ properties:
+ graph:
+ $ref: '#/definitions/waf.RuleGraph'
+ revision:
+ type: integer
+ type: object
waf.SiteRuleGroupsView:
properties:
applied_ids:
@@ -3918,17 +3869,24 @@ definitions:
type: array
applied_rule_groups:
items:
- $ref: '#/definitions/waf.RuleGroupView'
+ $ref: '#/definitions/waf.RuleView'
type: array
global_rule_group:
- $ref: '#/definitions/waf.RuleGroupView'
+ $ref: '#/definitions/waf.RuleView'
route_id:
type: integer
rule_groups:
items:
- $ref: '#/definitions/waf.RuleGroupView'
+ $ref: '#/definitions/waf.RuleView'
type: array
type: object
+ waf.UpdateRuleMetaInput:
+ properties:
+ enabled:
+ type: boolean
+ name:
+ type: string
+ type: object
zone.DomainInput:
properties:
cert_id:
@@ -10168,25 +10126,20 @@ paths:
- openflare-waf
/api/v1/d/waf/rule-groups:
get:
- description: 返回全部 WAF 规则组,需要管理员权限
produces:
- application/json
responses:
"200":
- description: 规则组列表
+ description: 规则列表
schema:
allOf:
- $ref: '#/definitions/response.Any'
- properties:
data:
items:
- $ref: '#/definitions/waf.RuleGroupView'
+ $ref: '#/definitions/waf.RuleView'
type: array
type: object
- "400":
- description: 参数错误
- schema:
- $ref: '#/definitions/response.Any'
"401":
description: 未登录
schema:
@@ -10201,31 +10154,30 @@ paths:
$ref: '#/definitions/response.Any'
security:
- SessionCookie: []
- summary: 列出 WAF 规则组
+ summary: 列出 WAF 规则
tags:
- openflare-waf
post:
consumes:
- application/json
- description: 创建新的 WAF 规则组,需要管理员权限
parameters:
- - description: 规则组参数
+ - description: 规则名称
in: body
name: request
required: true
schema:
- $ref: '#/definitions/waf.RuleGroupInput'
+ $ref: '#/definitions/waf.CreateRuleInput'
produces:
- application/json
responses:
"200":
- description: 创建成功的规则组
+ description: 创建成功
schema:
allOf:
- $ref: '#/definitions/response.Any'
- properties:
data:
- $ref: '#/definitions/waf.RuleGroupView'
+ $ref: '#/definitions/waf.RuleView'
type: object
"400":
description: 参数错误
@@ -10245,14 +10197,13 @@ paths:
$ref: '#/definitions/response.Any'
security:
- SessionCookie: []
- summary: 创建 WAF 规则组
+ summary: 创建 WAF 规则
tags:
- openflare-waf
/api/v1/d/waf/rule-groups/{id}:
get:
- description: 按 ID 返回 WAF 规则组详情,需要管理员权限
parameters:
- - description: 规则组 ID
+ - description: 规则 ID
in: path
name: id
required: true
@@ -10261,13 +10212,13 @@ paths:
- application/json
responses:
"200":
- description: 规则组详情
+ description: 规则详情
schema:
allOf:
- $ref: '#/definitions/response.Any'
- properties:
data:
- $ref: '#/definitions/waf.RuleGroupView'
+ $ref: '#/definitions/waf.RuleView'
type: object
"400":
description: 参数错误
@@ -10278,7 +10229,7 @@ paths:
schema:
$ref: '#/definitions/response.Any'
"404":
- description: 记录不存在
+ description: 无权限或不存在
schema:
$ref: '#/definitions/response.Any'
"500":
@@ -10287,14 +10238,13 @@ paths:
$ref: '#/definitions/response.Any'
security:
- SessionCookie: []
- summary: 获取 WAF 规则组详情
+ summary: 获取 WAF 规则详情
tags:
- openflare-waf
/api/v1/d/waf/rule-groups/{id}/delete:
post:
- description: 按 ID 删除 WAF 规则组,需要管理员权限
parameters:
- - description: 规则组 ID
+ - description: 规则 ID
in: path
name: id
required: true
@@ -10315,7 +10265,7 @@ paths:
schema:
$ref: '#/definitions/response.Any'
"404":
- description: 记录不存在
+ description: 无权限或不存在
schema:
$ref: '#/definitions/response.Any'
"500":
@@ -10324,37 +10274,89 @@ paths:
$ref: '#/definitions/response.Any'
security:
- SessionCookie: []
- summary: 删除 WAF 规则组
+ summary: 删除 WAF 规则
tags:
- openflare-waf
- /api/v1/d/waf/rule-groups/{id}/sites:
+ /api/v1/d/waf/rule-groups/{id}/graph:
post:
consumes:
- application/json
- description: 替换 WAF 规则组关联的代理站点列表,需要管理员权限
parameters:
- - description: 规则组 ID
+ - description: 规则 ID
in: path
name: id
required: true
type: integer
- - description: 站点 ID 列表
+ - description: 规则图和修订号
in: body
name: request
required: true
schema:
- $ref: '#/definitions/waf.IDsRequest'
+ $ref: '#/definitions/waf.SaveRuleGraphInput'
produces:
- application/json
responses:
"200":
- description: 更新后的规则组
+ description: 保存成功
schema:
allOf:
- $ref: '#/definitions/response.Any'
- properties:
data:
- $ref: '#/definitions/waf.RuleGroupView'
+ $ref: '#/definitions/waf.RuleView'
+ type: object
+ "400":
+ description: 参数或规则图错误
+ schema:
+ $ref: '#/definitions/response.Any'
+ "401":
+ description: 未登录
+ schema:
+ $ref: '#/definitions/response.Any'
+ "404":
+ description: 无权限或不存在
+ schema:
+ $ref: '#/definitions/response.Any'
+ "409":
+ description: 修订冲突
+ schema:
+ $ref: '#/definitions/response.Any'
+ "500":
+ description: 内部错误
+ schema:
+ $ref: '#/definitions/response.Any'
+ security:
+ - SessionCookie: []
+ summary: 保存 WAF 规则图
+ tags:
+ - openflare-waf
+ /api/v1/d/waf/rule-groups/{id}/meta:
+ post:
+ consumes:
+ - application/json
+ parameters:
+ - description: 规则 ID
+ in: path
+ name: id
+ required: true
+ type: integer
+ - description: 规则元数据
+ in: body
+ name: request
+ required: true
+ schema:
+ $ref: '#/definitions/waf.UpdateRuleMetaInput'
+ produces:
+ - application/json
+ responses:
+ "200":
+ description: 更新成功
+ schema:
+ allOf:
+ - $ref: '#/definitions/response.Any'
+ - properties:
+ data:
+ $ref: '#/definitions/waf.RuleView'
type: object
"400":
description: 参数错误
@@ -10365,7 +10367,7 @@ paths:
schema:
$ref: '#/definitions/response.Any'
"404":
- description: 记录不存在
+ description: 无权限或不存在
schema:
$ref: '#/definitions/response.Any'
"500":
@@ -10374,57 +10376,7 @@ paths:
$ref: '#/definitions/response.Any'
security:
- SessionCookie: []
- summary: 替换规则组站点绑定
- tags:
- - openflare-waf
- /api/v1/d/waf/rule-groups/{id}/update:
- post:
- consumes:
- - application/json
- description: 按 ID 更新 WAF 规则组,需要管理员权限
- parameters:
- - description: 规则组 ID
- in: path
- name: id
- required: true
- type: integer
- - description: 规则组参数
- in: body
- name: request
- required: true
- schema:
- $ref: '#/definitions/waf.RuleGroupInput'
- produces:
- - application/json
- responses:
- "200":
- description: 更新后的规则组
- schema:
- allOf:
- - $ref: '#/definitions/response.Any'
- - properties:
- data:
- $ref: '#/definitions/waf.RuleGroupView'
- type: object
- "400":
- description: 参数错误
- schema:
- $ref: '#/definitions/response.Any'
- "401":
- description: 未登录
- schema:
- $ref: '#/definitions/response.Any'
- "404":
- description: 记录不存在
- schema:
- $ref: '#/definitions/response.Any'
- "500":
- description: 内部错误
- schema:
- $ref: '#/definitions/response.Any'
- security:
- - SessionCookie: []
- summary: 更新 WAF 规则组
+ summary: 更新 WAF 规则元数据
tags:
- openflare-waf
/api/v1/d/waf/sites/{route_id}/rule-groups:
diff --git a/frontend/app/(main)/waf/components/rule-groups-table.tsx b/frontend/app/(main)/waf/components/rule-groups-table.tsx
index bfd62526..b03e1c75 100644
--- a/frontend/app/(main)/waf/components/rule-groups-table.tsx
+++ b/frontend/app/(main)/waf/components/rule-groups-table.tsx
@@ -1,6 +1,6 @@
'use client';
-import {Globe2, MoreHorizontal, Pencil, ShieldCheck, Trash2, Users} from 'lucide-react';
+import {Globe2, MoreHorizontal, Pencil, ShieldCheck, Trash2} from 'lucide-react';
import {Badge} from '@/components/ui/badge';
import {Button} from '@/components/ui/button';
@@ -20,14 +20,12 @@ interface RuleGroupsTableProps {
groups: WAFRule[];
onEdit: (group: WAFRule) => void;
onDelete: (group: WAFRule) => void;
- onBindSites: (group: WAFRule) => void;
}
export function RuleGroupsTable({
groups,
onEdit,
onDelete,
- onBindSites,
}: RuleGroupsTableProps) {
return (
@@ -85,12 +83,6 @@ export function RuleGroupsTable({
编排
- {!group.is_global ? (
- onBindSites(group)}>
-
- 绑定网站
-
- ) : null}
{!group.is_global ? (
<>
diff --git a/frontend/app/(main)/waf/components/site-binding-sheet.tsx b/frontend/app/(main)/waf/components/site-binding-sheet.tsx
deleted file mode 100644
index 456e393b..00000000
--- a/frontend/app/(main)/waf/components/site-binding-sheet.tsx
+++ /dev/null
@@ -1,139 +0,0 @@
-'use client';
-
-import {useEffect, useMemo, useState} from 'react';
-import {Check, Search} from 'lucide-react';
-
-import {Button} from '@/components/ui/button';
-import {Input} from '@/components/ui/input';
-import {Sheet, SheetContent, SheetDescription, SheetFooter, SheetHeader, SheetTitle,} from '@/components/ui/sheet';
-import {cn} from '@/lib/utils';
-import type {ProxyRouteItem, WAFRule} from '@/lib/services/openflare';
-
-interface SiteBindingSheetProps {
- group: WAFRule | null;
- routes: ProxyRouteItem[];
- open: boolean;
- pending: boolean;
- onOpenChange: (open: boolean) => void;
- onSave: (ids: number[]) => void;
-}
-
-export function SiteBindingSheet({
- group,
- routes,
- open,
- pending,
- onOpenChange,
- onSave,
-}: SiteBindingSheetProps) {
- const [keyword, setKeyword] = useState('');
- const [selectedIDs, setSelectedIDs] = useState([]);
-
- useEffect(() => {
- setSelectedIDs(group?.applied_site_ids ?? []);
- setKeyword('');
- }, [group, open]);
-
- const filteredRoutes = useMemo(() => {
- const normalized = keyword.trim().toLowerCase();
- if (!normalized) return routes;
- return routes.filter((route) =>
- [
- route.site_name,
- ...(route.zone_domains ?? []).map((item) => item.domain),
- ]
- .join(' ')
- .toLowerCase()
- .includes(normalized),
- );
- }, [keyword, routes]);
-
- const selectedSet = useMemo(() => new Set(selectedIDs), [selectedIDs]);
-
- const toggleID = (id: number) => {
- setSelectedIDs((current) =>
- current.includes(id)
- ? current.filter((item) => item !== id)
- : [...current, id],
- );
- };
-
- const selectFiltered = () => {
- const next = new Set(selectedIDs);
- filteredRoutes.forEach((route) => next.add(route.id));
- setSelectedIDs([...next]);
- };
-
- return (
-
-
-
- {group ? `绑定 ${group.name}` : '绑定规则组'}
-
- 选择这个自定义规则组要叠加到哪些网站。
-
-
-
-
-
-
- setKeyword(event.target.value)}
- />
-
-
-
-
- {filteredRoutes.map((route) => (
-
- ))}
-
-
-
-
-
-
-
-
-
- );
-}
diff --git a/frontend/app/(main)/waf/page.tsx b/frontend/app/(main)/waf/page.tsx
index 6591a66b..37793688 100644
--- a/frontend/app/(main)/waf/page.tsx
+++ b/frontend/app/(main)/waf/page.tsx
@@ -22,33 +22,25 @@ import {EmptyStateWithBorder} from '@/components/layout/empty';
import {ErrorInline} from '@/components/layout/error';
import {LoadingStateWithBorder} from '@/components/layout/loading';
import type {WAFRule} from '@/lib/services/openflare';
-import {ProxyRouteService, WafService} from '@/lib/services/openflare';
+import {WafService} from '@/lib/services/openflare';
import {CreateRuleDialog} from './components/create-rule-dialog';
import {getErrorMessage} from './components/helpers';
import {RuleGroupsTable} from './components/rule-groups-table';
-import {SiteBindingSheet} from './components/site-binding-sheet';
const ruleGroupsQueryKey = ['openflare', 'waf', 'rule-groups'];
-const routesQueryKey = ['openflare', 'proxy-routes'];
export default function WafPage() {
const router = useRouter();
const queryClient = useQueryClient();
const [createOpen, setCreateOpen] = useState(false);
const [deleteTarget, setDeleteTarget] = useState(null);
- const [bindingGroup, setBindingGroup] = useState(null);
const groupsQuery = useQuery({
queryKey: ruleGroupsQueryKey,
queryFn: () => WafService.listRuleGroups(),
});
- const routesQuery = useQuery({
- queryKey: routesQueryKey,
- queryFn: () => ProxyRouteService.list(),
- });
-
const invalidate = async () => {
await Promise.all([
queryClient.invalidateQueries({ queryKey: ruleGroupsQueryKey }),
@@ -80,26 +72,13 @@ export default function WafPage() {
},
});
- const bindMutation = useMutation({
- mutationFn: ({ id, ids }: { id: number; ids: number[] }) =>
- WafService.updateRuleGroupSites(id, ids),
- onSuccess: async () => {
- toast.success('规则组应用范围已更新');
- setBindingGroup(null);
- await invalidate();
- },
- onError: (error) => {
- toast.error(getErrorMessage(error));
- },
- });
-
const handleRefresh = () => {
void queryClient.invalidateQueries({ queryKey: ruleGroupsQueryKey });
};
const groups = groupsQuery.data ?? [];
- const loading = groupsQuery.isLoading || routesQuery.isLoading;
- const error = groupsQuery.error ?? routesQuery.error ?? null;
+ const loading = groupsQuery.isLoading;
+ const error = groupsQuery.error ?? null;
return (
@@ -162,7 +141,6 @@ export default function WafPage() {
groups={groups}
onEdit={(rule) => router.push(`/waf/rules/editor?id=${rule.id}`)}
onDelete={setDeleteTarget}
- onBindSites={setBindingGroup}
/>
)}
@@ -177,19 +155,6 @@ export default function WafPage() {
}}
/>
-
!open && setBindingGroup(null)}
- onSave={(ids) => {
- if (bindingGroup) {
- bindMutation.mutate({ id: bindingGroup.id, ids });
- }
- }}
- />
-
!open && setDeleteTarget(null)}
diff --git a/frontend/app/(main)/waf/rules/editor/components/editor-behavior.test.ts b/frontend/app/(main)/waf/rules/editor/components/editor-behavior.test.ts
new file mode 100644
index 00000000..d42b2559
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/editor-behavior.test.ts
@@ -0,0 +1,66 @@
+import {describe, expect, it} from 'vitest';
+
+import type {WAFRuleGraph} from '@/lib/services/openflare';
+
+import {acceptedNodeChanges, filterRemovableNodeIds, findGraphErrorTarget, getHistoryTransition, isConnectionAllowed, isPersistentEdgeChange, isPersistentNodeChange} from './editor-behavior';
+
+describe('React Flow persistence filtering', () => {
+ it('ignores dimensions and selection changes', () => {
+ expect(isPersistentNodeChange({type: 'dimensions', id: 'n', dimensions: {width: 10, height: 10}})).toBe(false);
+ expect(isPersistentNodeChange({type: 'select', id: 'n', selected: true})).toBe(false);
+ expect(isPersistentEdgeChange({type: 'select', id: 'e', selected: true})).toBe(false);
+ });
+
+ it('persists completed position changes and removals', () => {
+ expect(isPersistentNodeChange({type: 'position', id: 'n', position: {x: 1, y: 2}, dragging: false})).toBe(true);
+ expect(isPersistentNodeChange({type: 'remove', id: 'n'})).toBe(true);
+ expect(isPersistentEdgeChange({type: 'remove', id: 'e'})).toBe(true);
+ });
+});
+
+describe('editor safety constraints', () => {
+ const graph: WAFRuleGraph = {schema_version: 1, nodes: [
+ {id: 'start', type: 'start', position: {x: 0, y: 0}, config: {}},
+ {id: 'match', type: 'ip_match', position: {x: 0, y: 0}, config: {ips: [], cidrs: [], ip_group_ids: []}},
+ {id: 'allow', type: 'allow', position: {x: 0, y: 0}, config: {}},
+ ], edges: [{id: 'start-match', source: 'start', source_handle: 'next', target: 'match'}]};
+
+ it('protects start and allow from deletion', () => expect(filterRemovableNodeIds(graph.nodes, ['start', 'match', 'allow'])).toEqual(['match']));
+
+ it('does not persist a removal containing only protected nodes', () => {
+ expect(acceptedNodeChanges(graph.nodes, [{type: 'remove', id: 'start'}, {type: 'remove', id: 'allow'}])).toEqual({changes: [], persistent: false});
+ });
+
+ it('rejects invalid and already-used source ports', () => {
+ expect(isConnectionAllowed(graph, {source: 'match', sourceHandle: 'next', target: 'allow'})).toBe(false);
+ expect(isConnectionAllowed(graph, {source: 'start', sourceHandle: 'next', target: 'allow'})).toBe(false);
+ expect(isConnectionAllowed(graph, {source: 'match', sourceHandle: 'true', target: 'allow'})).toBe(true);
+ });
+});
+
+describe('server graph error targeting', () => {
+ const nodes = ['start', 'match-1'];
+ const edges = ['start-match'];
+
+ it('uses explicit node and edge ids from nested payloads', () => {
+ expect(findGraphErrorTarget({details: {node_id: 'match-1'}}, nodes, edges)).toEqual({kind: 'node', id: 'match-1'});
+ expect(findGraphErrorTarget({details: {edgeId: 'start-match'}}, nodes, edges)).toEqual({kind: 'edge', id: 'start-match'});
+ });
+
+ it('does not substring match unrelated error text', () => {
+ expect(findGraphErrorTarget({message: 'restart operation failed'}, nodes, edges)).toBeUndefined();
+ });
+
+ it('parses the real API envelope with strict ID boundaries', () => {
+ expect(findGraphErrorTarget({error_msg: '规则图无效: 节点 match-1 的 true 出口未连接', data: null}, nodes, edges)).toEqual({kind: 'node', id: 'match-1'});
+ expect(findGraphErrorTarget(new Error('规则图无效: 边 start-match 的目标节点不存在'), nodes, edges)).toEqual({kind: 'edge', id: 'start-match'});
+ expect(findGraphErrorTarget({error_msg: '节点 match-10 无效'}, nodes, edges)).toBeUndefined();
+ expect(findGraphErrorTarget({error_msg: '规则图无效: 边 ID start-match 重复'}, nodes, edges)).toEqual({kind: 'edge', id: 'start-match'});
+ expect(findGraphErrorTarget({error_msg: '边 ID start-matcher 重复'}, nodes, edges)).toBeUndefined();
+ });
+});
+
+it('calculates deterministic Back and Forward restoration deltas', () => {
+ expect(getHistoryTransition(4, 3)).toEqual({direction: 'back', restoreDelta: 1});
+ expect(getHistoryTransition(4, 6)).toEqual({direction: 'forward', restoreDelta: -2});
+});
diff --git a/frontend/app/(main)/waf/rules/editor/components/editor-behavior.ts b/frontend/app/(main)/waf/rules/editor/components/editor-behavior.ts
new file mode 100644
index 00000000..0f3b0c2d
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/editor-behavior.ts
@@ -0,0 +1,73 @@
+import type {EdgeChange, NodeChange} from '@xyflow/react';
+import type {WAFRuleGraph, WAFRuleNode} from '@/lib/services/openflare';
+
+import {wouldCreateCycle} from './graph-validation';
+
+export type GraphErrorTarget = {kind: 'node' | 'edge'; id: string};
+
+export function isPersistentNodeChange(change: NodeChange): boolean {
+ return change.type === 'remove' || change.type === 'add' || change.type === 'replace' || (change.type === 'position' && change.dragging === false && Boolean(change.position));
+}
+
+export function isPersistentEdgeChange(change: EdgeChange): boolean {
+ return change.type === 'remove' || change.type === 'add' || change.type === 'replace';
+}
+
+export function filterRemovableNodeIds(nodes: WAFRuleNode[], ids: string[]): string[] {
+ return ids.filter((id) => !['start', 'allow'].includes(nodes.find((node) => node.id === id)?.type ?? ''));
+}
+
+export function acceptedNodeChanges(nodes: WAFRuleNode[], changes: NodeChange[]): {changes: NodeChange[]; persistent: boolean} {
+ const accepted = changes.filter((change) => change.type === 'position' || (change.type === 'remove' && filterRemovableNodeIds(nodes, [change.id]).length === 1));
+ return {changes: accepted, persistent: accepted.some(isPersistentNodeChange)};
+}
+
+export function getHistoryTransition(current: number, target: number): {direction: 'back' | 'forward'; restoreDelta: number} {
+ return {direction: target < current ? 'back' : 'forward', restoreDelta: current - target};
+}
+
+export function isConnectionAllowed(graph: WAFRuleGraph, connection: {source?: string | null; sourceHandle?: string | null; target?: string | null}): boolean {
+ if (!connection.source || !connection.target || !connection.sourceHandle) return false;
+ const source = graph.nodes.find((node) => node.id === connection.source);
+ const handles: Partial> = {start: ['next'], ip_match: ['true', 'false'], geo_match: ['true', 'false'], pow: ['next']};
+ return Boolean(source && (handles[source.type] ?? []).includes(connection.sourceHandle)) && !graph.edges.some((edge) => edge.source === connection.source && edge.source_handle === connection.sourceHandle) && !wouldCreateCycle(graph, connection.source, connection.target);
+}
+
+export function findGraphErrorTarget(payload: unknown, nodeIds: string[], edgeIds: string[]): GraphErrorTarget | undefined {
+ const found = collectIdFields(payload);
+ for (const {key, value} of found) {
+ if ((key === 'node_id' || key === 'nodeId') && nodeIds.includes(value)) return {kind: 'node', id: value};
+ if ((key === 'edge_id' || key === 'edgeId') && edgeIds.includes(value)) return {kind: 'edge', id: value};
+ }
+ const messages = collectMessages(payload);
+ for (const message of messages) {
+ const nodeId = findMessageId(message, '节点', nodeIds);
+ if (nodeId) return {kind: 'node', id: nodeId};
+ const edgeId = findMessageId(message, '边', edgeIds);
+ if (edgeId) return {kind: 'edge', id: edgeId};
+ }
+ return undefined;
+}
+
+function collectMessages(value: unknown): string[] {
+ if (value instanceof Error) return [value.message, ...collectMessages(value.cause)];
+ if (!value || typeof value !== 'object') return [];
+ return Object.entries(value).flatMap(([key, child]) => key === 'error_msg' || key === 'message' ? typeof child === 'string' ? [child] : [] : collectMessages(child));
+}
+
+function findMessageId(message: string, prefix: '节点' | '边', ids: string[]): string | undefined {
+ const token = prefix === '边' ? '边(?:\\s+ID)?' : '节点';
+ return [...ids].sort((a, b) => b.length - a.length).find((id) => new RegExp(`${token}\\s+${escapeRegExp(id)}(?=$|[\\s,。,::的])`).test(message));
+}
+
+function escapeRegExp(value: string): string { return value.replace(/[.*+?^${}()|[\]\\]/g, '\\$&'); }
+
+function collectIdFields(value: unknown): {key: string; value: string}[] {
+ if (!value || typeof value !== 'object') return [];
+ const result: {key: string; value: string}[] = [];
+ for (const [key, child] of Object.entries(value)) {
+ if (typeof child === 'string') result.push({key, value: child});
+ else result.push(...collectIdFields(child));
+ }
+ return result;
+}
diff --git a/frontend/app/(main)/waf/rules/editor/components/graph-validation.test.ts b/frontend/app/(main)/waf/rules/editor/components/graph-validation.test.ts
new file mode 100644
index 00000000..a4a54cbb
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/graph-validation.test.ts
@@ -0,0 +1,116 @@
+import {describe, expect, it} from 'vitest';
+
+import type {WAFRuleGraph} from '@/lib/services/openflare';
+
+import {removeNodeFromGraph, validateGraph, wouldCreateCycle} from './graph-validation';
+
+const validGraph = (): WAFRuleGraph => ({
+ schema_version: 1,
+ nodes: [
+ {id: 'start', type: 'start', position: {x: 0, y: 0}, config: {}},
+ {id: 'match', type: 'ip_match', position: {x: 240, y: 0}, config: {ips: ['127.0.0.1'], cidrs: [], ip_group_ids: []}},
+ {id: 'allow', type: 'allow', position: {x: 520, y: -80}, config: {}},
+ {id: 'block', type: 'block', position: {x: 520, y: 100}, config: {status_code: 403, response_body: ''}},
+ ],
+ edges: [
+ {id: 'start-match', source: 'start', source_handle: 'next', target: 'match'},
+ {id: 'match-allow', source: 'match', source_handle: 'true', target: 'allow'},
+ {id: 'match-block', source: 'match', source_handle: 'false', target: 'block'},
+ ],
+});
+
+describe('validateGraph', () => {
+ it('accepts a complete terminating graph', () => expect(validateGraph(validGraph())).toEqual([]));
+
+ it('requires exactly one start and allow node', () => {
+ const graph = validGraph();
+ graph.nodes = graph.nodes.filter((node) => node.type !== 'allow');
+ expect(validateGraph(graph).map((issue) => issue.code)).toContain('allow_count');
+ });
+
+ it('requires every source handle', () => {
+ const graph = validGraph();
+ graph.edges = graph.edges.filter((edge) => edge.source_handle !== 'false');
+ expect(validateGraph(graph)).toContainEqual(expect.objectContaining({code: 'missing_handle', nodeId: 'match'}));
+ });
+
+ it('rejects cycles', () => {
+ const graph = validGraph();
+ graph.edges.push({id: 'cycle', source: 'block', source_handle: 'next', target: 'match'});
+ expect(validateGraph(graph).map((issue) => issue.code)).toContain('cycle');
+ });
+
+ it('reports unreachable nodes and paths without a terminal', () => {
+ const graph = validGraph();
+ graph.nodes.push({id: 'orphan', type: 'pow', position: {x: 0, y: 200}, config: {algorithm: 'fast', difficulty: 4, session_ttl: 60, challenge_ttl: 30}});
+ expect(validateGraph(graph)).toEqual(expect.arrayContaining([
+ expect.objectContaining({code: 'unreachable', nodeId: 'orphan'}),
+ expect.objectContaining({code: 'non_terminating', nodeId: 'orphan'}),
+ ]));
+ });
+
+ it('rejects duplicate identifiers and start incoming edges', () => {
+ const graph = validGraph();
+ graph.nodes.push({...graph.nodes[1]});
+ graph.edges.push({...graph.edges[0]}, {id: 'into-start', source: 'match', source_handle: 'true', target: 'start'});
+ expect(validateGraph(graph)).toEqual(expect.arrayContaining([
+ expect.objectContaining({code: 'duplicate_node_id', nodeId: 'match'}),
+ expect.objectContaining({code: 'duplicate_edge_id', edgeId: 'start-match'}),
+ expect.objectContaining({code: 'start_incoming', edgeId: 'into-start'}),
+ ]));
+ });
+
+ it('validates typed node configuration locally', () => {
+ const graph = validGraph();
+ graph.nodes = graph.nodes.map((node) => node.type === 'ip_match' ? {...node, config: {ips: ['999.1.1.1'], cidrs: ['broken'], ip_group_ids: [-1]}} : node.type === 'block' ? {...node, config: {status_code: 200, response_body: 'x'.repeat(65_537)}} : node);
+ expect(validateGraph(graph).filter((issue) => issue.code === 'invalid_config').map((issue) => issue.nodeId)).toEqual(expect.arrayContaining(['match', 'block']));
+ });
+
+ it('validates PoW bounds and geography codes', () => {
+ const graph = validGraph();
+ graph.nodes.push({id: 'pow', type: 'pow', position: {x: 0, y: 0}, config: {algorithm: 'fast', difficulty: 0, session_ttl: 0, challenge_ttl: 0}});
+ graph.nodes.push({id: 'geo', type: 'geo_match', position: {x: 0, y: 0}, config: {countries: ['china'], regions: ['']}});
+ expect(validateGraph(graph)).toEqual(expect.arrayContaining([
+ expect.objectContaining({code: 'invalid_config', nodeId: 'pow'}),
+ expect.objectContaining({code: 'invalid_config', nodeId: 'geo'}),
+ ]));
+ });
+
+ it('rejects non-finite and fractional integer configuration', () => {
+ const graph = validGraph();
+ graph.nodes.push({id: 'pow', type: 'pow', position: {x: 0, y: 0}, config: {algorithm: 'fast', difficulty: 4.5, session_ttl: Number.NaN, challenge_ttl: 30}});
+ graph.nodes.push({id: 'block-fraction', type: 'block', position: {x: 0, y: 0}, config: {status_code: 403.5, response_body: ''}});
+ expect(validateGraph(graph)).toEqual(expect.arrayContaining([
+ expect.objectContaining({code: 'invalid_config', nodeId: 'pow'}),
+ expect.objectContaining({code: 'invalid_config', nodeId: 'block-fraction'}),
+ ]));
+ });
+
+ it('matches server IP and prefix parsing semantics', () => {
+ const invalid = validGraph();
+ invalid.nodes = invalid.nodes.map((node) => node.type === 'ip_match' ? {...node, config: {ips: ['2001:db8::1', '::::'], cidrs: ['2001:db8::/32', '10.0.0.0/33'], ip_group_ids: []}} : node);
+ expect(validateGraph(invalid)).toContainEqual(expect.objectContaining({code: 'invalid_config', nodeId: 'match'}));
+ const valid = validGraph();
+ valid.nodes = valid.nodes.map((node) => node.type === 'ip_match' ? {...node, config: {ips: ['2001:db8::1', '192.0.2.1'], cidrs: ['2001:db8::/32', '10.0.0.0/8'], ip_group_ids: []}} : node);
+ expect(validateGraph(valid)).toEqual([]);
+ });
+
+ it('requires exactly one CIDR slash and rejects scoped IPv6 addresses', () => {
+ for (const value of ['10.0.0.0/8/extra', 'fe80::1%en0', 'fe80::%en0/64']) {
+ const graph = validGraph();
+ graph.nodes = graph.nodes.map((node) => node.type === 'ip_match' ? {...node, config: value.includes('/') ? {ips: [], cidrs: [value], ip_group_ids: []} : {ips: [value], cidrs: [], ip_group_ids: []}} : node);
+ expect(validateGraph(graph)).toContainEqual(expect.objectContaining({code: 'invalid_config', nodeId: 'match'}));
+ }
+ });
+});
+
+it('removes incident edges when deleting a node', () => {
+ const next = removeNodeFromGraph(validGraph(), 'match');
+ expect(next.nodes.some((node) => node.id === 'match')).toBe(false);
+ expect(next.edges).toEqual([]);
+});
+
+it('detects whether a new connection creates a cycle', () => {
+ expect(wouldCreateCycle(validGraph(), 'allow', 'start')).toBe(true);
+ expect(wouldCreateCycle(validGraph(), 'start', 'block')).toBe(false);
+});
diff --git a/frontend/app/(main)/waf/rules/editor/components/graph-validation.ts b/frontend/app/(main)/waf/rules/editor/components/graph-validation.ts
new file mode 100644
index 00000000..cd4fcad8
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/graph-validation.ts
@@ -0,0 +1,146 @@
+import type {WAFRuleGraph, WAFRuleNode} from '@/lib/services/openflare';
+
+export type GraphIssueCode = 'schema' | 'size_limit' | 'empty_id' | 'duplicate_node_id' | 'duplicate_edge_id' | 'start_count' | 'allow_count' | 'start_incoming' | 'missing_handle' | 'duplicate_handle' | 'invalid_edge' | 'invalid_config' | 'cycle' | 'unreachable' | 'non_terminating';
+
+export interface GraphIssue {
+ code: GraphIssueCode;
+ message: string;
+ nodeId?: string;
+ edgeId?: string;
+}
+
+const handles: Partial> = {
+ start: ['next'], ip_match: ['true', 'false'], geo_match: ['true', 'false'], pow: ['next'],
+};
+
+export function validateGraph(graph: WAFRuleGraph): GraphIssue[] {
+ const issues: GraphIssue[] = [];
+ const nodeMap = new Map(graph.nodes.map((node) => [node.id, node]));
+ if (graph.schema_version !== 1) issues.push({code: 'schema', message: '规则图 schema_version 必须为 1'});
+ if (graph.nodes.length > 128 || graph.edges.length > 256 || new TextEncoder().encode(JSON.stringify(graph)).length > 256 * 1024) issues.push({code: 'size_limit', message: '规则图超过大小限制'});
+ const nodeIds = new Set();
+ for (const node of graph.nodes) {
+ if (!node.id.trim()) issues.push({code: 'empty_id', message: '节点 ID 不能为空', nodeId: node.id});
+ if (nodeIds.has(node.id)) issues.push({code: 'duplicate_node_id', message: `节点 ID ${node.id} 重复`, nodeId: node.id});
+ nodeIds.add(node.id);
+ const configIssue = validateNodeConfig(node);
+ if (configIssue) issues.push({code: 'invalid_config', message: configIssue, nodeId: node.id});
+ }
+ const edgeIds = new Set();
+ for (const edge of graph.edges) {
+ if (!edge.id.trim()) issues.push({code: 'empty_id', message: '连线 ID 不能为空', edgeId: edge.id});
+ if (edgeIds.has(edge.id)) issues.push({code: 'duplicate_edge_id', message: `连线 ID ${edge.id} 重复`, edgeId: edge.id});
+ edgeIds.add(edge.id);
+ if (nodeMap.get(edge.target)?.type === 'start') issues.push({code: 'start_incoming', message: '开始节点不能有入边', edgeId: edge.id, nodeId: edge.target});
+ }
+ if (graph.nodes.filter((node) => node.type === 'start').length !== 1) issues.push({code: 'start_count', message: '规则图必须恰好有一个开始节点'});
+ if (graph.nodes.filter((node) => node.type === 'allow').length !== 1) issues.push({code: 'allow_count', message: '规则图必须恰好有一个通过节点'});
+
+ for (const edge of graph.edges) {
+ const source = nodeMap.get(edge.source);
+ if (!source || !nodeMap.has(edge.target) || !(handles[source.type] ?? []).includes(edge.source_handle)) {
+ issues.push({code: 'invalid_edge', message: `连线 ${edge.id} 的端点或出口无效`, edgeId: edge.id});
+ }
+ }
+ for (const node of graph.nodes) {
+ for (const handle of handles[node.type] ?? []) {
+ const outgoing = graph.edges.filter((edge) => edge.source === node.id && edge.source_handle === handle);
+ if (outgoing.length === 0) issues.push({code: 'missing_handle', message: `节点 ${node.id} 的 ${handle} 出口未连接`, nodeId: node.id});
+ if (outgoing.length > 1) issues.push({code: 'duplicate_handle', message: `节点 ${node.id} 的 ${handle} 出口只能连接一次`, nodeId: node.id});
+ }
+ }
+
+ const adjacency = new Map(graph.nodes.map((node) => [node.id, [] as string[]]));
+ const reverse = new Map(graph.nodes.map((node) => [node.id, [] as string[]]));
+ for (const edge of graph.edges) {
+ adjacency.get(edge.source)?.push(edge.target);
+ reverse.get(edge.target)?.push(edge.source);
+ }
+ const start = graph.nodes.find((node) => node.type === 'start');
+ const reachable = walk(start ? [start.id] : [], adjacency);
+ for (const node of graph.nodes) if (!reachable.has(node.id)) issues.push({code: 'unreachable', message: `节点 ${node.id} 无法从开始节点到达`, nodeId: node.id});
+ const terminals = graph.nodes.filter((node) => node.type === 'allow' || node.type === 'block').map((node) => node.id);
+ const canTerminate = walk(terminals, reverse);
+ for (const node of graph.nodes) if (!canTerminate.has(node.id)) issues.push({code: 'non_terminating', message: `节点 ${node.id} 无法抵达终止节点`, nodeId: node.id});
+ if (hasCycle(graph)) issues.push({code: 'cycle', message: '规则图不能包含循环'});
+ return issues;
+}
+
+function validateNodeConfig(node: WAFRuleNode): string | undefined {
+ if (node.type === 'ip_match') {
+ if (node.config.ips.some((value) => !isIP(value))) return `节点 ${node.id} 包含无效 IP`;
+ if (node.config.cidrs.some((value) => { const parts = value.split('/'); if (parts.length !== 2) return true; const [ip, bits] = parts; return !isIP(ip) || !/^\d+$/.test(bits) || Number(bits) > (ip.includes(':') ? 128 : 32); })) return `节点 ${node.id} 包含无效 CIDR`;
+ if (node.config.ip_group_ids.some((id) => !Number.isInteger(id) || id <= 0)) return `节点 ${node.id} 包含无效 IP 组`;
+ }
+ if (node.type === 'geo_match' && (node.config.countries.some((code) => !/^[A-Z]{2}$/.test(code)) || node.config.regions.some((code) => !/^[A-Z]{2}-[A-Z0-9]{1,3}$/.test(code)))) return `节点 ${node.id} 包含无效地域代码`;
+ if (node.type === 'pow' && (!['fast', 'slow'].includes(node.config.algorithm) || !isIntegerInRange(node.config.difficulty, 1, 16) || !isIntegerInRange(node.config.session_ttl, 60) || !isIntegerInRange(node.config.challenge_ttl, 30))) return `节点 ${node.id} 的 PoW 配置超出范围`;
+ if (node.type === 'block' && (!isIntegerInRange(node.config.status_code, 400, 599) || new TextEncoder().encode(node.config.response_body).length > 16 * 1024)) return `节点 ${node.id} 的阻止响应配置无效`;
+ return undefined;
+}
+
+function isIntegerInRange(value: number, min: number, max = Number.MAX_SAFE_INTEGER): boolean { return Number.isFinite(value) && Number.isInteger(value) && value >= min && value <= max; }
+
+function isIP(value: string): boolean {
+ if (value.includes(':')) return isIPv6(value);
+ const parts = value.split('.');
+ return parts.length === 4 && parts.every((part) => /^(0|[1-9]\d{0,2})$/.test(part) && Number(part) <= 255);
+}
+
+function isIPv6(value: string): boolean {
+ if (!/^[0-9a-f:.]+$/i.test(value) || value.includes(':::') || value.split('::').length > 2) return false;
+ const compressed = value.includes('::');
+ const sections = value.split('::');
+ const groups = sections.flatMap((section) => section ? section.split(':') : []);
+ let units = 0;
+ for (let index = 0; index < groups.length; index++) {
+ const group = groups[index];
+ if (group.includes('.')) {
+ if (index !== groups.length - 1 || !isIP(group)) return false;
+ units += 2;
+ } else {
+ if (!/^[0-9a-f]{1,4}$/i.test(group)) return false;
+ units++;
+ }
+ }
+ return compressed ? units < 8 : units === 8;
+}
+
+function walk(seeds: string[], links: Map): Set {
+ const seen = new Set();
+ const stack = [...seeds];
+ while (stack.length) {
+ const id = stack.pop()!;
+ if (seen.has(id)) continue;
+ seen.add(id);
+ stack.push(...(links.get(id) ?? []));
+ }
+ return seen;
+}
+
+function hasCycle(graph: WAFRuleGraph): boolean {
+ const indegree = new Map(graph.nodes.map((node) => [node.id, 0]));
+ for (const edge of graph.edges) if (indegree.has(edge.target) && indegree.has(edge.source)) indegree.set(edge.target, (indegree.get(edge.target) ?? 0) + 1);
+ const queue = [...indegree].filter(([, degree]) => degree === 0).map(([id]) => id);
+ let visited = 0;
+ while (queue.length) {
+ const id = queue.shift()!;
+ visited++;
+ for (const edge of graph.edges.filter((item) => item.source === id)) {
+ const next = (indegree.get(edge.target) ?? 0) - 1;
+ indegree.set(edge.target, next);
+ if (next === 0) queue.push(edge.target);
+ }
+ }
+ return visited !== graph.nodes.length;
+}
+
+export function wouldCreateCycle(graph: WAFRuleGraph, source: string, target: string): boolean {
+ if (source === target) return true;
+ const adjacency = new Map(graph.nodes.map((node) => [node.id, [] as string[]]));
+ for (const edge of graph.edges) adjacency.get(edge.source)?.push(edge.target);
+ return walk([target], adjacency).has(source);
+}
+
+export function removeNodeFromGraph(graph: WAFRuleGraph, nodeId: string): WAFRuleGraph {
+ return {...graph, nodes: graph.nodes.filter((node) => node.id !== nodeId), edges: graph.edges.filter((edge) => edge.source !== nodeId && edge.target !== nodeId)};
+}
diff --git a/frontend/app/(main)/waf/rules/editor/components/node-library.tsx b/frontend/app/(main)/waf/rules/editor/components/node-library.tsx
new file mode 100644
index 00000000..063f78de
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/node-library.tsx
@@ -0,0 +1,14 @@
+import {Ban, Fingerprint, Globe2, Plus, ShieldCheck} from 'lucide-react';
+
+import {Button} from '@/components/ui/button';
+import type {WAFRuleNode} from '@/lib/services/openflare';
+
+type AddableType = Extract;
+const items = [
+ {type: 'ip_match', label: 'IP 匹配', icon: Fingerprint}, {type: 'geo_match', label: '地域匹配', icon: Globe2},
+ {type: 'pow', label: 'PoW 挑战', icon: ShieldCheck}, {type: 'block', label: '阻止', icon: Ban},
+] satisfies {type: AddableType; label: string; icon: typeof Plus}[];
+
+export function NodeLibrary({onAdd}: {onAdd: (type: AddableType) => void}) {
+ return {items.map(({type, label, icon: Icon}) => )}
;
+}
diff --git a/frontend/app/(main)/waf/rules/editor/components/node-properties.test.tsx b/frontend/app/(main)/waf/rules/editor/components/node-properties.test.tsx
new file mode 100644
index 00000000..c88e2a45
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/node-properties.test.tsx
@@ -0,0 +1,34 @@
+import {fireEvent, render, screen} from '@testing-library/react';
+import {expect, it, vi} from 'vitest';
+
+import type {WAFIPGroup, WAFRuleNode} from '@/lib/services/openflare';
+
+import {NodeProperties} from './node-properties';
+
+it('edits IP group config through a typed multi-select', async () => {
+ const node: WAFRuleNode = {id: 'match', type: 'ip_match', position: {x: 0, y: 0}, config: {ips: [], cidrs: [], ip_group_ids: []}};
+ const group = {id: 7, name: '办公室出口'} as WAFIPGroup;
+ const onChange = vi.fn();
+ render();
+ fireEvent.click(screen.getByRole('button', {name: 'IP 组'}));
+ fireEvent.click(await screen.findByText('办公室出口'));
+ expect(onChange).toHaveBeenCalledWith(expect.objectContaining({config: expect.objectContaining({ip_group_ids: [7]})}));
+});
+
+it('associates numeric property labels and constrains server ranges', () => {
+ const node: WAFRuleNode = {id: 'pow', type: 'pow', position: {x: 0, y: 0}, config: {algorithm: 'fast', difficulty: 4, session_ttl: 60, challenge_ttl: 30}};
+ render();
+ expect(screen.getByLabelText('难度')).toHaveAttribute('min', '1');
+ expect(screen.getByLabelText('难度')).toHaveAttribute('max', '16');
+ expect(screen.getByLabelText('会话 TTL(秒)')).toHaveAttribute('min', '60');
+});
+
+it('creates any normalized valid geography code', async () => {
+ const node: WAFRuleNode = {id: 'geo', type: 'geo_match', position: {x: 0, y: 0}, config: {countries: [], regions: []}};
+ const onChange = vi.fn();
+ render();
+ fireEvent.click(screen.getByRole('button', {name: '国家代码'}));
+ fireEvent.change(await screen.findByPlaceholderText('输入代码并添加'), {target: {value: 'nz'}});
+ fireEvent.click(screen.getByRole('button', {name: '添加代码'}));
+ expect(onChange).toHaveBeenCalledWith(expect.objectContaining({config: expect.objectContaining({countries: ['NZ']})}));
+});
diff --git a/frontend/app/(main)/waf/rules/editor/components/node-properties.tsx b/frontend/app/(main)/waf/rules/editor/components/node-properties.tsx
new file mode 100644
index 00000000..7c6f7a8f
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/node-properties.tsx
@@ -0,0 +1,39 @@
+import {Settings2} from 'lucide-react';
+import {useState} from 'react';
+
+import {Button} from '@/components/ui/button';
+import {Checkbox} from '@/components/ui/checkbox';
+import {Field, FieldDescription, FieldGroup, FieldLabel} from '@/components/ui/field';
+import {Input} from '@/components/ui/input';
+import {Popover, PopoverContent, PopoverTrigger} from '@/components/ui/popover';
+import {ScrollArea} from '@/components/ui/scroll-area';
+import {Select, SelectContent, SelectGroup, SelectItem, SelectTrigger, SelectValue} from '@/components/ui/select';
+import {Separator} from '@/components/ui/separator';
+import {Textarea} from '@/components/ui/textarea';
+import type {WAFIPGroup, WAFRuleNode} from '@/lib/services/openflare';
+
+const countries = ['CN', 'US', 'JP', 'SG', 'DE', 'FR', 'GB', 'CA', 'AU', 'BR', 'IN', 'KR'].map((value) => ({value, label: value}));
+const regions = ['CN-BJ', 'CN-SH', 'CN-GD', 'CN-ZJ', 'US-CA', 'US-NY', 'US-TX', 'JP-13', 'DE-BE', 'GB-ENG'].map((value) => ({value, label: value}));
+
+export function NodeProperties({node, ipGroups, onChange}: {node?: WAFRuleNode; ipGroups: WAFIPGroup[]; onChange: (node: WAFRuleNode) => void}) {
+ return ;
+}
+
+function PropertyFields({node, ipGroups, onChange}: {node: WAFRuleNode; ipGroups: WAFIPGroup[]; onChange: (node: WAFRuleNode) => void}) {
+ if (node.type === 'start' || node.type === 'allow') return 系统节点无需配置。
;
+ if (node.type === 'ip_match') return onChange({...node, config: {...node.config, ips}})}/> onChange({...node, config: {...node.config, cidrs}})}/> ({value: String(group.id), label: group.name}))} value={node.config.ip_group_ids.map(String)} onChange={(values) => onChange({...node, config: {...node.config, ip_group_ids: values.map(Number)}})}/>;
+ if (node.type === 'geo_match') return onChange({...node, config: {...node.config, countries}})}/> onChange({...node, config: {...node.config, regions}})}/>;
+ if (node.type === 'pow') return 算法{(['difficulty', 'session_ttl', 'challenge_ttl'] as const).map((key) => onChange({...node, config: {...node.config, [key]: value}})}/>)};
+ return onChange({...node, config: {...node.config, status_code}})}/>HTML 响应体;
+}
+
+function CsvField({id, label, value, onChange}: {id: string; label: string; value: string[]; onChange: (value: string[]) => void}) { return {label}; }
+function NumberField({id, label, value, min, max, onChange}: {id: string; label: string; value: number; min?: number; max?: number; onChange: (value: number) => void}) { return {label} onChange(Number(event.target.value))}/>; }
+
+function MultiSelect({id, label, options, value, creatablePattern, onChange}: {id: string; label: string; options: {value: string; label: string}[]; value: string[]; creatablePattern?: RegExp; onChange: (value: string[]) => void}) {
+ const [draft, setDraft] = useState('');
+ const normalized = draft.trim().toUpperCase();
+ const visible = [...options, ...value.filter((selected) => !options.some((option) => option.value === selected)).map((selected) => ({value: selected, label: selected}))];
+ const canCreate = Boolean(creatablePattern?.test(normalized) && !value.includes(normalized));
+ return {label}{creatablePattern && setDraft(event.target.value)}/>
}{visible.length === 0 ? 暂无可选项
: visible.map((option) => )};
+}
diff --git a/frontend/app/(main)/waf/rules/editor/components/rule-flow-canvas.tsx b/frontend/app/(main)/waf/rules/editor/components/rule-flow-canvas.tsx
new file mode 100644
index 00000000..161ae35f
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/rule-flow-canvas.tsx
@@ -0,0 +1,68 @@
+'use client';
+
+import {useCallback, useEffect, useMemo, useRef} from 'react';
+import {addEdge, Background, Controls, MiniMap, ReactFlow, type Connection, type Edge, type Node, type NodeChange, applyNodeChanges, type EdgeChange, applyEdgeChanges, type ReactFlowInstance} from '@xyflow/react';
+import '@xyflow/react/dist/style.css';
+
+import type {WAFRuleEdge, WAFRuleGraph, WAFRuleNode} from '@/lib/services/openflare';
+
+import {type GraphIssue, removeNodeFromGraph} from './graph-validation';
+import {acceptedNodeChanges, filterRemovableNodeIds, type GraphErrorTarget, isConnectionAllowed, isPersistentEdgeChange} from './editor-behavior';
+import {NodeLibrary} from './node-library';
+import {RuleNode, type RuleFlowNodeData} from './rule-node';
+
+const nodeTypes = {rule: RuleNode};
+
+export function RuleFlowCanvas({graph, issues, selectedId, selectedEdgeId, focusTarget, onGraphChange, onSelect, onSelectEdge}: {graph: WAFRuleGraph; issues: GraphIssue[]; selectedId?: string; selectedEdgeId?: string; focusTarget?: GraphErrorTarget; onGraphChange: (graph: WAFRuleGraph, persistent?: boolean) => void; onSelect: (id?: string) => void; onSelectEdge: (id?: string) => void}) {
+ const instance = useRef, Edge> | null>(null);
+ const nodes = useMemo[]>(() => graph.nodes.map((rule) => ({id: rule.id, type: 'rule', position: rule.position, selected: rule.id === selectedId, data: {rule, issues: issues.filter((issue) => issue.nodeId === rule.id).length}})), [graph.nodes, issues, selectedId]);
+ const edges = useMemo(() => graph.edges.map((edge) => ({id: edge.id, source: edge.source, sourceHandle: edge.source_handle, target: edge.target, selected: edge.id === selectedEdgeId, animated: selectedId === edge.source || selectedId === edge.target})), [graph.edges, selectedEdgeId, selectedId]);
+
+ useEffect(() => {
+ if (!focusTarget || !instance.current) return;
+ if (focusTarget.kind === 'node') void instance.current.fitView({nodes: [{id: focusTarget.id}], duration: 350, maxZoom: 1.4});
+ else {
+ const edge = graph.edges.find((item) => item.id === focusTarget.id);
+ if (edge) void instance.current.fitView({nodes: [{id: edge.source}, {id: edge.target}], duration: 350, maxZoom: 1.4});
+ }
+ }, [focusTarget, graph.edges]);
+
+ const onNodesChange = useCallback((changes: NodeChange[]) => {
+ const accepted = acceptedNodeChanges(graph.nodes, changes);
+ if (accepted.changes.length === 0) return;
+ const removals = filterRemovableNodeIds(graph.nodes, accepted.changes.filter((change) => change.type === 'remove').map((change) => change.id));
+ let next = graph;
+ for (const id of removals) next = removeNodeFromGraph(next, id);
+ const positioned = applyNodeChanges(accepted.changes.filter((change) => change.type !== 'remove'), nodes);
+ const positions = new Map(positioned.map((node) => [node.id, node.position]));
+ onGraphChange({...next, nodes: next.nodes.map((node) => ({...node, position: positions.get(node.id) ?? node.position}))}, accepted.persistent);
+ }, [graph, nodes, onGraphChange]);
+
+ const onEdgesChange = useCallback((changes: EdgeChange[]) => {
+ const persistent = changes.filter(isPersistentEdgeChange);
+ if (persistent.length === 0) return;
+ const next = applyEdgeChanges(persistent, edges);
+ onGraphChange({...graph, edges: next.map(toRuleEdge)});
+ }, [edges, graph, onGraphChange]);
+
+ const isValidConnection = useCallback((connection: Edge | Connection) => {
+ return isConnectionAllowed(graph, connection);
+ }, [graph]);
+
+ const onConnect = useCallback((connection: Connection) => {
+ if (!isValidConnection(connection)) return;
+ const next = addEdge({...connection, id: `${connection.source}-${connection.sourceHandle}-${connection.target}`}, edges);
+ onGraphChange({...graph, edges: next.map(toRuleEdge)});
+ }, [edges, graph, isValidConnection, onGraphChange]);
+
+ const addNode = useCallback((type: 'ip_match' | 'geo_match' | 'pow' | 'block') => {
+ const id = `${type}-${crypto.randomUUID().slice(0, 8)}`;
+ const config = type === 'ip_match' ? {ips: [], cidrs: [], ip_group_ids: []} : type === 'geo_match' ? {countries: [], regions: []} : type === 'pow' ? {algorithm: 'fast' as const, difficulty: 4, session_ttl: 3600, challenge_ttl: 300} : {status_code: 403, response_body: ''};
+ onGraphChange({...graph, nodes: [...graph.nodes, {id, type, position: {x: 240, y: 140 + graph.nodes.length * 24}, config} as WAFRuleNode]});
+ onSelect(id);
+ }, [graph, onGraphChange, onSelect]);
+
+ return
{ instance.current = value; }} onNodesChange={onNodesChange} onEdgesChange={onEdgesChange} onConnect={onConnect} isValidConnection={isValidConnection} onNodeClick={(_, node) => { onSelectEdge(undefined); onSelect(node.id); }} onEdgeClick={(_, edge) => { onSelect(undefined); onSelectEdge(edge.id); }} onPaneClick={() => { onSelect(undefined); onSelectEdge(undefined); }} fitView deleteKeyCode={['Backspace', 'Delete']}>;
+}
+
+function toRuleEdge(edge: Edge): WAFRuleEdge { return {id: edge.id, source: edge.source, source_handle: edge.sourceHandle ?? '', target: edge.target}; }
diff --git a/frontend/app/(main)/waf/rules/editor/components/rule-node.tsx b/frontend/app/(main)/waf/rules/editor/components/rule-node.tsx
new file mode 100644
index 00000000..98688e3d
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/rule-node.tsx
@@ -0,0 +1,39 @@
+import {Handle, Position, type NodeProps} from '@xyflow/react';
+import {Ban, Fingerprint, Flag, Globe2, Play, ShieldCheck} from 'lucide-react';
+
+import {Badge} from '@/components/ui/badge';
+import {cn} from '@/lib/utils';
+import type {WAFRuleNode} from '@/lib/services/openflare';
+
+export interface RuleFlowNodeData extends Record { rule: WAFRuleNode; issues: number }
+
+const meta = {
+ start: {label: '开始', icon: Play}, ip_match: {label: 'IP 匹配', icon: Fingerprint}, geo_match: {label: '地域匹配', icon: Globe2},
+ pow: {label: 'PoW 挑战', icon: ShieldCheck}, allow: {label: '通过', icon: Flag}, block: {label: '阻止', icon: Ban},
+} as const;
+
+const outputHandles: Partial> = {start: ['next'], ip_match: ['true', 'false'], geo_match: ['true', 'false'], pow: ['next']};
+
+export function RuleNode({data, selected}: NodeProps) {
+ const value = data as RuleFlowNodeData;
+ const {rule, issues} = value;
+ const {label, icon: Icon} = meta[rule.type];
+ return (
+ 0 && 'border-destructive')}>
+ {rule.type !== 'start' &&
}
+
+
+
+ {label}
+ {rule.id}
+
+ {issues > 0 &&
{issues}}
+
+ {(outputHandles[rule.type] ?? []).map((handle, index, all) => (
+
+ {handle}
+
+ ))}
+
+ );
+}
diff --git a/frontend/app/(main)/waf/rules/editor/components/unsaved-changes.test.tsx b/frontend/app/(main)/waf/rules/editor/components/unsaved-changes.test.tsx
new file mode 100644
index 00000000..13954ed4
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/unsaved-changes.test.tsx
@@ -0,0 +1,46 @@
+import {render} from '@testing-library/react';
+import {afterEach, expect, it, vi} from 'vitest';
+
+import {UnsavedChanges} from './unsaved-changes';
+
+afterEach(() => vi.restoreAllMocks());
+
+it('blocks same-origin application links when dirty and confirmation is declined', () => {
+ vi.spyOn(window, 'confirm').mockReturnValue(false);
+ const {container} = render(<>WAF>);
+ const event = new MouseEvent('click', {bubbles: true, cancelable: true, button: 0});
+ container.querySelector('a')!.dispatchEvent(event);
+ expect(window.confirm).toHaveBeenCalledOnce();
+ expect(event.defaultPrevented).toBe(true);
+});
+
+it('does not block application links without changes', () => {
+ const confirm = vi.spyOn(window, 'confirm');
+ const {getByRole} = render(<>WAF>);
+ const event = new MouseEvent('click', {bubbles: true, cancelable: true, button: 0});
+ event.preventDefault();
+ getByRole('link').dispatchEvent(event);
+ expect(confirm).not.toHaveBeenCalled();
+});
+
+it('restores declined Back and Forward transitions by indexed delta', () => {
+ vi.spyOn(window, 'confirm').mockReturnValue(false);
+ const go = vi.spyOn(history, 'go').mockImplementation(() => undefined);
+ history.replaceState({__wafEditorIndex: 4}, '');
+ render();
+ window.dispatchEvent(new PopStateEvent('popstate', {state: {__wafEditorIndex: 3}}));
+ expect(go).toHaveBeenLastCalledWith(1);
+ window.dispatchEvent(new PopStateEvent('popstate', {state: {__wafEditorIndex: 4}}));
+ window.dispatchEvent(new PopStateEvent('popstate', {state: {__wafEditorIndex: 6}}));
+ expect(go).toHaveBeenLastCalledWith(-2);
+});
+
+it('prompts and restores the current URL for an unknown unindexed history entry', () => {
+ vi.spyOn(window, 'confirm').mockReturnValue(false);
+ history.replaceState({__wafEditorIndex: 4}, '', '/waf/rules/editor?id=9');
+ const push = vi.spyOn(history, 'pushState');
+ render();
+ window.dispatchEvent(new PopStateEvent('popstate', {state: {legacy: true}}));
+ expect(window.confirm).toHaveBeenCalledOnce();
+ expect(push).toHaveBeenCalledWith(expect.objectContaining({__wafEditorIndex: 4}), '', '/waf/rules/editor?id=9');
+});
diff --git a/frontend/app/(main)/waf/rules/editor/components/unsaved-changes.tsx b/frontend/app/(main)/waf/rules/editor/components/unsaved-changes.tsx
new file mode 100644
index 00000000..88f35f3c
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/components/unsaved-changes.tsx
@@ -0,0 +1,42 @@
+'use client';
+
+import {useEffect} from 'react';
+import {getHistoryTransition} from './editor-behavior';
+
+const historyIndexKey = '__wafEditorIndex';
+
+export function UnsavedChanges({dirty}: {dirty: boolean}) {
+ useEffect(() => {
+ if (!dirty) return;
+ const initialState = history.state && typeof history.state === 'object' ? history.state : {};
+ let currentIndex = Number.isInteger(initialState[historyIndexKey]) ? initialState[historyIndexKey] as number : 0;
+ const currentUrl = window.location.pathname + window.location.search + window.location.hash;
+ history.replaceState({...initialState, [historyIndexKey]: currentIndex}, '');
+ const originalPushState = history.pushState.bind(history);
+ const originalReplaceState = history.replaceState.bind(history);
+ history.pushState = (data, unused, url) => { currentIndex++; originalPushState({...data, [historyIndexKey]: currentIndex}, unused, url); };
+ history.replaceState = (data, unused, url) => originalReplaceState({...data, [historyIndexKey]: currentIndex}, unused, url);
+ let restoring = false;
+ const handler = (event: BeforeUnloadEvent) => { if (dirty) event.preventDefault(); };
+ const clickHandler = (event: MouseEvent) => {
+ if (!dirty || event.defaultPrevented || event.button !== 0 || event.metaKey || event.ctrlKey || event.shiftKey || event.altKey) return;
+ const link = (event.target as Element | null)?.closest('a[href]') as HTMLAnchorElement | null;
+ if (!link || link.target === '_blank' || new URL(link.href, window.location.href).origin !== window.location.origin) return;
+ if (!window.confirm('存在未保存的更改,确定离开吗?')) event.preventDefault();
+ };
+ const popstateHandler = (event: PopStateEvent) => {
+ const hasTargetIndex = Number.isInteger(event.state?.[historyIndexKey]);
+ const targetIndex = hasTargetIndex ? event.state[historyIndexKey] as number : currentIndex;
+ if (restoring) { restoring = false; currentIndex = targetIndex; return; }
+ if (window.confirm('存在未保存的更改,确定离开吗?')) { currentIndex = targetIndex; return; }
+ if (!hasTargetIndex) { originalPushState({...initialState, [historyIndexKey]: currentIndex}, '', currentUrl); return; }
+ restoring = true;
+ history.go(getHistoryTransition(currentIndex, targetIndex).restoreDelta);
+ };
+ window.addEventListener('beforeunload', handler);
+ document.addEventListener('click', clickHandler, true);
+ window.addEventListener('popstate', popstateHandler);
+ return () => { history.pushState = originalPushState; history.replaceState = originalReplaceState; window.removeEventListener('beforeunload', handler); document.removeEventListener('click', clickHandler, true); window.removeEventListener('popstate', popstateHandler); };
+ }, [dirty]);
+ return null;
+}
diff --git a/frontend/app/(main)/waf/rules/editor/page.test.tsx b/frontend/app/(main)/waf/rules/editor/page.test.tsx
new file mode 100644
index 00000000..5c3514c4
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/page.test.tsx
@@ -0,0 +1,124 @@
+import {QueryClient, QueryClientProvider} from '@tanstack/react-query';
+import {act, fireEvent, render, screen, waitFor} from '@testing-library/react';
+import {AxiosError, type AxiosResponse} from 'axios';
+import {beforeEach, expect, it, vi} from 'vitest';
+
+import type {WAFRule, WAFRuleGraph} from '@/lib/services';
+
+import WAFRuleEditorPage from './page';
+
+const {getRule, saveRuleGraph, updateRuleMeta, listIPGroups, toastError} = vi.hoisted(() => ({
+ getRule: vi.fn(),
+ saveRuleGraph: vi.fn(),
+ updateRuleMeta: vi.fn(),
+ listIPGroups: vi.fn().mockResolvedValue([]),
+ toastError: vi.fn(),
+}));
+
+vi.mock('next/navigation', () => ({useRouter: () => ({push: vi.fn()}), useSearchParams: () => new URLSearchParams('id=9')}));
+vi.mock('sonner', () => ({toast: {success: vi.fn(), error: toastError}}));
+vi.mock('@/lib/services', async (importOriginal) => {
+ const actual = await importOriginal();
+ return {...actual, services: {...actual.services, openflareWaf: {getRule, saveRuleGraph, updateRuleMeta, listIPGroups}}};
+});
+vi.mock('./components/rule-flow-canvas', () => ({RuleFlowCanvas: ({graph, focusTarget, onGraphChange}: {graph: WAFRuleGraph; focusTarget?: {kind: string; id: string}; onGraphChange: (graph: WAFRuleGraph) => void}) => {focusTarget && focus:{focusTarget.kind}:{focusTarget.id}}
}));
+vi.mock('./components/node-properties', () => ({NodeProperties: () => null}));
+
+const graph: WAFRuleGraph = {schema_version: 1, nodes: [
+ {id: 'start', type: 'start', position: {x: 0, y: 0}, config: {}},
+ {id: 'allow', type: 'allow', position: {x: 200, y: 0}, config: {}},
+], edges: [{id: 'start-allow', source: 'start', source_handle: 'next', target: 'allow'}]};
+const rule = {id: 9, name: '边缘防护', enabled: true, is_global: false, graph, revision: 1, applied_site_ids: [], applied_site_count: 0, created_at: '', updated_at: ''} satisfies WAFRule;
+
+function renderPage() {
+ const client = new QueryClient({defaultOptions: {queries: {retry: false}, mutations: {retry: false}}});
+ return render();
+}
+
+beforeEach(() => {
+ getRule.mockReset();
+ saveRuleGraph.mockReset();
+ updateRuleMeta.mockReset();
+ listIPGroups.mockClear();
+ toastError.mockReset();
+});
+
+it('renders a query error with a working retry action', async () => {
+ getRule.mockRejectedValueOnce(new Error('offline')).mockResolvedValueOnce(rule);
+ renderPage();
+ fireEvent.click(await screen.findByRole('button', {name: '重新加载'}));
+ expect(await screen.findByRole('heading', {name: '边缘防护'})).toBeInTheDocument();
+ expect(getRule).toHaveBeenCalledTimes(2);
+});
+
+it('exposes conflict reload and maps typed server node errors to canvas focus', async () => {
+ getRule.mockResolvedValue(rule);
+ const conflict = new AxiosError('conflict');
+ conflict.response = {status: 409, data: {}, headers: {}, config: {headers: {}}} as AxiosResponse;
+ saveRuleGraph.mockRejectedValueOnce(conflict).mockRejectedValueOnce(new Error('规则图无效: 节点 start 的 next 出口未连接'));
+ renderPage();
+ await screen.findByRole('heading', {name: '边缘防护'});
+ fireEvent.click(screen.getByRole('button', {name: '修改画布'}));
+ fireEvent.click(screen.getByRole('button', {name: '保存'}));
+ expect(await screen.findByRole('button', {name: '重新加载'})).toBeInTheDocument();
+ fireEvent.click(screen.getByRole('button', {name: '保存'}));
+ await waitFor(() => expect(screen.getByText('focus:node:start')).toBeInTheDocument());
+});
+
+it('optimistically enables a saved valid rule with its current name', async () => {
+ const disabledRule = {...rule, enabled: false};
+ let resolveUpdate: (value: WAFRule) => void = () => undefined;
+ updateRuleMeta.mockImplementation(() => new Promise((resolve) => {
+ resolveUpdate = resolve;
+ }));
+ getRule.mockResolvedValue(disabledRule);
+
+ renderPage();
+ expect(await screen.findByText('已停用')).toBeInTheDocument();
+ fireEvent.click(screen.getByRole('switch', {name: '启用规则'}));
+
+ await waitFor(() => {
+ expect(updateRuleMeta).toHaveBeenCalledWith(9, {
+ name: '边缘防护',
+ enabled: true,
+ });
+ expect(screen.getByText('已启用')).toBeInTheDocument();
+ });
+
+ await act(async () => resolveUpdate({...disabledRule, enabled: true}));
+});
+
+it('rolls back an optimistic enabled change when metadata update fails', async () => {
+ const disabledRule = {...rule, enabled: false};
+ let rejectUpdate: (reason: Error) => void = () => undefined;
+ updateRuleMeta.mockImplementation(() => new Promise((_resolve, reject) => {
+ rejectUpdate = reject;
+ }));
+ getRule.mockResolvedValue(disabledRule);
+
+ renderPage();
+ expect(await screen.findByText('已停用')).toBeInTheDocument();
+ fireEvent.click(screen.getByRole('switch', {name: '启用规则'}));
+ await waitFor(() => expect(screen.getByText('已启用')).toBeInTheDocument());
+
+ await act(async () => rejectUpdate(new Error('网络不可用')));
+
+ await waitFor(() => expect(screen.getByText('已停用')).toBeInTheDocument());
+ expect(toastError).toHaveBeenCalledWith('网络不可用');
+});
+
+it('requires graph changes to be saved before changing enabled state', async () => {
+ const disabledRule = {...rule, enabled: false};
+ getRule.mockResolvedValue(disabledRule);
+ saveRuleGraph.mockResolvedValue({...disabledRule, revision: 2});
+
+ renderPage();
+ const enabledSwitch = await screen.findByRole('switch', {name: '启用规则'});
+ expect(enabledSwitch).toBeEnabled();
+
+ fireEvent.click(screen.getByRole('button', {name: '修改画布'}));
+ expect(enabledSwitch).toBeDisabled();
+ fireEvent.click(screen.getByRole('button', {name: '保存'}));
+
+ await waitFor(() => expect(enabledSwitch).toBeEnabled());
+});
diff --git a/frontend/app/(main)/waf/rules/editor/page.tsx b/frontend/app/(main)/waf/rules/editor/page.tsx
new file mode 100644
index 00000000..8c982ae1
--- /dev/null
+++ b/frontend/app/(main)/waf/rules/editor/page.tsx
@@ -0,0 +1,123 @@
+'use client';
+
+import {Suspense, useCallback, useEffect, useMemo, useState} from 'react';
+import {useMutation, useQuery, useQueryClient} from '@tanstack/react-query';
+import axios from 'axios';
+import {ArrowLeft, GitBranch, Save} from 'lucide-react';
+import {useRouter, useSearchParams} from 'next/navigation';
+import {toast} from 'sonner';
+
+import {Badge} from '@/components/ui/badge';
+import {Button} from '@/components/ui/button';
+import {Label} from '@/components/ui/label';
+import {Skeleton} from '@/components/ui/skeleton';
+import {Switch} from '@/components/ui/switch';
+import {services, type WAFRule, type WAFRuleGraph, type WAFRuleNode} from '@/lib/services';
+
+import {getErrorMessage} from '../../components/helpers';
+import {validateGraph} from './components/graph-validation';
+import {NodeProperties} from './components/node-properties';
+import {RuleFlowCanvas} from './components/rule-flow-canvas';
+import {UnsavedChanges} from './components/unsaved-changes';
+import {findGraphErrorTarget, type GraphErrorTarget} from './components/editor-behavior';
+
+export default function WAFRuleEditorPage() {
+ return }>;
+}
+
+function EditorContent() {
+ const router = useRouter();
+ const searchParams = useSearchParams();
+ const queryClient = useQueryClient();
+ const id = Number(searchParams.get('id'));
+ const [graph, setGraph] = useState();
+ const [revision, setRevision] = useState(0);
+ const [selectedId, setSelectedId] = useState();
+ const [selectedEdgeId, setSelectedEdgeId] = useState();
+ const [focusTarget, setFocusTarget] = useState();
+ const [dirty, setDirty] = useState(false);
+ const [conflict, setConflict] = useState(false);
+ const ruleQueryKey = ['waf-rule', id] as const;
+
+ const ruleQuery = useQuery({queryKey: ruleQueryKey, queryFn: () => services.openflareWaf.getRule(id), enabled: Number.isFinite(id) && id > 0});
+ const ipGroupsQuery = useQuery({queryKey: ['waf-ip-groups'], queryFn: () => services.openflareWaf.listIPGroups()});
+ useEffect(() => { if (ruleQuery.data && !dirty) { setGraph(ruleQuery.data.graph); setRevision(ruleQuery.data.revision); } }, [dirty, ruleQuery.data]);
+ const issues = useMemo(() => graph ? validateGraph(graph) : [], [graph]);
+ const selected = graph?.nodes.find((node) => node.id === selectedId);
+
+ const saveMutation = useMutation({
+ mutationFn: () => services.openflareWaf.saveRuleGraph(id, {revision, graph: graph!}),
+ onSuccess: (rule) => { setGraph(rule.graph); setRevision(rule.revision); setDirty(false); setConflict(false); queryClient.setQueryData(['waf-rule', id], rule); toast.success('规则图已保存'); },
+ onError: (error) => {
+ if (axios.isAxiosError(error) && error.response?.status === 409) { setConflict(true); toast.error('规则已在其他页面更新,请重新加载'); return; }
+ const payload = axios.isAxiosError(error) ? error.response?.data : error;
+ const target = graph ? findGraphErrorTarget(payload, graph.nodes.map((node) => node.id), graph.edges.map((edge) => edge.id)) : undefined;
+ if (target?.kind === 'node') { setSelectedEdgeId(undefined); setSelectedId(target.id); }
+ if (target?.kind === 'edge') { setSelectedId(undefined); setSelectedEdgeId(target.id); }
+ setFocusTarget(target ? {...target} : undefined);
+ toast.error('保存失败,请检查标记的节点和连线');
+ },
+ });
+
+ const metaMutation = useMutation({
+ mutationFn: (enabled: boolean) => services.openflareWaf.updateRuleMeta(id, {
+ name: ruleQuery.data!.name,
+ enabled,
+ }),
+ onMutate: async (enabled) => {
+ await queryClient.cancelQueries({queryKey: ruleQueryKey});
+ const previous = queryClient.getQueryData(ruleQueryKey);
+ queryClient.setQueryData(ruleQueryKey, (current) =>
+ current ? {...current, enabled} : current,
+ );
+ return {previous};
+ },
+ onError: (error, _enabled, context) => {
+ if (context?.previous) {
+ queryClient.setQueryData(ruleQueryKey, context.previous);
+ }
+ toast.error(getErrorMessage(error));
+ },
+ onSuccess: (rule) => {
+ queryClient.setQueryData(ruleQueryKey, rule);
+ void queryClient.invalidateQueries({queryKey: ruleQueryKey, refetchType: 'none'});
+ void queryClient.invalidateQueries({queryKey: ['openflare', 'waf', 'rule-groups']});
+ void queryClient.invalidateQueries({queryKey: ['openflare', 'config-versions', 'diff']});
+ toast.success(rule.enabled ? '规则已启用' : '规则已停用');
+ },
+ });
+
+ const changeGraph = useCallback((next: WAFRuleGraph, persistent = true) => { setGraph(next); if (persistent) { setDirty(true); setConflict(false); } }, []);
+ const changeNode = useCallback((next: WAFRuleNode) => { if (!graph) return; changeGraph({...graph, nodes: graph.nodes.map((node) => node.id === next.id ? next : node)}); }, [changeGraph, graph]);
+ const leave = () => { if (!dirty || window.confirm('存在未保存的更改,确定离开吗?')) router.push('/waf'); };
+
+ if (!Number.isFinite(id) || id <= 0) return ;
+ if (ruleQuery.isError) return 规则加载失败,请重试。
;
+ if (ruleQuery.isLoading || !graph || !ruleQuery.data) return ;
+
+ return (
+
+ );
+}
+
+function EditorSkeleton() { return ; }
diff --git a/frontend/lib/services/openflare/index.ts b/frontend/lib/services/openflare/index.ts
index 4b975c0c..607dbfbf 100644
--- a/frontend/lib/services/openflare/index.ts
+++ b/frontend/lib/services/openflare/index.ts
@@ -76,6 +76,7 @@ export type {
WAFRuleGraph,
WAFRuleNode,
WAFSaveRuleGraphPayload,
+ WAFUpdateRuleMetaPayload,
WAFRuleGroup,
WAFRuleGroupPayload,
WAFSiteRuleGroups,
diff --git a/frontend/lib/services/openflare/types.ts b/frontend/lib/services/openflare/types.ts
index 200e96ce..67cdd38c 100644
--- a/frontend/lib/services/openflare/types.ts
+++ b/frontend/lib/services/openflare/types.ts
@@ -773,6 +773,11 @@ export interface WAFSaveRuleGraphPayload {
graph: WAFRuleGraph;
}
+export interface WAFUpdateRuleMetaPayload {
+ name: string;
+ enabled: boolean;
+}
+
export interface WAFRuleGroupPayload {
name: string;
enabled: boolean;
diff --git a/frontend/lib/services/openflare/waf.service.ts b/frontend/lib/services/openflare/waf.service.ts
index d7798637..9d7dadb0 100644
--- a/frontend/lib/services/openflare/waf.service.ts
+++ b/frontend/lib/services/openflare/waf.service.ts
@@ -9,6 +9,7 @@ import type {
WAFRule,
WAFSaveRuleGraphPayload,
WAFSiteRuleGroups,
+ WAFUpdateRuleMetaPayload,
} from './types';
export class WafService extends OpenFlareBaseService {
@@ -33,12 +34,15 @@ export class WafService extends OpenFlareBaseService {
return this.post(`/rule-groups/${id}/graph`, payload);
}
- static async deleteRuleGroup(id: number): Promise {
- return this.post(`/rule-groups/${id}/delete`);
+ static async updateRuleMeta(
+ id: number,
+ payload: WAFUpdateRuleMetaPayload,
+ ): Promise {
+ return this.post(`/rule-groups/${id}/meta`, payload);
}
- static async updateRuleGroupSites(id: number, ids: number[]): Promise {
- return this.post(`/rule-groups/${id}/sites`, { ids });
+ static async deleteRuleGroup(id: number): Promise {
+ return this.post(`/rule-groups/${id}/delete`);
}
static async listSiteRuleGroups(routeId: number): Promise {
diff --git a/frontend/tests/unit/waf-rule-service.test.ts b/frontend/tests/unit/waf-rule-service.test.ts
index c43b3969..014d3447 100644
--- a/frontend/tests/unit/waf-rule-service.test.ts
+++ b/frontend/tests/unit/waf-rule-service.test.ts
@@ -6,7 +6,7 @@ import {beforeEach, describe, expect, it, vi} from 'vitest';
import WafPage from '@/app/(main)/waf/page';
import apiClient from '@/lib/services/core/api-client';
-import {ProxyRouteService, WafService} from '@/lib/services/openflare';
+import {WafService} from '@/lib/services/openflare';
import type {WAFSiteRuleGroups} from '@/lib/services/openflare';
import {WafService as DirectWafService} from '@/lib/services/openflare/waf.service';
@@ -32,11 +32,6 @@ vi.mock('@/lib/services/openflare', async (importOriginal) => {
listIPGroups: vi.fn(),
createRule: vi.fn(),
deleteRuleGroup: vi.fn(),
- updateRuleGroupSites: vi.fn(),
- },
- ProxyRouteService: {
- ...actual.ProxyRouteService,
- list: vi.fn(),
},
};
});
@@ -109,6 +104,27 @@ describe('WafService rule graph API', () => {
undefined,
);
});
+
+ it('updates rule metadata with the current name and enabled state', async () => {
+ vi.mocked(apiClient.post).mockResolvedValue({
+ data: {error_msg: '', data: {...ruleSummary, enabled: true}},
+ } as AxiosResponse);
+
+ await DirectWafService.updateRuleMeta(17, {
+ name: '入口防护',
+ enabled: true,
+ });
+
+ expect(apiClient.post).toHaveBeenCalledWith(
+ '/api/v1/d/waf/rule-groups/17/meta',
+ {name: '入口防护', enabled: true},
+ undefined,
+ );
+ });
+
+ it('does not expose the retired rule-to-sites binding service', () => {
+ expect(DirectWafService).not.toHaveProperty('updateRuleGroupSites');
+ });
});
describe('WAF rule creation flow', () => {
@@ -118,8 +134,6 @@ describe('WAF rule creation flow', () => {
vi.mocked(WafService.listRuleGroups).mockResolvedValue([]);
vi.mocked(WafService.listIPGroups).mockReset();
vi.mocked(WafService.listIPGroups).mockResolvedValue([]);
- vi.mocked(ProxyRouteService.list).mockReset();
- vi.mocked(ProxyRouteService.list).mockResolvedValue([]);
vi.mocked(WafService.createRule).mockReset();
vi.mocked(WafService.createRule).mockResolvedValue(ruleSummary);
});
diff --git a/internal/apps/agent/config/config.go b/internal/apps/agent/config/config.go
index a7c23426..127fe060 100644
--- a/internal/apps/agent/config/config.go
+++ b/internal/apps/agent/config/config.go
@@ -24,6 +24,7 @@ const (
defaultRuntimeConfigDirRelativePath = "etc/openflare"
defaultPagesDirRelativePath = "var/lib/openflare/pages"
defaultMMDBRelativePath = "etc/openflare/GeoLite2-Country.mmdb"
+ defaultCityMMDBRelativePath = "etc/openflare/GeoLite2-City.mmdb"
defaultAccessLogRelativePath = "var/log/openflare/access.log"
defaultStateRelativePath = "var/lib/openflare/agent-state.json"
defaultObservabilityBufferRelativePath = "var/lib/openflare/observability-buffer.json"
@@ -31,6 +32,7 @@ const (
defaultObservabilityReplayMinutes = 15
defaultMMDBUpdateInterval = 24 * time.Hour
defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb"
+ defaultCityMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-City.mmdb"
defaultHeartbeatInterval = 10 * time.Second
defaultRequestTimeout = 10 * time.Second
configFilePerm = 0o600
@@ -58,8 +60,10 @@ type Config struct {
RuntimeConfigDir string `json:"runtime_config_dir"`
PagesDir string `json:"pages_dir"`
MMDBPath string `json:"mmdb_path"`
+ CityMMDBPath string `json:"city_mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"`
+ CityMMDBDownloadURL string `json:"city_mmdb_download_url"`
OpenrestyObservabilityPort int `json:"openresty_observability_port"`
ObservabilityBufferPath string `json:"observability_buffer_path"`
ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
@@ -89,8 +93,10 @@ type configFile struct {
RuntimeConfigDir string `json:"runtime_config_dir"`
PagesDir string `json:"pages_dir"`
MMDBPath string `json:"mmdb_path"`
+ CityMMDBPath string `json:"city_mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"`
+ CityMMDBDownloadURL string `json:"city_mmdb_download_url"`
OpenrestyObservabilityPort int `json:"openresty_observability_port"`
ObservabilityBufferPath string `json:"observability_buffer_path"`
ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
@@ -178,6 +184,7 @@ func applyAgentPathDefaults(cfg *Config, baseDir string) {
{&cfg.RuntimeConfigDir, defaultRuntimeConfigDirRelativePath},
{&cfg.PagesDir, defaultPagesDirRelativePath},
{&cfg.MMDBPath, defaultMMDBRelativePath},
+ {&cfg.CityMMDBPath, defaultCityMMDBRelativePath},
{&cfg.ObservabilityBufferPath, defaultObservabilityBufferRelativePath},
}
for _, item := range pathDefaults {
@@ -200,6 +207,9 @@ func applyAgentTimingDefaults(cfg *Config) {
if cfg.MMDBDownloadURL == "" {
cfg.MMDBDownloadURL = defaultMMDBDownloadURL
}
+ if cfg.CityMMDBDownloadURL == "" {
+ cfg.CityMMDBDownloadURL = defaultCityMMDBDownloadURL
+ }
if cfg.OpenrestyObservabilityPort <= 0 {
cfg.OpenrestyObservabilityPort = defaultOpenRestyObservabilityPort
}
@@ -232,6 +242,7 @@ func normalizeManagedPaths(cfg *Config) {
&cfg.StatePath,
&cfg.ObservabilityBufferPath,
&cfg.MMDBPath,
+ &cfg.CityMMDBPath,
}
for _, p := range paths {
if usesSlashPath(*p) {
@@ -256,6 +267,8 @@ func hasEnvConfig() bool {
"OPENFLARE_MMDB_PATH",
"OPENFLARE_MMDB_UPDATE_INTERVAL",
"OPENFLARE_MMDB_DOWNLOAD_URL",
+ "OPENFLARE_CITY_MMDB_PATH",
+ "OPENFLARE_CITY_MMDB_DOWNLOAD_URL",
} {
if strings.TrimSpace(os.Getenv(key)) != "" {
return true
@@ -283,6 +296,8 @@ func applyEnvOverrides(cfg *Config) {
overrideString("OPENFLARE_PAGES_DIR", &cfg.PagesDir)
overrideString("OPENFLARE_MMDB_PATH", &cfg.MMDBPath)
overrideString("OPENFLARE_MMDB_DOWNLOAD_URL", &cfg.MMDBDownloadURL)
+ overrideString("OPENFLARE_CITY_MMDB_PATH", &cfg.CityMMDBPath)
+ overrideString("OPENFLARE_CITY_MMDB_DOWNLOAD_URL", &cfg.CityMMDBDownloadURL)
if value := strings.TrimSpace(os.Getenv("OPENFLARE_HEARTBEAT_INTERVAL")); value != "" {
if duration, err := parseDurationValue(value); err == nil {
cfg.HeartbeatInterval = duration
diff --git a/internal/apps/agent/config/config_test.go b/internal/apps/agent/config/config_test.go
index d27013b4..b68e7bf6 100644
--- a/internal/apps/agent/config/config_test.go
+++ b/internal/apps/agent/config/config_test.go
@@ -61,6 +61,12 @@ func TestLoadDefaultsToManagedBinaryPaths(t *testing.T) {
if cfg.RuntimeConfigDir != filepath.Join(dir, "data", defaultRuntimeConfigDirRelativePath) {
t.Fatalf("unexpected runtime config dir: %s", cfg.RuntimeConfigDir)
}
+ if cfg.CityMMDBPath != filepath.Join(dir, "data", defaultCityMMDBRelativePath) {
+ t.Fatalf("unexpected city mmdb path: %s", cfg.CityMMDBPath)
+ }
+ if cfg.CityMMDBDownloadURL != defaultCityMMDBDownloadURL {
+ t.Fatalf("unexpected city mmdb download URL: %s", cfg.CityMMDBDownloadURL)
+ }
if cfg.OpenrestyCertDir != cfg.CertDir {
t.Fatalf("unexpected openresty cert dir: %s", cfg.OpenrestyCertDir)
}
@@ -333,6 +339,8 @@ func TestLoadEnvOverridesConfigFile(t *testing.T) {
t.Setenv("OPENFLARE_SERVER_URL", "http://new:3000")
t.Setenv("OPENFLARE_AGENT_TOKEN", "new-token")
t.Setenv("OPENFLARE_OPENRESTY_PATH", "/new/openresty")
+ t.Setenv("OPENFLARE_CITY_MMDB_PATH", "/new/GeoLite2-City.mmdb")
+ t.Setenv("OPENFLARE_CITY_MMDB_DOWNLOAD_URL", "https://geo.example/GeoLite2-City.mmdb")
cfg, err := Load(configPath)
if err != nil {
@@ -347,6 +355,25 @@ func TestLoadEnvOverridesConfigFile(t *testing.T) {
if cfg.OpenrestyPath != "/new/openresty" {
t.Fatalf("expected openresty path from env, got %s", cfg.OpenrestyPath)
}
+ if cfg.CityMMDBPath != "/new/GeoLite2-City.mmdb" || cfg.CityMMDBDownloadURL != "https://geo.example/GeoLite2-City.mmdb" {
+ t.Fatalf("unexpected City MMDB env overrides: %s / %s", cfg.CityMMDBPath, cfg.CityMMDBDownloadURL)
+ }
+}
+
+func TestLoadKeepsExplicitCityMMDBConfig(t *testing.T) {
+ dir := t.TempDir()
+ configPath := filepath.Join(dir, "agent.json")
+ payload := `{"server_url":"http://127.0.0.1:3000","agent_token":"token","node_name":"edge-01","node_ip":"10.0.0.8","city_mmdb_path":"/custom/GeoLite2-City.mmdb","city_mmdb_download_url":"https://custom.example/GeoLite2-City.mmdb"}`
+ if err := os.WriteFile(configPath, []byte(payload), 0o644); err != nil {
+ t.Fatalf("failed to write config: %v", err)
+ }
+ cfg, err := Load(configPath)
+ if err != nil {
+ t.Fatalf("Load failed: %v", err)
+ }
+ if cfg.CityMMDBPath != "/custom/GeoLite2-City.mmdb" || cfg.CityMMDBDownloadURL != "https://custom.example/GeoLite2-City.mmdb" {
+ t.Fatalf("explicit City MMDB config changed: %s / %s", cfg.CityMMDBPath, cfg.CityMMDBDownloadURL)
+ }
}
func TestLoadUsesMillisecondsForIntervals(t *testing.T) {
@@ -431,6 +458,9 @@ func TestSavePersistsMillisecondsAndOmitsRuntimeVersions(t *testing.T) {
if decoded["observability_replay_minutes"] != float64(defaultObservabilityReplayMinutes) {
t.Fatalf("unexpected observability replay minutes: %#v", decoded["observability_replay_minutes"])
}
+ if decoded["city_mmdb_path"] != cfg.CityMMDBPath || decoded["city_mmdb_download_url"] != cfg.CityMMDBDownloadURL {
+ t.Fatalf("City MMDB config was not persisted: %#v", decoded)
+ }
if _, ok := decoded["nginx_path"]; ok {
t.Fatal("legacy nginx_path should not be persisted")
}
diff --git a/internal/apps/agent/geoipupdate/updater.go b/internal/apps/agent/geoipupdate/updater.go
index e373617d..17eed6e1 100644
--- a/internal/apps/agent/geoipupdate/updater.go
+++ b/internal/apps/agent/geoipupdate/updater.go
@@ -3,6 +3,7 @@ package geoipupdate
import (
"context"
+ "errors"
"fmt"
"io/fs"
"log/slog"
@@ -22,9 +23,12 @@ const (
// Updater periodically downloads a fresh GeoIP MMDB file and seeds the
// initial embedded database when none is present on disk.
type Updater struct {
- MMDBPath string
- DownloadURL string
- UpdateInterval time.Duration
+ MMDBPath string
+ DownloadURL string
+ CityMMDBPath string
+ CityDownloadURL string
+ UpdateInterval time.Duration
+ downloadDatabase func(context.Context, string, string) error
}
// EnsureInitialDatabase seeds the MMDB file from the embedded database if it does not exist on disk.
@@ -52,13 +56,68 @@ func (u *Updater) EnsureInitialDatabase() error {
return nil
}
+// EnsureInitialDatabases retains the embedded Country seed and immediately
+// downloads City when it is absent so subdivision rules work before the first ticker interval.
+func (u *Updater) EnsureInitialDatabases(ctx context.Context) error {
+ var errs []error
+ if err := u.EnsureInitialDatabase(); err != nil {
+ errs = append(errs, err)
+ }
+ cityPath := filepath.Clean(u.CityMMDBPath)
+ if cityPath == "" || cityPath == "." || u.CityDownloadURL == "" {
+ return errors.Join(errs...)
+ }
+ if _, err := os.Stat(cityPath); err == nil {
+ return errors.Join(errs...)
+ } else if !os.IsNotExist(err) {
+ errs = append(errs, fmt.Errorf("stat City mmdb file failed: %w", err))
+ return errors.Join(errs...)
+ }
+ if err := u.download(ctx, cityPath, u.CityDownloadURL); err != nil {
+ errs = append(errs, fmt.Errorf("download initial City mmdb failed: %w", err))
+ } else {
+ slog.Info("initialized GeoIP City mmdb from provider", "path", cityPath)
+ }
+ return errors.Join(errs...)
+}
+
+func (u *Updater) download(ctx context.Context, path string, downloadURL string) error {
+ if u.downloadDatabase != nil {
+ return u.downloadDatabase(ctx, path, downloadURL)
+ }
+ return geoip.DownloadMaxMindDatabase(ctx, path, downloadURL)
+}
+
+func (u *Updater) updateDatabases(ctx context.Context) error {
+ databases := []struct {
+ name string
+ path string
+ downloadURL string
+ }{
+ {name: "Country", path: u.MMDBPath, downloadURL: u.DownloadURL},
+ {name: "City", path: u.CityMMDBPath, downloadURL: u.CityDownloadURL},
+ }
+ var errs []error
+ for _, database := range databases {
+ if database.path == "" || (database.name == "City" && database.downloadURL == "") {
+ continue
+ }
+ if err := u.download(ctx, database.path, database.downloadURL); err != nil {
+ errs = append(errs, fmt.Errorf("update GeoIP %s mmdb failed: %w", database.name, err))
+ continue
+ }
+ slog.Info("GeoIP mmdb updated", "database", database.name, "path", database.path)
+ }
+ return errors.Join(errs...)
+}
+
// Run starts the periodic GeoIP update loop and blocks until ctx is cancelled.
func (u *Updater) Run(ctx context.Context) {
if u == nil || u.MMDBPath == "" || u.UpdateInterval <= 0 {
return
}
- if err := u.EnsureInitialDatabase(); err != nil {
- slog.Warn("initialize GeoIP mmdb failed", "path", u.MMDBPath, "error", err)
+ if err := u.EnsureInitialDatabases(ctx); err != nil {
+ slog.Warn("initialize GeoIP databases failed", "country_path", u.MMDBPath, "city_path", u.CityMMDBPath, "error", err)
}
ticker := time.NewTicker(u.UpdateInterval)
defer ticker.Stop()
@@ -67,11 +126,9 @@ func (u *Updater) Run(ctx context.Context) {
case <-ctx.Done():
return
case <-ticker.C:
- if err := geoip.DownloadMaxMindDatabase(ctx, u.MMDBPath, u.DownloadURL); err != nil {
- slog.Warn("update GeoIP mmdb failed", "path", u.MMDBPath, "error", err)
- continue
+ if err := u.updateDatabases(ctx); err != nil {
+ slog.Warn("update GeoIP databases failed", "error", err)
}
- slog.Info("GeoIP mmdb updated", "path", u.MMDBPath)
}
}
}
diff --git a/internal/apps/agent/geoipupdate/updater_test.go b/internal/apps/agent/geoipupdate/updater_test.go
index 07dffda3..054c312e 100644
--- a/internal/apps/agent/geoipupdate/updater_test.go
+++ b/internal/apps/agent/geoipupdate/updater_test.go
@@ -1,8 +1,11 @@
package geoipupdate
import (
+ "context"
+ "errors"
"os"
"path/filepath"
+ "slices"
"testing"
)
@@ -22,3 +25,77 @@ func TestEnsureInitialDatabaseCopiesEmbeddedMMDB(t *testing.T) {
t.Fatal("expected copied mmdb to be non-empty")
}
}
+
+func TestEnsureInitialDatabasesDownloadsMissingCity(t *testing.T) {
+ tempDir := t.TempDir()
+ countryPath := filepath.Join(tempDir, "GeoLite2-Country.mmdb")
+ cityPath := filepath.Join(tempDir, "GeoLite2-City.mmdb")
+ updater := &Updater{
+ MMDBPath: countryPath,
+ CityMMDBPath: cityPath,
+ CityDownloadURL: "https://geo.example/GeoLite2-City.mmdb",
+ downloadDatabase: func(_ context.Context, path, downloadURL string) error {
+ if path != cityPath || downloadURL != "https://geo.example/GeoLite2-City.mmdb" {
+ t.Fatalf("unexpected initial download: %s / %s", path, downloadURL)
+ }
+ return os.WriteFile(path, []byte("city-mmdb"), 0o600)
+ },
+ }
+
+ if err := updater.EnsureInitialDatabases(context.Background()); err != nil {
+ t.Fatalf("EnsureInitialDatabases failed: %v", err)
+ }
+ if _, err := os.Stat(countryPath); err != nil {
+ t.Fatalf("expected embedded Country database: %v", err)
+ }
+ data, err := os.ReadFile(cityPath)
+ if err != nil || string(data) != "city-mmdb" {
+ t.Fatalf("expected downloaded City database, data=%q err=%v", data, err)
+ }
+}
+
+func TestEnsureInitialDatabasesKeepsCountryFallbackWhenCityDownloadFails(t *testing.T) {
+ tempDir := t.TempDir()
+ countryPath := filepath.Join(tempDir, "GeoLite2-Country.mmdb")
+ cityPath := filepath.Join(tempDir, "GeoLite2-City.mmdb")
+ updater := &Updater{
+ MMDBPath: countryPath,
+ CityMMDBPath: cityPath,
+ CityDownloadURL: "https://geo.example/GeoLite2-City.mmdb",
+ downloadDatabase: func(_ context.Context, _, _ string) error {
+ return errors.New("city unavailable")
+ },
+ }
+
+ if err := updater.EnsureInitialDatabases(context.Background()); err == nil {
+ t.Fatal("expected City download error to be reported")
+ }
+ if _, err := os.Stat(countryPath); err != nil {
+ t.Fatalf("expected Country fallback to remain available: %v", err)
+ }
+ if _, err := os.Stat(cityPath); !errors.Is(err, os.ErrNotExist) {
+ t.Fatalf("expected failed City download not to create a database, err=%v", err)
+ }
+}
+
+func TestUpdateDatabasesAttemptsCityAfterCountryFailure(t *testing.T) {
+ var paths []string
+ updater := &Updater{
+ MMDBPath: "/data/GeoLite2-Country.mmdb",
+ DownloadURL: "https://geo.example/GeoLite2-Country.mmdb",
+ CityMMDBPath: "/data/GeoLite2-City.mmdb",
+ CityDownloadURL: "https://geo.example/GeoLite2-City.mmdb",
+ downloadDatabase: func(_ context.Context, path, _ string) error {
+ paths = append(paths, path)
+ if path == "/data/GeoLite2-Country.mmdb" {
+ return errors.New("country unavailable")
+ }
+ return nil
+ },
+ }
+
+ err := updater.updateDatabases(context.Background())
+ if err == nil || !slices.Equal(paths, []string{"/data/GeoLite2-Country.mmdb", "/data/GeoLite2-City.mmdb"}) {
+ t.Fatalf("expected independent Country then City attempts, paths=%#v err=%v", paths, err)
+ }
+}
diff --git a/internal/apps/agent/nginx/manager.go b/internal/apps/agent/nginx/manager.go
index 52ae040f..4560b57e 100644
--- a/internal/apps/agent/nginx/manager.go
+++ b/internal/apps/agent/nginx/manager.go
@@ -18,9 +18,12 @@ import (
"path/filepath"
"regexp"
"sort"
+ "strconv"
"strings"
+ "sync"
"time"
+ sharedprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"github.com/Rain-kl/Wavelet/pkg/utils"
@@ -31,11 +34,23 @@ import (
// RuntimeConfigDirPlaceholder is substituted into generated configs at apply time.
const RuntimeConfigDirPlaceholder = "__OPENFLARE_RUNTIME_CONFIG_DIR__"
+// CountryMMDBPathPlaceholder is substituted with the configured Country database path.
+const CountryMMDBPathPlaceholder = "__OPENFLARE_COUNTRY_MMDB_PATH__"
+
+// CityMMDBPathPlaceholder is substituted with the configured City database path.
+const CityMMDBPathPlaceholder = "__OPENFLARE_CITY_MMDB_PATH__"
+
+// WAFIPGroupsMaxSnapshotBytesPlaceholder is replaced with the shared protocol aggregate snapshot limit.
+const WAFIPGroupsMaxSnapshotBytesPlaceholder = "__OPENFLARE_WAF_IP_GROUPS_MAX_SNAPSHOT_BYTES__"
+
// ResolverDirectivePlaceholder is substituted into generated configs at apply time.
const ResolverDirectivePlaceholder = "__OPENFLARE_RESOLVER_DIRECTIVE__"
// WAFIPGroupsConfigFileName is the runtime filename for synced WAF IP group data.
const WAFIPGroupsConfigFileName = "waf_ip_groups.json"
+
+// WAFIPGroupsChecksumFileName is atomically published after the IP group JSON snapshot.
+const WAFIPGroupsChecksumFileName = WAFIPGroupsConfigFileName + ".checksum"
const powConfigFileName = "pow_config.json"
const (
@@ -45,6 +60,7 @@ const (
stubStatusCheckTimeout = 1500 * time.Millisecond
nginxVersionSubmatchCount = 2
resolverAddressCapacity = 2
+ workerInitSubmatchCount = 2
)
// Executor controls OpenResty validation, reload, health, and lifecycle operations.
@@ -169,11 +185,15 @@ type Manager struct {
LuaDir string
NginxLuaDir string
RuntimeConfigDir string
+ MMDBPath string
+ CityMMDBPath string
PagesDir string
OpenrestyObservabilityListen string
OpenrestyObservabilityPort int
OpenrestyResolverDirective string
Executor Executor
+ atomicFileWriter func(path string, data []byte, perm os.FileMode) error
+ wafIPGroupsMu sync.Mutex
}
// ApplyStatus reports the outcome of an OpenResty configuration apply.
@@ -522,12 +542,21 @@ func (m *Manager) CurrentChecksum() (string, error) {
// WAFIPGroupChecksums returns checksums for locally synced WAF IP groups.
func (m *Manager) WAFIPGroupChecksums() (map[string]string, error) {
+ m.wafIPGroupsMu.Lock()
+ defer m.wafIPGroupsMu.Unlock()
+
config, err := m.readWAFIPGroupsRuntimeConfig()
if err != nil {
return nil, err
}
+ if err = m.ensureWAFIPGroupsChecksum(); err != nil {
+ return nil, err
+ }
result := make(map[string]string, len(config.Groups))
for id, group := range config.Groups {
+ if id != strconv.FormatUint(uint64(group.ID), 10) {
+ continue
+ }
if strings.TrimSpace(group.Checksum) != "" {
result[id] = strings.TrimSpace(group.Checksum)
}
@@ -535,39 +564,164 @@ func (m *Manager) WAFIPGroupChecksums() (map[string]string, error) {
return result, nil
}
-// SyncWAFIPGroups writes WAF IP group definitions to the runtime config directory.
-func (m *Manager) SyncWAFIPGroups(groups []protocol.WAFIPGroup) error {
- if m.RuntimeConfigDir == "" || len(groups) == 0 {
+// ReconcileWAFIPGroups atomically replaces the runtime snapshot with exactly the
+// authoritative target IDs, retaining local definitions that did not change.
+func (m *Manager) ReconcileWAFIPGroups(targetIDs []uint, changed []protocol.WAFIPGroup) error {
+ if m.RuntimeConfigDir == "" {
return nil
}
+ m.wafIPGroupsMu.Lock()
+ defer m.wafIPGroupsMu.Unlock()
+
config, err := m.readWAFIPGroupsRuntimeConfig()
if err != nil {
return err
}
- if config.Groups == nil {
- config.Groups = make(map[string]protocol.WAFIPGroup)
- }
- for _, group := range groups {
- if group.ID == 0 {
+ target := make(map[string]protocol.WAFIPGroup, len(targetIDs))
+ targetSet := make(map[uint]struct{}, len(targetIDs))
+ for _, id := range targetIDs {
+ if id == 0 {
continue
}
- config.Groups[fmt.Sprintf("%d", group.ID)] = group
+ targetSet[id] = struct{}{}
+ key := strconv.FormatUint(uint64(id), 10)
+ if group, ok := config.Groups[key]; ok && group.ID == id {
+ target[key] = group
+ }
}
- data, err := json.Marshal(config)
+ for _, group := range changed {
+ if _, ok := targetSet[group.ID]; !ok {
+ continue
+ }
+ target[strconv.FormatUint(uint64(group.ID), 10)] = group
+ }
+ for _, id := range targetIDs {
+ if id == 0 {
+ continue
+ }
+ if _, ok := target[strconv.FormatUint(uint64(id), 10)]; !ok {
+ return fmt.Errorf("missing referenced WAF IP group %d after synchronization", id)
+ }
+ }
+ return m.publishWAFIPGroups(target)
+}
+
+// UpdateExistingWAFIPGroups applies real-time changes only to groups already
+// present in the authoritative local snapshot.
+func (m *Manager) UpdateExistingWAFIPGroups(changed []protocol.WAFIPGroup) error {
+ if m.RuntimeConfigDir == "" || len(changed) == 0 {
+ return nil
+ }
+ m.wafIPGroupsMu.Lock()
+ defer m.wafIPGroupsMu.Unlock()
+
+ config, err := m.readWAFIPGroupsRuntimeConfig()
if err != nil {
return err
}
+ updated := false
+ for _, group := range changed {
+ key := strconv.FormatUint(uint64(group.ID), 10)
+ if _, ok := config.Groups[key]; !ok || group.ID == 0 {
+ continue
+ }
+ config.Groups[key] = group
+ updated = true
+ }
+ if !updated {
+ return nil
+ }
+ return m.publishWAFIPGroups(config.Groups)
+}
+
+func (m *Manager) publishWAFIPGroups(groups map[string]protocol.WAFIPGroup) error {
+ data, err := sharedprotocol.MarshalWAFIPGroupSnapshot(groups)
+ if err != nil {
+ return err
+ }
+ if len(data) > sharedprotocol.MaxWAFIPGroupSnapshotBytes {
+ return fmt.Errorf("WAF IP group snapshot size %d exceeds maximum %d bytes", len(data), sharedprotocol.MaxWAFIPGroupSnapshotBytes)
+ }
if err := os.MkdirAll(m.RuntimeConfigDir, nginxDirPerm); err != nil {
return err
}
path := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsConfigFileName)
- if err := os.WriteFile(path, data, nginxConfigFilePerm); err != nil {
+ if err := m.writeAtomicFile(path, data, nginxConfigFilePerm); err != nil {
return fmt.Errorf("write %s: %w", WAFIPGroupsConfigFileName, err)
}
+ checksumPath := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsChecksumFileName)
+ if err := m.writeAtomicFile(checksumPath, []byte(checksum(string(data))+"\n"), nginxConfigFilePerm); err != nil {
+ return fmt.Errorf("write %s: %w", WAFIPGroupsChecksumFileName, err)
+ }
slog.Info("synced waf ip groups", "path", path, "group_count", len(groups))
return nil
}
+func (m *Manager) ensureWAFIPGroupsChecksum() error {
+ if m.RuntimeConfigDir == "" {
+ return nil
+ }
+ jsonPath := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsConfigFileName)
+ data, err := os.ReadFile(jsonPath) //nolint:gosec // path is under managed RuntimeConfigDir
+ if err != nil {
+ if os.IsNotExist(err) {
+ return nil
+ }
+ return err
+ }
+ expected := checksum(string(data))
+ checksumPath := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsChecksumFileName)
+ current, err := os.ReadFile(checksumPath) //nolint:gosec // path is under managed RuntimeConfigDir
+ if err == nil && strings.TrimSpace(string(current)) == expected {
+ return nil
+ }
+ if err != nil && !os.IsNotExist(err) {
+ return err
+ }
+ return m.writeAtomicFile(checksumPath, []byte(expected+"\n"), nginxConfigFilePerm)
+}
+
+func (m *Manager) writeAtomicFile(path string, data []byte, perm os.FileMode) error {
+ if m.atomicFileWriter != nil {
+ return m.atomicFileWriter(path, data, perm)
+ }
+ return writeAtomicFile(path, data, perm)
+}
+
+func writeAtomicFile(path string, data []byte, perm os.FileMode) (resultErr error) {
+ tempFile, err := os.CreateTemp(filepath.Dir(path), "."+filepath.Base(path)+".tmp-*")
+ if err != nil {
+ return err
+ }
+ tempPath := tempFile.Name()
+ closed := false
+ defer func() {
+ if !closed {
+ if closeErr := tempFile.Close(); resultErr == nil && closeErr != nil {
+ resultErr = closeErr
+ }
+ }
+ _ = os.Remove(tempPath)
+ }()
+ if err = tempFile.Chmod(perm); err != nil {
+ return err
+ }
+ if _, err = tempFile.Write(data); err != nil {
+ return err
+ }
+ if err = tempFile.Sync(); err != nil {
+ return err
+ }
+ if err = tempFile.Close(); err != nil {
+ return err
+ }
+ closed = true
+ if err = os.Rename(tempPath, path); err != nil {
+ return err
+ }
+ return nil
+}
+
func (m *Manager) readWAFIPGroupsRuntimeConfig() (*wafIPGroupsRuntimeConfig, error) {
config := &wafIPGroupsRuntimeConfig{Groups: map[string]protocol.WAFIPGroup{}}
if m.RuntimeConfigDir == "" {
@@ -1276,6 +1430,9 @@ func (m *Manager) renderMainConfig(content string) string {
}
if luaDir := m.luaRuntimePath(); luaDir != "" {
rendered = strings.ReplaceAll(rendered, openrestyrender.LuaDirPlaceholder, luaDir)
+ if strings.Contains(rendered, "lua_shared_dict openflare_waf_config") {
+ rendered = injectWAFWorkerInit(rendered, luaDir)
+ }
}
if listen := strings.TrimSpace(m.OpenrestyObservabilityListen); listen != "" {
rendered = strings.ReplaceAll(rendered, openrestyrender.ObservabilityListenPlaceholder, listen)
@@ -1289,6 +1446,19 @@ func (m *Manager) renderMainConfig(content string) string {
return rendered
}
+func injectWAFWorkerInit(content string, luaDir string) string {
+ packagePath := fmt.Sprintf(" lua_package_path \"%s/?.lua;%s/?/init.lua;;\";\n", luaDir, luaDir)
+ if !strings.Contains(content, "lua_package_path ") {
+ content = strings.Replace(content, "http {", "http {\n"+packagePath, 1)
+ }
+ initPattern := regexp.MustCompile(`(?m)^[ \t]*init_worker_by_lua_file[ \t]+([^;]+);[ \t]*$`)
+ if match := initPattern.FindStringSubmatch(content); len(match) == workerInitSubmatchCount {
+ block := fmt.Sprintf(" init_worker_by_lua_block {\n require(\"waf.runtime\").init()\n dofile(%q)\n }", strings.TrimSpace(match[1]))
+ return initPattern.ReplaceAllString(content, block)
+ }
+ return strings.Replace(content, "http {", "http {\n init_worker_by_lua_block { require(\"waf.runtime\").init() }", 1)
+}
+
func (m *Manager) managedPowLuaFiles() []protocol.SupportFile {
files := ManagedPowLuaFiles()
runtimeConfigDir := filepath.ToSlash(strings.TrimSpace(m.RuntimeConfigDir))
@@ -1301,8 +1471,19 @@ func (m *Manager) managedPowLuaFiles() []protocol.SupportFile {
func (m *Manager) managedWAFLuaFiles() []protocol.SupportFile {
files := ManagedWAFLuaFiles()
runtimeConfigDir := filepath.ToSlash(strings.TrimSpace(m.RuntimeConfigDir))
+ countryMMDBPath := filepath.ToSlash(strings.TrimSpace(m.MMDBPath))
+ if countryMMDBPath == "" {
+ countryMMDBPath = filepath.ToSlash(filepath.Join(runtimeConfigDir, "GeoLite2-Country.mmdb"))
+ }
+ cityMMDBPath := filepath.ToSlash(strings.TrimSpace(m.CityMMDBPath))
+ if cityMMDBPath == "" {
+ cityMMDBPath = filepath.ToSlash(filepath.Join(runtimeConfigDir, "GeoLite2-City.mmdb"))
+ }
for index := range files {
files[index].Content = strings.ReplaceAll(files[index].Content, RuntimeConfigDirPlaceholder, runtimeConfigDir)
+ files[index].Content = strings.ReplaceAll(files[index].Content, CountryMMDBPathPlaceholder, countryMMDBPath)
+ files[index].Content = strings.ReplaceAll(files[index].Content, CityMMDBPathPlaceholder, cityMMDBPath)
+ files[index].Content = strings.ReplaceAll(files[index].Content, WAFIPGroupsMaxSnapshotBytesPlaceholder, strconv.Itoa(sharedprotocol.MaxWAFIPGroupSnapshotBytes))
}
return files
}
diff --git a/internal/apps/agent/nginx/manager_test.go b/internal/apps/agent/nginx/manager_test.go
index 06c805b7..ac5ece51 100644
--- a/internal/apps/agent/nginx/manager_test.go
+++ b/internal/apps/agent/nginx/manager_test.go
@@ -3,6 +3,7 @@ package nginx
import (
"context"
"errors"
+ "fmt"
"net"
"net/http"
"os"
@@ -13,6 +14,7 @@ import (
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
+ sharedprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
)
type runCall struct {
@@ -330,6 +332,29 @@ func TestManagerApplyWritesSupportFilesAndReplacesPlaceholder(t *testing.T) {
}
}
+func TestManagerRenderMainConfigInitializesWAFRuntimeInWorker(t *testing.T) {
+ manager := &Manager{NginxLuaDir: "/etc/nginx/openflare-lua"}
+ rendered := manager.renderMainConfig("events {}\nhttp {\n lua_shared_dict openflare_waf_config 1m;\n server {}\n}\n")
+ want := `init_worker_by_lua_block { require("waf.runtime").init() }`
+ if !strings.Contains(rendered, want) {
+ t.Fatalf("expected worker-time WAF initialization %q, got:\n%s", want, rendered)
+ }
+}
+
+func TestManagerRenderMainConfigMergesExistingWorkerInitializer(t *testing.T) {
+ manager := &Manager{NginxLuaDir: "/etc/nginx/openflare-lua"}
+ rendered := manager.renderMainConfig("http {\n lua_shared_dict openflare_waf_config 1m;\n init_worker_by_lua_file /etc/nginx/openflare-lua/observability/init.lua;\n}\n")
+ if strings.Count(rendered, "init_worker_by_lua_") != 1 {
+ t.Fatalf("expected one merged worker initializer, got:\n%s", rendered)
+ }
+ if !strings.Contains(rendered, `require("waf.runtime").init()`) {
+ t.Fatalf("expected WAF initialization in merged block, got:\n%s", rendered)
+ }
+ if !strings.Contains(rendered, `dofile("/etc/nginx/openflare-lua/observability/init.lua")`) {
+ t.Fatalf("expected existing worker initializer to be preserved, got:\n%s", rendered)
+ }
+}
+
func TestManagerCheckHealthUsesStubStatusInsteadOfConfigTest(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
@@ -545,12 +570,51 @@ func TestManagerEnsureLuaAssetsWritesReadableFiles(t *testing.T) {
if _, err := os.Stat(filepath.Join(manager.LuaDir, "pow", "check.lua")); err != nil {
t.Fatalf("failed to stat pow lua file: %v", err)
}
- data, err := os.ReadFile(filepath.Join(manager.LuaDir, "pow", "runtime.lua"))
+ data, err := os.ReadFile(filepath.Join(manager.LuaDir, "waf", "runtime.lua"))
if err != nil {
t.Fatalf("failed to read pow lua file: %v", err)
}
- if !strings.Contains(string(data), filepath.ToSlash(manager.RuntimeConfigDir)+"/waf_config.json") {
- t.Fatalf("expected pow lua to read runtime config dir, got %s", string(data))
+ if !strings.Contains(string(data), filepath.ToSlash(manager.RuntimeConfigDir)) || !strings.Contains(string(data), `runtime_dir .. "/waf_config.json"`) {
+ t.Fatalf("expected WAF runtime to load its worker snapshot from the runtime config dir, got %s", string(data))
+ }
+ ipGroupsData, err := os.ReadFile(filepath.Join(manager.LuaDir, "waf", "ip_groups.lua"))
+ if err != nil {
+ t.Fatalf("failed to read WAF IP group refresh module: %v", err)
+ }
+ if !strings.Contains(string(ipGroupsData), filepath.ToSlash(manager.RuntimeConfigDir)) || !strings.Contains(string(ipGroupsData), "waf_ip_groups.json.checksum") {
+ t.Fatalf("expected IP group module to use the managed runtime checksum path, got %s", string(ipGroupsData))
+ }
+ if !strings.Contains(string(ipGroupsData), "openflare_waf_ip_groups") ||
+ !strings.Contains(string(ipGroupsData), fmt.Sprintf("max_snapshot_bytes = options.max_snapshot_bytes or tonumber(\"%d\")", sharedprotocol.MaxWAFIPGroupSnapshotBytes)) {
+ t.Fatalf("expected dedicated dictionary and shared protocol size limit in deployed IP group module, got %s", string(ipGroupsData))
+ }
+ powData, err := os.ReadFile(filepath.Join(manager.LuaDir, "pow", "runtime.lua"))
+ if err != nil {
+ t.Fatalf("failed to read pow lua file: %v", err)
+ }
+ if strings.Contains(string(powData), "io.open") {
+ t.Fatalf("expected PoW node evaluation not to read configuration files, got %s", string(powData))
+ }
+}
+
+func TestManagerEnsureLuaAssetsUsesConfiguredGeoIPDatabasePaths(t *testing.T) {
+ tempDir := t.TempDir()
+ manager := &Manager{
+ LuaDir: filepath.Join(tempDir, "lua"),
+ MMDBPath: "/custom/GeoLite2-Country.mmdb",
+ CityMMDBPath: "/custom/GeoLite2-City.mmdb",
+ }
+ if err := manager.EnsureLuaAssets(); err != nil {
+ t.Fatalf("EnsureLuaAssets failed: %v", err)
+ }
+ data, err := os.ReadFile(filepath.Join(manager.LuaDir, "waf", "runtime.lua"))
+ if err != nil {
+ t.Fatalf("read WAF runtime: %v", err)
+ }
+ for _, path := range []string{manager.MMDBPath, manager.CityMMDBPath} {
+ if !strings.Contains(string(data), path) {
+ t.Fatalf("expected configured GeoIP path %q in WAF runtime", path)
+ }
}
}
@@ -667,7 +731,7 @@ func TestManagerCurrentChecksumIncludesPowConfig(t *testing.T) {
}
func TestManagedPowLuaFilesUseInternalChallengeFlow(t *testing.T) {
- if !strings.Contains(openRestyPowRuntimeLua, `return ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")`) {
+ if !strings.Contains(openRestyPowRuntimeLua, `ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")`) {
t.Fatal("expected pow runtime lua to internally execute make-challenge instead of issuing a 302 redirect")
}
if strings.Contains(openRestyPowRuntimeLua, "ngx.redirect(") {
@@ -699,12 +763,28 @@ func TestManagedPowLuaFilesUseInternalChallengeFlow(t *testing.T) {
}
}
-func TestManagedWAFLuaTreatsWhitelistAsBypass(t *testing.T) {
- if !strings.Contains(openRestyWAFRuntimeLua, "if ip_matches(group.ip_whitelist, ip)") {
- t.Fatal("expected waf runtime to bypass request when ip matches whitelist")
+func TestManagedPowLuaFilesPreserveConfigAcrossInternalRedirect(t *testing.T) {
+ for _, expected := range []string{
+ `pow_config_dict:set(config_key, cjson.encode(config)`,
+ `openflare_pow_config_key = config_key`,
+ `return false`,
+ } {
+ if !strings.Contains(openRestyPowRuntimeLua, expected) {
+ t.Fatalf("expected PoW runtime to contain %q", expected)
+ }
}
- if strings.Contains(openRestyWAFRuntimeLua, "first_allowlist_group") {
- t.Fatal("expected waf runtime not to block requests that miss configured whitelists")
+ if !strings.Contains(openRestyPowChallengeLua, `pow_config_dict:get(config_key)`) {
+ t.Fatal("expected challenge handler to restore reached PoW node config after internal redirect")
+ }
+}
+
+func TestManagedWAFLuaExecutesCompiledGraphWithoutRequestIO(t *testing.T) {
+ if !strings.Contains(openRestyWAFRuntimeLua, `node.type == "ip_match"`) {
+ t.Fatal("expected WAF runtime to execute compiled IP match nodes")
+ }
+ checkStart := strings.Index(openRestyWAFRuntimeLua, "function _M.check()")
+ if checkStart < 0 || strings.Contains(openRestyWAFRuntimeLua[checkStart:], "io.open") {
+ t.Fatal("expected WAF request path not to perform file I/O")
}
}
@@ -924,18 +1004,18 @@ func TestManagerApplyRejectsCertFilePathTraversal(t *testing.T) {
}
}
-func TestManagerSyncWAFIPGroupsWritesDeltaRuntimeFile(t *testing.T) {
+func TestManagerReconcileWAFIPGroupsRetainsUnchangedDeltaRuntimeFile(t *testing.T) {
manager := &Manager{RuntimeConfigDir: t.TempDir()}
- if err := manager.SyncWAFIPGroups([]protocol.WAFIPGroup{
+ if err := manager.ReconcileWAFIPGroups([]uint{1}, []protocol.WAFIPGroup{
{ID: 1, Enabled: true, IPList: []string{"203.0.113.10"}, Checksum: "sum-1"},
}); err != nil {
- t.Fatalf("SyncWAFIPGroups failed: %v", err)
+ t.Fatalf("ReconcileWAFIPGroups failed: %v", err)
}
- if err := manager.SyncWAFIPGroups([]protocol.WAFIPGroup{
+ if err := manager.ReconcileWAFIPGroups([]uint{1, 2}, []protocol.WAFIPGroup{
{ID: 2, Enabled: true, IPList: []string{"198.51.100.10"}, Checksum: "sum-2"},
}); err != nil {
- t.Fatalf("SyncWAFIPGroups second delta failed: %v", err)
+ t.Fatalf("ReconcileWAFIPGroups second delta failed: %v", err)
}
checksums, err := manager.WAFIPGroupChecksums()
@@ -955,6 +1035,252 @@ func TestManagerSyncWAFIPGroupsWritesDeltaRuntimeFile(t *testing.T) {
}
}
+func TestManagerReconcileWAFIPGroupsConvergesToAuthoritativeTarget(t *testing.T) {
+ runtimeDir := t.TempDir()
+ manager := &Manager{RuntimeConfigDir: runtimeDir}
+ initial, err := sharedprotocol.MarshalWAFIPGroupSnapshot(map[string]protocol.WAFIPGroup{
+ "1": {ID: 1, Name: "unchanged", Enabled: true, Checksum: "sum-1"},
+ "2": {ID: 2, Name: "old", Enabled: true, Checksum: "old-2"},
+ "99": {ID: 99, Name: strings.Repeat("x", sharedprotocol.MaxWAFIPGroupSnapshotBytes-1024), Enabled: true, Checksum: "stale"},
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err = os.WriteFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName), initial, 0o644); err != nil {
+ t.Fatal(err)
+ }
+
+ if err = manager.ReconcileWAFIPGroups([]uint{1, 2}, []protocol.WAFIPGroup{{
+ ID: 2, Name: "changed", Enabled: true, Checksum: "sum-2",
+ }}); err != nil {
+ t.Fatalf("ReconcileWAFIPGroups failed after pruning oversized stale data: %v", err)
+ }
+ config, err := manager.readWAFIPGroupsRuntimeConfig()
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(config.Groups) != 2 {
+ t.Fatalf("authoritative group count = %d, want 2: %#v", len(config.Groups), config.Groups)
+ }
+ if got := config.Groups["1"].Name; got != "unchanged" {
+ t.Fatalf("unchanged referenced group was not retained: %q", got)
+ }
+ if got := config.Groups["2"].Name; got != "changed" {
+ t.Fatalf("changed referenced group was not merged: %q", got)
+ }
+ if _, exists := config.Groups["99"]; exists {
+ t.Fatal("historical unreferenced group was not pruned")
+ }
+}
+
+func TestManagerUpdateExistingWAFIPGroupsIgnoresUnrelatedBroadcast(t *testing.T) {
+ runtimeDir := t.TempDir()
+ manager := &Manager{RuntimeConfigDir: runtimeDir}
+ if err := manager.ReconcileWAFIPGroups([]uint{1}, []protocol.WAFIPGroup{{ID: 1, Name: "old", Checksum: "old"}}); err != nil {
+ t.Fatal(err)
+ }
+ if err := manager.UpdateExistingWAFIPGroups([]protocol.WAFIPGroup{
+ {ID: 1, Name: "new", Checksum: "new"},
+ {ID: 2, Name: "unrelated", Checksum: "sum-2"},
+ }); err != nil {
+ t.Fatal(err)
+ }
+ config, err := manager.readWAFIPGroupsRuntimeConfig()
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(config.Groups) != 1 || config.Groups["1"].Name != "new" {
+ t.Fatalf("broadcast update escaped existing target: %#v", config.Groups)
+ }
+}
+
+func TestManagerReconcileWAFIPGroupsPublishesRemovalOnlyAndEmptyTargets(t *testing.T) {
+ runtimeDir := t.TempDir()
+ manager := &Manager{RuntimeConfigDir: runtimeDir}
+ if err := manager.ReconcileWAFIPGroups([]uint{1, 2}, []protocol.WAFIPGroup{
+ {ID: 1, Checksum: "sum-1"}, {ID: 2, Checksum: "sum-2"},
+ }); err != nil {
+ t.Fatal(err)
+ }
+ before, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName))
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err = manager.ReconcileWAFIPGroups([]uint{1}, nil); err != nil {
+ t.Fatalf("removal-only reconcile failed: %v", err)
+ }
+ after, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName))
+ if err != nil {
+ t.Fatal(err)
+ }
+ if string(before) == string(after) {
+ t.Fatal("removal-only reconcile did not publish a new checksum")
+ }
+ if err = manager.ReconcileWAFIPGroups(nil, nil); err != nil {
+ t.Fatalf("empty authoritative reconcile failed: %v", err)
+ }
+ config, err := manager.readWAFIPGroupsRuntimeConfig()
+ if err != nil {
+ t.Fatal(err)
+ }
+ if len(config.Groups) != 0 {
+ t.Fatalf("empty authoritative target retained groups: %#v", config.Groups)
+ }
+}
+
+func TestManagerReconcileWAFIPGroupsRejectsMissingReferencedGroup(t *testing.T) {
+ runtimeDir := t.TempDir()
+ data := []byte(`{"groups":{"7":{"id":8,"checksum":"mistaken-match"}}}`)
+ if err := os.WriteFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName), data, 0o644); err != nil {
+ t.Fatal(err)
+ }
+ manager := &Manager{RuntimeConfigDir: runtimeDir}
+ checksums, err := manager.WAFIPGroupChecksums()
+ if err != nil {
+ t.Fatal(err)
+ }
+ if _, ok := checksums["7"]; ok {
+ t.Fatalf("invalid local group was mistakenly reported matched: %#v", checksums)
+ }
+ err = manager.ReconcileWAFIPGroups([]uint{7}, nil)
+ if err == nil || !strings.Contains(err.Error(), "missing referenced WAF IP group 7") {
+ t.Fatalf("expected clear missing referenced group error, got %v", err)
+ }
+}
+
+func TestWAFIPGroupChecksumPublishesJSONBeforeSidecar(t *testing.T) {
+ runtimeDir := t.TempDir()
+ var writes []string
+ manager := &Manager{
+ RuntimeConfigDir: runtimeDir,
+ atomicFileWriter: func(path string, data []byte, perm os.FileMode) error {
+ writes = append(writes, filepath.Base(path))
+ return os.WriteFile(path, data, perm)
+ },
+ }
+
+ if err := manager.ReconcileWAFIPGroups([]uint{1}, []protocol.WAFIPGroup{{
+ ID: 1, Enabled: true, IPList: []string{"203.0.113.10"}, Checksum: "sum-1",
+ }}); err != nil {
+ t.Fatalf("SyncWAFIPGroups failed: %v", err)
+ }
+ if got, want := strings.Join(writes, ","), WAFIPGroupsConfigFileName+","+WAFIPGroupsChecksumFileName; got != want {
+ t.Fatalf("expected JSON then checksum publication, got %s", got)
+ }
+ jsonData, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName))
+ if err != nil {
+ t.Fatalf("read JSON snapshot: %v", err)
+ }
+ checksumData, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName))
+ if err != nil {
+ t.Fatalf("read checksum sidecar: %v", err)
+ }
+ if got, want := strings.TrimSpace(string(checksumData)), checksum(string(jsonData)); got != want {
+ t.Fatalf("checksum mismatch: got %q want %q", got, want)
+ }
+}
+
+func TestWAFIPGroupChecksumJSONFailurePreservesOldSidecar(t *testing.T) {
+ runtimeDir := t.TempDir()
+ checksumPath := filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName)
+ if err := os.WriteFile(checksumPath, []byte("old-checksum\n"), 0o644); err != nil {
+ t.Fatal(err)
+ }
+ manager := &Manager{
+ RuntimeConfigDir: runtimeDir,
+ atomicFileWriter: func(path string, _ []byte, _ os.FileMode) error {
+ if filepath.Base(path) == WAFIPGroupsConfigFileName {
+ return errors.New("json rename failed")
+ }
+ return errors.New("checksum must not be written")
+ },
+ }
+
+ if err := manager.ReconcileWAFIPGroups([]uint{1}, []protocol.WAFIPGroup{{ID: 1, Checksum: "sum-1"}}); err == nil {
+ t.Fatal("expected JSON publication failure")
+ }
+ data, err := os.ReadFile(checksumPath)
+ if err != nil || string(data) != "old-checksum\n" {
+ t.Fatalf("old checksum must remain unchanged, data=%q err=%v", data, err)
+ }
+}
+
+func TestWAFIPGroupChecksumBootstrapsLegacySnapshot(t *testing.T) {
+ runtimeDir := t.TempDir()
+ jsonData := []byte(`{"groups":{"7":{"id":7,"enabled":true,"checksum":"sum-7"}}}`)
+ if err := os.WriteFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName), jsonData, 0o644); err != nil {
+ t.Fatal(err)
+ }
+ manager := &Manager{RuntimeConfigDir: runtimeDir}
+
+ checksums, err := manager.WAFIPGroupChecksums()
+ if err != nil {
+ t.Fatalf("WAFIPGroupChecksums failed: %v", err)
+ }
+ if checksums["7"] != "sum-7" {
+ t.Fatalf("unexpected group checksums: %#v", checksums)
+ }
+ sidecar, err := os.ReadFile(filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName))
+ if err != nil {
+ t.Fatalf("legacy checksum sidecar was not created: %v", err)
+ }
+ if got, want := strings.TrimSpace(string(sidecar)), checksum(string(jsonData)); got != want {
+ t.Fatalf("legacy checksum mismatch: got %q want %q", got, want)
+ }
+}
+
+func TestWAFIPGroupChecksumRepairsStaleSidecar(t *testing.T) {
+ runtimeDir := t.TempDir()
+ jsonData := []byte(`{"groups":{"9":{"id":9,"enabled":true,"checksum":"sum-9"}}}`)
+ if err := os.WriteFile(filepath.Join(runtimeDir, WAFIPGroupsConfigFileName), jsonData, 0o644); err != nil {
+ t.Fatal(err)
+ }
+ checksumPath := filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName)
+ if err := os.WriteFile(checksumPath, []byte("stale-checksum\n"), 0o644); err != nil {
+ t.Fatal(err)
+ }
+ manager := &Manager{RuntimeConfigDir: runtimeDir}
+
+ if _, err := manager.WAFIPGroupChecksums(); err != nil {
+ t.Fatalf("WAFIPGroupChecksums failed: %v", err)
+ }
+ sidecar, err := os.ReadFile(checksumPath)
+ if err != nil {
+ t.Fatalf("read repaired checksum: %v", err)
+ }
+ if got, want := strings.TrimSpace(string(sidecar)), checksum(string(jsonData)); got != want {
+ t.Fatalf("stale checksum was not repaired: got %q want %q", got, want)
+ }
+}
+
+func TestWAFIPGroupSnapshotRejectsOversizeBeforePublication(t *testing.T) {
+ runtimeDir := t.TempDir()
+ jsonPath := filepath.Join(runtimeDir, WAFIPGroupsConfigFileName)
+ checksumPath := filepath.Join(runtimeDir, WAFIPGroupsChecksumFileName)
+ oldJSON := []byte(`{"groups":{"1":{"id":1,"enabled":true,"checksum":"old"}}}`)
+ oldChecksum := []byte("old-checksum\n")
+ if err := os.WriteFile(jsonPath, oldJSON, 0o644); err != nil {
+ t.Fatal(err)
+ }
+ if err := os.WriteFile(checksumPath, oldChecksum, 0o644); err != nil {
+ t.Fatal(err)
+ }
+ manager := &Manager{RuntimeConfigDir: runtimeDir}
+
+ err := manager.ReconcileWAFIPGroups([]uint{1, 2}, []protocol.WAFIPGroup{{
+ ID: 2, Name: strings.Repeat("x", sharedprotocol.MaxWAFIPGroupSnapshotBytes), Enabled: true, Checksum: "new",
+ }})
+ if err == nil || !strings.Contains(err.Error(), "exceeds maximum") {
+ t.Fatalf("expected clear oversized snapshot error, got %v", err)
+ }
+ if got, readErr := os.ReadFile(jsonPath); readErr != nil || string(got) != string(oldJSON) {
+ t.Fatalf("oversized snapshot touched committed JSON, got=%q err=%v", got, readErr)
+ }
+ if got, readErr := os.ReadFile(checksumPath); readErr != nil || string(got) != string(oldChecksum) {
+ t.Fatalf("oversized snapshot touched committed checksum, got=%q err=%v", got, readErr)
+ }
+}
+
func TestObservabilityListenAddress(t *testing.T) {
if got := ObservabilityListenAddress(18081); got != "127.0.0.1:18081" {
t.Fatalf("unexpected default observability listen address: %s", got)
diff --git a/internal/apps/agent/nginx/pow_assets.go b/internal/apps/agent/nginx/pow_assets.go
index e2eb137c..bda22b89 100644
--- a/internal/apps/agent/nginx/pow_assets.go
+++ b/internal/apps/agent/nginx/pow_assets.go
@@ -13,7 +13,81 @@ var powStaticFS embed.FS
const openRestyPowRuntimeLua = `local _M = {}
+local source = debug.getinfo(1, "S").source or ""
+if string.sub(source, 1, 1) == "@" then
+ local script_path = string.sub(source, 2)
+ local base_dir = string.match(script_path, "^(.*)/pow/[^/]+%.lua$")
+ if base_dir and base_dir ~= "" and not string.find(package.path, base_dir, 1, true) then
+ package.path = base_dir .. "/?.lua;" .. base_dir .. "/?/init.lua;" .. package.path
+ end
+end
+
+local policy = require "pow.policy"
+local pow_sessions = ngx.shared.openflare_pow_sessions
+local pow_config_dict = ngx.shared.openflare_pow_config
+local cjson = require "cjson.safe"
+
+local function session_cookie(value, ttl)
+ local cookie = "__openflare_pow=" .. value .. "; Path=/; HttpOnly; SameSite=Lax; Max-Age=" .. tostring(ttl)
+ if ngx.var.scheme == "https" then cookie = cookie .. "; Secure" end
+ return cookie
+end
+
+-- evaluate is called by a DAG pow node. true continues along its next edge;
+-- false means the challenge flow has taken ownership of the request.
+function _M.evaluate(config)
+ config = config or {}
+ ngx.ctx.openflare_pow_config = config
+
+ local host = ngx.var.host
+ if not host or host == "" then return true end
+ local session_ttl = config.session_ttl or 600
+ local uri = ngx.var.uri or ""
+ local ua = ngx.var.http_user_agent or ""
+ local remote_ip = ngx.var.remote_addr or ""
+
+ if policy.match_any(remote_ip, ua, uri, config.whitelist or {}) then return true end
+ local blacklist = config.blacklist or {}
+ if policy.has_entries(blacklist) and not policy.match_any(remote_ip, ua, uri, blacklist) then return true end
+
+ local cookie_val = ngx.var["cookie___openflare_pow"]
+ if cookie_val and cookie_val ~= "" then
+ local session_key = host .. ":" .. cookie_val
+ if pow_sessions:get(session_key) then
+ pow_sessions:set(session_key, "1", session_ttl)
+ ngx.header["Set-Cookie"] = session_cookie(cookie_val, session_ttl)
+ return true
+ end
+ end
+
+ local api_prefix = "/.within.website/x/cmd/anubis/api/"
+ local static_prefix = "/.within.website/x/cmd/anubis/static/"
+ if string.sub(uri, 1, #api_prefix) == api_prefix or string.sub(uri, 1, #static_prefix) == static_prefix then
+ return false
+ end
+
+ local config_key = "_request_config:" .. (ngx.var.request_id or ngx.md5(host .. uri .. tostring(ngx.now())))
+ pow_config_dict:set(config_key, cjson.encode(config), config.challenge_ttl or 300)
+ ngx.req.set_uri_args({
+ redir = ngx.var.scheme .. "://" .. host .. uri .. (ngx.var.args and ("?" .. ngx.var.args) or ""),
+ host = host,
+ openflare_pow_config_key = config_key,
+ })
+ ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")
+ return false
+end
+
+-- Compatibility entrypoint for old rendered routes. PoW selection now belongs
+-- exclusively to WAF graph nodes, so this function intentionally does nothing.
function _M.check()
+ return true
+end
+
+return _M
+`
+
+/* Removed legacy request-time configuration scanner. Graph execution now calls
+evaluate(config) with the reached node.
local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
local script_path = string.sub(source, 2)
@@ -198,7 +272,7 @@ return ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")
end
return _M
-`
+*/
const openRestyPowCheckLua = `local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
@@ -214,8 +288,8 @@ return require("pow.runtime").check()
const openRestyPowChallengeLua = `local cjson = require "cjson.safe"
-local pow_config_dict = ngx.shared.openflare_pow_config
local pow_challenges = ngx.shared.openflare_pow_challenges
+local pow_config_dict = ngx.shared.openflare_pow_config
local function generate_entropy()
local pieces = {
@@ -233,28 +307,20 @@ local args = ngx.req.get_uri_args()
local host = args["host"] or ngx.var.host or ""
local redir = args["redir"] or ""
-local site = ngx.var.openflare_waf_site or ""
-if site == "" then
+local config = ngx.ctx.openflare_pow_config
+local config_key = args["openflare_pow_config_key"] or ""
+if type(config) ~= "table" and config_key ~= "" then
+ local config_raw = pow_config_dict:get(config_key)
+ if config_raw then
+ config = cjson.decode(config_raw)
+ end
+end
+if config_key ~= "" then pow_config_dict:delete(config_key) end
+if type(config) ~= "table" then
ngx.status = 403
- ngx.say("PoW site not resolved; openflare_waf_site is required")
+ ngx.say("PoW graph node was not evaluated for this request")
return
end
-
-local config_raw = pow_config_dict:get(site)
-if not config_raw then
- ngx.status = 403
- ngx.say("PoW not configured for this site")
- return
-end
-
-local ok, route_config = pcall(cjson.decode, config_raw)
-if not ok or not route_config or not route_config.enabled then
- ngx.status = 403
- ngx.say("PoW not enabled for this site")
- return
-end
-
-local config = route_config.config or {}
local difficulty = config.difficulty or 4
local algorithm = config.algorithm or "fast"
local challenge_ttl = config.challenge_ttl or 300
diff --git a/internal/apps/agent/nginx/waf_assets.go b/internal/apps/agent/nginx/waf_assets.go
index b8d7e8fe..2149a9db 100644
--- a/internal/apps/agent/nginx/waf_assets.go
+++ b/internal/apps/agent/nginx/waf_assets.go
@@ -1,278 +1,16 @@
package nginx
-import "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
+import (
+ _ "embed"
-const openRestyWAFRuntimeLua = `local _M = {}
+ "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
+)
-function _M.check()
-local cjson = require "cjson.safe"
+//go:embed waf_runtime.lua
+var openRestyWAFRuntimeLua string
-local config_dict = ngx.shared.openflare_waf_config
-
-local function read_file(path)
- local f = io.open(path, "r")
- if not f then
- return nil
- end
- local content = f:read("*a")
- f:close()
- return content
-end
-
-local function load_config()
- local paths = {
- "__OPENFLARE_RUNTIME_CONFIG_DIR__/waf_config.json",
- "/etc/nginx/openflare-lua/waf_config.json",
- "/usr/local/openresty/nginx/conf/waf_config.json"
- }
- for _, path in ipairs(paths) do
- local content = read_file(path)
- if content and content ~= "" then
- local hash = ngx.md5(content)
- if config_dict:get("_config_hash") == hash then
- local cached = config_dict:get("_config_json")
- if cached then
- local decoded = cjson.decode(cached)
- if decoded then
- return decoded
- end
- end
- end
- local decoded = cjson.decode(content)
- if decoded then
- config_dict:set("_config_hash", hash, 0)
- config_dict:set("_config_json", content, 0)
- return decoded
- end
- end
- end
- return nil
-end
-
-local function load_ip_groups()
- local paths = {
- "__OPENFLARE_RUNTIME_CONFIG_DIR__/waf_ip_groups.json",
- "/etc/nginx/openflare-lua/waf_ip_groups.json",
- "/usr/local/openresty/nginx/conf/waf_ip_groups.json"
- }
- for _, path in ipairs(paths) do
- local content = read_file(path)
- if content and content ~= "" then
- local hash = ngx.md5(content)
- if config_dict:get("_ip_groups_hash") == hash then
- local cached = config_dict:get("_ip_groups_json")
- if cached then
- local decoded = cjson.decode(cached)
- if decoded then
- return decoded
- end
- end
- end
- local decoded = cjson.decode(content)
- if decoded then
- config_dict:set("_ip_groups_hash", hash, 0)
- config_dict:set("_ip_groups_json", content, 0)
- return decoded
- end
- end
- end
- return { groups = {} }
-end
-
-local function list_contains(items, value)
- if not items or type(items) ~= "table" or not value or value == "" then
- return false
- end
- for _, item in ipairs(items) do
- if item == value then
- return true
- end
- end
- return false
-end
-
-local function table_has_items(items)
- return type(items) == "table" and #items > 0
-end
-
-local function parse_ipv4(value)
- local a, b, c, d = string.match(value or "", "^(%d+)%.(%d+)%.(%d+)%.(%d+)$")
- if not a then
- return nil
- end
- a, b, c, d = tonumber(a), tonumber(b), tonumber(c), tonumber(d)
- if a > 255 or b > 255 or c > 255 or d > 255 then
- return nil
- end
- return ((a * 256 + b) * 256 + c) * 256 + d
-end
-
-local function ipv4_in_cidr(ip, cidr)
- local base, bits = string.match(cidr or "", "^([^/]+)/(%d+)$")
- if not base then
- return false
- end
- bits = tonumber(bits)
- if not bits or bits < 0 or bits > 32 then
- return false
- end
- local ip_num = parse_ipv4(ip)
- local base_num = parse_ipv4(base)
- if not ip_num or not base_num then
- return false
- end
- if bits == 0 then
- return true
- end
- local mask = 4294967295 - (2 ^ (32 - bits) - 1)
- return (ip_num - (ip_num % (2 ^ (32 - bits)))) == (base_num - (base_num % (2 ^ (32 - bits))))
-end
-
-local function ip_matches(items, ip)
- if not items or type(items) ~= "table" or not ip or ip == "" then
- return false
- end
- for _, item in ipairs(items) do
- if item == ip then
- return true
- end
- if string.find(item, "/", 1, true) and ipv4_in_cidr(ip, item) then
- return true
- end
- end
- return false
-end
-
-local function ip_matches_group_ids(group_ids, ip, ip_groups_config)
- if not group_ids or type(group_ids) ~= "table" or not ip or ip == "" then
- return false
- end
- local groups = (ip_groups_config or {}).groups or {}
- for _, id in ipairs(group_ids) do
- local group = groups[tostring(id)]
- if group and group.enabled and ip_matches(group.ip_list, ip) then
- return true
- end
- end
- return false
-end
-
-local function lookup_country(ip)
- local ok, maxminddb = pcall(require, "resty.maxminddb")
- if not ok or not maxminddb then
- return nil
- end
- local paths = {
- "__OPENFLARE_RUNTIME_CONFIG_DIR__/GeoLite2-Country.mmdb",
- "/etc/openflare/GeoLite2-Country.mmdb",
- "/usr/local/share/openflare/GeoLite2-Country.mmdb"
- }
- for _, path in ipairs(paths) do
- local opened = pcall(maxminddb.init, path)
- if opened then
- local res, err = maxminddb.lookup(ip)
- if res and res.country and res.country.iso_code then
- return string.upper(res.country.iso_code)
- end
- end
- end
- return nil
-end
-
-local function group_by_id(config)
- local result = {}
- for _, group in ipairs(config.rule_groups or {}) do
- result[tostring(group.id)] = group
- end
- return result
-end
-
-local function active_groups(config, groups)
- local site = ngx.var.openflare_waf_site or ""
- local ids = (config.site_rule_groups or {})[site]
- local result = {}
- for _, group in ipairs(config.rule_groups or {}) do
- if group.is_global then
- result[#result + 1] = group
- end
- end
- if ids then
- local by_id = group_by_id(config)
- for _, id in ipairs(ids) do
- local group = by_id[tostring(id)]
- if group and not group.is_global then
- result[#result + 1] = group
- end
- end
- end
- return result
-end
-
-local function exit_with_group(group)
- ngx.ctx.openflare_waf_blocked = true
- ngx.status = tonumber(group.block_status_code) or 418
- local body = group.block_response_body or ""
- if body ~= "" then
- ngx.header["Content-Type"] = "text/html; charset=utf-8"
- ngx.say(body)
- end
- return ngx.exit(ngx.status)
-end
-
-local config = load_config()
-if not config then
- if config_dict:add("_missing_config_logged", true, 60) then
- ngx.log(ngx.WARN, "openflare waf config is missing or invalid; requests will be allowed")
- end
- return
-end
-
-local ip = ngx.var.remote_addr or ""
-local groups = active_groups(config)
-local ip_groups_config = load_ip_groups()
-if #groups == 0 then
- if config_dict:add("_empty_groups_logged", true, 60) then
- ngx.log(ngx.WARN, "openflare waf has no active rule group for site: ", ngx.var.openflare_waf_site or "")
- end
- return
-end
-
-for _, group in ipairs(groups) do
- if ip_matches(group.ip_whitelist, ip) or ip_matches_group_ids(group.ip_whitelist_group_ids, ip, ip_groups_config) then
- return
- end
-end
-
-local country = nil
-for _, group in ipairs(groups) do
- if type(group.country_whitelist) == "table" and #group.country_whitelist > 0 then
- country = country or lookup_country(ip)
- if list_contains(group.country_whitelist, country) then
- return
- end
- end
-end
-
-for _, group in ipairs(groups) do
- if ip_matches(group.ip_blacklist, ip) or ip_matches_group_ids(group.ip_blacklist_group_ids, ip, ip_groups_config) then
- return exit_with_group(group)
- end
-end
-
-for _, group in ipairs(groups) do
- if type(group.country_blacklist) == "table" and #group.country_blacklist > 0 then
- country = country or lookup_country(ip)
- if list_contains(group.country_blacklist, country) then
- return exit_with_group(group)
- end
- end
-end
-
-return "ok"
-end
-
-return _M
-`
+//go:embed waf_ip_groups.lua
+var openRestyWAFIPGroupsLua string
const openRestyWAFCheckLua = `local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then
@@ -290,6 +28,7 @@ return require("waf.runtime").check()
func ManagedWAFLuaFiles() []protocol.SupportFile {
return []protocol.SupportFile{
{Path: "waf/runtime.lua", Content: openRestyWAFRuntimeLua},
+ {Path: "waf/ip_groups.lua", Content: openRestyWAFIPGroupsLua},
{Path: "waf/check.lua", Content: openRestyWAFCheckLua},
}
}
diff --git a/internal/apps/agent/nginx/waf_assets_test.go b/internal/apps/agent/nginx/waf_assets_test.go
new file mode 100644
index 00000000..289b119c
--- /dev/null
+++ b/internal/apps/agent/nginx/waf_assets_test.go
@@ -0,0 +1,44 @@
+package nginx
+
+import (
+ "path/filepath"
+ "testing"
+
+ lua "github.com/yuin/gopher-lua"
+)
+
+func TestWAFRuntime(t *testing.T) {
+ state := lua.NewState()
+ defer state.Close()
+
+ runtimePath, err := filepath.Abs("waf_runtime.lua")
+ if err != nil {
+ t.Fatal(err)
+ }
+ specPath, err := filepath.Abs("waf_runtime_spec.lua")
+ if err != nil {
+ t.Fatal(err)
+ }
+ state.SetGlobal("WAF_RUNTIME_PATH", lua.LString(runtimePath))
+ if err := state.DoFile(specPath); err != nil {
+ t.Fatalf("WAF runtime specification failed: %v", err)
+ }
+}
+
+func TestWAFIPGroupRefresh(t *testing.T) {
+ state := lua.NewState()
+ defer state.Close()
+
+ modulePath, err := filepath.Abs("waf_ip_groups.lua")
+ if err != nil {
+ t.Fatal(err)
+ }
+ specPath, err := filepath.Abs("waf_ip_groups_spec.lua")
+ if err != nil {
+ t.Fatal(err)
+ }
+ state.SetGlobal("WAF_IP_GROUPS_PATH", lua.LString(modulePath))
+ if err := state.DoFile(specPath); err != nil {
+ t.Fatalf("WAF IP group refresh specification failed: %v", err)
+ }
+}
diff --git a/internal/apps/agent/nginx/waf_ip_groups.lua b/internal/apps/agent/nginx/waf_ip_groups.lua
new file mode 100644
index 00000000..b1da22f3
--- /dev/null
+++ b/internal/apps/agent/nginx/waf_ip_groups.lua
@@ -0,0 +1,150 @@
+local _M = {}
+
+local current_groups = { groups = {} }
+local current_version
+local initialized = false
+local shared
+local read_checksum
+local read_json
+local decode
+local log_warning
+local max_snapshot_bytes
+
+local refresh_lock_key = "ip_groups_refresh_lock"
+local raw_snapshot_prefix = "ip_groups_raw:"
+local version_key = "ip_groups_version"
+local previous_version_key = "ip_groups_previous_version"
+
+local function warn(message, err, forcible)
+ local suffix = err and (": " .. tostring(err)) or ""
+ if forcible then suffix = suffix .. " (forcible eviction refused)" end
+ pcall(log_warning, "openflare WAF IP group refresh " .. message .. suffix)
+end
+
+local function safe_set(key, value, description)
+ local ok, err, forcible = shared:safe_set(key, value)
+ if ok ~= true or forcible == true then
+ warn(description, err, forcible)
+ return false
+ end
+ return true
+end
+
+local function read_file(path)
+ local file, err = io.open(path, "rb")
+ if not file then return nil, err end
+ local content = file:read("*a")
+ file:close()
+ return content
+end
+
+local function valid_snapshot(snapshot)
+ return type(snapshot) == "table" and type(snapshot.groups) == "table"
+end
+
+local function decode_snapshot(raw)
+ if type(raw) ~= "string" or raw == "" then return nil end
+ local called, snapshot = pcall(decode, raw)
+ if not called or not valid_snapshot(snapshot) then return nil end
+ return snapshot
+end
+
+local function refresh_from_checksum()
+ local called, checksum = pcall(read_checksum)
+ if not called or type(checksum) ~= "string" then return end
+ checksum = string.match(checksum, "^%s*(.-)%s*$")
+ local committed_version = shared:get(version_key)
+ if checksum == "" or checksum == committed_version then return end
+
+ local json_called, raw = pcall(read_json)
+ if not json_called then
+ warn("JSON read failed", raw)
+ return
+ end
+ if type(raw) ~= "string" or #raw > max_snapshot_bytes then
+ warn("snapshot exceeds maximum " .. tostring(max_snapshot_bytes) .. " bytes")
+ return
+ end
+ if not decode_snapshot(raw) then return end
+ local raw_key = raw_snapshot_prefix .. checksum
+ local existing_raw = shared:get(raw_key)
+ local published_new_raw = false
+ if existing_raw == nil then
+ if not safe_set(raw_key, raw, "raw publication failed") then return end
+ published_new_raw = true
+ elseif existing_raw ~= raw then
+ return
+ end
+ if not safe_set(version_key, checksum, "commit pointer publication failed") then
+ if published_new_raw then shared:delete(raw_key) end
+ return
+ end
+
+ local previous_version = shared:get(previous_version_key)
+ if type(committed_version) == "string" and committed_version ~= "" and committed_version ~= checksum then
+ if not safe_set(previous_version_key, committed_version, "previous version metadata publication failed") then return end
+ if type(previous_version) == "string" and previous_version ~= "" and
+ previous_version ~= committed_version and previous_version ~= checksum then
+ shared:delete(raw_snapshot_prefix .. previous_version)
+ end
+ end
+end
+
+local function adopt_shared_snapshot_if_changed()
+ local version = shared:get(version_key)
+ if type(version) ~= "string" or version == "" or version == current_version then return end
+ local snapshot = decode_snapshot(shared:get(raw_snapshot_prefix .. version))
+ if not snapshot then return end
+ current_groups = snapshot
+ current_version = version
+end
+
+local function tick(premature)
+ if premature then return end
+ local locked, lock_error, forcible = shared:safe_add(refresh_lock_key, true, 4)
+ if forcible == true then
+ warn("coordination lock refused forcible eviction", lock_error, true)
+ locked = false
+ elseif not locked and lock_error and lock_error ~= "exists" then
+ warn("coordination lock failed", lock_error)
+ end
+ if locked then refresh_from_checksum() end
+ adopt_shared_snapshot_if_changed()
+end
+
+function _M.init(options)
+ if initialized then return true end
+ options = options or {}
+ local runtime_dir = options.runtime_dir or "__OPENFLARE_RUNTIME_CONFIG_DIR__"
+ shared = options.shared or (ngx.shared and ngx.shared.openflare_waf_ip_groups)
+ assert(shared, "openflare_waf_ip_groups shared dictionary is required")
+ max_snapshot_bytes = options.max_snapshot_bytes or tonumber("__OPENFLARE_WAF_IP_GROUPS_MAX_SNAPSHOT_BYTES__")
+ assert(max_snapshot_bytes and max_snapshot_bytes > 0, "WAF IP group maximum snapshot size is required")
+ log_warning = options.log_warning or function(message)
+ if ngx and ngx.log then ngx.log(ngx.WARN, message) end
+ end
+ read_checksum = options.read_checksum or function()
+ return read_file(runtime_dir .. "/waf_ip_groups.json.checksum")
+ end
+ read_json = options.read_json or function()
+ return read_file(runtime_dir .. "/waf_ip_groups.json")
+ end
+ if options.decode then
+ decode = options.decode
+ else
+ local cjson = require("cjson.safe")
+ decode = cjson.decode
+ end
+ local timer_every = options.timer_every or ngx.timer.every
+ local ok, err = timer_every(5, tick)
+ if not ok then return nil, err end
+ initialized = true
+ tick(false)
+ return true
+end
+
+function _M.current()
+ return current_groups
+end
+
+return _M
diff --git a/internal/apps/agent/nginx/waf_ip_groups_spec.lua b/internal/apps/agent/nginx/waf_ip_groups_spec.lua
new file mode 100644
index 00000000..bb46aeb7
--- /dev/null
+++ b/internal/apps/agent/nginx/waf_ip_groups_spec.lua
@@ -0,0 +1,356 @@
+local module_path = assert(WAF_IP_GROUPS_PATH, "WAF_IP_GROUPS_PATH is required")
+
+local function assert_equal(actual, expected, message)
+ if actual ~= expected then
+ error((message or "values differ") .. ": expected " .. tostring(expected) .. ", got " .. tostring(actual), 2)
+ end
+end
+
+local shared_data = {}
+local locks = {}
+local shared = {}
+function shared:get(key) return shared_data[key] end
+function shared:set(key, value) shared_data[key] = value return true end
+function shared:delete(key) shared_data[key] = nil return true end
+function shared:safe_set(key, value) return shared:set(key, value) end
+function shared:add(key, value, ttl)
+ assert_equal(ttl, 4, "coordination lock TTL")
+ if locks[key] then return false end
+ locks[key] = value
+ return true
+end
+function shared:safe_add(key, value, ttl) return shared:add(key, value, ttl) end
+local function advance_time() locks = {} end
+
+local disk_checksum = "v1"
+local disk_json = "valid-v1"
+local checksum_reads = 0
+local json_reads = 0
+local timer_callbacks = {}
+
+local function decode(raw)
+ if raw == "valid-v1" then
+ return { groups = { ["1"] = { enabled = true, ip_list = { "192.0.2.1" } } } }
+ end
+ if raw == "valid-v2" then
+ return { groups = { ["2"] = { enabled = true, ip_list = { "198.51.100.2" } } } }
+ end
+ if raw == "valid-v3" then
+ return { groups = { ["3"] = { enabled = true, ip_list = { "203.0.113.3" } } } }
+ end
+ return nil, "invalid json"
+end
+
+local function load_worker()
+ local worker = assert(loadfile(module_path))()
+ worker.init({
+ shared = shared,
+ timer_every = function(interval, callback)
+ assert_equal(interval, 5, "refresh interval")
+ timer_callbacks[#timer_callbacks + 1] = callback
+ return true
+ end,
+ read_checksum = function()
+ checksum_reads = checksum_reads + 1
+ return disk_checksum
+ end,
+ read_json = function()
+ json_reads = json_reads + 1
+ return disk_json
+ end,
+ decode = decode,
+ max_snapshot_bytes = 20 * 1024 * 1024,
+ })
+ return worker
+end
+
+local first = load_worker()
+local second = load_worker()
+assert_equal(#timer_callbacks, 2, "each worker schedules a refresh timer")
+assert_equal(checksum_reads, 1, "one worker coordinates initial checksum read")
+assert_equal(json_reads, 1, "one worker reads initial JSON")
+assert_equal(first.current().groups["1"].ip_list[1], "192.0.2.1", "first worker adopts initial snapshot")
+assert_equal(second.current().groups["1"].ip_list[1], "192.0.2.1", "second worker adopts initial snapshot")
+
+local function tick_all()
+ advance_time()
+ for _, callback in ipairs(timer_callbacks) do callback(false) end
+end
+
+checksum_reads = 0
+json_reads = 0
+for _ = 1, 3 do tick_all() end
+assert_equal(checksum_reads, 3, "stable 15 seconds reads checksum once per interval")
+assert_equal(json_reads, 0, "unchanged checksum never reads JSON")
+
+disk_checksum = "v2"
+disk_json = "valid-v2"
+tick_all()
+assert_equal(json_reads, 1, "changed snapshot JSON is read once across workers")
+assert_equal(first.current().groups["2"].ip_list[1], "198.51.100.2", "first worker adopts v2")
+assert_equal(second.current().groups["2"].ip_list[1], "198.51.100.2", "second worker adopts v2")
+
+disk_checksum = "v3"
+disk_json = "valid-v3"
+tick_all()
+assert_equal(shared_data.ip_groups_previous_version, "v2", "previous pointer follows committed version")
+assert_equal(shared_data["ip_groups_raw:v1"], nil, "snapshot older than previous is cleaned")
+assert_equal(shared_data["ip_groups_raw:v2"], "valid-v2", "previous committed raw is retained")
+assert_equal(shared_data["ip_groups_raw:v3"], "valid-v3", "current committed raw is retained")
+
+disk_checksum = "v2"
+disk_json = "valid-v2"
+tick_all()
+assert_equal(shared_data.ip_groups_version, "v2", "rollback checksum becomes current commit")
+assert_equal(shared_data.ip_groups_previous_version, "v3", "rollback retains former current as previous")
+assert_equal(shared_data["ip_groups_raw:v2"], "valid-v2", "rollback must not clean its new current raw")
+assert_equal(shared_data["ip_groups_raw:v3"], "valid-v3", "rollback retains previous raw")
+
+disk_checksum = "v4"
+disk_json = "invalid-v4"
+tick_all()
+assert_equal(shared_data.ip_groups_version, "v2", "invalid update preserves shared version")
+assert_equal(first.current().groups["2"].ip_list[1], "198.51.100.2", "invalid update preserves first worker")
+assert_equal(second.current().groups["2"].ip_list[1], "198.51.100.2", "invalid update preserves second worker")
+
+local reads_before_requests = checksum_reads + json_reads
+for _ = 1, 20 do
+ assert_equal(first.current().groups["2"].enabled, true, "request reads worker-local object")
+end
+assert_equal(checksum_reads + json_reads, reads_before_requests, "current() performs zero file I/O")
+
+timer_callbacks[1](true)
+assert_equal(checksum_reads + json_reads, reads_before_requests, "premature timer performs zero file I/O")
+
+local function test_failed_commit_never_exposes_unpublished_raw_to_new_worker()
+ local data = {}
+ local held_locks = {}
+ local callbacks = {}
+ local checksum = "v1"
+ local raw = "valid-v1"
+ local reads = 0
+ local fail_commit = false
+ local interleaved_worker
+ local load_regression_worker
+ local regression_shared = {}
+
+ function regression_shared:get(key) return data[key] end
+ function regression_shared:add(key, value)
+ if held_locks[key] then return false end
+ held_locks[key] = value
+ return true
+ end
+ function regression_shared:delete(key) data[key] = nil return true end
+ local function set_regression_value(key, value)
+ if key == "ip_groups_version" and fail_commit then
+ return false, "shared dictionary full"
+ end
+ data[key] = value
+ if fail_commit and string.sub(key, 1, #"ip_groups_raw") == "ip_groups_raw" and not interleaved_worker then
+ interleaved_worker = load_regression_worker()
+ end
+ return true
+ end
+ function regression_shared:set(key, value) return set_regression_value(key, value) end
+ function regression_shared:safe_set(key, value) return set_regression_value(key, value) end
+ function regression_shared:safe_add(key, value) return regression_shared:add(key, value) end
+
+ load_regression_worker = function()
+ local worker = assert(loadfile(module_path))()
+ assert(worker.init({
+ shared = regression_shared,
+ timer_every = function(_, callback) callbacks[#callbacks + 1] = callback return true end,
+ read_checksum = function() return checksum end,
+ read_json = function() reads = reads + 1 return raw end,
+ decode = decode,
+ max_snapshot_bytes = 20 * 1024 * 1024,
+ }))
+ return worker
+ end
+
+ local established_worker = load_regression_worker()
+ assert_equal(established_worker.current().groups["1"].ip_list[1], "192.0.2.1", "v1 is committed before failure")
+
+ held_locks = {}
+ reads = 0
+ checksum = "v2"
+ raw = "valid-v2"
+ fail_commit = true
+ callbacks[1](false)
+
+ assert_equal(reads, 1, "failed commit still reads changed JSON only once")
+ assert_equal(data.ip_groups_version, "v1", "failed pointer write preserves committed version")
+ assert_equal(data["ip_groups_raw:v2"], nil, "failed commit cleans only unpublished v2 raw")
+ assert_equal(established_worker.current().groups["1"].ip_list[1], "192.0.2.1", "existing worker preserves committed v1")
+ assert(interleaved_worker, "raw publication must interleave a newly initialized worker")
+ assert_equal(interleaved_worker.current().groups["2"], nil, "new worker must not expose unpublished v2")
+ assert_equal(interleaved_worker.current().groups["1"].ip_list[1], "192.0.2.1", "new worker must never adopt unpublished v2 raw")
+end
+
+test_failed_commit_never_exposes_unpublished_raw_to_new_worker()
+
+local function test_capacity_failure_never_evicts_committed_snapshot()
+ local data = {
+ ip_groups_version = "v1",
+ ip_groups_previous_version = "v0",
+ ["ip_groups_raw:v1"] = "valid-v1",
+ ["ip_groups_raw:v0"] = "valid-v0",
+ }
+ local locks = {}
+ local callbacks = {}
+ local disk_checksum = "v1"
+ local disk_raw = "valid-v1"
+ local json_reads = 0
+ local ordinary_writes = 0
+ local warnings = {}
+ local dict = {}
+ function dict:get(key) return data[key] end
+ function dict:delete(key) data[key] = nil return true end
+ function dict:add(key, value)
+ if locks[key] then return false end
+ locks[key] = value
+ return true
+ end
+ function dict:safe_add(key, value) return dict:add(key, value) end
+ function dict:set(key, value)
+ ordinary_writes = ordinary_writes + 1
+ if key == "ip_groups_raw:v2" then
+ data = { [key] = value }
+ return true, nil, true
+ end
+ data[key] = value
+ return true, nil, false
+ end
+ function dict:safe_set(key, value)
+ if key == "ip_groups_raw:v2" then return nil, "no memory", false end
+ data[key] = value
+ return true, nil, false
+ end
+
+ local worker = assert(loadfile(module_path))()
+ assert(worker.init({
+ shared = dict,
+ timer_every = function(_, callback) callbacks[1] = callback return true end,
+ read_checksum = function() return disk_checksum end,
+ read_json = function() json_reads = json_reads + 1 return disk_raw end,
+ decode = decode,
+ max_snapshot_bytes = 20 * 1024 * 1024,
+ log_warning = function(message) warnings[#warnings + 1] = message end,
+ }))
+ assert_equal(worker.current().groups["1"].ip_list[1], "192.0.2.1", "worker starts from committed v1")
+
+ locks = {}
+ disk_checksum = "v2"
+ disk_raw = "valid-v2"
+ callbacks[1](false)
+
+ assert_equal(ordinary_writes, 0, "snapshot publication must never use evicting set")
+ assert_equal(json_reads, 1, "capacity failure reads changed JSON once")
+ assert_equal(data.ip_groups_version, "v1", "capacity failure preserves commit pointer")
+ assert_equal(data.ip_groups_previous_version, "v0", "capacity failure preserves previous metadata")
+ assert_equal(data["ip_groups_raw:v1"], "valid-v1", "capacity failure preserves current raw")
+ assert_equal(data["ip_groups_raw:v0"], "valid-v0", "capacity failure preserves previous raw")
+ assert_equal(data["ip_groups_raw:v2"], nil, "capacity failure does not publish new raw")
+ assert_equal(worker.current().groups["1"].ip_list[1], "192.0.2.1", "capacity failure preserves worker-local snapshot")
+ assert_equal(#warnings, 1, "capacity failure is logged")
+end
+
+local function test_previous_metadata_failure_keeps_committed_snapshot_without_cleanup()
+ local data = {
+ ip_groups_version = "v1",
+ ip_groups_previous_version = "v0",
+ ["ip_groups_raw:v1"] = "valid-v1",
+ ["ip_groups_raw:v0"] = "valid-v0",
+ }
+ local locks = {}
+ local callback
+ local checksum = "v1"
+ local raw = "valid-v1"
+ local deletes = 0
+ local warnings = {}
+ local dict = {}
+ function dict:get(key) return data[key] end
+ function dict:delete(key) deletes = deletes + 1 data[key] = nil return true end
+ function dict:add(key, value)
+ if locks[key] then return false end
+ locks[key] = value
+ return true
+ end
+ function dict:safe_add(key, value) return dict:add(key, value) end
+ function dict:set(key, value) data[key] = value return true end
+ function dict:safe_set(key, value)
+ if key == "ip_groups_previous_version" then return nil, "no memory", false end
+ data[key] = value
+ return true, nil, false
+ end
+
+ local worker = assert(loadfile(module_path))()
+ assert(worker.init({
+ shared = dict,
+ timer_every = function(_, value) callback = value return true end,
+ read_checksum = function() return checksum end,
+ read_json = function() return raw end,
+ decode = decode,
+ max_snapshot_bytes = 20 * 1024 * 1024,
+ log_warning = function(message) warnings[#warnings + 1] = message end,
+ }))
+
+ locks = {}
+ checksum = "v2"
+ raw = "valid-v2"
+ callback(false)
+
+ assert_equal(data.ip_groups_version, "v2", "successful commit pointer remains authoritative")
+ assert_equal(data.ip_groups_previous_version, "v0", "failed previous metadata write is not forced")
+ assert_equal(data["ip_groups_raw:v2"], "valid-v2", "new committed raw remains")
+ assert_equal(data["ip_groups_raw:v1"], "valid-v1", "old current raw remains when cleanup is skipped")
+ assert_equal(data["ip_groups_raw:v0"], "valid-v0", "old previous raw remains when cleanup is skipped")
+ assert_equal(deletes, 0, "previous metadata failure skips all cleanup")
+ assert_equal(worker.current().groups["2"].ip_list[1], "198.51.100.2", "worker adopts valid committed v2")
+ assert_equal(#warnings, 1, "previous metadata failure is logged")
+end
+
+local function test_oversized_raw_is_rejected_before_shared_publication()
+ local data = { ip_groups_version = "v1", ["ip_groups_raw:v1"] = "valid-v1" }
+ local locks = {}
+ local callback
+ local checksum = "v1"
+ local raw = "valid-v1"
+ local shared_writes = 0
+ local warnings = {}
+ local dict = {}
+ function dict:get(key) return data[key] end
+ function dict:delete(key) data[key] = nil return true end
+ function dict:add(key, value) if locks[key] then return false end locks[key] = value return true end
+ function dict:safe_add(key, value) return dict:add(key, value) end
+ function dict:set(key, value) shared_writes = shared_writes + 1 data[key] = value return true end
+ function dict:safe_set(key, value) shared_writes = shared_writes + 1 data[key] = value return true, nil, false end
+
+ local worker = assert(loadfile(module_path))()
+ assert(worker.init({
+ shared = dict,
+ timer_every = function(_, value) callback = value return true end,
+ read_checksum = function() return checksum end,
+ read_json = function() return raw end,
+ decode = decode,
+ max_snapshot_bytes = 4,
+ log_warning = function(message) warnings[#warnings + 1] = message end,
+ }))
+
+ locks = {}
+ checksum = "v2"
+ raw = "valid-v2"
+ callback(false)
+
+ assert_equal(shared_writes, 0, "oversized raw is rejected before shared writes")
+ assert_equal(data.ip_groups_version, "v1", "oversized raw preserves commit pointer")
+ assert_equal(data["ip_groups_raw:v1"], "valid-v1", "oversized raw preserves committed data")
+ assert_equal(worker.current().groups["1"].ip_list[1], "192.0.2.1", "oversized raw preserves worker-local snapshot")
+ assert_equal(#warnings, 1, "oversized raw rejection is logged")
+end
+
+test_capacity_failure_never_evicts_committed_snapshot()
+test_previous_metadata_failure_keeps_committed_snapshot_without_cleanup()
+test_oversized_raw_is_rejected_before_shared_publication()
+
+return true
diff --git a/internal/apps/agent/nginx/waf_runtime.lua b/internal/apps/agent/nginx/waf_runtime.lua
new file mode 100644
index 00000000..1c2b28c9
--- /dev/null
+++ b/internal/apps/agent/nginx/waf_runtime.lua
@@ -0,0 +1,404 @@
+local _M = {}
+
+local rules_config
+local ip_groups_config
+local ip_groups_runtime
+local pow_runtime
+local geo_lookup
+local geo_module
+local geo_profiles = { city = false, country = false }
+
+local function read_file(path)
+ local file, err = io.open(path, "r")
+ if not file then
+ return nil, err
+ end
+ local content = file:read("*a")
+ file:close()
+ return content
+end
+
+local function load_json(path)
+ local content, err = read_file(path)
+ if not content or content == "" then
+ return nil, err or "empty file"
+ end
+ local decoded, decode_err = require("cjson.safe").decode(content)
+ if not decoded then
+ return nil, decode_err or "invalid JSON"
+ end
+ return decoded
+end
+
+local function warn_rate_limited(key, ...)
+ local dict = ngx.shared and ngx.shared.openflare_waf_config
+ if not dict or not dict.add or dict:add(key, true, 60) then
+ ngx.log(ngx.WARN, ...)
+ end
+end
+
+local function file_exists(path)
+ local file = io.open(path, "rb")
+ if not file then return false end
+ file:close()
+ return true
+end
+
+local function init_geo_databases(country_path, city_path, path_exists, region_required)
+ local ok, module_or_error = pcall(require, "resty.maxminddb")
+ if not ok or not module_or_error then
+ warn_rate_limited("_geo_module_unavailable", "openflare waf GeoIP module unavailable: ", module_or_error)
+ return
+ end
+ geo_module = module_or_error
+ local profiles = {}
+ if path_exists(city_path) then profiles.city = city_path end
+ if path_exists(country_path) then profiles.country = country_path end
+ if not profiles.city and region_required then
+ warn_rate_limited("_geo_city_unavailable", "openflare waf GeoLite2 City database unavailable; region match takes false branch")
+ end
+ if not profiles.country and not profiles.city then
+ warn_rate_limited("_geo_database_unavailable", "openflare waf GeoIP databases unavailable")
+ return
+ end
+ local function initialize_profile(profile, path)
+ local called, init_result, init_error = pcall(geo_module.init, { [profile] = path })
+ if not called or init_result ~= true then
+ return false, init_error or init_result
+ end
+ geo_profiles[profile] = true
+ return true
+ end
+ local city_initialized, city_error = false, nil
+ if profiles.city then
+ city_initialized, city_error = initialize_profile("city", profiles.city)
+ if not city_initialized and region_required then
+ warn_rate_limited("_geo_city_unavailable", "openflare waf GeoLite2 City database initialization failed; region match takes false branch: ", city_error)
+ end
+ end
+ local country_initialized, country_error = false, nil
+ if profiles.country then
+ country_initialized, country_error = initialize_profile("country", profiles.country)
+ end
+ if not city_initialized and not country_initialized then
+ warn_rate_limited("_geo_database_unavailable", "openflare waf GeoIP database initialization failed: ", country_error or city_error)
+ end
+end
+
+local function lookup_geo_profile(ip, profile)
+ if not geo_module or not geo_profiles[profile] then return nil end
+ local ok, result, lookup_error = pcall(geo_module.lookup, ip, nil, profile)
+ if not ok or not result then
+ warn_rate_limited("_geo_lookup_failed_" .. profile, "openflare waf GeoIP ", profile, " lookup failed: ", lookup_error or result)
+ return nil
+ end
+ return result
+end
+
+local function default_geo_lookup(ip, region_required)
+ local result = lookup_geo_profile(ip, "city")
+ local from_city = result ~= nil
+ if not result then result = lookup_geo_profile(ip, "country") end
+ if not result then return nil, nil end
+ local country = result.country and result.country.iso_code or nil
+ local subdivision
+ if from_city then
+ subdivision = result.most_specific_subdivision and result.most_specific_subdivision.iso_code or nil
+ if not subdivision and result.subdivisions and result.subdivisions[1] then
+ subdivision = result.subdivisions[1].iso_code
+ end
+ elseif region_required then
+ warn_rate_limited("_geo_city_unavailable", "openflare waf GeoLite2 City database unavailable; region match takes false branch")
+ end
+ country = country and string.upper(country) or nil
+ subdivision = subdivision and string.upper(subdivision) or nil
+ local region = subdivision
+ if country and subdivision and not string.match(subdivision, "^[A-Z][A-Z]%-") then
+ region = country .. "-" .. subdivision
+ end
+ return country, region
+end
+
+local function config_geo_requirements(config)
+ local uses_geo, uses_region = false, false
+ for _, rule in ipairs(config.rule_groups or {}) do
+ for _, node in pairs((rule.graph or {}).nodes or {}) do
+ if node.type == "geo_match" then
+ uses_geo = true
+ local node_config = node.config or {}
+ if type(node_config.regions) == "table" and #node_config.regions > 0 then uses_region = true end
+ end
+ end
+ end
+ return uses_geo, uses_region
+end
+
+function _M.init(options)
+ if rules_config then
+ return true
+ end
+ options = options or {}
+ local runtime_dir = options.runtime_dir or "__OPENFLARE_RUNTIME_CONFIG_DIR__"
+ if options.config then
+ rules_config = options.config
+ else
+ local err
+ rules_config, err = load_json(runtime_dir .. "/waf_config.json")
+ assert(rules_config, "load waf_config.json failed: " .. tostring(err))
+ end
+ if options.ip_groups then
+ ip_groups_config = options.ip_groups
+ else
+ ip_groups_runtime = options.ip_groups_runtime or require("waf.ip_groups")
+ local initialized, init_error = ip_groups_runtime.init({ runtime_dir = runtime_dir })
+ assert(initialized, "initialize WAF IP groups failed: " .. tostring(init_error))
+ end
+ pow_runtime = options.pow or require("pow.runtime")
+ if options.geo_lookup then
+ geo_lookup = options.geo_lookup
+ else
+ local uses_geo, uses_region = config_geo_requirements(rules_config)
+ if uses_geo then
+ init_geo_databases(
+ options.country_mmdb_path or "__OPENFLARE_COUNTRY_MMDB_PATH__",
+ options.city_mmdb_path or "__OPENFLARE_CITY_MMDB_PATH__",
+ options.geo_file_exists or file_exists,
+ uses_region
+ )
+ end
+ geo_lookup = default_geo_lookup
+ end
+ return true
+end
+
+-- Task 7 can atomically replace the worker-local IP group snapshot through this seam.
+function _M.replace_ip_groups(snapshot)
+ ip_groups_config = snapshot or { groups = {} }
+end
+
+local function list_contains(items, value)
+ if type(items) ~= "table" or not value then return false end
+ value = string.upper(value)
+ for _, item in ipairs(items) do
+ if string.upper(tostring(item)) == value then return true end
+ end
+ return false
+end
+
+local function parse_ipv4(value)
+ local a, b, c, d = string.match(value or "", "^(%d+)%.(%d+)%.(%d+)%.(%d+)$")
+ if not a then return nil end
+ a, b, c, d = tonumber(a), tonumber(b), tonumber(c), tonumber(d)
+ if a > 255 or b > 255 or c > 255 or d > 255 then return nil end
+ return ((a * 256 + b) * 256 + c) * 256 + d
+end
+
+local function split_ipv6_side(value)
+ local result = {}
+ if value == "" then return result end
+ for part in string.gmatch(value, "[^:]+") do
+ if string.find(part, ".", 1, true) then
+ local ipv4 = parse_ipv4(part)
+ if not ipv4 then return nil end
+ result[#result + 1] = math.floor(ipv4 / 65536)
+ result[#result + 1] = ipv4 % 65536
+ else
+ if #part > 4 or not string.match(part, "^[%x]+$") then return nil end
+ local number = tonumber(part, 16)
+ if not number or number > 65535 then return nil end
+ result[#result + 1] = number
+ end
+ end
+ return result
+end
+
+local function parse_ipv6(value)
+ value = string.lower(value or "")
+ local compressed_at = string.find(value, "::", 1, true)
+ if compressed_at and string.find(value, "::", compressed_at + 2, true) then return nil end
+ local left, right
+ if compressed_at then
+ left = split_ipv6_side(string.sub(value, 1, compressed_at - 1))
+ right = split_ipv6_side(string.sub(value, compressed_at + 2))
+ else
+ if string.sub(value, 1, 1) == ":" or string.sub(value, -1) == ":" then return nil end
+ left, right = split_ipv6_side(value), {}
+ end
+ if not left or not right then return nil end
+ local missing = 8 - #left - #right
+ if (compressed_at and missing < 1) or (not compressed_at and missing ~= 0) then return nil end
+ local result = {}
+ for _, number in ipairs(left) do result[#result + 1] = number end
+ for _ = 1, missing do result[#result + 1] = 0 end
+ for _, number in ipairs(right) do result[#result + 1] = number end
+ if #result ~= 8 then return nil end
+ return result
+end
+
+local function ipv6_equal(left, right)
+ left, right = parse_ipv6(left), parse_ipv6(right)
+ if not left or not right then return false end
+ for index = 1, 8 do
+ if left[index] ~= right[index] then return false end
+ end
+ return true
+end
+
+local function ip_in_cidr(ip, cidr)
+ local base, bits = string.match(cidr or "", "^([^/]+)/(%d+)$")
+ bits = tonumber(bits)
+ if not base or not bits then return false end
+ local ip_number, base_number = parse_ipv4(ip), parse_ipv4(base)
+ if ip_number and base_number then
+ if bits < 0 or bits > 32 then return false end
+ if bits == 0 then return true end
+ local size = 2 ^ (32 - bits)
+ return ip_number - (ip_number % size) == base_number - (base_number % size)
+ end
+ local ip_groups, base_groups = parse_ipv6(ip), parse_ipv6(base)
+ if not ip_groups or not base_groups or bits < 0 or bits > 128 then return false end
+ local full_groups, remaining_bits = math.floor(bits / 16), bits % 16
+ for index = 1, full_groups do
+ if ip_groups[index] ~= base_groups[index] then return false end
+ end
+ if remaining_bits > 0 then
+ local size = 2 ^ (16 - remaining_bits)
+ local index = full_groups + 1
+ if math.floor(ip_groups[index] / size) ~= math.floor(base_groups[index] / size) then return false end
+ end
+ return true
+end
+
+local function matches_ip_values(config, ip)
+ for _, item in ipairs(config.ips or {}) do
+ if item == ip or ipv6_equal(item, ip) then return true end
+ end
+ for _, cidr in ipairs(config.cidrs or {}) do
+ if ip_in_cidr(ip, cidr) then return true end
+ end
+ local snapshot = ip_groups_config or ip_groups_runtime.current()
+ local groups = (snapshot or {}).groups or {}
+ for _, id in ipairs(config.ip_group_ids or {}) do
+ local group = groups[tostring(id)]
+ if group and group.enabled then
+ for _, item in ipairs(group.ip_list or {}) do
+ if item == ip or ipv6_equal(item, ip) or ip_in_cidr(ip, item) then return true end
+ end
+ end
+ end
+ return false
+end
+
+local function fail_closed(reason)
+ local dict = ngx.shared and ngx.shared.openflare_waf_config
+ if not dict or not dict.add or dict:add("_damaged_graph_logged", true, 60) then
+ ngx.log(ngx.ERR, "openflare waf damaged runtime graph: ", reason)
+ end
+ ngx.ctx.openflare_waf_blocked = true
+ ngx.status = 500
+ ngx.header["Content-Type"] = "text/plain; charset=utf-8"
+ ngx.say("OpenFlare WAF runtime error")
+ return ngx.exit(500)
+end
+
+local function render_block(config)
+ config = config or {}
+ local status = tonumber(config.status_code) or 403
+ ngx.ctx.openflare_waf_blocked = true
+ ngx.status = status
+ local body = config.response_body or ""
+ if body ~= "" then
+ ngx.header["Content-Type"] = "text/html; charset=utf-8"
+ ngx.say(body)
+ end
+ return ngx.exit(status)
+end
+
+local function execute_graph(graph)
+ if type(graph) ~= "table" or type(graph.nodes) ~= "table" or type(graph.entry) ~= "string" then
+ return nil, "invalid graph"
+ end
+ local node_count = 0
+ for _ in pairs(graph.nodes) do node_count = node_count + 1 end
+ local current = graph.entry
+ for _ = 1, node_count do
+ local node = graph.nodes[current]
+ if type(node) ~= "table" or type(node.type) ~= "string" then
+ return nil, "missing node " .. tostring(current)
+ end
+ if node.type == "allow" then
+ return { kind = "allow" }
+ end
+ if node.type == "block" then
+ return { kind = "block", config = node.config }
+ end
+ local handle
+ if node.type == "start" then
+ handle = "next"
+ elseif node.type == "ip_match" then
+ handle = matches_ip_values(node.config or {}, ngx.var.remote_addr or "") and "true" or "false"
+ elseif node.type == "geo_match" then
+ local config = node.config or {}
+ local region_required = type(config.regions) == "table" and #config.regions > 0
+ local country, region = geo_lookup(ngx.var.remote_addr or "", region_required)
+ handle = (list_contains(config.countries, country) or list_contains(config.regions, region)) and "true" or "false"
+ elseif node.type == "pow" then
+ if pow_runtime.evaluate(node.config or {}) ~= true then
+ return { kind = "takeover" }
+ end
+ handle = "next"
+ else
+ return nil, "unknown node type " .. node.type
+ end
+ if type(node.next) ~= "table" or type(node.next[handle]) ~= "string" then
+ return nil, "missing " .. handle .. " edge from " .. current
+ end
+ current = node.next[handle]
+ end
+ return nil, "graph exceeded maximum steps"
+end
+
+local function active_rules(site)
+ local by_id, result = {}, {}
+ for _, rule in ipairs(rules_config.rule_groups or {}) do
+ by_id[tostring(rule.id)] = rule
+ if rule.enabled and rule.is_global then result[#result + 1] = rule end
+ end
+ for _, binding in ipairs(rules_config.bindings or {}) do
+ if binding.site_name == site then
+ for _, id in ipairs(binding.rule_group_ids or {}) do
+ local rule = by_id[tostring(id)]
+ if rule and rule.enabled and not rule.is_global then result[#result + 1] = rule end
+ end
+ break
+ end
+ end
+ return result
+end
+
+local function is_internal_pow_continuation()
+ if not ngx.req or not ngx.req.is_internal or not ngx.req.is_internal() then return false end
+ local uri = ngx.var.uri or ""
+ local api_prefix = "/.within.website/x/cmd/anubis/api/"
+ local static_prefix = "/.within.website/x/cmd/anubis/static/"
+ return string.sub(uri, 1, #api_prefix) == api_prefix or string.sub(uri, 1, #static_prefix) == static_prefix
+end
+
+function _M.check()
+ if not rules_config then
+ return fail_closed("runtime not initialized")
+ end
+ if is_internal_pow_continuation() then
+ ngx.ctx.openflare_pow_takeover = true
+ return
+ end
+ for _, rule in ipairs(active_rules(ngx.var.openflare_waf_site or "")) do
+ local decision, err = execute_graph(rule.graph)
+ if not decision then return fail_closed(err) end
+ if decision.kind == "block" then return render_block(decision.config) end
+ if decision.kind == "takeover" then return end
+ end
+ return "ok"
+end
+
+return _M
diff --git a/internal/apps/agent/nginx/waf_runtime_spec.lua b/internal/apps/agent/nginx/waf_runtime_spec.lua
new file mode 100644
index 00000000..2e841060
--- /dev/null
+++ b/internal/apps/agent/nginx/waf_runtime_spec.lua
@@ -0,0 +1,535 @@
+local runtime_path = assert(WAF_RUNTIME_PATH, "WAF_RUNTIME_PATH is required")
+
+local function assert_equal(actual, expected, message)
+ if actual ~= expected then
+ error((message or "values differ") .. ": expected " .. tostring(expected) .. ", got " .. tostring(actual), 2)
+ end
+end
+
+local output
+local pow_calls
+local pow_results
+local shared_keys = {}
+local logs = {}
+
+ngx = {
+ WARN = "WARN",
+ ERR = "ERR",
+ var = {},
+ ctx = {},
+ header = {},
+ shared = {
+ openflare_waf_config = {
+ add = function(_, key)
+ if shared_keys[key] then return false end
+ shared_keys[key] = true
+ return true
+ end,
+ },
+ },
+ req = { is_internal = function() return ngx.var.openflare_internal == true end },
+ say = function(body) output.body = body end,
+ exit = function(status) output.exit = status return status end,
+ log = function(_, ...)
+ local parts = { ... }
+ for index, value in ipairs(parts) do parts[index] = tostring(value) end
+ output.log = table.concat(parts)
+ logs[#logs + 1] = output.log
+ end,
+}
+
+local pow_stub = {}
+function pow_stub.evaluate(config)
+ pow_calls[#pow_calls + 1] = config.difficulty
+ local result = pow_results[1]
+ table.remove(pow_results, 1)
+ return result
+end
+
+local function node(node_type, config, next_nodes)
+ return { type = node_type, config = config or {}, next = next_nodes }
+end
+
+local function graph(nodes, entry)
+ return { entry = entry or "start", nodes = nodes }
+end
+
+local function rule(id, is_global, rule_graph)
+ return { id = id, enabled = true, is_global = is_global or false, graph = rule_graph }
+end
+
+local function start_to(target)
+ return node("start", {}, { next = target })
+end
+
+local function load_runtime(config, options)
+ local chunk = assert(loadfile(runtime_path))
+ local runtime = chunk()
+ options = options or {}
+ runtime.init({
+ config = config,
+ ip_groups = options.ip_groups or { groups = {} },
+ pow = pow_stub,
+ geo_lookup = options.geo_lookup,
+ runtime_dir = options.runtime_dir,
+ geo_file_exists = options.geo_file_exists,
+ country_mmdb_path = options.country_mmdb_path or (options.runtime_dir and (options.runtime_dir .. "/GeoLite2-Country.mmdb") or nil),
+ city_mmdb_path = options.city_mmdb_path or (options.runtime_dir and (options.runtime_dir .. "/GeoLite2-City.mmdb") or nil),
+ })
+ return runtime
+end
+
+local function reset_request(site, ip, uri, is_internal)
+ ngx.var = { openflare_waf_site = site, remote_addr = ip or "192.0.2.1", uri = uri or "/", request_id = "request-1", openflare_internal = is_internal == true }
+ ngx.ctx = {}
+ ngx.header = {}
+ output = {}
+ pow_calls = {}
+ pow_results = {}
+end
+
+local function binding(site, ids)
+ return { site_name = site, rule_group_ids = ids }
+end
+
+local function test_ip_true_and_false()
+ local config = {
+ rule_groups = { rule(1, false, graph({
+ start = start_to("match"),
+ match = node("ip_match", { ips = { "192.0.2.1" }, cidrs = { "198.51.100.0/24" }, ip_group_ids = { 7 } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 451, response_body = "ip blocked" }),
+ allow = node("allow"),
+ })) },
+ bindings = { binding("ip-site", { 1 }) },
+ }
+ local runtime = load_runtime(config, { ip_groups = { groups = { ["7"] = { enabled = true, ip_list = { "203.0.113.7" } } } } })
+
+ reset_request("ip-site", "192.0.2.1")
+ runtime.check()
+ assert_equal(output.exit, 451, "exact IP true branch")
+
+ reset_request("ip-site", "198.51.100.8")
+ runtime.check()
+ assert_equal(output.exit, 451, "CIDR true branch")
+
+ reset_request("ip-site", "203.0.113.7")
+ runtime.check()
+ assert_equal(output.exit, 451, "IP group true branch")
+
+ reset_request("ip-site", "203.0.113.8")
+ runtime.check()
+ assert_equal(output.exit, nil, "IP false branch")
+end
+
+local function test_ipv6_exact_cidr_and_group()
+ local config = {
+ rule_groups = { rule(8, false, graph({
+ start = start_to("match"),
+ match = node("ip_match", { ips = { "2001:db8::1" }, cidrs = { "2001:db8:abcd::/48" }, ip_group_ids = { 9 } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 451, response_body = "ipv6 blocked" }),
+ allow = node("allow"),
+ })) },
+ bindings = { binding("ipv6-site", { 8 }) },
+ }
+ local runtime = load_runtime(config, { ip_groups = { groups = { ["9"] = { enabled = true, ip_list = { "2001:db8:ffff::/48" } } } } })
+
+ reset_request("ipv6-site", "2001:0db8:0:0:0:0:0:1")
+ runtime.check()
+ assert_equal(output.exit, 451, "canonical-equivalent IPv6 exact match")
+
+ reset_request("ipv6-site", "2001:db8:abcd:12::9")
+ runtime.check()
+ assert_equal(output.exit, 451, "IPv6 CIDR true branch")
+
+ reset_request("ipv6-site", "2001:db8:ffff:beef::9")
+ runtime.check()
+ assert_equal(output.exit, 451, "IP group IPv6 CIDR true branch")
+
+ reset_request("ipv6-site", "2001:db9::1")
+ runtime.check()
+ assert_equal(output.exit, nil, "IPv6 false branch")
+end
+
+local function test_geo_true_and_false()
+ local config = {
+ rule_groups = { rule(2, false, graph({
+ start = start_to("geo"),
+ geo = node("geo_match", { countries = { "US" }, regions = { "DE-BE" } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 403, response_body = "geo blocked" }),
+ allow = node("allow"),
+ })) },
+ bindings = { binding("geo-site", { 2 }) },
+ }
+ local country, region = "US", "NY"
+ local runtime = load_runtime(config, { geo_lookup = function() return country, region end })
+
+ reset_request("geo-site")
+ runtime.check()
+ assert_equal(output.exit, 403, "country true branch")
+
+ country, region = "DE", "DE-BE"
+ reset_request("geo-site")
+ runtime.check()
+ assert_equal(output.exit, 403, "region true branch")
+
+ country, region = "DE", "BE"
+ reset_request("geo-site")
+ runtime.check()
+ assert_equal(output.exit, nil, "geo false branch")
+end
+
+local function test_geo_module_is_initialized_once_and_composes_region()
+ local init_calls, lookup_calls = 0, 0
+ local initialized_profiles = {}
+ package.loaded["resty.maxminddb"] = nil
+ package.preload["resty.maxminddb"] = function()
+ return {
+ init = function(profiles)
+ init_calls = init_calls + 1
+ for profile, path in pairs(profiles) do initialized_profiles[profile] = path end
+ return true
+ end,
+ has_profile = function(profile) return initialized_profiles[profile] ~= nil end,
+ lookup = function(_, _, profile)
+ lookup_calls = lookup_calls + 1
+ assert_equal(profile, "city", "subdivision lookup uses City profile")
+ return { country = { iso_code = "US" }, subdivisions = { { iso_code = "CA" } } }
+ end,
+ }
+ end
+ local config = {
+ rule_groups = { rule(12, false, graph({
+ start = start_to("geo"),
+ geo = node("geo_match", { regions = { "US-CA" } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 403 }),
+ allow = node("allow"),
+ })) },
+ bindings = { binding("geo-cache", { 12 }) },
+ }
+ local runtime = load_runtime(config, { runtime_dir = "/runtime", geo_file_exists = function() return true end })
+ assert_equal(init_calls, 2, "each MaxMind profile initializes independently during worker init")
+ assert_equal(initialized_profiles.city, "/runtime/GeoLite2-City.mmdb", "City profile path")
+ assert_equal(initialized_profiles.country, "/runtime/GeoLite2-Country.mmdb", "Country profile path")
+ for _ = 1, 3 do
+ reset_request("geo-cache")
+ runtime.check()
+ assert_equal(output.exit, 403, "MaxMind subdivision composes validator-compatible region")
+ end
+ assert_equal(init_calls, 2, "MaxMind database is not initialized on requests")
+ assert_equal(lookup_calls, 3, "requests only perform lookup")
+end
+
+local function test_geo_country_fallback_does_not_fake_region()
+ local profiles
+ package.loaded["resty.maxminddb"] = nil
+ package.preload["resty.maxminddb"] = function()
+ return {
+ init = function(value) profiles = value return true end,
+ has_profile = function(profile) return profiles[profile] ~= nil end,
+ lookup = function(_, _, profile)
+ assert_equal(profile, "country", "fallback lookup uses Country profile")
+ return { country = { iso_code = "US" }, subdivisions = { { iso_code = "CA" } } }
+ end,
+ }
+ end
+ shared_keys = {}
+ logs = {}
+ local country_graph = graph({
+ start = start_to("geo"),
+ geo = node("geo_match", { countries = { "US" } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 403 }), allow = node("allow"),
+ })
+ local region_graph = graph({
+ start = start_to("geo"),
+ geo = node("geo_match", { regions = { "US-CA" } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 451 }), allow = node("allow"),
+ })
+ local runtime = load_runtime({
+ rule_groups = { rule(15, false, country_graph), rule(16, false, region_graph) },
+ bindings = { binding("country-only", { 15 }), binding("region-without-city", { 16 }) },
+ }, {
+ runtime_dir = "/runtime",
+ geo_file_exists = function(path) return string.find(path, "Country", 1, true) ~= nil end,
+ })
+
+ reset_request("country-only")
+ runtime.check()
+ assert_equal(output.exit, 403, "Country fallback remains available")
+
+ reset_request("region-without-city")
+ runtime.check()
+ assert_equal(output.exit, nil, "Country subdivisions must not satisfy region")
+ runtime.check()
+ assert_equal(#logs, 1, "missing City warning is rate limited")
+end
+
+local function test_geo_city_init_failure_retries_country_profile()
+ local init_calls = {}
+ local profiles = {}
+ package.loaded["resty.maxminddb"] = nil
+ package.preload["resty.maxminddb"] = function()
+ return {
+ init = function(value)
+ init_calls[#init_calls + 1] = value
+ if value.city then return false end
+ profiles = value
+ return true
+ end,
+ has_profile = function(profile) return profiles[profile] ~= nil end,
+ lookup = function(_, _, profile)
+ assert_equal(profile, "country", "corrupt City fallback uses Country")
+ return { country = { iso_code = "DE" } }
+ end,
+ }
+ end
+ shared_keys = {}
+ logs = {}
+ local runtime = load_runtime({
+ rule_groups = { rule(17, false, graph({
+ start = start_to("geo"),
+ geo = node("geo_match", { countries = { "DE" }, regions = { "DE-BE" } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 403 }), allow = node("allow"),
+ })) },
+ bindings = { binding("corrupt-city", { 17 }) },
+ }, { runtime_dir = "/runtime", geo_file_exists = function() return true end })
+
+ reset_request("corrupt-city")
+ runtime.check()
+ assert_equal(#init_calls, 2, "Country profile is retried after City profile init failure")
+ assert_equal(output.exit, 403, "Country remains available after corrupt City init")
+end
+
+local function test_geo_partial_init_never_looks_up_corrupt_city()
+ local opened = {}
+ local lookups = {}
+ package.loaded["resty.maxminddb"] = nil
+ package.preload["resty.maxminddb"] = function()
+ return {
+ init = function(profiles)
+ if profiles.country then opened.country = true end
+ if profiles.city then return nil, "corrupt City" end
+ return true
+ end,
+ initted = function() return next(opened) ~= nil end,
+ lookup = function(_, _, profile)
+ lookups[#lookups + 1] = profile
+ assert_equal(opened[profile], true, "lookup must only use an opened profile")
+ return { country = { iso_code = "DE" } }
+ end,
+ }
+ end
+ local runtime = load_runtime({
+ rule_groups = { rule(18, false, graph({
+ start = start_to("geo"),
+ geo = node("geo_match", { countries = { "DE" } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 403 }), allow = node("allow"),
+ })) },
+ bindings = { binding("partial-corrupt-city", { 18 }) },
+ }, { runtime_dir = "/runtime", geo_file_exists = function() return true end })
+
+ reset_request("partial-corrupt-city")
+ runtime.check()
+ assert_equal(table.concat(lookups, ","), "country", "corrupt City is never looked up")
+ assert_equal(output.exit, 403, "valid Country remains available")
+end
+
+local function test_geo_partial_init_never_looks_up_corrupt_country()
+ local opened = {}
+ local lookups = {}
+ package.loaded["resty.maxminddb"] = nil
+ package.preload["resty.maxminddb"] = function()
+ return {
+ init = function(profiles)
+ if profiles.city then opened.city = true end
+ if profiles.country then return nil, "corrupt Country" end
+ return true
+ end,
+ initted = function() return next(opened) ~= nil end,
+ lookup = function(_, _, profile)
+ lookups[#lookups + 1] = profile
+ assert_equal(opened[profile], true, "lookup must only use an opened profile")
+ return nil, "address absent"
+ end,
+ }
+ end
+ local runtime = load_runtime({
+ rule_groups = { rule(19, false, graph({
+ start = start_to("geo"),
+ geo = node("geo_match", { countries = { "DE" } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 403 }), allow = node("allow"),
+ })) },
+ bindings = { binding("partial-corrupt-country", { 19 }) },
+ }, { runtime_dir = "/runtime", geo_file_exists = function() return true end })
+
+ reset_request("partial-corrupt-country")
+ runtime.check()
+ assert_equal(table.concat(lookups, ","), "city", "corrupt Country is never used as fallback")
+ assert_equal(output.exit, nil, "missing City result takes false branch without corrupt fallback")
+end
+
+local function test_geo_unavailable_warning_is_rate_limited()
+ package.loaded["resty.maxminddb"] = nil
+ package.preload["resty.maxminddb"] = function() error("module unavailable") end
+ shared_keys = {}
+ logs = {}
+ local config = {
+ rule_groups = { rule(13, false, graph({
+ start = start_to("geo"),
+ geo = node("geo_match", { countries = { "US" } }, { ["true"] = "blocked", ["false"] = "allow" }),
+ blocked = node("block", { status_code = 403 }),
+ allow = node("allow"),
+ })) },
+ bindings = { binding("geo-missing", { 13 }) },
+ }
+ local first = load_runtime(config)
+ local second = load_runtime(config)
+ reset_request("geo-missing")
+ first.check()
+ second.check()
+ assert_equal(#logs, 1, "missing MaxMind warning is rate limited across workers")
+end
+
+local function test_pow_takeover_and_completion()
+ local config = {
+ rule_groups = { rule(3, false, graph({
+ start = start_to("pow"),
+ pow = node("pow", { algorithm = "fast", difficulty = 5, session_ttl = 600, challenge_ttl = 300 }, { next = "blocked" }),
+ blocked = node("block", { status_code = 429, response_body = "after pow" }),
+ allow = node("allow"),
+ })) },
+ bindings = { binding("pow-site", { 3 }) },
+ }
+ local runtime = load_runtime(config)
+
+ reset_request("pow-site")
+ pow_results = { false }
+ runtime.check()
+ assert_equal(output.exit, nil, "PoW takeover must stop graph execution")
+ assert_equal(#pow_calls, 1, "PoW evaluated once")
+
+ reset_request("pow-site")
+ pow_results = { true }
+ runtime.check()
+ assert_equal(output.exit, 429, "completed PoW follows next edge")
+end
+
+local function test_pow_internal_redirect_bypasses_graph_as_takeover()
+ local config = {
+ rule_groups = { rule(14, false, graph({
+ start = start_to("pow"),
+ pow = node("pow", { difficulty = 4 }, { next = "blocked" }),
+ blocked = node("block", { status_code = 429 }),
+ allow = node("allow"),
+ })) },
+ bindings = { binding("pow-internal", { 14 }) },
+ }
+ local runtime = load_runtime(config)
+ reset_request("pow-internal", "192.0.2.1", "/.within.website/x/cmd/anubis/api/make-challenge", true)
+ pow_results = { true }
+ runtime.check()
+ assert_equal(#pow_calls, 0, "internal challenge continuation must not re-enter DAG")
+ assert_equal(output.exit, nil, "internal challenge continuation must not follow pow next")
+end
+
+local function test_block_config_and_rule_order()
+ local function pow_allow(difficulty)
+ return graph({
+ start = start_to("pow"),
+ pow = node("pow", { algorithm = "fast", difficulty = difficulty, session_ttl = 600, challenge_ttl = 300 }, { next = "allow" }),
+ allow = node("allow"),
+ })
+ end
+ local config = {
+ rule_groups = {
+ rule(10, true, pow_allow(10)),
+ rule(20, false, pow_allow(20)),
+ rule(30, false, pow_allow(30)),
+ rule(40, false, graph({
+ start = start_to("blocked"),
+ blocked = node("block", { status_code = 418, response_body = "custom block" }),
+ allow = node("allow"),
+ })),
+ },
+ bindings = { binding("ordered-site", { 30, 20, 40 }) },
+ }
+ local runtime = load_runtime(config)
+
+ reset_request("ordered-site")
+ pow_results = { true, true, true }
+ runtime.check()
+ assert_equal(table.concat(pow_calls, ","), "10,30,20", "global rule precedes binding order")
+ assert_equal(output.exit, 418, "block status comes from reached block node")
+ assert_equal(output.body, "custom block", "block body comes from reached block node")
+ assert_equal(ngx.header["Content-Type"], "text/html; charset=utf-8", "block content type")
+end
+
+local function test_damaged_graphs_fail_closed()
+ local configs = {
+ graph({ start = start_to("unknown"), unknown = node("future_node"), allow = node("allow") }),
+ graph({ start = start_to("missing"), allow = node("allow") }),
+ graph({ start = start_to("loop"), loop = node("start", {}, { next = "loop" }), allow = node("allow") }),
+ }
+ for index, damaged in ipairs(configs) do
+ local runtime = load_runtime({ rule_groups = { rule(index, false, damaged) }, bindings = { binding("damaged", { index }) } })
+ reset_request("damaged")
+ runtime.check()
+ assert_equal(output.exit, 500, "damaged graph " .. index .. " must fail closed")
+ end
+end
+
+local function test_request_path_has_no_file_io()
+ local opens = 0
+ local original_open = io.open
+ io.open = function(path, mode)
+ opens = opens + 1
+ local value = path:match("waf_ip_groups%.json$") and "IP_GROUPS" or "CONFIG"
+ return {
+ read = function() return value end,
+ close = function() end,
+ }
+ end
+ package.loaded["cjson.safe"] = nil
+ package.preload["cjson.safe"] = function()
+ return { decode = function(value)
+ if value == "IP_GROUPS" then return { groups = {} } end
+ return {
+ rule_groups = { rule(1, false, graph({ start = start_to("allow"), allow = node("allow") })) },
+ bindings = { binding("io-site", { 1 }) },
+ }
+ end }
+ end
+ local chunk = assert(loadfile(runtime_path))
+ local runtime = chunk()
+ runtime.init({
+ runtime_dir = "/runtime",
+ pow = pow_stub,
+ ip_groups_runtime = {
+ init = function() return true end,
+ current = function() return { groups = {} } end,
+ },
+ })
+ local init_opens = opens
+ assert_equal(init_opens, 1, "WAF graph initializes once; IP groups are owned by refresh module")
+
+ reset_request("io-site")
+ for _ = 1, 3 do runtime.check() end
+ assert_equal(opens, init_opens, "request execution performs no file I/O")
+ io.open = original_open
+end
+
+test_ip_true_and_false()
+test_ipv6_exact_cidr_and_group()
+test_geo_true_and_false()
+test_geo_module_is_initialized_once_and_composes_region()
+test_geo_country_fallback_does_not_fake_region()
+test_geo_city_init_failure_retries_country_profile()
+test_geo_partial_init_never_looks_up_corrupt_city()
+test_geo_partial_init_never_looks_up_corrupt_country()
+test_geo_unavailable_warning_is_rate_limited()
+test_pow_takeover_and_completion()
+test_pow_internal_redirect_bypasses_graph_as_takeover()
+test_block_config_and_rule_order()
+test_damaged_graphs_fail_closed()
+test_request_path_has_no_file_io()
+
+return true
diff --git a/internal/apps/agent/sync/service.go b/internal/apps/agent/sync/service.go
index 67e5b254..80ca9c6b 100644
--- a/internal/apps/agent/sync/service.go
+++ b/internal/apps/agent/sync/service.go
@@ -9,6 +9,7 @@ import (
"fmt"
"log/slog"
"sort"
+ "strconv"
"strings"
"sync"
@@ -42,7 +43,8 @@ type NginxManager interface {
EnsureSafeFallbackRuntime(ctx context.Context, reason string) error
CurrentChecksum() (string, error)
WAFIPGroupChecksums() (map[string]string, error)
- SyncWAFIPGroups(groups []protocol.WAFIPGroup) error
+ ReconcileWAFIPGroups(targetIDs []uint, changed []protocol.WAFIPGroup) error
+ UpdateExistingWAFIPGroups(changed []protocol.WAFIPGroup) error
EnsureWorkerReadAccess() error
}
@@ -132,12 +134,12 @@ func (s *Service) WAFIPGroupChecksums() (map[string]string, error) {
return s.nginxManager.WAFIPGroupChecksums()
}
-// ApplyWAFIPGroups writes the given WAF IP groups to the nginx manager.
+// ApplyWAFIPGroups applies real-time changes only to groups already in the local authoritative snapshot.
func (s *Service) ApplyWAFIPGroups(_ context.Context, groups []protocol.WAFIPGroup) error {
if len(groups) == 0 || s.nginxManager == nil {
return nil
}
- return s.nginxManager.SyncWAFIPGroups(groups)
+ return s.nginxManager.UpdateExistingWAFIPGroups(groups)
}
func (s *Service) applyIfNeeded(ctx context.Context, mode string, startup bool, snapshot *state.Snapshot, currentChecksum string, target *protocol.ActiveConfigMeta, config *protocol.ActiveConfigResponse) error {
@@ -218,25 +220,42 @@ func (s *Service) applyRenderedConfig(ctx context.Context, mode string, snapshot
}
func (s *Service) syncReferencedWAFIPGroups(ctx context.Context, supportFiles []protocol.SupportFile) error {
- ids := referencedWAFIPGroupIDs(supportFiles)
+ ids, err := referencedWAFIPGroupIDs(supportFiles)
+ if err != nil {
+ return err
+ }
if len(ids) == 0 {
- return nil
+ if s.nginxManager == nil {
+ return nil
+ }
+ return s.nginxManager.ReconcileWAFIPGroups([]uint{}, nil)
}
checksums, err := s.WAFIPGroupChecksums()
if err != nil {
return err
}
+ targetChecksums := make(map[string]string, len(ids))
+ for _, id := range ids {
+ key := strconv.FormatUint(uint64(id), 10)
+ if value := strings.TrimSpace(checksums[key]); value != "" {
+ targetChecksums[key] = value
+ }
+ }
response, err := s.client.SyncWAFIPGroups(ctx, protocol.WAFIPGroupSyncRequest{
IDs: ids,
- Checksums: checksums,
+ Checksums: targetChecksums,
})
if err != nil {
return err
}
- if response == nil || len(response.Groups) == 0 {
+ if s.nginxManager == nil {
return nil
}
- return s.ApplyWAFIPGroups(ctx, response.Groups)
+ var changed []protocol.WAFIPGroup
+ if response != nil {
+ changed = response.Groups
+ }
+ return s.nginxManager.ReconcileWAFIPGroups(ids, changed)
}
type renderedActiveConfig struct {
@@ -288,7 +307,7 @@ func fromOpenRestySupportFiles(files []openrestyrender.SupportFile) []protocol.S
return result
}
-func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) []uint {
+func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) ([]uint, error) {
var content string
for _, file := range supportFiles {
if file.Path == "waf_config.json" {
@@ -297,28 +316,38 @@ func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) []uint {
}
}
if content == "" {
- return []uint{}
- }
- var payload struct {
- RuleGroups []struct {
- IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"`
- IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"`
- } `json:"rule_groups"`
+ return []uint{}, nil
}
+ var payload openrestyrender.WAFDocument
if err := json.Unmarshal([]byte(content), &payload); err != nil {
- slog.Debug("decode waf_config.json for ip group references failed", "error", err)
- return []uint{}
+ return nil, fmt.Errorf("decode waf_config.json for ip group references: %w", err)
}
seen := make(map[uint]struct{})
for _, group := range payload.RuleGroups {
- for _, id := range group.IPWhitelistGroups {
- if id > 0 {
- seen[id] = struct{}{}
+ for _, legacyIDs := range [][]uint{group.IPWhitelistGroups, group.IPBlacklistGroups} {
+ for _, id := range legacyIDs {
+ if id > 0 {
+ seen[id] = struct{}{}
+ }
}
}
- for _, id := range group.IPBlacklistGroups {
- if id > 0 {
- seen[id] = struct{}{}
+ for nodeID, node := range group.Graph.Nodes {
+ if node.Type != "ip_match" {
+ continue
+ }
+ var config *struct {
+ IPGroupIDs []uint `json:"ip_group_ids"`
+ }
+ if err := json.Unmarshal(node.Config, &config); err != nil || config == nil {
+ if err == nil {
+ err = errors.New("config must be a JSON object")
+ }
+ return nil, fmt.Errorf("decode ip_match config for rule group %d node %s: %w", group.ID, nodeID, err)
+ }
+ for _, id := range config.IPGroupIDs {
+ if id > 0 {
+ seen[id] = struct{}{}
+ }
}
}
}
@@ -327,7 +356,7 @@ func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) []uint {
ids = append(ids, id)
}
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
- return ids
+ return ids, nil
}
func shouldReportNoopApply(snapshot *state.Snapshot, version string, checksum string) bool {
diff --git a/internal/apps/agent/sync/service_test.go b/internal/apps/agent/sync/service_test.go
index 3846c050..2797a990 100644
--- a/internal/apps/agent/sync/service_test.go
+++ b/internal/apps/agent/sync/service_test.go
@@ -30,8 +30,10 @@ func testPagesSourceConfigJSON(deploymentID uint, checksum string) string {
type fakeClient struct {
config protocol.ActiveConfigResponse
reports []protocol.ApplyLogPayload
+ wafSyncCalls []protocol.WAFIPGroupSyncRequest
pagesPackages map[uint][]byte
pagesHashes map[uint]string
+ wafSyncResult protocol.WAFIPGroupSyncResponse
fetchCalls int
hashCalls int
}
@@ -47,6 +49,12 @@ type fakeManager struct {
applyMainContents []string
applyRouteContents []string
applyFiles [][]protocol.SupportFile
+ wafChecksums map[string]string
+ wafReconcileIDs []uint
+ wafReconcileGroups []protocol.WAFIPGroup
+ wafUpdatedGroups []protocol.WAFIPGroup
+ wafReconcileErr error
+ wafReconcileCalls int
}
func testSourceConfigJSON(workerProcesses string, listen int) string {
@@ -106,7 +114,9 @@ func (f *fakeClient) ReportApplyLog(ctx context.Context, payload protocol.ApplyL
}
func (f *fakeClient) SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error) {
- return &protocol.WAFIPGroupSyncResponse{}, nil
+ f.wafSyncCalls = append(f.wafSyncCalls, payload)
+ result := f.wafSyncResult
+ return &result, nil
}
func (m *fakeManager) Apply(ctx context.Context, mainConfig string, routeConfig string, supportFiles []protocol.SupportFile) nginx.ApplyOutcome {
@@ -134,10 +144,18 @@ func (m *fakeManager) CurrentChecksum() (string, error) {
}
func (m *fakeManager) WAFIPGroupChecksums() (map[string]string, error) {
- return map[string]string{}, nil
+ return m.wafChecksums, nil
}
-func (m *fakeManager) SyncWAFIPGroups(groups []protocol.WAFIPGroup) error {
+func (m *fakeManager) ReconcileWAFIPGroups(ids []uint, groups []protocol.WAFIPGroup) error {
+ m.wafReconcileCalls++
+ m.wafReconcileIDs = append([]uint(nil), ids...)
+ m.wafReconcileGroups = append([]protocol.WAFIPGroup(nil), groups...)
+ return m.wafReconcileErr
+}
+
+func (m *fakeManager) UpdateExistingWAFIPGroups(groups []protocol.WAFIPGroup) error {
+ m.wafUpdatedGroups = append([]protocol.WAFIPGroup(nil), groups...)
return nil
}
@@ -145,6 +163,143 @@ func (m *fakeManager) EnsureWorkerReadAccess() error {
return nil
}
+func TestReferencedWAFIPGroupIDsSyncsCompiledDAGReferences(t *testing.T) {
+ client := &fakeClient{wafSyncResult: protocol.WAFIPGroupSyncResponse{Groups: []protocol.WAFIPGroup{{ID: 7, Checksum: "new-7"}}}}
+ manager := &fakeManager{wafChecksums: map[string]string{"2": "sum-2", "7": "old-7", "99": "stale"}}
+ service := New(client, manager, nil)
+ supportFiles := []protocol.SupportFile{{Path: "waf_config.json", Content: `{
+ "rule_groups":[
+ {"id":1,"ip_whitelist_group_ids":[0,11,7],"graph":{"entry":"start","nodes":{
+ "start":{"type":"start","config":{}},
+ "first":{"type":"ip_match","config":{"ip_group_ids":[7,0,2,7]}},
+ "geo":{"type":"geo_match","config":{"countries":["US"],"ip_group_ids":[700]}},
+ "pow":{"type":"pow","config":{"difficulty":4,"ip_group_ids":[800]}},
+ "second":{"type":"ip_match","config":{"ip_group_ids":[9,2]}}
+ }}}
+ ],
+ "ip_groups":[{"id":2},{"id":7},{"id":9},{"id":11},{"id":404}],
+ "bindings":[]
+ }`}}
+
+ if err := service.syncReferencedWAFIPGroups(context.Background(), supportFiles); err != nil {
+ t.Fatalf("syncReferencedWAFIPGroups failed: %v", err)
+ }
+ if len(client.wafSyncCalls) != 1 {
+ t.Fatalf("expected one WAF IP group sync request, got %d", len(client.wafSyncCalls))
+ }
+ if got, want := fmt.Sprint(client.wafSyncCalls[0].IDs), "[2 7 9 11]"; got != want {
+ t.Fatalf("referenced IDs = %s, want %s", got, want)
+ }
+ if got, want := fmt.Sprint(client.wafSyncCalls[0].Checksums), "map[2:sum-2 7:old-7]"; got != want {
+ t.Fatalf("request checksums = %s, want target-only %s", got, want)
+ }
+ if got, want := fmt.Sprint(manager.wafReconcileIDs), "[2 7 9 11]"; got != want {
+ t.Fatalf("reconcile IDs = %s, want %s", got, want)
+ }
+ if len(manager.wafReconcileGroups) != 1 || manager.wafReconcileGroups[0].ID != 7 {
+ t.Fatalf("changed groups not passed to reconcile: %#v", manager.wafReconcileGroups)
+ }
+}
+
+func TestReferencedWAFIPGroupIDsReconcilesEmptyResponseAndSurfacesMissingLocal(t *testing.T) {
+ client := &fakeClient{}
+ manager := &fakeManager{
+ wafChecksums: map[string]string{"7": "mistaken-match"},
+ wafReconcileErr: fmt.Errorf("missing referenced WAF IP group 7"),
+ }
+ service := New(client, manager, nil)
+ err := service.syncReferencedWAFIPGroups(context.Background(), []protocol.SupportFile{{
+ Path: "waf_config.json", Content: `{"rule_groups":[{"graph":{"nodes":{"match":{"type":"ip_match","config":{"ip_group_ids":[7]}}}}}]}`,
+ }})
+ if err == nil || !strings.Contains(err.Error(), "missing referenced WAF IP group 7") {
+ t.Fatalf("expected missing local group error after empty delta, got %v", err)
+ }
+ if len(client.wafSyncCalls) != 1 || len(manager.wafReconcileIDs) != 1 || manager.wafReconcileIDs[0] != 7 {
+ t.Fatalf("empty response did not reach authoritative reconcile: calls=%#v ids=%#v", client.wafSyncCalls, manager.wafReconcileIDs)
+ }
+}
+
+func TestReferencedWAFIPGroupIDsEmptyTargetClearsWithoutRequest(t *testing.T) {
+ client := &fakeClient{}
+ manager := &fakeManager{}
+ service := New(client, manager, nil)
+ if err := service.syncReferencedWAFIPGroups(context.Background(), nil); err != nil {
+ t.Fatal(err)
+ }
+ if len(client.wafSyncCalls) != 0 {
+ t.Fatalf("empty target performed request I/O: %#v", client.wafSyncCalls)
+ }
+ if manager.wafReconcileCalls != 1 {
+ t.Fatal("empty authoritative target was not reconciled")
+ }
+}
+
+func TestApplyWAFIPGroupsUsesExistingOnlyUpdatePath(t *testing.T) {
+ manager := &fakeManager{}
+ service := New(nil, manager, nil)
+ groups := []protocol.WAFIPGroup{{ID: 99, Checksum: "broadcast"}}
+ if err := service.ApplyWAFIPGroups(context.Background(), groups); err != nil {
+ t.Fatal(err)
+ }
+ if len(manager.wafUpdatedGroups) != 1 || manager.wafUpdatedGroups[0].ID != 99 {
+ t.Fatalf("broadcast did not use existing-only update path: %#v", manager.wafUpdatedGroups)
+ }
+ if manager.wafReconcileCalls != 0 {
+ t.Fatal("broadcast must not use authoritative reconciliation")
+ }
+}
+
+func TestReferencedWAFIPGroupIDsRejectsMalformedRuntimeConfig(t *testing.T) {
+ tests := []struct {
+ name string
+ content string
+ wantErr string
+ }{
+ {name: "document", content: `{`, wantErr: "decode waf_config.json"},
+ {
+ name: "ip match config",
+ content: `{"rule_groups":[{"id":1,"graph":{"nodes":{"match":{"type":"ip_match","config":{"ip_group_ids":"bad"}}}}}],"bindings":[]}`,
+ wantErr: "decode ip_match config",
+ },
+ }
+ for _, test := range tests {
+ t.Run(test.name, func(t *testing.T) {
+ client := &fakeClient{}
+ service := New(client, &fakeManager{}, nil)
+ err := service.syncReferencedWAFIPGroups(context.Background(), []protocol.SupportFile{{
+ Path: "waf_config.json", Content: test.content,
+ }})
+ if err == nil || !strings.Contains(err.Error(), test.wantErr) {
+ t.Fatalf("expected %q error, got %v", test.wantErr, err)
+ }
+ if len(client.wafSyncCalls) != 0 {
+ t.Fatalf("malformed WAF config must not send sync request, got %#v", client.wafSyncCalls)
+ }
+ })
+ }
+}
+
+func TestWAFIPGroupChecksumServicePublishesSidecar(t *testing.T) {
+ runtimeDir := t.TempDir()
+ manager := &nginx.Manager{RuntimeConfigDir: runtimeDir}
+ if err := manager.ReconcileWAFIPGroups([]uint{3}, []protocol.WAFIPGroup{{
+ ID: 3, Enabled: true, IPList: []string{"203.0.113.3"}, Checksum: "sum-3",
+ }}); err != nil {
+ t.Fatalf("ReconcileWAFIPGroups failed: %v", err)
+ }
+ jsonData, err := os.ReadFile(filepath.Join(runtimeDir, nginx.WAFIPGroupsConfigFileName))
+ if err != nil {
+ t.Fatalf("read IP group JSON: %v", err)
+ }
+ checksumData, err := os.ReadFile(filepath.Join(runtimeDir, nginx.WAFIPGroupsChecksumFileName))
+ if err != nil {
+ t.Fatalf("read IP group checksum: %v", err)
+ }
+ if got, want := strings.TrimSpace(string(checksumData)), testBytesChecksum(jsonData); got != want {
+ t.Fatalf("published checksum mismatch: got %q want %q", got, want)
+ }
+}
+
func TestSyncOnceSuccess(t *testing.T) {
client := &fakeClient{
config: protocol.ActiveConfigResponse{
diff --git a/internal/apps/openflare/agent/waf_ip_group.go b/internal/apps/openflare/agent/waf_ip_group.go
index bd2ed516..4b5cbe21 100644
--- a/internal/apps/openflare/agent/waf_ip_group.go
+++ b/internal/apps/openflare/agent/waf_ip_group.go
@@ -4,6 +4,7 @@
package agent
import (
+ "bytes"
"context"
"crypto/sha256"
"encoding/hex"
@@ -11,43 +12,32 @@ import (
"errors"
"fmt"
"sort"
+ "strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/pkg/protocol"
+ openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
)
-type snapshotWAFRuleGroupRef struct {
- IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
- IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
-}
-
-type snapshotWAFSection struct {
- RuleGroups []snapshotWAFRuleGroupRef `json:"rule_groups"`
-}
-
type activeConfigSnapshot struct {
- WAF snapshotWAFSection `json:"waf"`
+ WAF openrestyrender.WAFDocument `json:"waf"`
+}
+
+type runtimeIPMatchConfig struct {
+ IPs []string `json:"ips,omitempty"`
+ CIDRs []string `json:"cidrs,omitempty"`
+ IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
}
// WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
- return buildAgentWAFIPGroups(ctx, ids)
+ return validatedAgentWAFIPGroups(ctx, ids, false)
}
// ChangedWAFIPGroupsForAgent returns WAF IP groups whose checksums differ from the agent state.
func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) {
- targetIDs := uniqueUintIDs(ids)
- if len(targetIDs) == 0 {
- activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
- if err != nil {
- return nil, err
- }
- targetIDs = activeIDs
- }
- if len(targetIDs) == 0 {
- return []WAFIPGroup{}, nil
- }
- groups, err := buildAgentWAFIPGroups(ctx, targetIDs)
+ groups, err := validatedAgentWAFIPGroups(ctx, ids, true)
if err != nil {
return nil, err
}
@@ -61,6 +51,45 @@ func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[s
return changed, nil
}
+func validatedAgentWAFIPGroups(ctx context.Context, ids []uint, fallbackToActive bool) ([]WAFIPGroup, error) {
+ targetIDs := uniqueUintIDs(ids)
+ activeIDs, err := activeConfigWAFIPGroupIDs(ctx)
+ if err != nil {
+ return nil, err
+ }
+ if len(targetIDs) == 0 && fallbackToActive {
+ targetIDs = activeIDs
+ }
+ if len(targetIDs) == 0 {
+ return []WAFIPGroup{}, nil
+ }
+
+ validationIDs := uniqueUintIDs(append(append([]uint{}, activeIDs...), targetIDs...))
+ allGroups, err := buildAgentWAFIPGroups(ctx, validationIDs)
+ if err != nil {
+ return nil, err
+ }
+ runtimeGroups := make(map[string]protocol.WAFIPGroup, len(allGroups))
+ for _, group := range allGroups {
+ runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = group
+ }
+ if err = protocol.ValidateWAFIPGroupSnapshotSize(runtimeGroups); err != nil {
+ return nil, err
+ }
+
+ targetSet := make(map[uint]struct{}, len(targetIDs))
+ for _, id := range targetIDs {
+ targetSet[id] = struct{}{}
+ }
+ result := make([]WAFIPGroup, 0, len(targetIDs))
+ for _, group := range allGroups {
+ if _, ok := targetSet[group.ID]; ok {
+ result = append(result, group)
+ }
+ }
+ return result, nil
+}
+
func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
ids = uniqueUintIDs(ids)
if len(ids) == 0 {
@@ -142,6 +171,8 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
}
idSet := make(map[uint]struct{})
for _, group := range snapshot.WAF.RuleGroups {
+ // Retain legacy flattened references while older active snapshots may
+ // still exist during a rolling Server upgrade.
for _, id := range group.IPWhitelistGroups {
if id > 0 {
idSet[id] = struct{}{}
@@ -152,6 +183,20 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
idSet[id] = struct{}{}
}
}
+ for nodeID, node := range group.Graph.Nodes {
+ if node.Type != "ip_match" {
+ continue
+ }
+ ids, err := runtimeIPMatchGroupIDs(node.Config)
+ if err != nil {
+ return nil, fmt.Errorf("活动配置 WAF 规则 %d 节点 %s 的 IP 匹配配置无效: %w", group.ID, nodeID, err)
+ }
+ for _, id := range ids {
+ if id > 0 {
+ idSet[id] = struct{}{}
+ }
+ }
+ }
}
ids := make([]uint, 0, len(idSet))
for id := range idSet {
@@ -161,6 +206,16 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
return ids, nil
}
+func runtimeIPMatchGroupIDs(raw json.RawMessage) ([]uint, error) {
+ var config runtimeIPMatchConfig
+ decoder := json.NewDecoder(bytes.NewReader(raw))
+ decoder.DisallowUnknownFields()
+ if err := decoder.Decode(&config); err != nil {
+ return nil, err
+ }
+ return config.IPGroupIDs, nil
+}
+
func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) {
text := strings.TrimSpace(snapshotJSON)
if text == "" {
@@ -171,7 +226,7 @@ func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, erro
return nil, err
}
if snapshot.WAF.RuleGroups == nil {
- snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{}
+ snapshot.WAF.RuleGroups = []openrestyrender.WAFRuleGroup{}
}
return &snapshot, nil
}
diff --git a/internal/apps/openflare/agent/waf_ip_group_test.go b/internal/apps/openflare/agent/waf_ip_group_test.go
index 41e8bd37..f1ea210d 100644
--- a/internal/apps/openflare/agent/waf_ip_group_test.go
+++ b/internal/apps/openflare/agent/waf_ip_group_test.go
@@ -7,10 +7,12 @@ import (
"context"
"encoding/json"
"strconv"
+ "strings"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/pkg/protocol"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -63,6 +65,115 @@ func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID
}).Error)
}
+func seedActiveConfigWithWAFGraphIPGroup(t *testing.T, ctx context.Context, ipGroupID uint) {
+ t.Helper()
+
+ snapshot := map[string]any{
+ "routes": []any{},
+ "waf": map[string]any{
+ "rule_groups": []map[string]any{
+ {
+ "id": 1,
+ "name": "graph refs",
+ "enabled": true,
+ "graph": map[string]any{
+ "entry": "start",
+ "nodes": map[string]any{
+ "start": map[string]any{
+ "type": "start",
+ "config": map[string]any{},
+ "next": map[string]string{"next": "match"},
+ },
+ "match": map[string]any{
+ "type": "ip_match",
+ "config": map[string]any{
+ "ip_group_ids": []uint{ipGroupID},
+ },
+ "next": map[string]string{"true": "allow", "false": "allow"},
+ },
+ "allow": map[string]any{
+ "type": "allow",
+ "config": map[string]any{},
+ },
+ },
+ },
+ },
+ },
+ "bindings": []any{},
+ },
+ }
+ snapshotJSON, err := json.Marshal(snapshot)
+ require.NoError(t, err)
+
+ require.NoError(t, db.DB(ctx).Create(&model.ConfigVersion{
+ Version: "20260713-graph-001",
+ SnapshotJSON: string(snapshotJSON),
+ Checksum: "graph-test-checksum",
+ IsActive: true,
+ }).Error)
+}
+
+func TestChangedWAFIPGroupsForAgentDiscoversGraphReferences(t *testing.T) {
+ cleanup := setupWAFIPGroupTestDB(t)
+ defer cleanup()
+
+ ctx := context.Background()
+ ipGroup := &model.OpenFlareWAFIPGroup{
+ Name: "graph runtime group",
+ Type: "manual",
+ Enabled: true,
+ IPList: `["192.0.2.88"]`,
+ }
+ require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
+ seedActiveConfigWithWAFGraphIPGroup(t, ctx, ipGroup.ID)
+
+ groups, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
+ require.NoError(t, err)
+ require.Len(t, groups, 1)
+ assert.Equal(t, ipGroup.ID, groups[0].ID)
+ assert.Equal(t, []string{"192.0.2.88"}, groups[0].IPList)
+}
+
+func TestChangedWAFIPGroupsForAgentRejectsMalformedIPMatchConfig(t *testing.T) {
+ cleanup := setupWAFIPGroupTestDB(t)
+ defer cleanup()
+
+ ctx := context.Background()
+ require.NoError(t, db.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":[]}}`,
+ Checksum: "malformed-test-checksum",
+ IsActive: true,
+ }).Error)
+
+ _, err := ChangedWAFIPGroupsForAgent(ctx, nil, nil)
+ require.ErrorContains(t, err, "规则 7 节点 match")
+ require.ErrorContains(t, err, "IP 匹配配置无效")
+}
+
+func TestChangedWAFIPGroupsForAgentRejectsOversizedSnapshotBeforeChecksumDelta(t *testing.T) {
+ cleanup := setupWAFIPGroupTestDB(t)
+ defer cleanup()
+
+ ctx := context.Background()
+ ipGroup := &model.OpenFlareWAFIPGroup{
+ Name: strings.Repeat("x", protocol.MaxWAFIPGroupSnapshotBytes),
+ Type: "manual",
+ Enabled: true,
+ IPList: `[]`,
+ }
+ require.NoError(t, model.CreateOpenFlareWAFIPGroup(ctx, ipGroup))
+ agentGroup, err := buildAgentWAFIPGroup(ipGroup)
+ require.NoError(t, err)
+
+ _, err = ChangedWAFIPGroupsForAgent(ctx, []uint{ipGroup.ID}, map[string]string{
+ strconv.FormatUint(uint64(ipGroup.ID), 10): agentGroup.Checksum,
+ })
+ require.ErrorContains(t, err, "WAF IP 组快照大小")
+ require.ErrorContains(t, err, "超过上限")
+}
+
func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t)
defer cleanup()
diff --git a/internal/apps/openflare/config_version/logics_test.go b/internal/apps/openflare/config_version/logics_test.go
index d2127b1f..f0be1642 100644
--- a/internal/apps/openflare/config_version/logics_test.go
+++ b/internal/apps/openflare/config_version/logics_test.go
@@ -13,6 +13,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
+ openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -151,13 +152,7 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err)
- customGroup := &model.OpenFlareWAFRuleGroup{
- Name: "pow-group",
- Enabled: true,
- PoWEnabled: true,
- PoWConfig: `{"difficulty":4,"algorithm":"fast","session_ttl":600,"challenge_ttl":300}`,
- }
- require.NoError(t, model.CreateOpenFlareWAFRuleGroup(ctx, customGroup))
+ customGroup := createSnapshotRule(t, ctx, "pow-group", waf.DefaultRuleGraph())
require.NoError(t, model.ReplaceOpenFlareWAFRuleGroupBindings(ctx, customGroup.ID, []uint{route.ID}))
bundle, err := buildCurrentConfigBundle(ctx, true)
@@ -177,20 +172,24 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
}
assert.True(t, found, "expected WAF binding for enabled route")
- var wafRuntime struct {
- SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
- }
+ var wafRuntime openrestyrender.WAFDocument
+ foundWAFConfig := false
for _, file := range bundle.SupportFiles {
if file.Path != "waf_config.json" {
continue
}
+ foundWAFConfig = true
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
}
- require.Contains(t, wafRuntime.SiteRuleGroups, "example.com")
- require.Contains(t, wafRuntime.SiteRuleGroups["example.com"], customGroup.ID)
- require.Contains(t, wafRuntime.SiteRuleGroups["example.com"], globalGroup.ID)
+ require.True(t, foundWAFConfig, "expected rendered WAF support file")
+ require.NotEmpty(t, wafRuntime.RuleGroups)
+ assert.Equal(t, globalGroup.ID, wafRuntime.RuleGroups[0].ID)
+ assert.True(t, wafRuntime.RuleGroups[0].IsGlobal)
+ require.Len(t, wafRuntime.Bindings, 1)
+ assert.Equal(t, route.ID, wafRuntime.Bindings[0].RouteID)
+ assert.Equal(t, "example.com", wafRuntime.Bindings[0].SiteName)
+ assert.Equal(t, []uint{customGroup.ID}, wafRuntime.Bindings[0].RuleGroupIDs)
assert.Contains(t, bundle.RouteConfig, `set $openflare_waf_site "example.com"`)
- assert.Contains(t, bundle.RouteConfig, `require("pow.runtime").check()`)
}
func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testing.T) {
@@ -210,34 +209,30 @@ func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testi
require.NoError(t, waf.EnsureDefaultRuleGroup(ctx))
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err)
- globalGroup.PoWEnabled = true
- globalGroup.PoWConfig = `{"difficulty":4,"algorithm":"fast","session_ttl":600,"challenge_ttl":300}`
- require.NoError(t, model.UpdateOpenFlareWAFRuleGroup(ctx, globalGroup))
+ 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)
bundle, err := buildCurrentConfigBundle(ctx, true)
require.NoError(t, err)
- assert.Contains(t, bundle.RouteConfig, `require("pow.runtime").check()`)
-
- var wafRuntime struct {
- RuleGroups []struct {
- ID uint `json:"id"`
- PoWEnabled bool `json:"pow_enabled"`
- PoWConfig *struct {
- Difficulty int `json:"difficulty"`
- } `json:"pow_config"`
- } `json:"rule_groups"`
- SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
- }
+ var wafRuntime openrestyrender.WAFDocument
+ foundWAFConfig := false
for _, file := range bundle.SupportFiles {
if file.Path != "waf_config.json" {
continue
}
+ foundWAFConfig = true
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
}
- require.Contains(t, wafRuntime.SiteRuleGroups, "pow-global.example.com")
- require.Contains(t, wafRuntime.SiteRuleGroups["pow-global.example.com"], globalGroup.ID)
+ require.True(t, foundWAFConfig, "expected rendered WAF support file")
require.NotEmpty(t, wafRuntime.RuleGroups)
- assert.True(t, wafRuntime.RuleGroups[0].PoWEnabled)
- require.NotNil(t, wafRuntime.RuleGroups[0].PoWConfig)
- assert.Equal(t, 4, wafRuntime.RuleGroups[0].PoWConfig.Difficulty)
+ assert.Equal(t, globalGroup.ID, wafRuntime.RuleGroups[0].ID)
+ assert.True(t, wafRuntime.RuleGroups[0].IsGlobal)
+ assert.Equal(t, string(waf.RuleNodePoW), wafRuntime.RuleGroups[0].Graph.Nodes["pow"].Type)
+ require.Len(t, wafRuntime.Bindings, 1)
+ assert.Equal(t, "pow-global.example.com", wafRuntime.Bindings[0].SiteName)
+ assert.Empty(t, wafRuntime.Bindings[0].RuleGroupIDs)
+ require.NotEmpty(t, bundle.WAFSnapshot.RuleGroups)
+ assert.Equal(t, waf.RuleNodePoW, bundle.WAFSnapshot.RuleGroups[0].Graph.Nodes["pow"].Type)
}
diff --git a/internal/apps/openflare/config_version/snapshot.go b/internal/apps/openflare/config_version/snapshot.go
index e5c28d8e..2dff4d46 100644
--- a/internal/apps/openflare/config_version/snapshot.go
+++ b/internal/apps/openflare/config_version/snapshot.go
@@ -9,17 +9,21 @@ import (
"errors"
"fmt"
"sort"
+ "strconv"
"strings"
oftls "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
+ "github.com/Rain-kl/Wavelet/pkg/protocol"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
+ "gorm.io/gorm"
)
const (
- supportFilesPerCertificate = 2
+ supportFilesPerCertificate = 2
+ wafIPGroupChecksumHexLength = 64
// OpenResty 默认配置值
defaultOpenRestyReturnStatus = 421
@@ -66,22 +70,11 @@ type snapshotRoute struct {
}
type snapshotWAFRuleGroup struct {
- ID uint `json:"id"`
- Name string `json:"name"`
- Enabled bool `json:"enabled"`
- IsGlobal bool `json:"is_global"`
- BlockStatusCode int `json:"block_status_code"`
- BlockResponseBody string `json:"block_response_body,omitempty"`
- IPWhitelist []string `json:"ip_whitelist,omitempty"`
- IPBlacklist []string `json:"ip_blacklist,omitempty"`
- IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
- IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
- CountryWhitelist []string `json:"country_whitelist,omitempty"`
- CountryBlacklist []string `json:"country_blacklist,omitempty"`
- RegionWhitelist []string `json:"region_whitelist,omitempty"`
- RegionBlacklist []string `json:"region_blacklist,omitempty"`
- PoWEnabled bool `json:"pow_enabled,omitempty"`
- PoWConfig *openrestyrender.PoWConfig `json:"pow_config,omitempty"`
+ ID uint `json:"id"`
+ Name string `json:"name"`
+ Enabled bool `json:"enabled"`
+ IsGlobal bool `json:"is_global"`
+ Graph waf.RuntimeRuleGraph `json:"graph"`
}
type snapshotWAFIPGroup struct {
@@ -311,35 +304,37 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
if err := waf.EnsureDefaultRuleGroup(ctx); err != nil {
return snapshotWAFDocument{}, err
}
- views, err := waf.ListRuleGroups(ctx)
+ groups, err := model.ListOpenFlareWAFRuleGroups(ctx)
if err != nil {
return snapshotWAFDocument{}, err
}
- ruleGroups := make([]snapshotWAFRuleGroup, 0, len(views))
- for _, view := range views {
- if !view.Enabled {
+ ruleGroups := make([]snapshotWAFRuleGroup, 0, len(groups))
+ referencedIPGroupIDs := make(map[uint]struct{})
+ enabledRuleIDs := make(map[uint]struct{})
+ for _, group := range groups {
+ if !group.Enabled {
continue
}
+ var editorGraph waf.RuleGraph
+ if err = json.Unmarshal([]byte(group.Graph), &editorGraph); err != nil {
+ return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 的图数据无效: %w", group.Name, err)
+ }
+ if err = waf.ValidateRuleGraph(ctx, editorGraph, snapshotWAFIPGroupExists); err != nil {
+ return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 的图无效: %w", group.Name, err)
+ }
+ runtimeGraph, compileErr := waf.CompileRuleGraph(editorGraph)
+ if compileErr != nil {
+ return snapshotWAFDocument{}, fmt.Errorf("WAF 规则 %s 编译失败: %w", group.Name, compileErr)
+ }
ruleGroups = append(ruleGroups, snapshotWAFRuleGroup{
- ID: view.ID,
- Name: view.Name,
- Enabled: view.Enabled,
- IsGlobal: view.IsGlobal,
- BlockStatusCode: view.BlockStatusCode,
- BlockResponseBody: view.BlockResponseBody,
- IPWhitelist: view.IPWhitelist,
- IPBlacklist: view.IPBlacklist,
- IPWhitelistGroups: view.IPWhitelistGroups,
- IPBlacklistGroups: view.IPBlacklistGroups,
- CountryWhitelist: view.CountryWhitelist,
- CountryBlacklist: view.CountryBlacklist,
- RegionWhitelist: view.RegionWhitelist,
- RegionBlacklist: view.RegionBlacklist,
- PoWEnabled: view.PoWEnabled,
- PoWConfig: convertPoWConfig(view.PoWEnabled, view.PoWConfig),
+ ID: group.ID, Name: group.Name, Enabled: group.Enabled, IsGlobal: group.IsGlobal, Graph: runtimeGraph,
})
+ enabledRuleIDs[group.ID] = struct{}{}
+ for _, id := range waf.ReferencedIPGroupIDs(editorGraph) {
+ referencedIPGroupIDs[id] = struct{}{}
+ }
}
- ipGroups, err := buildSnapshotWAFIPGroups(ctx, ruleGroups)
+ ipGroups, err := buildSnapshotWAFIPGroups(ctx, referencedIPGroupIDs)
if err != nil {
return snapshotWAFDocument{}, err
}
@@ -366,16 +361,16 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
if _, ok := enabledRouteSiteNames[binding.ProxyRouteID]; !ok {
continue
}
- groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID)
+ if _, enabled := enabledRuleIDs[binding.RuleGroupID]; enabled {
+ groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID)
+ }
}
bindings := make([]snapshotWAFBinding, 0, len(enabledRouteSiteNames))
for routeID, siteName := range enabledRouteSiteNames {
- groupIDs := groupIDsByRoute[routeID]
- sort.Slice(groupIDs, func(i, j int) bool { return groupIDs[i] < groupIDs[j] })
bindings = append(bindings, snapshotWAFBinding{
RouteID: routeID,
SiteName: siteName,
- RuleGroupIDs: groupIDs,
+ RuleGroupIDs: groupIDsByRoute[routeID],
})
}
sort.Slice(bindings, func(i, j int) bool {
@@ -387,16 +382,26 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
return snapshotWAFDocument{RuleGroups: ruleGroups, IPGroups: ipGroups, Bindings: bindings}, nil
}
-func buildSnapshotWAFIPGroups(ctx context.Context, ruleGroups []snapshotWAFRuleGroup) ([]snapshotWAFIPGroup, error) {
- idSet := make(map[uint]struct{})
- for _, group := range ruleGroups {
- for _, id := range group.IPWhitelistGroups {
- idSet[id] = struct{}{}
+func validateSnapshotWAFIPGroupSize(groups []snapshotWAFIPGroup) error {
+ runtimeGroups := make(map[string]protocol.WAFIPGroup, len(groups))
+ for _, group := range groups {
+ ipList := group.IPList
+ if !group.Enabled {
+ ipList = []string{}
}
- for _, id := range group.IPBlacklistGroups {
- idSet[id] = struct{}{}
+ runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = protocol.WAFIPGroup{
+ ID: group.ID,
+ Name: group.Name,
+ Type: group.Type,
+ Enabled: group.Enabled,
+ IPList: ipList,
+ Checksum: strings.Repeat("0", wafIPGroupChecksumHexLength),
}
}
+ return protocol.ValidateWAFIPGroupSnapshotSize(runtimeGroups)
+}
+
+func buildSnapshotWAFIPGroups(ctx context.Context, idSet map[uint]struct{}) ([]snapshotWAFIPGroup, error) {
if len(idSet) == 0 {
return []snapshotWAFIPGroup{}, nil
}
@@ -431,9 +436,20 @@ func buildSnapshotWAFIPGroups(ctx context.Context, ruleGroups []snapshotWAFRuleG
IPList: ipList,
})
}
+ if err = validateSnapshotWAFIPGroupSize(snapshots); err != nil {
+ return nil, err
+ }
return snapshots, nil
}
+func snapshotWAFIPGroupExists(ctx context.Context, id uint) (bool, error) {
+ group, err := model.GetOpenFlareWAFIPGroupByID(ctx, id)
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return false, nil
+ }
+ return group != nil, err
+}
+
func decodeIPList(raw string) ([]string, error) {
text := strings.TrimSpace(raw)
if text == "" {
@@ -446,36 +462,6 @@ func decodeIPList(raw string) ([]string, error) {
return items, nil
}
-func convertPoWConfig(enabled bool, config *waf.PoWConfig) *openrestyrender.PoWConfig {
- if !enabled {
- return nil
- }
- if config == nil {
- defaultConfig := openrestyrender.DefaultPoWConfig()
- return &defaultConfig
- }
- return &openrestyrender.PoWConfig{
- Difficulty: config.Difficulty,
- Algorithm: config.Algorithm,
- SessionTTL: config.SessionTTL,
- ChallengeTTL: config.ChallengeTTL,
- Whitelist: openrestyrender.PoWListConfig{
- IPs: config.Whitelist.IPs,
- IPCidrs: config.Whitelist.IPCidrs,
- Paths: config.Whitelist.Paths,
- PathRegexes: config.Whitelist.PathRegexes,
- UserAgents: config.Whitelist.UserAgents,
- },
- Blacklist: openrestyrender.PoWListConfig{
- IPs: config.Blacklist.IPs,
- IPCidrs: config.Blacklist.IPCidrs,
- Paths: config.Blacklist.Paths,
- PathRegexes: config.Blacklist.PathRegexes,
- UserAgents: config.Blacklist.UserAgents,
- },
- }
-}
-
func buildOpenRestyConfigSnapshot(ctx context.Context) openRestyConfigSnapshot {
// 读取所有 OpenResty 配置,使用默认值作为降级
getIntConfig := func(key string, defaultVal int) int {
diff --git a/internal/apps/openflare/config_version/waf_graph_snapshot_test.go b/internal/apps/openflare/config_version/waf_graph_snapshot_test.go
new file mode 100644
index 00000000..1daf0aed
--- /dev/null
+++ b/internal/apps/openflare/config_version/waf_graph_snapshot_test.go
@@ -0,0 +1,135 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package config_version
+
+import (
+ "context"
+ "encoding/json"
+ "strings"
+ "testing"
+
+ "github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
+ "github.com/Rain-kl/Wavelet/internal/db"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+)
+
+func TestBuildSnapshotRejectsOversizedAggregateWAFIPGroups(t *testing.T) {
+ cleanup := setupConfigVersionTestDB(t)
+ defer cleanup()
+ ctx := context.Background()
+
+ // Each group remains below the existing 2 MiB per-subscription ceiling,
+ // while the complete Agent runtime document exceeds the aggregate limit.
+ ipList, err := json.Marshal(strings.Fields(strings.Repeat("192.0.2.1 ", 165000)))
+ require.NoError(t, err)
+ require.Less(t, len(ipList), 2<<20)
+
+ groupIDs := make([]uint, 0, 12)
+ for index := 0; index < 12; index++ {
+ group := &model.OpenFlareWAFIPGroup{
+ Name: "aggregate-" + strings.Repeat("x", index),
+ Type: "manual",
+ Enabled: true,
+ IPList: string(ipList),
+ }
+ require.NoError(t, db.DB(ctx).Create(group).Error)
+ groupIDs = append(groupIDs, group.ID)
+ }
+ createSnapshotRule(t, ctx, "oversized-aggregate", snapshotIPMatchGraphForGroups(groupIDs))
+
+ _, err = buildSnapshotWAFDocument(ctx, nil)
+ require.ErrorContains(t, err, "WAF IP 组快照大小")
+ require.ErrorContains(t, err, "超过上限")
+}
+
+func TestWAFGraphSnapshotPreservesOrderAndGraphReferences(t *testing.T) {
+ cleanup := setupConfigVersionTestDB(t)
+ defer cleanup()
+ ctx := context.Background()
+
+ route := &model.ProxyRoute{SiteName: "ordered.example.com", OriginURL: "http://origin:8080", Upstreams: `["http://origin:8080"]`, Enabled: true}
+ require.NoError(t, model.CreateProxyRouteRecord(ctx, route))
+ createSnapshotZoneDomains(t, ctx, route, route.SiteName)
+
+ 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)
+
+ customA := createSnapshotRule(t, ctx, "custom-a", waf.DefaultRuleGraph())
+ customB := createSnapshotRule(t, ctx, "custom-b", snapshotIPMatchGraph(referenced.ID))
+ require.NoError(t, model.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, route.ID, []uint{customB.ID, customA.ID}))
+
+ snapshot, err := buildSnapshotWAFDocument(ctx, []*model.ProxyRoute{route})
+ require.NoError(t, err)
+ require.Len(t, snapshot.Bindings, 1)
+ assert.Equal(t, []uint{customB.ID, customA.ID}, snapshot.Bindings[0].RuleGroupIDs)
+ require.Len(t, snapshot.IPGroups, 1)
+ assert.Equal(t, referenced.ID, snapshot.IPGroups[0].ID)
+
+ var customBSnapshot *snapshotWAFRuleGroup
+ for index := range snapshot.RuleGroups {
+ if snapshot.RuleGroups[index].ID == customB.ID {
+ customBSnapshot = &snapshot.RuleGroups[index]
+ }
+ }
+ require.NotNil(t, customBSnapshot)
+ assert.Equal(t, "start", customBSnapshot.Graph.Entry)
+ assert.Equal(t, waf.RuleNodeIPMatch, customBSnapshot.Graph.Nodes["match"].Type)
+ raw, err := json.Marshal(customBSnapshot)
+ require.NoError(t, err)
+ assert.NotContains(t, string(raw), "position")
+ assert.NotContains(t, string(raw), "ip_whitelist")
+}
+
+func TestBuildSnapshotRejectsInvalidWAFGraph(t *testing.T) {
+ cleanup := setupConfigVersionTestDB(t)
+ defer cleanup()
+ 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)
+ _, err := buildSnapshotWAFDocument(ctx, nil)
+ require.ErrorContains(t, err, "invalid")
+}
+
+func createSnapshotRule(t *testing.T, ctx context.Context, name string, graph waf.RuleGraph) *model.OpenFlareWAFRuleGroup {
+ t.Helper()
+ 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)
+ return rule
+}
+
+func snapshotIPMatchGraph(ipGroupID uint) waf.RuleGraph {
+ return snapshotIPMatchGraphForGroups([]uint{ipGroupID})
+}
+
+func snapshotIPMatchGraphForGroups(ipGroupIDs []uint) waf.RuleGraph {
+ config, _ := json.Marshal(waf.IPMatchConfig{IPGroupIDs: ipGroupIDs})
+ return waf.RuleGraph{SchemaVersion: waf.RuleGraphSchemaVersion, Nodes: []waf.RuleNode{
+ {ID: "start", Type: waf.RuleNodeStart, Position: waf.RulePosition{X: 1, Y: 2}, Config: json.RawMessage(`{}`)},
+ {ID: "match", Type: waf.RuleNodeIPMatch, Position: waf.RulePosition{X: 3, Y: 4}, Config: config},
+ {ID: "allow", Type: waf.RuleNodeAllow, Position: waf.RulePosition{X: 5, Y: 6}, Config: json.RawMessage(`{}`)},
+ }, Edges: []waf.RuleEdge{
+ {ID: "e1", Source: "start", SourceHandle: "next", Target: "match"},
+ {ID: "e2", Source: "match", SourceHandle: "true", Target: "allow"},
+ {ID: "e3", Source: "match", SourceHandle: "false", Target: "allow"},
+ }}
+}
+
+func snapshotPoWGraph() waf.RuleGraph {
+ config, _ := json.Marshal(waf.PoWNodeConfig{Algorithm: "fast", Difficulty: 4, SessionTTL: 600, ChallengeTTL: 300})
+ return waf.RuleGraph{SchemaVersion: waf.RuleGraphSchemaVersion, Nodes: []waf.RuleNode{
+ {ID: "start", Type: waf.RuleNodeStart, Config: json.RawMessage(`{}`)},
+ {ID: "pow", Type: waf.RuleNodePoW, Config: config},
+ {ID: "allow", Type: waf.RuleNodeAllow, Config: json.RawMessage(`{}`)},
+ }, Edges: []waf.RuleEdge{
+ {ID: "e1", Source: "start", SourceHandle: "next", Target: "pow"},
+ {ID: "e2", Source: "pow", SourceHandle: "next", Target: "allow"},
+ }}
+}
diff --git a/internal/apps/openflare/integration/security_test.go b/internal/apps/openflare/integration/security_test.go
index 837bc361..238c92fb 100644
--- a/internal/apps/openflare/integration/security_test.go
+++ b/internal/apps/openflare/integration/security_test.go
@@ -109,12 +109,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
t.Run("WAF rule group create", func(t *testing.T) {
rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/waf/rule-groups"), map[string]any{
- "name": "edge-security",
- "enabled": true,
- "block_status_code": 403,
- "ip_whitelist": []string{"192.0.2.1"},
- "ip_blacklist": []string{"203.0.113.10"},
- "country_blacklist": []string{"CN"},
+ "name": "edge-security",
}, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code)
@@ -124,7 +119,8 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
assert.NotZero(t, ruleGroupID)
assert.Equal(t, "edge-security", data["name"])
assert.Equal(t, false, data["is_global"])
- assert.Equal(t, float64(403), data["block_status_code"])
+ assert.Equal(t, float64(1), data["revision"])
+ assert.NotNil(t, data["graph"])
})
t.Run("WAF rule group list includes global and custom groups", func(t *testing.T) {
@@ -174,11 +170,9 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
t,
engine,
http.MethodPost,
- fmt.Sprintf("%s/waf/rule-groups/%d/update", apiPath(""), ruleGroupID),
+ fmt.Sprintf("%s/waf/rule-groups/%d/meta", apiPath(""), ruleGroupID),
map[string]any{
- "name": "edge-security-updated",
- "enabled": true,
- "block_status_code": 451,
+ "name": "edge-security-updated", "enabled": true,
},
adminAuthHeaders(seed.Token),
)
@@ -187,7 +181,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data)
assert.Equal(t, "edge-security-updated", data["name"])
- assert.Equal(t, float64(451), data["block_status_code"])
+ assert.Equal(t, true, data["enabled"])
})
t.Run("WAF IP group create", func(t *testing.T) {
diff --git a/internal/apps/openflare/waf/graph_types.go b/internal/apps/openflare/waf/graph_types.go
index 332adeab..4a924b49 100644
--- a/internal/apps/openflare/waf/graph_types.go
+++ b/internal/apps/openflare/waf/graph_types.go
@@ -5,25 +5,35 @@ package waf
import "encoding/json"
+// RuleGraphSchemaVersion is the current persisted rule graph schema version.
const RuleGraphSchemaVersion = 1
+// RuleNodeType identifies the behavior of a rule graph node.
type RuleNodeType string
const (
- RuleNodeStart RuleNodeType = "start"
- RuleNodeAllow RuleNodeType = "allow"
- RuleNodeBlock RuleNodeType = "block"
- RuleNodeIPMatch RuleNodeType = "ip_match"
+ // RuleNodeStart begins graph execution.
+ RuleNodeStart RuleNodeType = "start"
+ // RuleNodeAllow terminates execution with an allow decision.
+ RuleNodeAllow RuleNodeType = "allow"
+ // RuleNodeBlock terminates execution with a blocking response.
+ RuleNodeBlock RuleNodeType = "block"
+ // RuleNodeIPMatch branches on an IP match.
+ RuleNodeIPMatch RuleNodeType = "ip_match"
+ // RuleNodeGeoMatch branches on a geographic match.
RuleNodeGeoMatch RuleNodeType = "geo_match"
- RuleNodePoW RuleNodeType = "pow"
+ // RuleNodePoW runs a proof-of-work challenge before continuing.
+ RuleNodePoW RuleNodeType = "pow"
)
+// RuleGraph is the editor-facing representation of an executable WAF graph.
type RuleGraph struct {
SchemaVersion int `json:"schema_version"`
Nodes []RuleNode `json:"nodes"`
Edges []RuleEdge `json:"edges"`
}
+// RuleNode stores one editor node and its type-specific configuration.
type RuleNode struct {
ID string `json:"id"`
Type RuleNodeType `json:"type"`
@@ -32,11 +42,13 @@ type RuleNode struct {
Config json.RawMessage `json:"config"`
}
+// RulePosition stores a node's editor canvas coordinates.
type RulePosition struct {
X float64 `json:"x"`
Y float64 `json:"y"`
}
+// RuleEdge connects one source handle to a target node.
type RuleEdge struct {
ID string `json:"id"`
Source string `json:"source"`
@@ -44,17 +56,20 @@ type RuleEdge struct {
Target string `json:"target"`
}
+// IPMatchConfig configures literal, CIDR, and managed-group IP matching.
type IPMatchConfig struct {
IPs []string `json:"ips,omitempty"`
CIDRs []string `json:"cidrs,omitempty"`
IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
}
+// GeoMatchConfig configures country and region matching.
type GeoMatchConfig struct {
Countries []string `json:"countries,omitempty"`
Regions []string `json:"regions,omitempty"`
}
+// PoWNodeConfig configures a proof-of-work challenge node.
type PoWNodeConfig struct {
Algorithm string `json:"algorithm"`
Difficulty int `json:"difficulty"`
@@ -62,11 +77,13 @@ type PoWNodeConfig struct {
ChallengeTTL int `json:"challenge_ttl"`
}
+// BlockNodeConfig configures a terminal blocking response.
type BlockNodeConfig struct {
StatusCode int `json:"status_code"`
ResponseBody string `json:"response_body,omitempty"`
}
+// DefaultRuleGraph returns the minimal start-to-allow graph.
func DefaultRuleGraph() RuleGraph {
return RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
{ID: "start", Type: RuleNodeStart, Position: RulePosition{X: 0, Y: 0}, Config: json.RawMessage(`{}`)},
diff --git a/internal/apps/openflare/waf/graph_validate.go b/internal/apps/openflare/waf/graph_validate.go
index 7d1af57f..45c74096 100644
--- a/internal/apps/openflare/waf/graph_validate.go
+++ b/internal/apps/openflare/waf/graph_validate.go
@@ -26,7 +26,33 @@ var (
regionCodePattern = regexp.MustCompile(`^[A-Z]{2}-[A-Z0-9]{1,3}$`)
)
+// ValidateRuleGraph validates graph structure, node configuration, references,
+// reachability, and termination before compilation.
func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(context.Context, uint) (bool, error)) error {
+ if err := validateRuleGraphLimits(graph); err != nil {
+ return err
+ }
+ nodes, startID, err := validateRuleGraphNodes(ctx, graph.Nodes, ipGroupExists)
+ if err != nil {
+ return err
+ }
+ outgoing, incoming, handleTargets, err := validateRuleGraphEdges(nodes, graph.Edges)
+ if err != nil {
+ return err
+ }
+ if hasRuleGraphCycle(nodes, outgoing, incoming) {
+ return errors.New("规则图不能包含循环")
+ }
+ if err := validateRequiredHandles(graph.Nodes, handleTargets); err != nil {
+ return err
+ }
+ if err := validateRuleGraphConnectivity(graph.Nodes, startID, outgoing, incoming); err != nil {
+ return err
+ }
+ return validateTerminalPaths(graph.Nodes, outgoing)
+}
+
+func validateRuleGraphLimits(graph RuleGraph) error {
if graph.SchemaVersion != RuleGraphSchemaVersion {
return fmt.Errorf("规则图 schema_version 必须为 %d", RuleGraphSchemaVersion)
}
@@ -41,15 +67,18 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
} else if len(raw) > maxRuleGraphBytes {
return fmt.Errorf("规则图大小不能超过 256 KiB")
}
+ return nil
+}
- nodes := make(map[string]RuleNode, len(graph.Nodes))
+func validateRuleGraphNodes(ctx context.Context, graphNodes []RuleNode, ipGroupExists func(context.Context, uint) (bool, error)) (map[string]RuleNode, string, error) {
+ nodes := make(map[string]RuleNode, len(graphNodes))
startCount, allowCount, startID := 0, 0, ""
- for _, node := range graph.Nodes {
+ for _, node := range graphNodes {
if strings.TrimSpace(node.ID) == "" {
- return errors.New("节点 ID 不能为空")
+ return nil, "", errors.New("节点 ID 不能为空")
}
if _, exists := nodes[node.ID]; exists {
- return fmt.Errorf("节点 ID %s 重复", node.ID)
+ return nil, "", fmt.Errorf("节点 ID %s 重复", node.ID)
}
nodes[node.ID] = node
switch node.Type {
@@ -60,67 +89,74 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
allowCount++
case RuleNodeBlock, RuleNodeIPMatch, RuleNodeGeoMatch, RuleNodePoW:
default:
- return fmt.Errorf("节点 %s 的类型 %s 未知", node.ID, node.Type)
+ return nil, "", fmt.Errorf("节点 %s 的类型 %s 未知", node.ID, node.Type)
}
if err := validateRuleNodeConfig(ctx, node, ipGroupExists); err != nil {
- return err
+ return nil, "", err
}
}
if startCount != 1 {
- return errors.New("规则图必须恰好包含一个开始节点")
+ return nil, "", errors.New("规则图必须恰好包含一个开始节点")
}
if allowCount != 1 {
- return errors.New("规则图必须恰好包含一个通过节点")
+ return nil, "", errors.New("规则图必须恰好包含一个通过节点")
}
+ return nodes, startID, nil
+}
- edgeIDs := make(map[string]struct{}, len(graph.Edges))
+func validateRuleGraphEdges(nodes map[string]RuleNode, graphEdges []RuleEdge) (map[string][]RuleEdge, map[string]int, map[string]int, error) {
+ edgeIDs := make(map[string]struct{}, len(graphEdges))
outgoing := make(map[string][]RuleEdge)
incoming := make(map[string]int)
handleTargets := make(map[string]int)
- for _, edge := range graph.Edges {
+ for _, edge := range graphEdges {
if strings.TrimSpace(edge.ID) == "" {
- return errors.New("边 ID 不能为空")
+ return nil, nil, nil, errors.New("边 ID 不能为空")
}
if _, exists := edgeIDs[edge.ID]; exists {
- return fmt.Errorf("边 ID %s 重复", edge.ID)
+ return nil, nil, nil, fmt.Errorf("边 ID %s 重复", edge.ID)
}
edgeIDs[edge.ID] = struct{}{}
source, ok := nodes[edge.Source]
if !ok {
- return fmt.Errorf("边 %s 的源节点 %s 不存在", edge.ID, edge.Source)
+ return nil, nil, nil, fmt.Errorf("边 %s 的源节点 %s 不存在", edge.ID, edge.Source)
}
if _, ok := nodes[edge.Target]; !ok {
- return fmt.Errorf("边 %s 的目标节点 %s 不存在", edge.ID, edge.Target)
+ return nil, nil, nil, fmt.Errorf("边 %s 的目标节点 %s 不存在", edge.ID, edge.Target)
}
if !validSourceHandle(source.Type, edge.SourceHandle) {
- return fmt.Errorf("边 %s 的源端口 %s 不适用于节点 %s", edge.ID, edge.SourceHandle, edge.Source)
+ return nil, nil, nil, fmt.Errorf("边 %s 的源端口 %s 不适用于节点 %s", edge.ID, edge.SourceHandle, edge.Source)
}
key := edge.Source + "\x00" + edge.SourceHandle
handleTargets[key]++
if handleTargets[key] > 1 {
- return fmt.Errorf("节点 %s 的 %s 出口连接了多个目标", edge.Source, edge.SourceHandle)
+ return nil, nil, nil, fmt.Errorf("节点 %s 的 %s 出口连接了多个目标", edge.Source, edge.SourceHandle)
}
outgoing[edge.Source] = append(outgoing[edge.Source], edge)
incoming[edge.Target]++
}
- if hasRuleGraphCycle(nodes, outgoing, incoming) {
- return errors.New("规则图不能包含循环")
- }
- for _, node := range graph.Nodes {
+ return outgoing, incoming, handleTargets, nil
+}
+
+func validateRequiredHandles(nodes []RuleNode, handleTargets map[string]int) error {
+ for _, node := range nodes {
for _, handle := range requiredHandles(node.Type) {
if handleTargets[node.ID+"\x00"+handle] == 0 {
return fmt.Errorf("节点 %s 的 %s 出口未连接", node.ID, handle)
}
}
}
+ return nil
+}
+func validateRuleGraphConnectivity(nodes []RuleNode, startID string, outgoing map[string][]RuleEdge, incoming map[string]int) error {
reachable := walkRuleGraph(startID, outgoing)
- for _, node := range graph.Nodes {
+ for _, node := range nodes {
if !reachable[node.ID] {
return fmt.Errorf("节点 %s 无法从开始节点到达", node.ID)
}
}
- for _, node := range graph.Nodes {
+ for _, node := range nodes {
if node.Type == RuleNodeStart && incoming[node.ID] != 0 {
return fmt.Errorf("开始节点 %s 不能有入边", node.ID)
}
@@ -131,96 +167,129 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
return fmt.Errorf("终止节点 %s 不能有出口", node.ID)
}
}
- if err := validateTerminalPaths(graph.Nodes, outgoing); err != nil {
- return err
- }
return nil
}
func validateRuleNodeConfig(ctx context.Context, node RuleNode, exists func(context.Context, uint) (bool, error)) error {
switch node.Type {
case RuleNodeStart, RuleNodeAllow:
- var cfg struct{}
- if err := decodeStrictConfig(node.Config, &cfg); err != nil {
- return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
- }
+ return validateEmptyNodeConfig(node)
case RuleNodeIPMatch:
- var cfg IPMatchConfig
- if err := decodeStrictConfig(node.Config, &cfg); err != nil {
- return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
- }
- for _, raw := range cfg.IPs {
- if _, err := netip.ParseAddr(raw); err != nil {
- return fmt.Errorf("节点 %s 的 IP %s 无效", node.ID, raw)
- }
- }
- for _, raw := range cfg.CIDRs {
- if _, err := netip.ParsePrefix(raw); err != nil {
- return fmt.Errorf("节点 %s 的 CIDR %s 无效", node.ID, raw)
- }
- }
- for _, id := range cfg.IPGroupIDs {
- if id == 0 {
- return fmt.Errorf("节点 %s 引用的 IP 组 ID 无效", node.ID)
- }
- if exists == nil {
- return fmt.Errorf("节点 %s 无法校验 IP 组 %d", node.ID, id)
- }
- ok, err := exists(ctx, id)
- if err != nil {
- return fmt.Errorf("节点 %s 校验 IP 组 %d 失败: %w", node.ID, id, err)
- }
- if !ok {
- return fmt.Errorf("节点 %s 引用的 IP 组 %d 不存在", node.ID, id)
- }
- }
+ return validateIPMatchNodeConfig(ctx, node, exists)
case RuleNodeGeoMatch:
- var cfg GeoMatchConfig
- if err := decodeStrictConfig(node.Config, &cfg); err != nil {
- return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
- }
- for _, code := range cfg.Countries {
- if !countryCodePattern.MatchString(code) {
- return fmt.Errorf("节点 %s 的国家代码 %s 无效", node.ID, code)
- }
- }
- for _, code := range cfg.Regions {
- if !regionCodePattern.MatchString(code) {
- return fmt.Errorf("节点 %s 的地区代码 %s 无效", node.ID, code)
- }
- }
+ return validateGeoMatchNodeConfig(node)
case RuleNodePoW:
- var cfg PoWNodeConfig
- if err := decodeStrictConfig(node.Config, &cfg); err != nil {
- return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
- }
- if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
- return fmt.Errorf("节点 %s 的 PoW 难度必须在 1-16 之间", node.ID)
- }
- if cfg.Algorithm != "fast" && cfg.Algorithm != "slow" {
- return fmt.Errorf("节点 %s 的 PoW 算法必须为 fast 或 slow", node.ID)
- }
- if cfg.SessionTTL < 60 {
- return fmt.Errorf("节点 %s 的 PoW 会话 TTL 不能小于 60 秒", node.ID)
- }
- if cfg.ChallengeTTL < 30 {
- return fmt.Errorf("节点 %s 的 PoW 挑战 TTL 不能小于 30 秒", node.ID)
- }
+ return validatePoWNodeConfig(node)
case RuleNodeBlock:
- var cfg BlockNodeConfig
- if err := decodeStrictConfig(node.Config, &cfg); err != nil {
- return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
+ return validateBlockNodeConfig(node)
+ }
+ return nil
+}
+
+func validateEmptyNodeConfig(node RuleNode) error {
+ var cfg struct{}
+ return decodeNodeConfig(node, &cfg)
+}
+
+func validateIPMatchNodeConfig(ctx context.Context, node RuleNode, exists func(context.Context, uint) (bool, error)) error {
+ var cfg IPMatchConfig
+ if err := decodeNodeConfig(node, &cfg); err != nil {
+ return err
+ }
+ for _, raw := range cfg.IPs {
+ if _, err := netip.ParseAddr(raw); err != nil {
+ return fmt.Errorf("节点 %s 的 IP %s 无效", node.ID, raw)
}
- if cfg.StatusCode < 400 || cfg.StatusCode > 599 {
- return fmt.Errorf("节点 %s 的阻止状态码必须在 400-599 之间", node.ID)
+ }
+ for _, raw := range cfg.CIDRs {
+ if _, err := netip.ParsePrefix(raw); err != nil {
+ return fmt.Errorf("节点 %s 的 CIDR %s 无效", node.ID, raw)
}
- if len([]byte(cfg.ResponseBody)) > maxWAFBlockBodyBytes {
- return fmt.Errorf("节点 %s 的阻止响应体不能超过 %d 字节", node.ID, maxWAFBlockBodyBytes)
+ }
+ for _, id := range cfg.IPGroupIDs {
+ if err := validateIPGroupReference(ctx, node.ID, id, exists); err != nil {
+ return err
}
}
return nil
}
+func validateIPGroupReference(ctx context.Context, nodeID string, id uint, exists func(context.Context, uint) (bool, error)) error {
+ if id == 0 {
+ return fmt.Errorf("节点 %s 引用的 IP 组 ID 无效", nodeID)
+ }
+ if exists == nil {
+ return fmt.Errorf("节点 %s 无法校验 IP 组 %d", nodeID, id)
+ }
+ ok, err := exists(ctx, id)
+ if err != nil {
+ return fmt.Errorf("节点 %s 校验 IP 组 %d 失败: %w", nodeID, id, err)
+ }
+ if !ok {
+ return fmt.Errorf("节点 %s 引用的 IP 组 %d 不存在", nodeID, id)
+ }
+ return nil
+}
+
+func validateGeoMatchNodeConfig(node RuleNode) error {
+ var cfg GeoMatchConfig
+ if err := decodeNodeConfig(node, &cfg); err != nil {
+ return err
+ }
+ for _, code := range cfg.Countries {
+ if !countryCodePattern.MatchString(code) {
+ return fmt.Errorf("节点 %s 的国家代码 %s 无效", node.ID, code)
+ }
+ }
+ for _, code := range cfg.Regions {
+ if !regionCodePattern.MatchString(code) {
+ return fmt.Errorf("节点 %s 的地区代码 %s 无效", node.ID, code)
+ }
+ }
+ return nil
+}
+
+func validatePoWNodeConfig(node RuleNode) error {
+ var cfg PoWNodeConfig
+ if err := decodeNodeConfig(node, &cfg); err != nil {
+ return err
+ }
+ if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
+ return fmt.Errorf("节点 %s 的 PoW 难度必须在 1-16 之间", node.ID)
+ }
+ if cfg.Algorithm != powAlgorithmFast && cfg.Algorithm != powAlgorithmSlow {
+ return fmt.Errorf("节点 %s 的 PoW 算法必须为 fast 或 slow", node.ID)
+ }
+ if cfg.SessionTTL < minPoWSessionTTLSeconds {
+ return fmt.Errorf("节点 %s 的 PoW 会话 TTL 不能小于 60 秒", node.ID)
+ }
+ if cfg.ChallengeTTL < minPoWChallengeTTLSeconds {
+ return fmt.Errorf("节点 %s 的 PoW 挑战 TTL 不能小于 30 秒", node.ID)
+ }
+ return nil
+}
+
+func validateBlockNodeConfig(node RuleNode) error {
+ var cfg BlockNodeConfig
+ if err := decodeNodeConfig(node, &cfg); err != nil {
+ return err
+ }
+ if cfg.StatusCode < 400 || cfg.StatusCode > 599 {
+ return fmt.Errorf("节点 %s 的阻止状态码必须在 400-599 之间", node.ID)
+ }
+ if len([]byte(cfg.ResponseBody)) > maxWAFBlockBodyBytes {
+ return fmt.Errorf("节点 %s 的阻止响应体不能超过 %d 字节", node.ID, maxWAFBlockBodyBytes)
+ }
+ return nil
+}
+
+func decodeNodeConfig(node RuleNode, dst any) error {
+ if err := decodeStrictConfig(node.Config, dst); err != nil {
+ return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
+ }
+ return nil
+}
+
func decodeStrictConfig(raw json.RawMessage, dst any) error {
trimmed := bytes.TrimSpace(raw)
if bytes.Equal(trimmed, []byte("null")) {
diff --git a/internal/apps/openflare/waf/ip_group_sync.go b/internal/apps/openflare/waf/ip_group_sync.go
index 0b9acc37..b8979119 100644
--- a/internal/apps/openflare/waf/ip_group_sync.go
+++ b/internal/apps/openflare/waf/ip_group_sync.go
@@ -90,7 +90,7 @@ func syncOpenFlareWAFIPGroup(ctx context.Context, group *model.OpenFlareWAFIPGro
case wafIPGroupTypeAutomatic:
return syncIPGroupAutomatic(ctx, group, now)
default:
- return nil, errors.New("只有自动和订阅类型 IP 组支持同步")
+ return nil, &RuleValidationError{Err: errors.New("只有自动和订阅类型 IP 组支持同步")}
}
}
diff --git a/internal/apps/openflare/waf/logics.go b/internal/apps/openflare/waf/logics.go
index 4c5d0b84..53c115d3 100644
--- a/internal/apps/openflare/waf/logics.go
+++ b/internal/apps/openflare/waf/logics.go
@@ -13,7 +13,6 @@ import (
"sort"
"strings"
"time"
- "unicode"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -22,8 +21,7 @@ import (
)
const (
- defaultWAFBlockStatusCode = 418
- maxWAFBlockBodyBytes = 16 * 1024
+ maxWAFBlockBodyBytes = 16 * 1024
wafIPGroupTypeManual = "manual"
wafIPGroupTypeAutomatic = "automatic"
@@ -36,79 +34,19 @@ const (
defaultWAFIPGroupAutoLookbackMinutes = 60
minWAFIPGroupSyncIntervalMinutes = 5
maxWAFIPGroupSyncIntervalMinutes = 43200
-
- minPoWSessionTTLSeconds = 60
- minPoWChallengeTTLSeconds = 30
+ minPoWSessionTTLSeconds = 60
+ minPoWChallengeTTLSeconds = 30
+ powAlgorithmFast = "fast"
+ powAlgorithmSlow = "slow"
)
-// RuleGroupInput is the create/update payload for WAF rule groups.
-type RuleGroupInput struct {
- Name string `json:"name"`
- Enabled bool `json:"enabled"`
- BlockStatusCode int `json:"block_status_code"`
- BlockResponseBody string `json:"block_response_body"`
- IPWhitelist []string `json:"ip_whitelist"`
- IPBlacklist []string `json:"ip_blacklist"`
- IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"`
- IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"`
- CountryWhitelist []string `json:"country_whitelist"`
- CountryBlacklist []string `json:"country_blacklist"`
- RegionWhitelist []string `json:"region_whitelist"`
- RegionBlacklist []string `json:"region_blacklist"`
- PoWEnabled bool `json:"pow_enabled"`
- PoWConfig json.RawMessage `json:"pow_config"`
-}
-
-// PoWListConfig stores PoW whitelist/blacklist dimensions.
-type PoWListConfig struct {
- IPs []string `json:"ips"`
- IPCidrs []string `json:"ip_cidrs"`
- Paths []string `json:"paths"`
- PathRegexes []string `json:"path_regexes"`
- UserAgents []string `json:"user_agents"`
-}
-
-// PoWConfig stores proof-of-work settings for a rule group.
-type PoWConfig struct {
- Difficulty int `json:"difficulty"`
- Algorithm string `json:"algorithm"`
- SessionTTL int `json:"session_ttl"`
- ChallengeTTL int `json:"challenge_ttl"`
- Whitelist PoWListConfig `json:"whitelist"`
- Blacklist PoWListConfig `json:"blacklist"`
-}
-
-// RuleGroupView is the API view for a WAF rule group.
-type RuleGroupView struct {
- ID uint `json:"id"`
- Name string `json:"name"`
- Enabled bool `json:"enabled"`
- IsGlobal bool `json:"is_global"`
- BlockStatusCode int `json:"block_status_code"`
- BlockResponseBody string `json:"block_response_body"`
- IPWhitelist []string `json:"ip_whitelist"`
- IPBlacklist []string `json:"ip_blacklist"`
- IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"`
- IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"`
- CountryWhitelist []string `json:"country_whitelist"`
- CountryBlacklist []string `json:"country_blacklist"`
- RegionWhitelist []string `json:"region_whitelist"`
- RegionBlacklist []string `json:"region_blacklist"`
- PoWEnabled bool `json:"pow_enabled"`
- PoWConfig *PoWConfig `json:"pow_config"`
- AppliedSiteIDs []uint `json:"applied_site_ids"`
- AppliedSiteCount int `json:"applied_site_count"`
- CreatedAt string `json:"created_at"`
- UpdatedAt string `json:"updated_at"`
-}
-
// SiteRuleGroupsView is the site-level WAF binding view.
type SiteRuleGroupsView struct {
- RouteID uint `json:"route_id"`
- GlobalRuleGroup *RuleGroupView `json:"global_rule_group"`
- RuleGroups []RuleGroupView `json:"rule_groups"`
- AppliedRuleGroups []RuleGroupView `json:"applied_rule_groups"`
- AppliedIDs []uint `json:"applied_ids"`
+ RouteID uint `json:"route_id"`
+ GlobalRuleGroup *RuleView `json:"global_rule_group"`
+ RuleGroups []RuleView `json:"rule_groups"`
+ AppliedRuleGroups []RuleView `json:"applied_rule_groups"`
+ AppliedIDs []uint `json:"applied_ids"`
}
// IDsRequest carries a list of numeric ids.
@@ -197,120 +135,12 @@ type ipGroupExtIP struct {
CapturedAt time.Time `json:"captured_at"`
}
-var powAlgorithmValues = map[string]bool{"fast": true, "slow": true}
-
-// ListRuleGroups returns all WAF rule groups.
-func ListRuleGroups(ctx context.Context) ([]RuleGroupView, error) {
- if err := EnsureDefaultRuleGroup(ctx); err != nil {
- return nil, err
- }
- groups, err := model.ListOpenFlareWAFRuleGroups(ctx)
- if err != nil {
- return nil, err
- }
- bindings, err := loadRuleGroupBindings(ctx)
- if err != nil {
- return nil, err
- }
- views := make([]RuleGroupView, 0, len(groups))
- for _, group := range groups {
- view, buildErr := buildRuleGroupView(group, bindings[group.ID])
- if buildErr != nil {
- return nil, buildErr
- }
- views = append(views, view)
- }
- return views, nil
-}
-
-// GetRuleGroup returns a WAF rule group by id.
-func GetRuleGroup(ctx context.Context, id uint) (*RuleGroupView, error) {
- group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
- if err != nil {
- return nil, err
- }
- bindings, err := loadRuleGroupBindings(ctx)
- if err != nil {
- return nil, err
- }
- view, err := buildRuleGroupView(group, bindings[group.ID])
- if err != nil {
- return nil, err
- }
- return &view, nil
-}
-
-// CreateRuleGroup creates a custom WAF rule group.
-func CreateRuleGroup(ctx context.Context, input RuleGroupInput) (*RuleGroupView, error) {
- group, err := buildRuleGroup(ctx, nil, input)
- if err != nil {
- return nil, err
- }
- group.IsGlobal = false
- if err = model.CreateOpenFlareWAFRuleGroup(ctx, group); err != nil {
- return nil, err
- }
- return GetRuleGroup(ctx, group.ID)
-}
-
-// UpdateRuleGroup updates a WAF rule group.
-func UpdateRuleGroup(ctx context.Context, id uint, input RuleGroupInput) (*RuleGroupView, error) {
- group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
- if err != nil {
- return nil, err
- }
- isGlobal := group.IsGlobal
- group, err = buildRuleGroup(ctx, group, input)
- if err != nil {
- return nil, err
- }
- group.IsGlobal = isGlobal
- if isGlobal && strings.TrimSpace(group.Name) == "" {
- group.Name = "全局规则组"
- }
- if err = model.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil {
- return nil, err
- }
- return GetRuleGroup(ctx, group.ID)
-}
-
-// DeleteRuleGroup deletes a non-global WAF rule group.
-func DeleteRuleGroup(ctx context.Context, id uint) error {
- group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
- if err != nil {
- return err
- }
- if group.IsGlobal {
- return errors.New("全局 WAF 规则组不能删除")
- }
- return model.DeleteOpenFlareWAFRuleGroupWithBindings(ctx, group.ID)
-}
-
-// ReplaceRuleGroupSites replaces site bindings for a rule group.
-func ReplaceRuleGroupSites(ctx context.Context, groupID uint, routeIDs []uint) (*RuleGroupView, error) {
- group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, groupID)
- if err != nil {
- return nil, err
- }
- if group.IsGlobal {
- return nil, errors.New("全局 WAF 规则组默认应用到所有网站,不能手动绑定")
- }
- normalized, err := normalizeRouteIDs(ctx, routeIDs)
- if err != nil {
- return nil, err
- }
- if err = model.ReplaceOpenFlareWAFRuleGroupBindings(ctx, groupID, normalized); err != nil {
- return nil, err
- }
- return GetRuleGroup(ctx, groupID)
-}
-
// GetSiteRuleGroups returns WAF rule groups for a proxy route.
func GetSiteRuleGroups(ctx context.Context, routeID uint) (*SiteRuleGroupsView, error) {
if _, err := model.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil {
return nil, err
}
- groups, err := ListRuleGroups(ctx)
+ groups, err := ListRules(ctx)
if err != nil {
return nil, err
}
@@ -318,13 +148,10 @@ func GetSiteRuleGroups(ctx context.Context, routeID uint) (*SiteRuleGroupsView,
if err != nil {
return nil, err
}
- appliedSet := make(map[uint]struct{}, len(appliedIDs))
- for _, id := range appliedIDs {
- appliedSet[id] = struct{}{}
- }
- var global *RuleGroupView
- custom := make([]RuleGroupView, 0, len(groups))
- applied := make([]RuleGroupView, 0, len(appliedIDs))
+ var global *RuleView
+ custom := make([]RuleView, 0, len(groups))
+ applied := make([]RuleView, 0, len(appliedIDs))
+ groupByID := make(map[uint]RuleView, len(groups))
for index := range groups {
group := groups[index]
if group.IsGlobal {
@@ -333,7 +160,10 @@ func GetSiteRuleGroups(ctx context.Context, routeID uint) (*SiteRuleGroupsView,
continue
}
custom = append(custom, group)
- if _, ok := appliedSet[group.ID]; ok {
+ groupByID[group.ID] = group
+ }
+ for _, id := range appliedIDs {
+ if group, ok := groupByID[id]; ok {
applied = append(applied, group)
}
}
@@ -353,7 +183,7 @@ func ReplaceSiteRuleGroups(ctx context.Context, routeID uint, groupIDs []uint) (
}
normalized, err := normalizeRuleGroupIDs(ctx, groupIDs)
if err != nil {
- return nil, err
+ return nil, &RuleValidationError{Err: err}
}
if err = model.ReplaceOpenFlareWAFSiteRuleGroupBindings(ctx, routeID, normalized); err != nil {
return nil, err
@@ -383,22 +213,12 @@ func EnsureDefaultRuleGroup(ctx context.Context) error {
if !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
+ graph, marshalErr := json.Marshal(DefaultRuleGraph())
+ if marshalErr != nil {
+ return marshalErr
+ }
group := &model.OpenFlareWAFRuleGroup{
- Name: "全局规则组",
- Enabled: true,
- IsGlobal: true,
- BlockStatusCode: defaultWAFBlockStatusCode,
- IPWhitelist: "[]",
- IPBlacklist: "[]",
- IPWhitelistGroups: "[]",
- IPBlacklistGroups: "[]",
- CountryWhitelist: "[]",
- CountryBlacklist: "[]",
- RegionWhitelist: "[]",
- RegionBlacklist: "[]",
- PoWEnabled: false,
- PoWConfig: "{}",
- BlockResponseBody: "",
+ Name: "全局规则组", Enabled: true, IsGlobal: true, Graph: string(graph), Revision: 1,
}
return model.CreateOpenFlareWAFRuleGroup(ctx, group)
}
@@ -445,7 +265,7 @@ func GetIPGroup(ctx context.Context, id uint) (*IPGroupView, error) {
func CreateIPGroup(ctx context.Context, input IPGroupInput) (*IPGroupView, error) {
group, err := buildIPGroup(nil, input)
if err != nil {
- return nil, err
+ return nil, &RuleValidationError{Err: err}
}
if err = model.CreateOpenFlareWAFIPGroup(ctx, group); err != nil {
return nil, err
@@ -462,7 +282,7 @@ func UpdateIPGroup(ctx context.Context, id uint, input IPGroupInput) (*IPGroupVi
}
group, err = buildIPGroup(group, input)
if err != nil {
- return nil, err
+ return nil, &RuleValidationError{Err: err}
}
if err = model.UpdateOpenFlareWAFIPGroup(ctx, group); err != nil {
return nil, err
@@ -482,7 +302,7 @@ func DeleteIPGroup(ctx context.Context, id uint) error {
return err
}
if counts[group.ID] > 0 {
- return errors.New("IP 组已被 WAF 规则组引用,请先移除引用")
+ return &RuleValidationError{Err: errors.New("IP 组已被 WAF 规则引用,请先移除引用")}
}
return model.DeleteOpenFlareWAFIPGroup(ctx, group.ID)
}
@@ -500,7 +320,7 @@ func SyncIPGroup(ctx context.Context, id uint) (*IPGroupSyncResult, error) {
func TestIPGroupAutoConfig(ctx context.Context, input IPGroupAutoTestInput) (*IPGroupAutoTestResult, error) {
config, err := parseIPGroupAutoConfig(input.AutoConfig)
if err != nil {
- return nil, err
+ return nil, &RuleValidationError{Err: err}
}
now := time.Now().UTC()
ips, err := evaluateParsedIPGroupAutoConfig(ctx, config, now)
@@ -516,131 +336,259 @@ func TestIPGroupAutoConfig(ctx context.Context, input IPGroupAutoTestInput) (*IP
}, nil
}
-func buildRuleGroup(ctx context.Context, group *model.OpenFlareWAFRuleGroup, input RuleGroupInput) (*model.OpenFlareWAFRuleGroup, error) {
- name := strings.TrimSpace(input.Name)
- if name == "" {
- return nil, errors.New("规则组名称不能为空")
- }
- statusCode := input.BlockStatusCode
- if statusCode == 0 {
- statusCode = defaultWAFBlockStatusCode
- }
- if statusCode < 400 || statusCode > 599 {
- return nil, errors.New("拦截状态码必须在 400-599 之间")
- }
- if len([]byte(input.BlockResponseBody)) > maxWAFBlockBodyBytes {
- return nil, fmt.Errorf("拦截页面内容不能超过 %d 字节", maxWAFBlockBodyBytes)
- }
- ipWhitelist, err := normalizeIPList(input.IPWhitelist)
- if err != nil {
- return nil, fmt.Errorf("IP 白名单无效: %w", err)
- }
- ipBlacklist, err := normalizeIPList(input.IPBlacklist)
- if err != nil {
- return nil, fmt.Errorf("IP 黑名单无效: %w", err)
- }
- ipWhitelistGroups, err := normalizeIPGroupIDs(ctx, input.IPWhitelistGroups)
- if err != nil {
- return nil, fmt.Errorf("IP 白名单引用无效: %w", err)
- }
- ipBlacklistGroups, err := normalizeIPGroupIDs(ctx, input.IPBlacklistGroups)
- if err != nil {
- return nil, fmt.Errorf("IP 黑名单引用无效: %w", err)
- }
- countryWhitelist, err := normalizeCountryList(input.CountryWhitelist)
- if err != nil {
- return nil, fmt.Errorf("地域白名单无效: %w", err)
- }
- countryBlacklist, err := normalizeCountryList(input.CountryBlacklist)
- if err != nil {
- return nil, fmt.Errorf("地域黑名单无效: %w", err)
- }
- regionWhitelist := normalizeStringList(input.RegionWhitelist)
- regionBlacklist := normalizeStringList(input.RegionBlacklist)
- powConfigRaw := strings.TrimSpace(string(input.PoWConfig))
- if powConfigRaw == "" {
- powConfigRaw = "{}"
- }
- powConfig, err := normalizePoWConfig(input.PoWEnabled, powConfigRaw)
+func loadRuleGroupBindings(ctx context.Context) (map[uint][]uint, error) {
+ bindings, err := model.ListOpenFlareWAFRuleGroupBindings(ctx)
if err != nil {
return nil, err
}
- powConfigJSON, _ := json.Marshal(powConfig)
-
- ipWhitelistJSON, _ := json.Marshal(ipWhitelist)
- ipBlacklistJSON, _ := json.Marshal(ipBlacklist)
- ipWhitelistGroupsJSON, _ := json.Marshal(ipWhitelistGroups)
- ipBlacklistGroupsJSON, _ := json.Marshal(ipBlacklistGroups)
- countryWhitelistJSON, _ := json.Marshal(countryWhitelist)
- countryBlacklistJSON, _ := json.Marshal(countryBlacklist)
- regionWhitelistJSON, _ := json.Marshal(regionWhitelist)
- regionBlacklistJSON, _ := json.Marshal(regionBlacklist)
-
- if group == nil {
- group = &model.OpenFlareWAFRuleGroup{}
+ result := make(map[uint][]uint, len(bindings))
+ for _, binding := range bindings {
+ result[binding.RuleGroupID] = append(result[binding.RuleGroupID], binding.ProxyRouteID)
}
- group.Name = name
- group.Enabled = input.Enabled
- group.BlockStatusCode = statusCode
- group.BlockResponseBody = input.BlockResponseBody
- group.IPWhitelist = string(ipWhitelistJSON)
- group.IPBlacklist = string(ipBlacklistJSON)
- group.IPWhitelistGroups = string(ipWhitelistGroupsJSON)
- group.IPBlacklistGroups = string(ipBlacklistGroupsJSON)
- group.CountryWhitelist = string(countryWhitelistJSON)
- group.CountryBlacklist = string(countryBlacklistJSON)
- group.RegionWhitelist = string(regionWhitelistJSON)
- group.RegionBlacklist = string(regionBlacklistJSON)
- group.PoWEnabled = input.PoWEnabled
- group.PoWConfig = string(powConfigJSON)
- return group, nil
+ return result, nil
}
-func buildRuleGroupView(group *model.OpenFlareWAFRuleGroup, appliedSiteIDs []uint) (RuleGroupView, error) {
- if group == nil {
- return RuleGroupView{}, errors.New("waf rule group is nil")
+func loadIPGroupReferenceCounts(ctx context.Context) (map[uint]int, error) {
+ groups, err := model.ListOpenFlareWAFRuleGroups(ctx)
+ if err != nil {
+ return nil, err
}
- sort.Slice(appliedSiteIDs, func(i, j int) bool { return appliedSiteIDs[i] < appliedSiteIDs[j] })
- view := RuleGroupView{
- ID: group.ID,
- Name: group.Name,
- Enabled: group.Enabled,
- IsGlobal: group.IsGlobal,
- BlockStatusCode: group.BlockStatusCode,
- BlockResponseBody: group.BlockResponseBody,
- PoWEnabled: group.PoWEnabled,
- AppliedSiteIDs: appliedSiteIDs,
- AppliedSiteCount: len(appliedSiteIDs),
- CreatedAt: group.CreatedAt.Format(time.RFC3339),
- UpdatedAt: group.UpdatedAt.Format(time.RFC3339),
+ counts := make(map[uint]int)
+ for _, group := range groups {
+ graph := DefaultRuleGraph()
+ if strings.TrimSpace(group.Graph) != "" {
+ if err = json.Unmarshal([]byte(group.Graph), &graph); err != nil {
+ return nil, fmt.Errorf("decode WAF rule %d graph: %w", group.ID, err)
+ }
+ }
+ for _, id := range ReferencedIPGroupIDs(graph) {
+ counts[id]++
+ }
}
- var err error
- if view.IPWhitelist, err = decodeStringList(group.IPWhitelist); err != nil {
- return view, err
- }
- if view.IPBlacklist, err = decodeStringList(group.IPBlacklist); err != nil {
- return view, err
- }
- view.IPWhitelistGroups = mustDecodeUintList(group.IPWhitelistGroups)
- view.IPBlacklistGroups = mustDecodeUintList(group.IPBlacklistGroups)
- if view.CountryWhitelist, err = decodeStringList(group.CountryWhitelist); err != nil {
- return view, err
- }
- if view.CountryBlacklist, err = decodeStringList(group.CountryBlacklist); err != nil {
- return view, err
- }
- if view.RegionWhitelist, err = decodeStringList(group.RegionWhitelist); err != nil {
- return view, err
- }
- if view.RegionBlacklist, err = decodeStringList(group.RegionBlacklist); err != nil {
- return view, err
- }
- if view.PoWConfig, err = decodeStoredPoWConfig(group.PoWEnabled, group.PoWConfig); err != nil {
- return view, err
- }
- return view, nil
+ return counts, nil
}
+func pruneIPGroupExtIPs(group *model.OpenFlareWAFIPGroup, ipList []string) error {
+ if group == nil {
+ return nil
+ }
+ allowed := make(map[string]struct{}, len(ipList))
+ for _, ip := range ipList {
+ allowed[ip] = struct{}{}
+ }
+ var extIPs []ipGroupExtIP
+ if group.ExtIPs != "" && group.ExtIPs != "[]" {
+ if err := json.Unmarshal([]byte(group.ExtIPs), &extIPs); err != nil {
+ return err
+ }
+ }
+ pruned := make([]ipGroupExtIP, 0, len(extIPs))
+ for _, extIP := range extIPs {
+ if _, ok := allowed[extIP.IP]; ok {
+ pruned = append(pruned, extIP)
+ }
+ }
+ extIPsJSON, err := json.Marshal(pruned)
+ if err != nil {
+ return err
+ }
+ group.ExtIPs = string(extIPsJSON)
+ return nil
+}
+
+func normalizeIPList(items []string) ([]string, error) {
+ normalized := make([]string, 0, len(items))
+ for _, raw := range items {
+ item := strings.TrimSpace(raw)
+ if item == "" {
+ continue
+ }
+ if strings.Contains(item, "/") {
+ prefix, err := netip.ParsePrefix(item)
+ if err != nil {
+ return nil, fmt.Errorf("%s 不是合法 IP 段", item)
+ }
+ item = prefix.Masked().String()
+ } else {
+ addr, err := netip.ParseAddr(item)
+ if err != nil {
+ return nil, fmt.Errorf("%s 不是合法 IP", item)
+ }
+ item = addr.String()
+ }
+ normalized = append(normalized, item)
+ }
+ normalized = uniqueStrings(normalized)
+ sort.Strings(normalized)
+ return normalized, nil
+}
+
+func decodeStringList(raw string) ([]string, error) {
+ text := strings.TrimSpace(raw)
+ if text == "" {
+ return []string{}, nil
+ }
+ var items []string
+ if err := json.Unmarshal([]byte(text), &items); err != nil {
+ return nil, err
+ }
+ return items, nil
+}
+
+func normalizeRuleGroupIDs(ctx context.Context, groupIDs []uint) ([]uint, error) {
+ normalized := uniqueUintIDsInOrder(groupIDs)
+ for _, groupID := range normalized {
+ group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, groupID)
+ if err != nil {
+ return nil, fmt.Errorf("WAF 规则组 %d 不存在", groupID)
+ }
+ if group.IsGlobal {
+ return nil, errors.New("全局 WAF 规则组不需要手动绑定")
+ }
+ }
+ return normalized, nil
+}
+
+func uniqueUintIDsInOrder(ids []uint) []uint {
+ seen := make(map[uint]struct{}, len(ids))
+ result := make([]uint, 0, len(ids))
+ for _, id := range ids {
+ if id == 0 {
+ continue
+ }
+ if _, exists := seen[id]; exists {
+ continue
+ }
+ seen[id] = struct{}{}
+ result = append(result, id)
+ }
+ return result
+}
+
+func uniqueStrings(items []string) []string {
+ seen := make(map[string]struct{}, len(items))
+ result := make([]string, 0, len(items))
+ for _, item := range items {
+ if _, ok := seen[item]; ok {
+ continue
+ }
+ seen[item] = struct{}{}
+ result = append(result, item)
+ }
+ return result
+}
+
+func normalizeIPGroupAutoConfig(raw json.RawMessage) (string, error) {
+ text := strings.TrimSpace(string(raw))
+ if text == "" {
+ text = "{}"
+ }
+ config, err := parseIPGroupAutoConfig(json.RawMessage(text))
+ if err != nil {
+ return "", err
+ }
+ normalized, _ := json.Marshal(config)
+ return string(normalized), nil
+}
+
+func parseIPGroupAutoConfig(raw json.RawMessage) (ipGroupAutoConfig, error) {
+ text := strings.TrimSpace(string(raw))
+ if text == "" {
+ text = "{}"
+ }
+ var config ipGroupAutoConfig
+ if err := json.Unmarshal([]byte(text), &config); err != nil {
+ return ipGroupAutoConfig{}, errors.New("自动 IP 组配置必须是 JSON 对象")
+ }
+ var object map[string]any
+ if err := json.Unmarshal([]byte(text), &object); err != nil || object == nil {
+ return ipGroupAutoConfig{}, errors.New("自动 IP 组配置必须是 JSON 对象")
+ }
+ if config.LookbackMinutes <= 0 {
+ config.LookbackMinutes = defaultWAFIPGroupAutoLookbackMinutes
+ }
+ if config.LookbackMinutes < minWAFIPGroupSyncIntervalMinutes {
+ config.LookbackMinutes = minWAFIPGroupSyncIntervalMinutes
+ }
+ if config.LookbackMinutes > maxWAFIPGroupSyncIntervalMinutes {
+ config.LookbackMinutes = maxWAFIPGroupSyncIntervalMinutes
+ }
+ if config.TTL == 0 {
+ config.TTL = -1
+ }
+ if config.Rules == nil {
+ config.Rules = []ipGroupAutoRule{}
+ }
+ for i, rule := range config.Rules {
+ rule.Name = strings.TrimSpace(rule.Name)
+ rule.Expr = strings.TrimSpace(rule.Expr)
+ if rule.Expr == "" {
+ return ipGroupAutoConfig{}, fmt.Errorf("自动规则 %d 的 Expr 表达式不能为空", i+1)
+ }
+ if _, err := exprlang.Compile(rule.Expr, exprlang.Env(ipGroupAutoRuleEnv{}), exprlang.AsBool()); err != nil {
+ return ipGroupAutoConfig{}, fmt.Errorf("自动规则 %s Expr 无效: %w", displayIPGroupAutoRuleName(rule, i), err)
+ }
+ config.Rules[i] = rule
+ }
+ return config, nil
+}
+
+func validateSubscriptionURL(rawURL string) error {
+ parsed, err := url.Parse(strings.TrimSpace(rawURL))
+ if err != nil || parsed.Host == "" {
+ return errors.New("订阅 URL 无效")
+ }
+ if parsed.Scheme != "http" && parsed.Scheme != "https" {
+ return errors.New("订阅 URL 仅支持 http 或 https")
+ }
+ return nil
+}
+
+func normalizeIPGroupType(value string) string {
+ switch strings.TrimSpace(value) {
+ case wafIPGroupTypeManual, "":
+ return wafIPGroupTypeManual
+ case wafIPGroupTypeAutomatic:
+ return wafIPGroupTypeAutomatic
+ case wafIPGroupTypeSubscription:
+ return wafIPGroupTypeSubscription
+ default:
+ return ""
+ }
+}
+
+func normalizeIPGroupSubscriptionFormat(value string) string {
+ switch strings.TrimSpace(value) {
+ case wafIPGroupSubscriptionFormatJSON:
+ return wafIPGroupSubscriptionFormatJSON
+ default:
+ return wafIPGroupSubscriptionFormatText
+ }
+}
+
+func normalizeIPGroupSyncInterval(value int) int {
+ if value <= 0 {
+ return defaultWAFIPGroupSyncIntervalMinutes
+ }
+ if value < minWAFIPGroupSyncIntervalMinutes {
+ return minWAFIPGroupSyncIntervalMinutes
+ }
+ if value > maxWAFIPGroupSyncIntervalMinutes {
+ return maxWAFIPGroupSyncIntervalMinutes
+ }
+ return value
+}
+
+func nextIPGroupSyncAt(groupType string, enabled bool, interval int, current *time.Time) *time.Time {
+ if (groupType != wafIPGroupTypeSubscription && groupType != wafIPGroupTypeAutomatic) || !enabled {
+ return nil
+ }
+ if current != nil && current.After(time.Now().UTC()) {
+ return current
+ }
+ next := time.Now().UTC().Add(time.Duration(normalizeIPGroupSyncInterval(interval)) * time.Minute)
+ return &next
+}
func buildIPGroup(group *model.OpenFlareWAFIPGroup, input IPGroupInput) (*model.OpenFlareWAFIPGroup, error) {
name := strings.TrimSpace(input.Name)
if name == "" {
@@ -755,388 +703,3 @@ func buildIPGroupView(group *model.OpenFlareWAFIPGroup, referenceCount int) (IPG
}
return view, nil
}
-
-func loadRuleGroupBindings(ctx context.Context) (map[uint][]uint, error) {
- bindings, err := model.ListOpenFlareWAFRuleGroupBindings(ctx)
- if err != nil {
- return nil, err
- }
- result := make(map[uint][]uint, len(bindings))
- for _, binding := range bindings {
- result[binding.RuleGroupID] = append(result[binding.RuleGroupID], binding.ProxyRouteID)
- }
- return result, nil
-}
-
-func loadIPGroupReferenceCounts(ctx context.Context) (map[uint]int, error) {
- groups, err := model.ListOpenFlareWAFRuleGroups(ctx)
- if err != nil {
- return nil, err
- }
- counts := make(map[uint]int)
- for _, group := range groups {
- for _, id := range mustDecodeUintList(group.IPWhitelistGroups) {
- counts[id]++
- }
- for _, id := range mustDecodeUintList(group.IPBlacklistGroups) {
- counts[id]++
- }
- }
- return counts, nil
-}
-
-func pruneIPGroupExtIPs(group *model.OpenFlareWAFIPGroup, ipList []string) error {
- if group == nil {
- return nil
- }
- allowed := make(map[string]struct{}, len(ipList))
- for _, ip := range ipList {
- allowed[ip] = struct{}{}
- }
- var extIPs []ipGroupExtIP
- if group.ExtIPs != "" && group.ExtIPs != "[]" {
- if err := json.Unmarshal([]byte(group.ExtIPs), &extIPs); err != nil {
- return err
- }
- }
- pruned := make([]ipGroupExtIP, 0, len(extIPs))
- for _, extIP := range extIPs {
- if _, ok := allowed[extIP.IP]; ok {
- pruned = append(pruned, extIP)
- }
- }
- extIPsJSON, err := json.Marshal(pruned)
- if err != nil {
- return err
- }
- group.ExtIPs = string(extIPsJSON)
- return nil
-}
-
-func normalizeIPList(items []string) ([]string, error) {
- normalized := make([]string, 0, len(items))
- for _, raw := range items {
- item := strings.TrimSpace(raw)
- if item == "" {
- continue
- }
- if strings.Contains(item, "/") {
- prefix, err := netip.ParsePrefix(item)
- if err != nil {
- return nil, fmt.Errorf("%s 不是合法 IP 段", item)
- }
- item = prefix.Masked().String()
- } else {
- addr, err := netip.ParseAddr(item)
- if err != nil {
- return nil, fmt.Errorf("%s 不是合法 IP", item)
- }
- item = addr.String()
- }
- normalized = append(normalized, item)
- }
- normalized = uniqueStrings(normalized)
- sort.Strings(normalized)
- return normalized, nil
-}
-
-func normalizeCountryList(items []string) ([]string, error) {
- normalized := make([]string, 0, len(items))
- for _, raw := range items {
- item := strings.ToUpper(strings.TrimSpace(raw))
- if item == "" {
- continue
- }
- if len(item) != 2 || !unicode.IsLetter(rune(item[0])) || !unicode.IsLetter(rune(item[1])) {
- return nil, fmt.Errorf("%s 不是合法国家代码", item)
- }
- normalized = append(normalized, item)
- }
- normalized = uniqueStrings(normalized)
- sort.Strings(normalized)
- return normalized, nil
-}
-
-func normalizeStringList(items []string) []string {
- normalized := make([]string, 0, len(items))
- for _, raw := range items {
- item := strings.TrimSpace(raw)
- if item == "" {
- continue
- }
- normalized = append(normalized, item)
- }
- normalized = uniqueStrings(normalized)
- sort.Strings(normalized)
- return normalized
-}
-
-func decodeStringList(raw string) ([]string, error) {
- text := strings.TrimSpace(raw)
- if text == "" {
- return []string{}, nil
- }
- var items []string
- if err := json.Unmarshal([]byte(text), &items); err != nil {
- return nil, err
- }
- return items, nil
-}
-
-func normalizeRouteIDs(ctx context.Context, routeIDs []uint) ([]uint, error) {
- normalized := uniqueUintIDs(routeIDs)
- for _, routeID := range normalized {
- if _, err := model.GetOpenFlareProxyRouteByID(ctx, routeID); err != nil {
- return nil, fmt.Errorf("网站 %d 不存在", routeID)
- }
- }
- return normalized, nil
-}
-
-func normalizeRuleGroupIDs(ctx context.Context, groupIDs []uint) ([]uint, error) {
- normalized := uniqueUintIDs(groupIDs)
- for _, groupID := range normalized {
- group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, groupID)
- if err != nil {
- return nil, fmt.Errorf("WAF 规则组 %d 不存在", groupID)
- }
- if group.IsGlobal {
- return nil, errors.New("全局 WAF 规则组不需要手动绑定")
- }
- }
- return normalized, nil
-}
-
-func normalizeIPGroupIDs(ctx context.Context, ids []uint) ([]uint, error) {
- normalized := uniqueUintIDs(ids)
- for _, id := range normalized {
- if _, err := model.GetOpenFlareWAFIPGroupByID(ctx, id); err != nil {
- if errors.Is(err, gorm.ErrRecordNotFound) {
- return nil, fmt.Errorf("IP 组 %d 不存在", id)
- }
- return nil, err
- }
- }
- return normalized, nil
-}
-
-func uniqueUintIDs(ids []uint) []uint {
- normalized := make([]uint, 0, len(ids))
- for _, id := range ids {
- if id == 0 {
- continue
- }
- normalized = append(normalized, id)
- }
- normalized = uniqueUints(normalized)
- sort.Slice(normalized, func(i, j int) bool { return normalized[i] < normalized[j] })
- return normalized
-}
-
-func uniqueUints(items []uint) []uint {
- seen := make(map[uint]struct{}, len(items))
- result := make([]uint, 0, len(items))
- for _, item := range items {
- if _, ok := seen[item]; ok {
- continue
- }
- seen[item] = struct{}{}
- result = append(result, item)
- }
- return result
-}
-
-func uniqueStrings(items []string) []string {
- seen := make(map[string]struct{}, len(items))
- result := make([]string, 0, len(items))
- for _, item := range items {
- if _, ok := seen[item]; ok {
- continue
- }
- seen[item] = struct{}{}
- result = append(result, item)
- }
- return result
-}
-
-func mustDecodeUintList(raw string) []uint {
- var values []uint
- if err := json.Unmarshal([]byte(strings.TrimSpace(raw)), &values); err != nil {
- return []uint{}
- }
- values = uniqueUintIDs(values)
- sort.Slice(values, func(i, j int) bool { return values[i] < values[j] })
- return values
-}
-
-func defaultPoWConfig() PoWConfig {
- return PoWConfig{
- Difficulty: 4,
- Algorithm: "fast",
- SessionTTL: 600,
- ChallengeTTL: 300,
- Whitelist: PoWListConfig{IPs: []string{}, IPCidrs: []string{}, Paths: []string{}, PathRegexes: []string{}, UserAgents: []string{}},
- Blacklist: PoWListConfig{IPs: []string{}, IPCidrs: []string{}, Paths: []string{}, PathRegexes: []string{}, UserAgents: []string{}},
- }
-}
-
-func normalizePoWConfig(enabled bool, raw string) (PoWConfig, error) {
- cfg, err := parsePoWConfigRaw(enabled, raw)
- if err != nil {
- return cfg, err
- }
- if err := validatePoWCoreSettings(cfg); err != nil {
- return cfg, err
- }
- if err := validatePoWCIDRs(cfg.Whitelist.IPCidrs, "白名单"); err != nil {
- return cfg, err
- }
- if err := validatePoWCIDRs(cfg.Blacklist.IPCidrs, "黑名单"); err != nil {
- return cfg, err
- }
- if err := validatePoWPathRegexes(cfg.Whitelist.PathRegexes, "白名单"); err != nil {
- return cfg, err
- }
- if err := validatePoWPathRegexes(cfg.Blacklist.PathRegexes, "黑名单"); err != nil {
- return cfg, err
- }
- if err := validatePoWIPs(cfg.Whitelist.IPs, "白名单"); err != nil {
- return cfg, err
- }
- if err := validatePoWIPs(cfg.Blacklist.IPs, "黑名单"); err != nil {
- return cfg, err
- }
- if err := validatePoWListMutualExclusion(cfg); err != nil {
- return cfg, err
- }
- return cfg, nil
-}
-
-func decodeStoredPoWConfig(enabled bool, raw string) (*PoWConfig, error) {
- if !enabled {
- cfg := defaultPoWConfig()
- return &cfg, nil
- }
- text := strings.TrimSpace(raw)
- if text == "" || text == "{}" {
- cfg := defaultPoWConfig()
- return &cfg, nil
- }
- var cfg PoWConfig
- if err := json.Unmarshal([]byte(text), &cfg); err != nil {
- return nil, errors.New("pow_config 格式无效")
- }
- return &cfg, nil
-}
-
-func normalizeIPGroupAutoConfig(raw json.RawMessage) (string, error) {
- text := strings.TrimSpace(string(raw))
- if text == "" {
- text = "{}"
- }
- config, err := parseIPGroupAutoConfig(json.RawMessage(text))
- if err != nil {
- return "", err
- }
- normalized, _ := json.Marshal(config)
- return string(normalized), nil
-}
-
-func parseIPGroupAutoConfig(raw json.RawMessage) (ipGroupAutoConfig, error) {
- text := strings.TrimSpace(string(raw))
- if text == "" {
- text = "{}"
- }
- var config ipGroupAutoConfig
- if err := json.Unmarshal([]byte(text), &config); err != nil {
- return ipGroupAutoConfig{}, errors.New("自动 IP 组配置必须是 JSON 对象")
- }
- var object map[string]any
- if err := json.Unmarshal([]byte(text), &object); err != nil || object == nil {
- return ipGroupAutoConfig{}, errors.New("自动 IP 组配置必须是 JSON 对象")
- }
- if config.LookbackMinutes <= 0 {
- config.LookbackMinutes = defaultWAFIPGroupAutoLookbackMinutes
- }
- if config.LookbackMinutes < minWAFIPGroupSyncIntervalMinutes {
- config.LookbackMinutes = minWAFIPGroupSyncIntervalMinutes
- }
- if config.LookbackMinutes > maxWAFIPGroupSyncIntervalMinutes {
- config.LookbackMinutes = maxWAFIPGroupSyncIntervalMinutes
- }
- if config.TTL == 0 {
- config.TTL = -1
- }
- if config.Rules == nil {
- config.Rules = []ipGroupAutoRule{}
- }
- for i, rule := range config.Rules {
- rule.Name = strings.TrimSpace(rule.Name)
- rule.Expr = strings.TrimSpace(rule.Expr)
- if rule.Expr == "" {
- return ipGroupAutoConfig{}, fmt.Errorf("自动规则 %d 的 Expr 表达式不能为空", i+1)
- }
- if _, err := exprlang.Compile(rule.Expr, exprlang.Env(ipGroupAutoRuleEnv{}), exprlang.AsBool()); err != nil {
- return ipGroupAutoConfig{}, fmt.Errorf("自动规则 %s Expr 无效: %w", displayIPGroupAutoRuleName(rule, i), err)
- }
- config.Rules[i] = rule
- }
- return config, nil
-}
-
-func validateSubscriptionURL(rawURL string) error {
- parsed, err := url.Parse(strings.TrimSpace(rawURL))
- if err != nil || parsed.Host == "" {
- return errors.New("订阅 URL 无效")
- }
- if parsed.Scheme != "http" && parsed.Scheme != "https" {
- return errors.New("订阅 URL 仅支持 http 或 https")
- }
- return nil
-}
-
-func normalizeIPGroupType(value string) string {
- switch strings.TrimSpace(value) {
- case wafIPGroupTypeManual, "":
- return wafIPGroupTypeManual
- case wafIPGroupTypeAutomatic:
- return wafIPGroupTypeAutomatic
- case wafIPGroupTypeSubscription:
- return wafIPGroupTypeSubscription
- default:
- return ""
- }
-}
-
-func normalizeIPGroupSubscriptionFormat(value string) string {
- switch strings.TrimSpace(value) {
- case wafIPGroupSubscriptionFormatJSON:
- return wafIPGroupSubscriptionFormatJSON
- default:
- return wafIPGroupSubscriptionFormatText
- }
-}
-
-func normalizeIPGroupSyncInterval(value int) int {
- if value <= 0 {
- return defaultWAFIPGroupSyncIntervalMinutes
- }
- if value < minWAFIPGroupSyncIntervalMinutes {
- return minWAFIPGroupSyncIntervalMinutes
- }
- if value > maxWAFIPGroupSyncIntervalMinutes {
- return maxWAFIPGroupSyncIntervalMinutes
- }
- return value
-}
-
-func nextIPGroupSyncAt(groupType string, enabled bool, interval int, current *time.Time) *time.Time {
- if (groupType != wafIPGroupTypeSubscription && groupType != wafIPGroupTypeAutomatic) || !enabled {
- return nil
- }
- if current != nil && current.After(time.Now().UTC()) {
- return current
- }
- next := time.Now().UTC().Add(time.Duration(normalizeIPGroupSyncInterval(interval)) * time.Minute)
- return &next
-}
diff --git a/internal/apps/openflare/waf/logics_test.go b/internal/apps/openflare/waf/logics_test.go
index 8a5618f3..c9ad8c02 100644
--- a/internal/apps/openflare/waf/logics_test.go
+++ b/internal/apps/openflare/waf/logics_test.go
@@ -26,6 +26,7 @@ func setupWAFTestDB(t *testing.T) func() {
&model.OpenFlareWAFRuleGroup{},
&model.OpenFlareWAFIPGroup{},
&model.OpenFlareWAFRuleGroupBinding{},
+ &model.OriginProxyRoute{},
))
db.SetDB(sqliteDB)
@@ -34,38 +35,6 @@ func setupWAFTestDB(t *testing.T) func() {
}
}
-func TestCreateRuleGroup(t *testing.T) {
- cleanup := setupWAFTestDB(t)
- defer cleanup()
- ctx := context.Background()
-
- group, err := CreateRuleGroup(ctx, RuleGroupInput{
- Name: "edge guard",
- Enabled: true,
- BlockStatusCode: 451,
- IPWhitelist: []string{" 192.0.2.1 ", "192.0.2.1", "198.51.100.0/24"},
- IPBlacklist: []string{"203.0.113.10"},
- CountryBlacklist: []string{" cn ", "CN", "us"},
- })
- require.NoError(t, err)
- assert.NotZero(t, group.ID)
- assert.False(t, group.IsGlobal)
- assert.Equal(t, "edge guard", group.Name)
- require.Len(t, group.IPWhitelist, 2)
- assert.Equal(t, "192.0.2.1", group.IPWhitelist[0])
- assert.Equal(t, "198.51.100.0/24", group.IPWhitelist[1])
- require.Len(t, group.CountryBlacklist, 2)
- assert.Equal(t, "CN", group.CountryBlacklist[0])
- assert.Equal(t, "US", group.CountryBlacklist[1])
-
- _, err = CreateRuleGroup(ctx, RuleGroupInput{
- Name: "bad ip",
- Enabled: true,
- IPBlacklist: []string{"not-an-ip"},
- })
- require.Error(t, err)
-}
-
func TestPruneIPGroupExtIPs(t *testing.T) {
group := &model.OpenFlareWAFIPGroup{
ExtIPs: `[{"ip":"203.0.113.10","captured_at":"2026-06-18T10:00:00Z"},{"ip":"203.0.113.11","captured_at":"2026-06-18T11:00:00Z"}]`,
diff --git a/internal/apps/openflare/waf/pow_helpers.go b/internal/apps/openflare/waf/pow_helpers.go
deleted file mode 100644
index 242100df..00000000
--- a/internal/apps/openflare/waf/pow_helpers.go
+++ /dev/null
@@ -1,92 +0,0 @@
-// Copyright 2026 Arctel.net
-// SPDX-License-Identifier: Apache-2.0
-
-package waf
-
-import (
- "encoding/json"
- "errors"
- "fmt"
- "net"
- "regexp"
- "strings"
-)
-
-func parsePoWConfigRaw(enabled bool, raw string) (PoWConfig, error) {
- if !enabled {
- return defaultPoWConfig(), nil
- }
- cfg := defaultPoWConfig()
- text := strings.TrimSpace(raw)
- if text == "" || text == "{}" {
- return cfg, nil
- }
- if err := json.Unmarshal([]byte(text), &cfg); err != nil {
- return cfg, errors.New("pow_config 格式无效")
- }
- return cfg, nil
-}
-
-func validatePoWCoreSettings(cfg PoWConfig) error {
- if cfg.Difficulty < 1 || cfg.Difficulty > 16 {
- return errors.New("pow_config.difficulty 必须在 1-16 之间")
- }
- if !powAlgorithmValues[cfg.Algorithm] {
- return errors.New("pow_config.algorithm 必须为 fast 或 slow")
- }
- if cfg.SessionTTL < minPoWSessionTTLSeconds {
- return errors.New("pow_config.session_ttl 不能小于 60 秒")
- }
- if cfg.ChallengeTTL < minPoWChallengeTTLSeconds {
- return errors.New("pow_config.challenge_ttl 不能小于 30 秒")
- }
- return nil
-}
-
-func validatePoWCIDRs(cidrs []string, listName string) error {
- for _, cidr := range cidrs {
- if _, _, err := net.ParseCIDR(cidr); err != nil {
- return fmt.Errorf("pow_config %s IP CIDR 格式无效: %s", listName, cidr)
- }
- }
- return nil
-}
-
-func validatePoWPathRegexes(regexes []string, listName string) error {
- for _, re := range regexes {
- if _, err := regexp.Compile(re); err != nil {
- return fmt.Errorf("pow_config %s路径正则格式无效: %s", listName, re)
- }
- }
- return nil
-}
-
-func validatePoWIPs(ips []string, listName string) error {
- for _, ip := range ips {
- if net.ParseIP(ip) == nil {
- return fmt.Errorf("pow_config %s IP 格式无效: %s", listName, ip)
- }
- }
- return nil
-}
-
-func validatePoWListMutualExclusion(cfg PoWConfig) error {
- type dimension struct {
- name string
- wl []string
- bl []string
- }
- dimensions := []dimension{
- {"IP", cfg.Whitelist.IPs, cfg.Blacklist.IPs},
- {"IP CIDR", cfg.Whitelist.IPCidrs, cfg.Blacklist.IPCidrs},
- {"路径", cfg.Whitelist.Paths, cfg.Blacklist.Paths},
- {"路径正则", cfg.Whitelist.PathRegexes, cfg.Blacklist.PathRegexes},
- {"User-Agent", cfg.Whitelist.UserAgents, cfg.Blacklist.UserAgents},
- }
- for _, dim := range dimensions {
- if len(dim.wl) > 0 && len(dim.bl) > 0 {
- return fmt.Errorf("pow_config %s 不能同时配置白名单和黑名单", dim.name)
- }
- }
- return nil
-}
diff --git a/internal/apps/openflare/waf/routers.go b/internal/apps/openflare/waf/routers.go
index ab121de0..cfc1c930 100644
--- a/internal/apps/openflare/waf/routers.go
+++ b/internal/apps/openflare/waf/routers.go
@@ -12,14 +12,6 @@ import (
"github.com/gin-gonic/gin"
)
-
-func handleLogicError(c *gin.Context, err error) bool {
- if err == nil {
- return false
- }
- return apiutil.AbortNotFoundIfMissing(c, err, "记录不存在")
-}
-
func routeIDParam(c *gin.Context) (uint, bool) {
raw := c.Param("route_id")
if raw == "" {
@@ -34,167 +26,6 @@ func routeIDParam(c *gin.Context) (uint, bool) {
return uint(id64), true
}
-// ListRuleGroupsHandler 列出全部 WAF 规则组。
-// @Summary 列出 WAF 规则组
-// @Description 返回全部 WAF 规则组,需要管理员权限
-// @Tags openflare-waf
-// @Produce json
-// @Security SessionCookie
-// @Success 200 {object} response.Any{data=[]waf.RuleGroupView} "规则组列表"
-// @Failure 400 {object} response.Any "参数错误"
-// @Failure 401 {object} response.Any "未登录"
-// @Failure 404 {object} response.Any "无权限或不存在"
-// @Failure 500 {object} response.Any "内部错误"
-// @Router /api/v1/d/waf/rule-groups [get]
-func ListRuleGroupsHandler(c *gin.Context) {
- groups, err := ListRuleGroups(c.Request.Context())
- if handleLogicError(c, err) {
- return
- }
- c.JSON(http.StatusOK, response.OK(groups))
-}
-
-// GetRuleGroupHandler 获取 WAF 规则组详情。
-// @Summary 获取 WAF 规则组详情
-// @Description 按 ID 返回 WAF 规则组详情,需要管理员权限
-// @Tags openflare-waf
-// @Produce json
-// @Security SessionCookie
-// @Param id path int true "规则组 ID"
-// @Success 200 {object} response.Any{data=waf.RuleGroupView} "规则组详情"
-// @Failure 400 {object} response.Any "参数错误"
-// @Failure 401 {object} response.Any "未登录"
-// @Failure 404 {object} response.Any "无权限或不存在"
-// @Failure 404 {object} response.Any "记录不存在"
-// @Failure 500 {object} response.Any "内部错误"
-// @Router /api/v1/d/waf/rule-groups/{id} [get]
-func GetRuleGroupHandler(c *gin.Context) {
- id, ok := apiutil.IDParam(c)
- if !ok {
- return
- }
- group, err := GetRuleGroup(c.Request.Context(), id)
- if handleLogicError(c, err) {
- return
- }
- c.JSON(http.StatusOK, response.OK(group))
-}
-
-// CreateRuleGroupHandler 创建 WAF 规则组。
-// @Summary 创建 WAF 规则组
-// @Description 创建新的 WAF 规则组,需要管理员权限
-// @Tags openflare-waf
-// @Accept json
-// @Produce json
-// @Security SessionCookie
-// @Param request body waf.RuleGroupInput true "规则组参数"
-// @Success 200 {object} response.Any{data=waf.RuleGroupView} "创建成功的规则组"
-// @Failure 400 {object} response.Any "参数错误"
-// @Failure 401 {object} response.Any "未登录"
-// @Failure 404 {object} response.Any "无权限或不存在"
-// @Failure 500 {object} response.Any "内部错误"
-// @Router /api/v1/d/waf/rule-groups [post]
-func CreateRuleGroupHandler(c *gin.Context) {
- var input RuleGroupInput
- if !apiutil.BindJSON(c, &input) {
- return
- }
- group, err := CreateRuleGroup(c.Request.Context(), input)
- if handleLogicError(c, err) {
- return
- }
- c.JSON(http.StatusOK, response.OK(group))
-}
-
-// UpdateRuleGroupHandler 更新 WAF 规则组。
-// @Summary 更新 WAF 规则组
-// @Description 按 ID 更新 WAF 规则组,需要管理员权限
-// @Tags openflare-waf
-// @Accept json
-// @Produce json
-// @Security SessionCookie
-// @Param id path int true "规则组 ID"
-// @Param request body waf.RuleGroupInput true "规则组参数"
-// @Success 200 {object} response.Any{data=waf.RuleGroupView} "更新后的规则组"
-// @Failure 400 {object} response.Any "参数错误"
-// @Failure 401 {object} response.Any "未登录"
-// @Failure 404 {object} response.Any "无权限或不存在"
-// @Failure 404 {object} response.Any "记录不存在"
-// @Failure 500 {object} response.Any "内部错误"
-// @Router /api/v1/d/waf/rule-groups/{id}/update [post]
-func UpdateRuleGroupHandler(c *gin.Context) {
- id, ok := apiutil.IDParam(c)
- if !ok {
- return
- }
- var input RuleGroupInput
- if !apiutil.BindJSON(c, &input) {
- return
- }
- group, err := UpdateRuleGroup(c.Request.Context(), id, input)
- if handleLogicError(c, err) {
- return
- }
- c.JSON(http.StatusOK, response.OK(group))
-}
-
-// DeleteRuleGroupHandler 删除 WAF 规则组。
-// @Summary 删除 WAF 规则组
-// @Description 按 ID 删除 WAF 规则组,需要管理员权限
-// @Tags openflare-waf
-// @Produce json
-// @Security SessionCookie
-// @Param id path int true "规则组 ID"
-// @Success 200 {object} response.Any "删除成功"
-// @Failure 400 {object} response.Any "参数错误"
-// @Failure 401 {object} response.Any "未登录"
-// @Failure 404 {object} response.Any "无权限或不存在"
-// @Failure 404 {object} response.Any "记录不存在"
-// @Failure 500 {object} response.Any "内部错误"
-// @Router /api/v1/d/waf/rule-groups/{id}/delete [post]
-func DeleteRuleGroupHandler(c *gin.Context) {
- id, ok := apiutil.IDParam(c)
- if !ok {
- return
- }
- if err := DeleteRuleGroup(c.Request.Context(), id); handleLogicError(c, err) {
- return
- }
- c.JSON(http.StatusOK, response.OKNil())
-}
-
-// ReplaceRuleGroupSitesHandler 替换规则组绑定的站点。
-// @Summary 替换规则组站点绑定
-// @Description 替换 WAF 规则组关联的代理站点列表,需要管理员权限
-// @Tags openflare-waf
-// @Accept json
-// @Produce json
-// @Security SessionCookie
-// @Param id path int true "规则组 ID"
-// @Param request body waf.IDsRequest true "站点 ID 列表"
-// @Success 200 {object} response.Any{data=waf.RuleGroupView} "更新后的规则组"
-// @Failure 400 {object} response.Any "参数错误"
-// @Failure 401 {object} response.Any "未登录"
-// @Failure 404 {object} response.Any "无权限或不存在"
-// @Failure 404 {object} response.Any "记录不存在"
-// @Failure 500 {object} response.Any "内部错误"
-// @Router /api/v1/d/waf/rule-groups/{id}/sites [post]
-func ReplaceRuleGroupSitesHandler(c *gin.Context) {
- id, ok := apiutil.IDParam(c)
- if !ok {
- return
- }
- var request IDsRequest
- if !apiutil.BindJSON(c, &request) {
- return
- }
- group, err := ReplaceRuleGroupSites(c.Request.Context(), id, request.IDs)
- if handleLogicError(c, err) {
- return
- }
- c.JSON(http.StatusOK, response.OK(group))
-}
-
// GetSiteRuleGroupsHandler 获取站点的 WAF 规则组绑定。
// @Summary 获取站点 WAF 规则组
// @Description 返回代理站点关联的 WAF 规则组绑定,需要管理员权限
@@ -215,7 +46,7 @@ func GetSiteRuleGroupsHandler(c *gin.Context) {
return
}
view, err := GetSiteRuleGroups(c.Request.Context(), routeID)
- if handleLogicError(c, err) {
+ if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
@@ -247,7 +78,7 @@ func ReplaceSiteRuleGroupsHandler(c *gin.Context) {
return
}
view, err := ReplaceSiteRuleGroups(c.Request.Context(), routeID, request.IDs)
- if handleLogicError(c, err) {
+ if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(view))
@@ -267,7 +98,7 @@ func ReplaceSiteRuleGroupsHandler(c *gin.Context) {
// @Router /api/v1/d/waf/ip-groups [get]
func ListIPGroupsHandler(c *gin.Context) {
groups, err := ListIPGroups(c.Request.Context())
- if handleLogicError(c, err) {
+ if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(groups))
@@ -293,7 +124,7 @@ func GetIPGroupHandler(c *gin.Context) {
return
}
group, err := GetIPGroup(c.Request.Context(), id)
- if handleLogicError(c, err) {
+ if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
@@ -319,7 +150,7 @@ func CreateIPGroupHandler(c *gin.Context) {
return
}
group, err := CreateIPGroup(c.Request.Context(), input)
- if handleLogicError(c, err) {
+ if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
@@ -351,7 +182,7 @@ func UpdateIPGroupHandler(c *gin.Context) {
return
}
group, err := UpdateIPGroup(c.Request.Context(), id, input)
- if handleLogicError(c, err) {
+ if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(group))
@@ -376,7 +207,7 @@ func DeleteIPGroupHandler(c *gin.Context) {
if !ok {
return
}
- if err := DeleteIPGroup(c.Request.Context(), id); handleLogicError(c, err) {
+ if err := DeleteIPGroup(c.Request.Context(), id); handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OKNil())
@@ -402,7 +233,7 @@ func SyncIPGroupHandler(c *gin.Context) {
return
}
result, err := SyncIPGroup(c.Request.Context(), id)
- if handleLogicError(c, err) {
+ if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
@@ -428,8 +259,8 @@ func TestIPGroupAutoConfigHandler(c *gin.Context) {
return
}
result, err := TestIPGroupAutoConfig(c.Request.Context(), input)
- if handleLogicError(c, err) {
+ if handleRuleError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
-}
\ No newline at end of file
+}
diff --git a/internal/apps/openflare/waf/rule_logics.go b/internal/apps/openflare/waf/rule_logics.go
new file mode 100644
index 00000000..deae54fa
--- /dev/null
+++ b/internal/apps/openflare/waf/rule_logics.go
@@ -0,0 +1,185 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package waf
+
+import (
+ "context"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "strings"
+ "time"
+
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "gorm.io/gorm"
+)
+
+// CreateRuleInput is the minimal payload used to create an orchestrated rule.
+type CreateRuleInput struct {
+ Name string `json:"name"`
+}
+
+// SaveRuleGraphInput atomically replaces a rule graph at the supplied revision.
+type SaveRuleGraphInput struct {
+ Revision uint64 `json:"revision"`
+ Graph RuleGraph `json:"graph"`
+}
+
+// UpdateRuleMetaInput updates metadata without replacing the graph.
+type UpdateRuleMetaInput struct {
+ Name string `json:"name"`
+ Enabled bool `json:"enabled"`
+}
+
+// RuleValidationError represents a safe user-facing validation failure.
+type RuleValidationError struct{ Err error }
+
+func (err *RuleValidationError) Error() string { return err.Err.Error() }
+func (err *RuleValidationError) Unwrap() error { return err.Err }
+
+// RuleView is the API representation of an orchestrated WAF rule.
+type RuleView struct {
+ ID uint `json:"id"`
+ Name string `json:"name"`
+ Enabled bool `json:"enabled"`
+ IsGlobal bool `json:"is_global"`
+ Graph RuleGraph `json:"graph"`
+ Revision uint64 `json:"revision"`
+ AppliedSiteIDs []uint `json:"applied_site_ids"`
+ AppliedSiteCount int `json:"applied_site_count"`
+ CreatedAt string `json:"created_at"`
+ UpdatedAt string `json:"updated_at"`
+}
+
+// ListRules returns all orchestrated WAF rules.
+func ListRules(ctx context.Context) ([]RuleView, error) {
+ if err := EnsureDefaultRuleGroup(ctx); err != nil {
+ return nil, err
+ }
+ groups, err := model.ListOpenFlareWAFRuleGroups(ctx)
+ if err != nil {
+ return nil, err
+ }
+ bindings, err := loadRuleGroupBindings(ctx)
+ if err != nil {
+ return nil, err
+ }
+ views := make([]RuleView, 0, len(groups))
+ for _, group := range groups {
+ view, buildErr := buildRuleView(group, bindings[group.ID])
+ if buildErr != nil {
+ return nil, buildErr
+ }
+ views = append(views, view)
+ }
+ return views, nil
+}
+
+// GetRule returns one orchestrated WAF rule.
+func GetRule(ctx context.Context, id uint) (*RuleView, error) {
+ group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
+ if err != nil {
+ return nil, err
+ }
+ bindings, err := loadRuleGroupBindings(ctx)
+ if err != nil {
+ return nil, err
+ }
+ view, err := buildRuleView(group, bindings[group.ID])
+ return &view, err
+}
+
+// CreateRule creates a disabled custom rule with the safe default graph.
+func CreateRule(ctx context.Context, input CreateRuleInput) (*RuleView, error) {
+ name := strings.TrimSpace(input.Name)
+ if name == "" {
+ return nil, &RuleValidationError{Err: errors.New("WAF 规则名称不能为空")}
+ }
+ raw, err := json.Marshal(DefaultRuleGraph())
+ if err != nil {
+ return nil, err
+ }
+ group := &model.OpenFlareWAFRuleGroup{Name: name, Enabled: false, IsGlobal: false, Graph: string(raw), Revision: 1}
+ if err = model.CreateOpenFlareWAFRuleGroup(ctx, group); err != nil {
+ return nil, err
+ }
+ // GORM applies the model's database default to a false bool on Create, so
+ // explicitly persist the safe disabled state after the row has an ID.
+ group.Enabled = false
+ if err = model.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil {
+ return nil, err
+ }
+ return GetRule(ctx, group.ID)
+}
+
+// UpdateRuleMeta updates rule metadata without touching its graph revision.
+func UpdateRuleMeta(ctx context.Context, id uint, input UpdateRuleMetaInput) (*RuleView, error) {
+ group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
+ if err != nil {
+ return nil, err
+ }
+ name := strings.TrimSpace(input.Name)
+ if name == "" {
+ return nil, &RuleValidationError{Err: errors.New("WAF 规则名称不能为空")}
+ }
+ group.Name, group.Enabled = name, input.Enabled
+ if err = model.UpdateOpenFlareWAFRuleGroup(ctx, group); err != nil {
+ return nil, err
+ }
+ return GetRule(ctx, id)
+}
+
+// DeleteRuleGroup deletes a non-global orchestrated WAF rule.
+func DeleteRuleGroup(ctx context.Context, id uint) error {
+ group, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id)
+ if err != nil {
+ return err
+ }
+ if group.IsGlobal {
+ return &RuleValidationError{Err: errors.New("全局 WAF 规则不能删除")}
+ }
+ return model.DeleteOpenFlareWAFRuleGroupWithBindings(ctx, id)
+}
+
+// SaveRuleGraph validates and atomically replaces a rule graph.
+func SaveRuleGraph(ctx context.Context, id uint, input SaveRuleGraphInput) (*RuleView, error) {
+ if _, err := model.GetOpenFlareWAFRuleGroupByID(ctx, id); err != nil {
+ return nil, err
+ }
+ if err := ValidateRuleGraph(ctx, input.Graph, ruleIPGroupExists); err != nil {
+ return nil, &RuleValidationError{Err: fmt.Errorf("规则图无效: %w", err)}
+ }
+ raw, err := json.Marshal(input.Graph)
+ if err != nil {
+ return nil, err
+ }
+ if _, err = model.UpdateOpenFlareWAFRuleGraph(ctx, id, input.Revision, string(raw)); err != nil {
+ return nil, err
+ }
+ return GetRule(ctx, id)
+}
+
+func ruleIPGroupExists(ctx context.Context, id uint) (bool, error) {
+ _, err := model.GetOpenFlareWAFIPGroupByID(ctx, id)
+ if errors.Is(err, gorm.ErrRecordNotFound) {
+ return false, nil
+ }
+ return err == nil, err
+}
+
+func buildRuleView(group *model.OpenFlareWAFRuleGroup, appliedSiteIDs []uint) (RuleView, error) {
+ if group == nil {
+ return RuleView{}, errors.New("waf rule is nil")
+ }
+ graph := DefaultRuleGraph()
+ if strings.TrimSpace(group.Graph) != "" {
+ if err := json.Unmarshal([]byte(group.Graph), &graph); err != nil {
+ return RuleView{}, err
+ }
+ }
+ ids := append([]uint(nil), appliedSiteIDs...)
+ return RuleView{ID: group.ID, Name: group.Name, Enabled: group.Enabled, IsGlobal: group.IsGlobal,
+ Graph: graph, Revision: group.Revision, AppliedSiteIDs: ids, AppliedSiteCount: len(ids),
+ CreatedAt: group.CreatedAt.Format(time.RFC3339), UpdatedAt: group.UpdatedAt.Format(time.RFC3339)}, nil
+}
diff --git a/internal/apps/openflare/waf/rule_logics_test.go b/internal/apps/openflare/waf/rule_logics_test.go
new file mode 100644
index 00000000..cf4be3b4
--- /dev/null
+++ b/internal/apps/openflare/waf/rule_logics_test.go
@@ -0,0 +1,181 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package waf
+
+import (
+ "bytes"
+ "context"
+ "encoding/json"
+ "errors"
+ "net/http"
+ "net/http/httptest"
+ "strconv"
+ "testing"
+
+ "github.com/Rain-kl/Wavelet/internal/common/response"
+ "github.com/Rain-kl/Wavelet/internal/db"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/gin-gonic/gin"
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+)
+
+func TestDeleteIPGroupRejectsGraphReference(t *testing.T) {
+ cleanup := setupWAFTestDB(t)
+ defer cleanup()
+ ctx := context.Background()
+
+ group, err := CreateIPGroup(ctx, IPGroupInput{Name: "trusted", Type: wafIPGroupTypeManual, Enabled: true})
+ require.NoError(t, err)
+ rule, err := CreateRule(ctx, CreateRuleInput{Name: "guard"})
+ require.NoError(t, err)
+ graph := RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
+ {ID: "start", Type: RuleNodeStart, Config: json.RawMessage(`{}`)},
+ {ID: "match", Type: RuleNodeIPMatch, Config: json.RawMessage(`{"ip_group_ids":[` + strconv.FormatUint(uint64(group.ID), 10) + `]}`)},
+ {ID: "allow", Type: RuleNodeAllow, Config: json.RawMessage(`{}`)},
+ }, Edges: []RuleEdge{
+ {ID: "e1", Source: "start", SourceHandle: "next", Target: "match"},
+ {ID: "e2", Source: "match", SourceHandle: "true", Target: "allow"},
+ {ID: "e3", Source: "match", SourceHandle: "false", Target: "allow"},
+ }}
+ _, err = SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: graph})
+ require.NoError(t, err)
+
+ view, err := GetIPGroup(ctx, group.ID)
+ require.NoError(t, err)
+ assert.Equal(t, 1, view.ReferencedByRuleCount)
+ require.ErrorContains(t, DeleteIPGroup(ctx, group.ID), "已被 WAF 规则引用")
+}
+
+func TestRuleHandlersMapFailures(t *testing.T) {
+ gin.SetMode(gin.TestMode)
+ tests := []struct {
+ name string
+ method string
+ path string
+ body string
+ setup func(t *testing.T) func()
+ want int
+ }{
+ {name: "invalid id", method: http.MethodGet, path: "/rules/nope", setup: setupWAFTestDB, want: http.StatusBadRequest},
+ {name: "malformed json", method: http.MethodPost, path: "/rules", body: `{`, setup: setupWAFTestDB, want: http.StatusBadRequest},
+ {name: "invalid graph", method: http.MethodPost, path: "/rules/1/graph", body: `{"revision":1,"graph":{"schema_version":1,"nodes":[],"edges":[]}}`, setup: func(t *testing.T) func() {
+ cleanup := setupWAFTestDB(t)
+ _, err := CreateRule(context.Background(), CreateRuleInput{Name: "one"})
+ require.NoError(t, err)
+ return cleanup
+ }, want: http.StatusBadRequest},
+ {name: "manual IP group sync", method: http.MethodPost, path: "/ip-groups/1/sync", setup: func(t *testing.T) func() {
+ cleanup := setupWAFTestDB(t)
+ _, err := CreateIPGroup(context.Background(), IPGroupInput{Name: "manual", Type: wafIPGroupTypeManual, Enabled: true})
+ require.NoError(t, err)
+ return cleanup
+ }, want: http.StatusBadRequest},
+ {name: "missing", method: http.MethodGet, path: "/rules/999", setup: setupWAFTestDB, want: http.StatusNotFound},
+ {name: "conflict", method: http.MethodPost, path: "/rules/1/graph", body: mustGraphRequest(t, 0), setup: func(t *testing.T) func() {
+ cleanup := setupWAFTestDB(t)
+ _, err := CreateRule(context.Background(), CreateRuleInput{Name: "one"})
+ require.NoError(t, err)
+ return cleanup
+ }, want: http.StatusConflict},
+ {name: "database failure", method: http.MethodGet, path: "/rules", setup: func(t *testing.T) func() { db.SetDB(nil); return func() {} }, want: http.StatusInternalServerError},
+ }
+ for _, tt := range tests {
+ t.Run(tt.name, func(t *testing.T) {
+ cleanup := tt.setup(t)
+ defer cleanup()
+ router := gin.New()
+ router.Use(response.ErrorHandlerMiddleware())
+ router.GET("/rules", ListRulesHandler)
+ router.POST("/rules", CreateRuleHandler)
+ router.GET("/rules/:id", GetRuleHandler)
+ router.POST("/rules/:id/graph", SaveRuleGraphHandler)
+ router.POST("/ip-groups/:id/sync", SyncIPGroupHandler)
+ rec := httptest.NewRecorder()
+ req := httptest.NewRequest(tt.method, tt.path, bytes.NewBufferString(tt.body))
+ req.Header.Set("Content-Type", "application/json")
+ router.ServeHTTP(rec, req)
+ assert.Equal(t, tt.want, rec.Code, rec.Body.String())
+ })
+ }
+}
+
+func mustGraphRequest(t *testing.T, revision uint64) string {
+ t.Helper()
+ raw, err := json.Marshal(SaveRuleGraphInput{Revision: revision, Graph: DefaultRuleGraph()})
+ require.NoError(t, err)
+ return string(raw)
+}
+
+func TestCreateRuleCreatesDefaultGraph(t *testing.T) {
+ cleanup := setupWAFTestDB(t)
+ defer cleanup()
+
+ rule, err := CreateRule(context.Background(), CreateRuleInput{Name: " edge guard "})
+ require.NoError(t, err)
+ assert.Equal(t, "edge guard", rule.Name)
+ assert.False(t, rule.Enabled)
+ assert.Equal(t, uint64(1), rule.Revision)
+ assert.Equal(t, DefaultRuleGraph(), rule.Graph)
+}
+
+func TestCreateRuleRejectsEmptyName(t *testing.T) {
+ cleanup := setupWAFTestDB(t)
+ defer cleanup()
+
+ _, err := CreateRule(context.Background(), CreateRuleInput{Name: " "})
+ require.ErrorContains(t, err, "名称不能为空")
+}
+
+func TestSaveRuleGraphValidationAndRevisionConflict(t *testing.T) {
+ cleanup := setupWAFTestDB(t)
+ defer cleanup()
+ ctx := context.Background()
+
+ rule, err := CreateRule(ctx, CreateRuleInput{Name: "guard"})
+ require.NoError(t, err)
+ invalid := DefaultRuleGraph()
+ invalid.Edges = nil
+ _, err = SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: invalid})
+ require.Error(t, err)
+
+ updated, err := SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: DefaultRuleGraph()})
+ require.NoError(t, err)
+ assert.Equal(t, uint64(2), updated.Revision)
+ _, err = SaveRuleGraph(ctx, rule.ID, SaveRuleGraphInput{Revision: rule.Revision, Graph: DefaultRuleGraph()})
+ assert.ErrorIs(t, err, model.ErrWAFRuleRevisionConflict)
+}
+
+func TestReplaceSiteRuleGroupsPreservesOrderAndRejectsGlobal(t *testing.T) {
+ cleanup := setupWAFTestDB(t)
+ defer cleanup()
+ ctx := context.Background()
+
+ require.NoError(t, db.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"})
+ require.NoError(t, err)
+ third, err := CreateRule(ctx, CreateRuleInput{Name: "third"})
+ require.NoError(t, err)
+
+ view, err := ReplaceSiteRuleGroups(ctx, 7, []uint{third.ID, first.ID, second.ID, first.ID})
+ require.NoError(t, err)
+ assert.Equal(t, []uint{third.ID, first.ID, second.ID}, view.AppliedIDs)
+
+ require.NoError(t, EnsureDefaultRuleGroup(ctx))
+ global, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
+ require.NoError(t, err)
+ _, err = ReplaceSiteRuleGroups(ctx, 7, []uint{global.ID, second.ID})
+ require.Error(t, err)
+ assert.False(t, errors.Is(err, model.ErrWAFRuleRevisionConflict))
+ assert.Equal(t, []uint{third.ID, first.ID, second.ID}, mustListSiteRuleGroupIDs(t, ctx, 7))
+}
+
+func mustListSiteRuleGroupIDs(t *testing.T, ctx context.Context, routeID uint) []uint {
+ t.Helper()
+ ids, err := ListSiteRuleGroupIDs(ctx, routeID)
+ require.NoError(t, err)
+ return ids
+}
diff --git a/internal/apps/openflare/waf/rule_routers.go b/internal/apps/openflare/waf/rule_routers.go
new file mode 100644
index 00000000..f6c8aa72
--- /dev/null
+++ b/internal/apps/openflare/waf/rule_routers.go
@@ -0,0 +1,186 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package waf
+
+import (
+ "errors"
+ "net/http"
+
+ "github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
+ "github.com/Rain-kl/Wavelet/internal/common/response"
+ "github.com/Rain-kl/Wavelet/internal/model"
+ "github.com/Rain-kl/Wavelet/pkg/logger"
+ "github.com/gin-gonic/gin"
+ "gorm.io/gorm"
+)
+
+func handleRuleError(c *gin.Context, err error) bool {
+ if err == nil {
+ return false
+ }
+ var validation *RuleValidationError
+ switch {
+ case errors.As(err, &validation):
+ response.AbortBadRequest(c, validation.Error())
+ case errors.Is(err, model.ErrWAFRuleRevisionConflict):
+ response.AbortConflict(c, "规则已被其他操作更新,请重新加载")
+ case errors.Is(err, gorm.ErrRecordNotFound):
+ response.AbortNotFound(c, "WAF 规则不存在")
+ default:
+ logger.ErrorF(c.Request.Context(), "[OpenFlareWAF] rule API failed: %v", err)
+ response.AbortInternal(c, "WAF 规则操作失败")
+ }
+ return true
+}
+
+// ListRulesHandler lists orchestrated WAF rules.
+// @Summary 列出 WAF 规则
+// @Tags openflare-waf
+// @Produce json
+// @Security SessionCookie
+// @Success 200 {object} response.Any{data=[]waf.RuleView} "规则列表"
+// @Failure 401 {object} response.Any "未登录"
+// @Failure 404 {object} response.Any "无权限或不存在"
+// @Failure 500 {object} response.Any "内部错误"
+// @Router /api/v1/d/waf/rule-groups [get]
+func ListRulesHandler(c *gin.Context) {
+ rules, err := ListRules(c.Request.Context())
+ if handleRuleError(c, err) {
+ return
+ }
+ c.JSON(http.StatusOK, response.OK(rules))
+}
+
+// GetRuleHandler gets an orchestrated WAF rule.
+// @Summary 获取 WAF 规则详情
+// @Tags openflare-waf
+// @Produce json
+// @Security SessionCookie
+// @Param id path int true "规则 ID"
+// @Success 200 {object} response.Any{data=waf.RuleView} "规则详情"
+// @Failure 400 {object} response.Any "参数错误"
+// @Failure 401 {object} response.Any "未登录"
+// @Failure 404 {object} response.Any "无权限或不存在"
+// @Failure 500 {object} response.Any "内部错误"
+// @Router /api/v1/d/waf/rule-groups/{id} [get]
+func GetRuleHandler(c *gin.Context) {
+ id, ok := apiutil.IDParam(c)
+ if !ok {
+ return
+ }
+ rule, err := GetRule(c.Request.Context(), id)
+ if handleRuleError(c, err) {
+ return
+ }
+ c.JSON(http.StatusOK, response.OK(rule))
+}
+
+// CreateRuleHandler creates an orchestrated WAF rule from a name only.
+// @Summary 创建 WAF 规则
+// @Tags openflare-waf
+// @Accept json
+// @Produce json
+// @Security SessionCookie
+// @Param request body waf.CreateRuleInput true "规则名称"
+// @Success 200 {object} response.Any{data=waf.RuleView} "创建成功"
+// @Failure 400 {object} response.Any "参数错误"
+// @Failure 401 {object} response.Any "未登录"
+// @Failure 404 {object} response.Any "无权限或不存在"
+// @Failure 500 {object} response.Any "内部错误"
+// @Router /api/v1/d/waf/rule-groups [post]
+func CreateRuleHandler(c *gin.Context) {
+ var input CreateRuleInput
+ if !apiutil.BindJSON(c, &input) {
+ return
+ }
+ rule, err := CreateRule(c.Request.Context(), input)
+ if handleRuleError(c, err) {
+ return
+ }
+ c.JSON(http.StatusOK, response.OK(rule))
+}
+
+// UpdateRuleMetaHandler updates rule name and enabled state.
+// @Summary 更新 WAF 规则元数据
+// @Tags openflare-waf
+// @Accept json
+// @Produce json
+// @Security SessionCookie
+// @Param id path int true "规则 ID"
+// @Param request body waf.UpdateRuleMetaInput true "规则元数据"
+// @Success 200 {object} response.Any{data=waf.RuleView} "更新成功"
+// @Failure 400 {object} response.Any "参数错误"
+// @Failure 401 {object} response.Any "未登录"
+// @Failure 404 {object} response.Any "无权限或不存在"
+// @Failure 500 {object} response.Any "内部错误"
+// @Router /api/v1/d/waf/rule-groups/{id}/meta [post]
+func UpdateRuleMetaHandler(c *gin.Context) {
+ id, ok := apiutil.IDParam(c)
+ if !ok {
+ return
+ }
+ var input UpdateRuleMetaInput
+ if !apiutil.BindJSON(c, &input) {
+ return
+ }
+ rule, err := UpdateRuleMeta(c.Request.Context(), id, input)
+ if handleRuleError(c, err) {
+ return
+ }
+ c.JSON(http.StatusOK, response.OK(rule))
+}
+
+// SaveRuleGraphHandler saves a complete versioned rule graph.
+// @Summary 保存 WAF 规则图
+// @Tags openflare-waf
+// @Accept json
+// @Produce json
+// @Security SessionCookie
+// @Param id path int true "规则 ID"
+// @Param request body waf.SaveRuleGraphInput true "规则图和修订号"
+// @Success 200 {object} response.Any{data=waf.RuleView} "保存成功"
+// @Failure 400 {object} response.Any "参数或规则图错误"
+// @Failure 401 {object} response.Any "未登录"
+// @Failure 404 {object} response.Any "无权限或不存在"
+// @Failure 409 {object} response.Any "修订冲突"
+// @Failure 500 {object} response.Any "内部错误"
+// @Router /api/v1/d/waf/rule-groups/{id}/graph [post]
+func SaveRuleGraphHandler(c *gin.Context) {
+ id, ok := apiutil.IDParam(c)
+ if !ok {
+ return
+ }
+ var input SaveRuleGraphInput
+ if !apiutil.BindJSON(c, &input) {
+ return
+ }
+ rule, err := SaveRuleGraph(c.Request.Context(), id, input)
+ if handleRuleError(c, err) {
+ return
+ }
+ c.JSON(http.StatusOK, response.OK(rule))
+}
+
+// DeleteRuleHandler deletes a non-global WAF rule.
+// @Summary 删除 WAF 规则
+// @Tags openflare-waf
+// @Produce json
+// @Security SessionCookie
+// @Param id path int true "规则 ID"
+// @Success 200 {object} response.Any "删除成功"
+// @Failure 400 {object} response.Any "参数错误"
+// @Failure 401 {object} response.Any "未登录"
+// @Failure 404 {object} response.Any "无权限或不存在"
+// @Failure 500 {object} response.Any "内部错误"
+// @Router /api/v1/d/waf/rule-groups/{id}/delete [post]
+func DeleteRuleHandler(c *gin.Context) {
+ id, ok := apiutil.IDParam(c)
+ if !ok {
+ return
+ }
+ if err := DeleteRuleGroup(c.Request.Context(), id); handleRuleError(c, err) {
+ return
+ }
+ c.JSON(http.StatusOK, response.OKNil())
+}
diff --git a/internal/db/migrator/goose/postgres/202607150003_drop_legacy_waf_rule_fields.sql b/internal/db/migrator/goose/postgres/202607150003_drop_legacy_waf_rule_fields.sql
new file mode 100644
index 00000000..70268652
--- /dev/null
+++ b/internal/db/migrator/goose/postgres/202607150003_drop_legacy_waf_rule_fields.sql
@@ -0,0 +1,27 @@
+-- +goose Up
+ALTER TABLE of_waf_rule_groups DROP COLUMN block_status_code;
+ALTER TABLE of_waf_rule_groups DROP COLUMN block_response_body;
+ALTER TABLE of_waf_rule_groups DROP COLUMN ip_whitelist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN ip_blacklist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN ip_whitelist_groups;
+ALTER TABLE of_waf_rule_groups DROP COLUMN ip_blacklist_groups;
+ALTER TABLE of_waf_rule_groups DROP COLUMN country_whitelist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN country_blacklist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN region_whitelist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN region_blacklist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN pow_enabled;
+ALTER TABLE of_waf_rule_groups DROP COLUMN pow_config;
+
+-- +goose Down
+ALTER TABLE of_waf_rule_groups ADD COLUMN block_status_code INTEGER NOT NULL DEFAULT 418;
+ALTER TABLE of_waf_rule_groups ADD COLUMN block_response_body TEXT NOT NULL DEFAULT '';
+ALTER TABLE of_waf_rule_groups ADD COLUMN ip_whitelist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN ip_blacklist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN ip_whitelist_groups TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN ip_blacklist_groups TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN country_whitelist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN country_blacklist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN region_whitelist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN region_blacklist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT FALSE;
+ALTER TABLE of_waf_rule_groups ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}';
diff --git a/internal/db/migrator/goose/sqlite/202607150003_drop_legacy_waf_rule_fields.sql b/internal/db/migrator/goose/sqlite/202607150003_drop_legacy_waf_rule_fields.sql
new file mode 100644
index 00000000..164eb190
--- /dev/null
+++ b/internal/db/migrator/goose/sqlite/202607150003_drop_legacy_waf_rule_fields.sql
@@ -0,0 +1,27 @@
+-- +goose Up
+ALTER TABLE of_waf_rule_groups DROP COLUMN block_status_code;
+ALTER TABLE of_waf_rule_groups DROP COLUMN block_response_body;
+ALTER TABLE of_waf_rule_groups DROP COLUMN ip_whitelist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN ip_blacklist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN ip_whitelist_groups;
+ALTER TABLE of_waf_rule_groups DROP COLUMN ip_blacklist_groups;
+ALTER TABLE of_waf_rule_groups DROP COLUMN country_whitelist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN country_blacklist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN region_whitelist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN region_blacklist;
+ALTER TABLE of_waf_rule_groups DROP COLUMN pow_enabled;
+ALTER TABLE of_waf_rule_groups DROP COLUMN pow_config;
+
+-- +goose Down
+ALTER TABLE of_waf_rule_groups ADD COLUMN block_status_code INTEGER NOT NULL DEFAULT 418;
+ALTER TABLE of_waf_rule_groups ADD COLUMN block_response_body TEXT NOT NULL DEFAULT '';
+ALTER TABLE of_waf_rule_groups ADD COLUMN ip_whitelist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN ip_blacklist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN ip_whitelist_groups TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN ip_blacklist_groups TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN country_whitelist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN country_blacklist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN region_whitelist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN region_blacklist TEXT NOT NULL DEFAULT '[]';
+ALTER TABLE of_waf_rule_groups ADD COLUMN pow_enabled BOOLEAN NOT NULL DEFAULT 0;
+ALTER TABLE of_waf_rule_groups ADD COLUMN pow_config TEXT NOT NULL DEFAULT '{}';
diff --git a/internal/model/openflare_waf.go b/internal/model/openflare_waf.go
index 15e5fc56..188cd8fd 100644
--- a/internal/model/openflare_waf.go
+++ b/internal/model/openflare_waf.go
@@ -14,26 +14,14 @@ import (
// OpenFlareWAFRuleGroup stores a WAF rule group.
type OpenFlareWAFRuleGroup struct {
- ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
- Name string `json:"name" gorm:"size:255;not null"`
- Enabled bool `json:"enabled" gorm:"not null;default:true"`
- IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"`
- BlockStatusCode int `json:"block_status_code" gorm:"not null;default:418"`
- BlockResponseBody string `json:"block_response_body" gorm:"type:text;not null;default:''"`
- IPWhitelist string `json:"ip_whitelist" gorm:"type:text;not null;default:'[]'"`
- IPBlacklist string `json:"ip_blacklist" gorm:"type:text;not null;default:'[]'"`
- IPWhitelistGroups string `json:"ip_whitelist_group_ids" gorm:"column:ip_whitelist_groups;type:text;not null;default:'[]'"`
- IPBlacklistGroups string `json:"ip_blacklist_group_ids" gorm:"column:ip_blacklist_groups;type:text;not null;default:'[]'"`
- CountryWhitelist string `json:"country_whitelist" gorm:"type:text;not null;default:'[]'"`
- CountryBlacklist string `json:"country_blacklist" gorm:"type:text;not null;default:'[]'"`
- RegionWhitelist string `json:"region_whitelist" gorm:"type:text;not null;default:'[]'"`
- RegionBlacklist string `json:"region_blacklist" gorm:"type:text;not null;default:'[]'"`
- PoWEnabled bool `json:"pow_enabled" gorm:"column:pow_enabled;not null;default:false"`
- PoWConfig string `json:"pow_config" gorm:"column:pow_config;type:text;not null;default:'{}'"`
- Graph string `json:"graph" gorm:"type:text;not null;default:''"`
- Revision uint64 `json:"revision" gorm:"not null;default:1"`
- CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
- UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
+ ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
+ Name string `json:"name" gorm:"size:255;not null"`
+ Enabled bool `json:"enabled" gorm:"not null;default:true"`
+ IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"`
+ Graph string `json:"graph" gorm:"type:text;not null;default:''"`
+ Revision uint64 `json:"revision" gorm:"not null;default:1"`
+ CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
+ UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName returns the GORM table name.
@@ -147,21 +135,9 @@ func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGro
return err
}
return conn.Model(&OpenFlareWAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
- "name": group.Name,
- colEnabled: group.Enabled,
- "is_global": group.IsGlobal,
- "block_status_code": group.BlockStatusCode,
- "block_response_body": group.BlockResponseBody,
- "ip_whitelist": group.IPWhitelist,
- "ip_blacklist": group.IPBlacklist,
- "ip_whitelist_groups": group.IPWhitelistGroups,
- "ip_blacklist_groups": group.IPBlacklistGroups,
- "country_whitelist": group.CountryWhitelist,
- "country_blacklist": group.CountryBlacklist,
- "region_whitelist": group.RegionWhitelist,
- "region_blacklist": group.RegionBlacklist,
- "pow_enabled": group.PoWEnabled,
- "pow_config": group.PoWConfig,
+ "name": group.Name,
+ colEnabled: group.Enabled,
+ "is_global": group.IsGlobal,
}).Error
}
diff --git a/internal/model/openflare_waf_graph_test.go b/internal/model/openflare_waf_graph_test.go
index 9255f18c..32819d1a 100644
--- a/internal/model/openflare_waf_graph_test.go
+++ b/internal/model/openflare_waf_graph_test.go
@@ -31,6 +31,7 @@ func wafMigrationFS(t *testing.T) fs.FS {
for _, name := range []string{
"202607150001_orchestrate_waf_rules.sql",
"202607150002_reset_waf_rule_graphs.sql",
+ "202607150003_drop_legacy_waf_rule_fields.sql",
} {
contents, err := os.ReadFile(filepath.Join(dir, name))
require.NoError(t, err)
@@ -45,7 +46,7 @@ func TestOpenFlareWAFGraphMigrationResetsGraphsAndOrdersBindings(t *testing.T) {
sqlDB, err := conn.DB()
require.NoError(t, err)
- require.NoError(t, conn.Exec(`CREATE TABLE of_waf_rule_groups (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL)`).Error)
+ require.NoError(t, conn.Exec(`CREATE TABLE of_waf_rule_groups (id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, block_status_code INTEGER NOT NULL DEFAULT 418, block_response_body TEXT NOT NULL DEFAULT '', ip_whitelist TEXT NOT NULL DEFAULT '[]', ip_blacklist TEXT NOT NULL DEFAULT '[]', ip_whitelist_groups TEXT NOT NULL DEFAULT '[]', ip_blacklist_groups TEXT NOT NULL DEFAULT '[]', country_whitelist TEXT NOT NULL DEFAULT '[]', country_blacklist TEXT NOT NULL DEFAULT '[]', region_whitelist TEXT NOT NULL DEFAULT '[]', region_blacklist TEXT NOT NULL DEFAULT '[]', pow_enabled BOOLEAN NOT NULL DEFAULT 0, pow_config TEXT NOT NULL DEFAULT '{}')`).Error)
require.NoError(t, conn.Exec(`CREATE TABLE of_waf_rule_group_bindings (id INTEGER PRIMARY KEY AUTOINCREMENT, rule_group_id INTEGER NOT NULL, proxy_route_id INTEGER NOT NULL)`).Error)
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (id, name) VALUES (1, 'one'), (2, 'two')`).Error)
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_group_bindings (id, rule_group_id, proxy_route_id) VALUES (20, 2, 7), (10, 1, 7)`).Error)
@@ -61,6 +62,9 @@ func TestOpenFlareWAFGraphMigrationResetsGraphsAndOrdersBindings(t *testing.T) {
require.JSONEq(t, defaultWAFRuleGraph, group.Graph)
assert.Equal(t, uint64(1), group.Revision)
}
+ for _, column := range []string{"block_status_code", "ip_whitelist", "pow_enabled"} {
+ assert.False(t, conn.Migrator().HasColumn("of_waf_rule_groups", column))
+ }
require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (name) VALUES ('new')`).Error)
var newGroup OpenFlareWAFRuleGroup
@@ -105,3 +109,15 @@ func TestReplaceOpenFlareWAFRuleGroupBindingsPreservesInputOrder(t *testing.T) {
assert.Equal(t, []uint{30, 10, 20}, []uint{bindings[0].RuleGroupID, bindings[1].RuleGroupID, bindings[2].RuleGroupID})
assert.Equal(t, []int{0, 1, 2}, []int{bindings[0].Sequence, bindings[1].Sequence, bindings[2].Sequence})
}
+
+func TestLegacyWAFColumnsRemoved(t *testing.T) {
+ conn, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{})
+ require.NoError(t, err)
+ require.NoError(t, conn.AutoMigrate(&OpenFlareWAFRuleGroup{}))
+ legacy := []string{"block_status_code", "block_response_body", "ip_whitelist", "ip_blacklist", "ip_whitelist_groups", "ip_blacklist_groups", "country_whitelist", "country_blacklist", "region_whitelist", "region_blacklist", "pow_enabled", "pow_config"}
+ for _, column := range legacy {
+ if conn.Migrator().HasColumn(&OpenFlareWAFRuleGroup{}, column) {
+ t.Fatalf("legacy WAF column %s still exists", column)
+ }
+ }
+}
diff --git a/internal/router/v1/openflare/register_waf.go b/internal/router/v1/openflare/register_waf.go
index bea01a41..5340eaea 100644
--- a/internal/router/v1/openflare/register_waf.go
+++ b/internal/router/v1/openflare/register_waf.go
@@ -21,12 +21,12 @@ func registerWAFRoutes(apiGroup *gin.RouterGroup) {
wafRoute.POST("/ip-groups/:id/delete", waf.DeleteIPGroupHandler)
wafRoute.POST("/ip-groups/:id/sync", waf.SyncIPGroupHandler)
- wafRoute.GET("/rule-groups", waf.ListRuleGroupsHandler)
- wafRoute.GET("/rule-groups/:id", waf.GetRuleGroupHandler)
- wafRoute.POST("/rule-groups", waf.CreateRuleGroupHandler)
- wafRoute.POST("/rule-groups/:id/update", waf.UpdateRuleGroupHandler)
- wafRoute.POST("/rule-groups/:id/delete", waf.DeleteRuleGroupHandler)
- wafRoute.POST("/rule-groups/:id/sites", waf.ReplaceRuleGroupSitesHandler)
+ wafRoute.GET("/rule-groups", waf.ListRulesHandler)
+ wafRoute.GET("/rule-groups/:id", waf.GetRuleHandler)
+ wafRoute.POST("/rule-groups", waf.CreateRuleHandler)
+ wafRoute.POST("/rule-groups/:id/meta", waf.UpdateRuleMetaHandler)
+ wafRoute.POST("/rule-groups/:id/graph", waf.SaveRuleGraphHandler)
+ wafRoute.POST("/rule-groups/:id/delete", waf.DeleteRuleHandler)
wafRoute.GET("/sites/:route_id/rule-groups", waf.GetSiteRuleGroupsHandler)
wafRoute.POST("/sites/:route_id/rule-groups", waf.ReplaceSiteRuleGroupsHandler)
diff --git a/pkg/protocol/waf_ip_group_snapshot.go b/pkg/protocol/waf_ip_group_snapshot.go
new file mode 100644
index 00000000..6cb23a06
--- /dev/null
+++ b/pkg/protocol/waf_ip_group_snapshot.go
@@ -0,0 +1,39 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package protocol
+
+import (
+ "encoding/json"
+ "fmt"
+)
+
+// MaxWAFIPGroupSnapshotBytes is the maximum serialized size accepted for the
+// complete Agent/OpenResty WAF IP group runtime document.
+const MaxWAFIPGroupSnapshotBytes = 20 << 20
+
+type wafIPGroupSnapshot struct {
+ Groups map[string]WAFIPGroup `json:"groups"`
+}
+
+// MarshalWAFIPGroupSnapshot serializes the exact document written by the
+// Agent to waf_ip_groups.json.
+func MarshalWAFIPGroupSnapshot(groups map[string]WAFIPGroup) ([]byte, error) {
+ if groups == nil {
+ groups = map[string]WAFIPGroup{}
+ }
+ return json.Marshal(wafIPGroupSnapshot{Groups: groups})
+}
+
+// ValidateWAFIPGroupSnapshotSize rejects a complete runtime document that
+// cannot be published safely to the OpenResty shared-memory snapshot.
+func ValidateWAFIPGroupSnapshotSize(groups map[string]WAFIPGroup) error {
+ data, err := MarshalWAFIPGroupSnapshot(groups)
+ if err != nil {
+ return err
+ }
+ if len(data) > MaxWAFIPGroupSnapshotBytes {
+ return fmt.Errorf("WAF IP 组快照大小 %d 字节超过上限 %d 字节", len(data), MaxWAFIPGroupSnapshotBytes)
+ }
+ return nil
+}
diff --git a/pkg/protocol/waf_ip_group_snapshot_test.go b/pkg/protocol/waf_ip_group_snapshot_test.go
new file mode 100644
index 00000000..990f3df5
--- /dev/null
+++ b/pkg/protocol/waf_ip_group_snapshot_test.go
@@ -0,0 +1,57 @@
+// Copyright 2026 Arctel.net
+// SPDX-License-Identifier: Apache-2.0
+
+package protocol
+
+import (
+ "strings"
+ "testing"
+)
+
+func TestMarshalWAFIPGroupSnapshotMatchesAgentRuntimeDocument(t *testing.T) {
+ data, err := MarshalWAFIPGroupSnapshot(map[string]WAFIPGroup{
+ "7": {ID: 7, Name: "deny", Type: "manual", Enabled: true, IPList: []string{"192.0.2.7"}, Checksum: "sum"},
+ })
+ if err != nil {
+ t.Fatalf("MarshalWAFIPGroupSnapshot failed: %v", err)
+ }
+ want := `{"groups":{"7":{"id":7,"name":"deny","type":"manual","enabled":true,"ip_list":["192.0.2.7"],"checksum":"sum"}}}`
+ if string(data) != want {
+ t.Fatalf("snapshot = %s, want %s", data, want)
+ }
+}
+
+func TestValidateWAFIPGroupSnapshotSizeBoundary(t *testing.T) {
+ groups := map[string]WAFIPGroup{
+ "1": {ID: 1, Type: "manual", Enabled: true, IPList: []string{"192.0.2.1"}, Checksum: strings.Repeat("a", 64)},
+ }
+ base, err := MarshalWAFIPGroupSnapshot(groups)
+ if err != nil {
+ t.Fatalf("marshal base snapshot: %v", err)
+ }
+ groups["1"] = WAFIPGroup{
+ ID: 1,
+ Name: strings.Repeat("x", MaxWAFIPGroupSnapshotBytes-len(base)),
+ Type: "manual",
+ Enabled: true,
+ IPList: []string{"192.0.2.1"},
+ Checksum: strings.Repeat("a", 64),
+ }
+ atLimit, err := MarshalWAFIPGroupSnapshot(groups)
+ if err != nil {
+ t.Fatalf("marshal boundary snapshot: %v", err)
+ }
+ if len(atLimit) != MaxWAFIPGroupSnapshotBytes {
+ t.Fatalf("boundary snapshot size = %d, want %d", len(atLimit), MaxWAFIPGroupSnapshotBytes)
+ }
+ if err := ValidateWAFIPGroupSnapshotSize(groups); err != nil {
+ t.Fatalf("boundary snapshot rejected: %v", err)
+ }
+
+ group := groups["1"]
+ group.Name += "x"
+ groups["1"] = group
+ if err := ValidateWAFIPGroupSnapshotSize(groups); err == nil {
+ t.Fatal("oversized snapshot was accepted")
+ }
+}
diff --git a/pkg/render/openresty/render.go b/pkg/render/openresty/render.go
index d85a922f..f7d6b980 100644
--- a/pkg/render/openresty/render.go
+++ b/pkg/render/openresty/render.go
@@ -112,96 +112,10 @@ func RenderRouteConfig(doc Document, certificateFiles []SupportFile) (string, er
// RenderWAFConfig serialises the WAF runtime configuration (rule groups and
// per-site bindings) as a JSON string consumed by the OpenResty Lua runtime.
func RenderWAFConfig(snapshot WAFDocument) (string, error) {
- type wafRuntimeRuleGroup struct {
- ID uint `json:"id"`
- Name string `json:"name"`
- IsGlobal bool `json:"is_global"`
- BlockStatusCode int `json:"block_status_code"`
- BlockResponseBody string `json:"block_response_body"`
- IPWhitelist []string `json:"ip_whitelist"`
- IPBlacklist []string `json:"ip_blacklist"`
- IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
- IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
- CountryWhitelist []string `json:"country_whitelist"`
- CountryBlacklist []string `json:"country_blacklist"`
- RegionWhitelist []string `json:"region_whitelist"`
- RegionBlacklist []string `json:"region_blacklist"`
- PoWEnabled bool `json:"pow_enabled"`
- PoWConfig *PoWConfig `json:"pow_config,omitempty"`
- }
- type wafRuntimeConfig struct {
- DefaultBlockStatusCode int `json:"default_block_status_code"`
- RuleGroups []wafRuntimeRuleGroup `json:"rule_groups"`
- SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
- }
- groups := make([]wafRuntimeRuleGroup, 0, len(snapshot.RuleGroups))
- globalGroupIDs := make([]uint, 0)
- enabledGroupIDs := make(map[uint]struct{}, len(snapshot.RuleGroups))
- for _, group := range snapshot.RuleGroups {
- if !group.Enabled {
- continue
- }
- statusCode := group.BlockStatusCode
- if statusCode == 0 {
- statusCode = defaultWAFBlockStatus
- }
- if group.IsGlobal {
- globalGroupIDs = append(globalGroupIDs, group.ID)
- }
- enabledGroupIDs[group.ID] = struct{}{}
- powConfig := ensurePoWConfig(group.PoWEnabled, group.PoWConfig)
- groups = append(groups, wafRuntimeRuleGroup{
- ID: group.ID,
- Name: group.Name,
- IsGlobal: group.IsGlobal,
- BlockStatusCode: statusCode,
- BlockResponseBody: group.BlockResponseBody,
- IPWhitelist: sortedUniqueStrings(group.IPWhitelist),
- IPBlacklist: sortedUniqueStrings(group.IPBlacklist),
- IPWhitelistGroups: sortedUniqueUintIDs(group.IPWhitelistGroups),
- IPBlacklistGroups: sortedUniqueUintIDs(group.IPBlacklistGroups),
- CountryWhitelist: group.CountryWhitelist,
- CountryBlacklist: group.CountryBlacklist,
- RegionWhitelist: group.RegionWhitelist,
- RegionBlacklist: group.RegionBlacklist,
- PoWEnabled: group.PoWEnabled,
- PoWConfig: powConfig,
- })
- }
- sort.Slice(groups, func(i, j int) bool {
- if groups[i].IsGlobal != groups[j].IsGlobal {
- return groups[i].IsGlobal
- }
- return groups[i].ID < groups[j].ID
- })
- sort.Slice(globalGroupIDs, func(i, j int) bool { return globalGroupIDs[i] < globalGroupIDs[j] })
- siteRuleGroups := make(map[string][]uint, len(snapshot.Bindings))
- for _, binding := range snapshot.Bindings {
- ids := append([]uint{}, globalGroupIDs...)
- for _, id := range binding.RuleGroupIDs {
- if _, ok := enabledGroupIDs[id]; ok {
- ids = append(ids, id)
- }
- }
- siteRuleGroups[binding.SiteName] = uniqueUintIDs(ids)
- }
- data, err := json.Marshal(wafRuntimeConfig{DefaultBlockStatusCode: defaultWAFBlockStatus, RuleGroups: groups, SiteRuleGroups: siteRuleGroups})
+ data, err := json.Marshal(snapshot)
return string(data), err
}
-func sortedUniqueStrings(values []string) []string {
- items := append([]string{}, values...)
- items = uniqueStrings(items)
- sort.Strings(items)
- return items
-}
-
-func sortedUniqueUintIDs(values []uint) []uint {
- items := uniqueUintIDs(values)
- sort.Slice(items, func(i, j int) bool { return items[i] < items[j] })
- return items
-}
-
// ChecksumBundle returns a stable SHA-256 hex digest over the combined content
// of the main config, route config, and deduplicated support files, excluding
// the source config JSON file itself.
@@ -313,7 +227,7 @@ func renderOpenRestyLimitZoneBlock() string {
}
func renderOpenRestyObservabilityTemplateBlock() string {
- return fmt.Sprintf(" lua_shared_dict openflare_observability 10m;\n lua_shared_dict openflare_pow_challenges 10m;\n lua_shared_dict openflare_pow_sessions 10m;\n lua_shared_dict openflare_pow_config 1m;\n lua_shared_dict openflare_waf_config 1m;\n init_worker_by_lua_file %s/observability/init.lua;\n log_by_lua_file %s/observability/log.lua;\n\n server {\n listen %s;\n server_name openflare-observability;\n access_log off;\n\n location = /openflare/stub_status {\n stub_status;\n }\n\n location = /openflare/observability {\n default_type application/json;\n content_by_lua_file %s/observability/read.lua;\n }\n }\n\n", LuaDirPlaceholder, LuaDirPlaceholder, ObservabilityListenPlaceholder, LuaDirPlaceholder)
+ return fmt.Sprintf(" lua_shared_dict openflare_observability 10m;\n lua_shared_dict openflare_pow_challenges 10m;\n lua_shared_dict openflare_pow_sessions 10m;\n lua_shared_dict openflare_pow_config 1m;\n lua_shared_dict openflare_waf_config 1m;\n lua_shared_dict openflare_waf_ip_groups 64m;\n init_worker_by_lua_file %s/observability/init.lua;\n log_by_lua_file %s/observability/log.lua;\n\n server {\n listen %s;\n server_name openflare-observability;\n access_log off;\n\n location = /openflare/stub_status {\n stub_status;\n }\n\n location = /openflare/observability {\n default_type application/json;\n content_by_lua_file %s/observability/read.lua;\n }\n }\n\n", LuaDirPlaceholder, LuaDirPlaceholder, ObservabilityListenPlaceholder, LuaDirPlaceholder)
}
func renderHTTPProxyServer(serverNames string, siteName string, originURL string, originHost string, customHeaders []CustomHeader, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg ConfigSnapshot) string {
@@ -771,7 +685,6 @@ func getPoWConfigForRoute(routeID uint, snapshot WAFDocument) (bool, *PoWConfig)
globalGroupIDs = append(globalGroupIDs, group.ID)
}
}
- sort.Slice(globalGroupIDs, func(i, j int) bool { return globalGroupIDs[i] < globalGroupIDs[j] })
var boundGroupIDs []uint
for _, binding := range snapshot.Bindings {
@@ -789,23 +702,20 @@ func getPoWConfigForRoute(routeID uint, snapshot WAFDocument) (bool, *PoWConfig)
activeGroupIDs := uniqueUintIDs(append(append([]uint{}, globalGroupIDs...), boundGroupIDs...))
for _, groupID := range activeGroupIDs {
group := enabledGroups[groupID]
- if group.PoWEnabled {
- config := ensurePoWConfig(true, group.PoWConfig)
- return true, config
+ if graphContainsNodeType(group.Graph, "pow") {
+ return true, nil
}
}
return false, nil
}
-func ensurePoWConfig(enabled bool, config *PoWConfig) *PoWConfig {
- if !enabled {
- return nil
+func graphContainsNodeType(graph WAFRuleGraph, nodeType string) bool {
+ for _, node := range graph.Nodes {
+ if node.Type == nodeType {
+ return true
+ }
}
- if config != nil {
- return config
- }
- defaultConfig := DefaultPoWConfig()
- return &defaultConfig
+ return false
}
func uniqueUintIDs(values []uint) []uint {
@@ -824,23 +734,6 @@ func uniqueUintIDs(values []uint) []uint {
return result
}
-func uniqueStrings(values []string) []string {
- seen := make(map[string]struct{}, len(values))
- result := make([]string, 0, len(values))
- for _, value := range values {
- item := strings.TrimSpace(value)
- if item == "" {
- continue
- }
- if _, ok := seen[item]; ok {
- continue
- }
- seen[item] = struct{}{}
- result = append(result, item)
- }
- return result
-}
-
func resolveUpstreamServerName(originURL string, originHost string) string {
parsed, err := url.Parse(originURL)
if err != nil || !strings.EqualFold(parsed.Scheme, "https") {
diff --git a/pkg/render/openresty/render_test.go b/pkg/render/openresty/render_test.go
index 4d58ebb9..445d2159 100644
--- a/pkg/render/openresty/render_test.go
+++ b/pkg/render/openresty/render_test.go
@@ -6,6 +6,16 @@ import (
"testing"
)
+func TestRenderOpenRestyUsesDedicatedWAFIPGroupSharedDict(t *testing.T) {
+ block := renderOpenRestyObservabilityTemplateBlock()
+ if !strings.Contains(block, "lua_shared_dict openflare_waf_config 1m;") {
+ t.Fatal("expected general WAF coordination dictionary to remain available")
+ }
+ if !strings.Contains(block, "lua_shared_dict openflare_waf_ip_groups 64m;") {
+ t.Fatalf("expected dedicated 64m WAF IP group dictionary, got:\n%s", block)
+ }
+}
+
func TestRenderWAFConfigIncludesAllRouteSiteNames(t *testing.T) {
doc := Document{
Routes: []Route{
@@ -15,11 +25,8 @@ func TestRenderWAFConfigIncludesAllRouteSiteNames(t *testing.T) {
WAF: WAFDocument{
RuleGroups: []WAFRuleGroup{
{
- ID: 1,
- Name: "pow-group",
- Enabled: true,
- PoWEnabled: true,
- PoWConfig: &PoWConfig{Difficulty: 4, Algorithm: "fast", SessionTTL: 600, ChallengeTTL: 300},
+ ID: 1, Name: "pow-group", Enabled: true,
+ Graph: WAFRuleGraph{Entry: "pow", Nodes: map[string]WAFRuleNode{"pow": {Type: "pow"}}},
},
},
Bindings: []WAFBinding{
@@ -34,18 +41,13 @@ func TestRenderWAFConfigIncludesAllRouteSiteNames(t *testing.T) {
t.Fatalf("RenderWAFConfig() error = %v", err)
}
- var decoded struct {
- SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
- }
+ var decoded WAFDocument
if err := json.Unmarshal([]byte(wafConfig), &decoded); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err)
}
- for _, route := range doc.Routes {
- siteName := resolveRouteSiteName(route)
- if _, ok := decoded.SiteRuleGroups[siteName]; !ok {
- t.Fatalf("site_rule_groups missing site %q, got %#v", siteName, decoded.SiteRuleGroups)
- }
+ if len(decoded.Bindings) != 2 || decoded.Bindings[0].SiteName != "example.com" || decoded.Bindings[1].SiteName != "named-site" {
+ t.Fatalf("bindings did not preserve route site names: %#v", decoded.Bindings)
}
routeConfig, err := RenderRouteConfig(doc, nil)
@@ -60,7 +62,7 @@ func TestRenderWAFConfigIncludesAllRouteSiteNames(t *testing.T) {
}
}
-func TestRenderWAFConfigUsesDefaultPoWConfigWhenEnabledWithoutPayload(t *testing.T) {
+func TestRenderWAFConfigDoesNotSynthesizeLegacyPoWConfig(t *testing.T) {
doc := WAFDocument{
RuleGroups: []WAFRuleGroup{
{
@@ -81,26 +83,15 @@ func TestRenderWAFConfigUsesDefaultPoWConfigWhenEnabledWithoutPayload(t *testing
t.Fatalf("RenderWAFConfig() error = %v", err)
}
- var decoded struct {
- RuleGroups []struct {
- PoWEnabled bool `json:"pow_enabled"`
- PoWConfig *PoWConfig `json:"pow_config"`
- } `json:"rule_groups"`
- }
+ var decoded WAFDocument
if err := json.Unmarshal([]byte(wafConfig), &decoded); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err)
}
if len(decoded.RuleGroups) != 1 {
t.Fatalf("expected 1 rule group, got %d", len(decoded.RuleGroups))
}
- if !decoded.RuleGroups[0].PoWEnabled {
- t.Fatal("expected pow_enabled=true")
- }
- if decoded.RuleGroups[0].PoWConfig == nil {
- t.Fatal("expected default pow_config to be emitted")
- }
- if decoded.RuleGroups[0].PoWConfig.Difficulty != 4 {
- t.Fatalf("expected default difficulty 4, got %d", decoded.RuleGroups[0].PoWConfig.Difficulty)
+ if decoded.RuleGroups[0].PoWConfig != nil {
+ t.Fatalf("expected renderer not to synthesize legacy PoW config, got %#v", decoded.RuleGroups[0].PoWConfig)
}
}
@@ -108,11 +99,8 @@ func TestGetPoWConfigForRouteUsesGlobalGroupWithoutExplicitBinding(t *testing.T)
snapshot := WAFDocument{
RuleGroups: []WAFRuleGroup{
{
- ID: 1,
- Name: "global",
- Enabled: true,
- IsGlobal: true,
- PoWEnabled: true,
+ ID: 1, Name: "global", Enabled: true, IsGlobal: true,
+ Graph: WAFRuleGraph{Entry: "pow", Nodes: map[string]WAFRuleNode{"pow": {Type: "pow"}}},
},
},
Bindings: []WAFBinding{
@@ -124,8 +112,67 @@ func TestGetPoWConfigForRouteUsesGlobalGroupWithoutExplicitBinding(t *testing.T)
if !enabled {
t.Fatal("expected pow to be enabled via global rule group")
}
- if config == nil || config.Difficulty != 4 {
- t.Fatalf("expected default pow config, got %#v", config)
+ if config != nil {
+ t.Fatalf("expected node config to stay in runtime graph, got legacy config %#v", config)
+ }
+}
+
+func TestRenderRouteConfigEnablesPoWLocationsFromRuntimeGraph(t *testing.T) {
+ doc := Document{
+ Routes: []Route{{ID: 1, SiteName: "pow.example.com", Domains: []string{"pow.example.com"}, OriginURL: "http://127.0.0.1:8080", Enabled: true}},
+ WAF: WAFDocument{
+ RuleGroups: []WAFRuleGroup{{
+ ID: 1, Name: "graph-pow", Enabled: true, IsGlobal: true,
+ Graph: WAFRuleGraph{Entry: "start", Nodes: map[string]WAFRuleNode{
+ "start": {Type: "start", Next: map[string]string{"next": "pow"}},
+ "pow": {Type: "pow", Config: json.RawMessage(`{"algorithm":"fast","difficulty":4,"session_ttl":600,"challenge_ttl":300}`), Next: map[string]string{"next": "allow"}},
+ "allow": {Type: "allow"},
+ }},
+ }},
+ Bindings: []WAFBinding{{RouteID: 1, SiteName: "pow.example.com", RuleGroupIDs: []uint{}}},
+ },
+ }
+
+ rendered, err := RenderRouteConfig(doc, nil)
+ if err != nil {
+ t.Fatalf("RenderRouteConfig() error = %v", err)
+ }
+ for _, expected := range []string{
+ `location = /.within.website/x/cmd/anubis/api/make-challenge`,
+ `location = /.within.website/x/cmd/anubis/api/pass-challenge`,
+ `location /.within.website/x/cmd/anubis/static/`,
+ } {
+ if !strings.Contains(rendered, expected) {
+ t.Fatalf("expected graph PoW route to contain %q, got:\n%s", expected, rendered)
+ }
+ }
+}
+
+func TestRenderWAFConfigPreservesRuntimeGraphAndBindingOrder(t *testing.T) {
+ doc := WAFDocument{
+ RuleGroups: []WAFRuleGroup{{
+ ID: 9, Name: "graph", Enabled: true,
+ Graph: WAFRuleGraph{Entry: "start", Nodes: map[string]WAFRuleNode{
+ "start": {Type: "start", Next: map[string]string{"next": "allow"}},
+ "allow": {Type: "allow"},
+ }},
+ }},
+ Bindings: []WAFBinding{{RouteID: 3, SiteName: "ordered.example.com", RuleGroupIDs: []uint{9, 4, 7}}},
+ }
+
+ raw, err := RenderWAFConfig(doc)
+ if err != nil {
+ t.Fatalf("RenderWAFConfig() error = %v", err)
+ }
+ var decoded WAFDocument
+ if err := json.Unmarshal([]byte(raw), &decoded); err != nil {
+ t.Fatalf("json.Unmarshal() error = %v", err)
+ }
+ if decoded.RuleGroups[0].Graph.Entry != "start" {
+ t.Fatalf("runtime graph was not preserved: %#v", decoded.RuleGroups[0].Graph)
+ }
+ if got := decoded.Bindings[0].RuleGroupIDs; len(got) != 3 || got[0] != 9 || got[1] != 4 || got[2] != 7 {
+ t.Fatalf("binding order changed: %#v", got)
}
}
diff --git a/pkg/render/openresty/types.go b/pkg/render/openresty/types.go
index bb7f23cc..1dee3df5 100644
--- a/pkg/render/openresty/types.go
+++ b/pkg/render/openresty/types.go
@@ -1,5 +1,7 @@
package openresty
+import "encoding/json"
+
// Placeholder constants used as sentinel values in rendered OpenResty config
// files; the deploy process replaces them with real paths before reload.
const (
@@ -177,25 +179,39 @@ type PagesDeployment struct {
LocalRoot string `json:"local_root"`
}
-// WAFRuleGroup defines a WAF rule group with IP/country/region lists, PoW
-// integration, and per-group block status configuration.
+// WAFRuleGraph is the compact graph executed by the OpenResty WAF runtime.
+type WAFRuleGraph struct {
+ Entry string `json:"entry"`
+ Nodes map[string]WAFRuleNode `json:"nodes"`
+}
+
+// WAFRuleNode contains one compiled node and its handle-to-target edges.
+type WAFRuleNode struct {
+ Type string `json:"type"`
+ Config json.RawMessage `json:"config,omitempty"`
+ Next map[string]string `json:"next,omitempty"`
+}
+
+// WAFRuleGroup defines one enabled runtime graph. Legacy flattened fields are
+// retained only for decoding older stored snapshots during rolling upgrades.
type WAFRuleGroup struct {
- ID uint `json:"id"`
- Name string `json:"name"`
- Enabled bool `json:"enabled"`
- IsGlobal bool `json:"is_global"`
- BlockStatusCode int `json:"block_status_code"`
- BlockResponseBody string `json:"block_response_body,omitempty"`
- IPWhitelist []string `json:"ip_whitelist,omitempty"`
- IPBlacklist []string `json:"ip_blacklist,omitempty"`
- IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
- IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
- CountryWhitelist []string `json:"country_whitelist,omitempty"`
- CountryBlacklist []string `json:"country_blacklist,omitempty"`
- RegionWhitelist []string `json:"region_whitelist,omitempty"`
- RegionBlacklist []string `json:"region_blacklist,omitempty"`
- PoWEnabled bool `json:"pow_enabled,omitempty"`
- PoWConfig *PoWConfig `json:"pow_config,omitempty"`
+ ID uint `json:"id"`
+ Name string `json:"name"`
+ Enabled bool `json:"enabled"`
+ IsGlobal bool `json:"is_global"`
+ BlockStatusCode int `json:"block_status_code"`
+ BlockResponseBody string `json:"block_response_body,omitempty"`
+ IPWhitelist []string `json:"ip_whitelist,omitempty"`
+ IPBlacklist []string `json:"ip_blacklist,omitempty"`
+ IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
+ IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
+ CountryWhitelist []string `json:"country_whitelist,omitempty"`
+ CountryBlacklist []string `json:"country_blacklist,omitempty"`
+ RegionWhitelist []string `json:"region_whitelist,omitempty"`
+ RegionBlacklist []string `json:"region_blacklist,omitempty"`
+ PoWEnabled bool `json:"pow_enabled,omitempty"`
+ PoWConfig *PoWConfig `json:"pow_config,omitempty"`
+ Graph WAFRuleGraph `json:"graph"`
}
// WAFIPGroup is a named, reusable list of IP addresses or CIDRs that can be