feat(waf): complete composable rule orchestration

Add the React Flow rule editor, ordered graph APIs and runtime DAG execution.\n\nPublish rules only on OpenResty reload and reconcile checksum-driven IP group snapshots in bounded shared memory.
This commit is contained in:
ryan
2026-07-13 14:16:55 +08:00
parent d36409fbf9
commit a1a997bcda
72 changed files with 5897 additions and 3080 deletions
+16 -4
View File
@@ -66,6 +66,7 @@ func main() {
"lua_dir", cfg.LuaDir, "lua_dir", cfg.LuaDir,
"runtime_config_dir", cfg.RuntimeConfigDir, "runtime_config_dir", cfg.RuntimeConfigDir,
"mmdb_path", cfg.MMDBPath, "mmdb_path", cfg.MMDBPath,
"city_mmdb_path", cfg.CityMMDBPath,
) )
client := httpclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration()) client := httpclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration())
@@ -81,6 +82,8 @@ func main() {
LuaDir: cfg.LuaDir, LuaDir: cfg.LuaDir,
NginxLuaDir: cfg.OpenrestyLuaDir, NginxLuaDir: cfg.OpenrestyLuaDir,
RuntimeConfigDir: cfg.RuntimeConfigDir, RuntimeConfigDir: cfg.RuntimeConfigDir,
MMDBPath: cfg.MMDBPath,
CityMMDBPath: cfg.CityMMDBPath,
PagesDir: cfg.PagesDir, PagesDir: cfg.PagesDir,
OpenrestyObservabilityListen: nginx.ObservabilityListenAddress(cfg.OpenrestyObservabilityPort), OpenrestyObservabilityListen: nginx.ObservabilityListenAddress(cfg.OpenrestyObservabilityPort),
OpenrestyObservabilityPort: cfg.OpenrestyObservabilityPort, OpenrestyObservabilityPort: cfg.OpenrestyObservabilityPort,
@@ -122,10 +125,9 @@ func main() {
} }
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM)
geoIPUpdater := &geoipupdate.Updater{ geoIPUpdater := newGeoIPUpdater(cfg)
MMDBPath: cfg.MMDBPath, if err = geoIPUpdater.EnsureInitialDatabases(ctx); err != nil {
DownloadURL: cfg.MMDBDownloadURL, slog.Warn("failed to prepare GeoIP databases before agent startup", "error", err)
UpdateInterval: cfg.MMDBUpdateInterval.Duration(),
} }
go geoIPUpdater.Run(ctx) go geoIPUpdater.Run(ctx)
slog.Info("agent process started") slog.Info("agent process started")
@@ -138,3 +140,13 @@ func main() {
stop() stop()
slog.Info("agent process stopped") 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(),
}
}
+24
View File
@@ -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)
}
}
+12
View File
@@ -21,8 +21,20 @@ sidebar: false
## [unreleased] ## [unreleased]
### 新增
- WAF 规则支持可视化 DAG 编排、版本冲突保护和有序路由绑定,发布时编译为 OpenResty 纯内存运行图。
- WAF IP 组支持 checksum 驱动的 Worker 内存热刷新,并补充 City MMDB 地区匹配数据源。
### 变更
- 移除 WAF 规则旧固定黑白名单、地域名单与 PoW 数据库字段;升级后需在发布前重新编排规则。
### 修复 ### 修复
- WAF 规则编辑器新增启用/停用控制,并移除已废弃的规则组侧站点绑定入口;站点规则顺序统一在反代路由详情中管理。
- 修复 Agent 心跳无法从活动 WAF 运行图发现 IP 组引用的问题,并对发布与同步的完整 IP 组快照增加 20 MiB 聚合容量保护。
- 修复 Agent 增量同步长期保留已取消引用的 WAF IP 组、最终阻塞合法热更新的问题;配置同步现按活动引用集合权威收敛,实时广播仅更新本地现存组。
- 修复 Pages 部署文件清单前端请求路径与后端路由不一致导致的 404 错误。 - 修复 Pages 部署文件清单前端请求路径与后端路由不一致导致的 404 错误。
## [v3.2.0] - 2026-07-12 ## [v3.2.0] - 2026-07-12
+8 -125
View File
@@ -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),会导致: 地域节点使用 Country 与 City MMDB。数据库不可用时地域匹配返回 `false` 并限频告警,不允许因数据损坏意外放行其它执行错误。
1. **源站负载飙升**:高频的数据库查询与 CPU 运算极易耗尽服务器资源。
2. **敏感接口被刷**:登录、注册、短信验证码接口容易被恶意滥用导致财产损失。
3. **数据泄露风险**:恶意的通用漏洞探测行为无法被提前拦截。
因此,OpenFlare 需要在最前端的数据面(OpenResty)构建一套 **高性能、可弹性伸缩的 WAF 过滤引擎**。该引擎能够在最接近用户的边缘层以毫秒级的极低开销对恶意请求进行深度过滤,减轻源站压力,并提供防 CC(PoW 挑战)、IP 黑白名单与地域级别拦截等核心安全防护能力。 ## 安全顺序
--- 启用的全局规则固定前置;路由规则按绑定 sequence 执行。阻止节点立即终止,通过节点仅结束当前规则,全部规则通过后才进入回源链路。未知节点、缺失出口或步数超限一律阻止请求。
## 核心功能
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)` 平滑退出请求,防止请求继续向后透传。
+2 -2
View File
@@ -70,11 +70,11 @@ IP 组采用协调 Worker、共享快照和 Worker 本地对象的两级缓存
1. 请求始终读取当前 Worker 内存中的 IP 组对象,不访问文件或共享字典中的 JSON。 1. 请求始终读取当前 Worker 内存中的 IP 组对象,不访问文件或共享字典中的 JSON。
2. 每 5 秒只有一个取得共享锁的 Worker 读取轻量 checksum 文件。 2. 每 5 秒只有一个取得共享锁的 Worker 读取轻量 checksum 文件。
3. checksum 未变化时立即结束,不读取完整 `waf_ip_groups.json`。 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 发现共享版本变化后,从共享内存取得快照、解析并原子替换各自的本地对象,不重复读取磁盘。 5. 其他 Worker 发现共享版本变化后,从共享内存取得快照、解析并原子替换各自的本地对象,不重复读取磁盘。
6. 刷新失败时继续使用上一份有效对象,限频记录错误,并在下一周期重试。 6. 刷新失败时继续使用上一份有效对象,限频记录错误,并在下一周期重试。
Agent 必须先原子替换 IP 组 JSON,最后原子更新 checksum,使 Worker 永远不会把半写入文件识别为新版本。 Agent 必须先原子替换 IP 组 JSON,最后原子更新 checksum,使 Worker 永远不会把半写入文件识别为新版本。Server 发布/同步和 Agent 落盘共同执行 20 MiB 聚合快照上限;共享字典使用不会强制淘汰旧键的安全写入,失败时保留当前与上一代不可变快照。
## API 与编辑器 ## API 与编辑器
+204 -276
View File
@@ -10249,17 +10249,16 @@ const docTemplate = `{
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "返回全部 WAF 规则组,需要管理员权限",
"produces": [ "produces": [
"application/json" "application/json"
], ],
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "列出 WAF 规则组", "summary": "列出 WAF 规则",
"responses": { "responses": {
"200": { "200": {
"description": "规则组列表", "description": "规则列表",
"schema": { "schema": {
"allOf": [ "allOf": [
{ {
@@ -10271,7 +10270,7 @@ const docTemplate = `{
"data": { "data": {
"type": "array", "type": "array",
"items": { "items": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
} }
} }
@@ -10279,12 +10278,6 @@ const docTemplate = `{
] ]
} }
}, },
"400": {
"description": "参数错误",
"schema": {
"$ref": "#/definitions/response.Any"
}
},
"401": { "401": {
"description": "未登录", "description": "未登录",
"schema": { "schema": {
@@ -10311,7 +10304,6 @@ const docTemplate = `{
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "创建新的 WAF 规则组,需要管理员权限",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -10321,21 +10313,21 @@ const docTemplate = `{
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "创建 WAF 规则组", "summary": "创建 WAF 规则",
"parameters": [ "parameters": [
{ {
"description": "规则组参数", "description": "规则名称",
"name": "request", "name": "request",
"in": "body", "in": "body",
"required": true, "required": true,
"schema": { "schema": {
"$ref": "#/definitions/waf.RuleGroupInput" "$ref": "#/definitions/waf.CreateRuleInput"
} }
} }
], ],
"responses": { "responses": {
"200": { "200": {
"description": "创建成功的规则组", "description": "创建成功",
"schema": { "schema": {
"allOf": [ "allOf": [
{ {
@@ -10345,7 +10337,7 @@ const docTemplate = `{
"type": "object", "type": "object",
"properties": { "properties": {
"data": { "data": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
} }
} }
@@ -10386,18 +10378,17 @@ const docTemplate = `{
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "按 ID 返回 WAF 规则组详情,需要管理员权限",
"produces": [ "produces": [
"application/json" "application/json"
], ],
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "获取 WAF 规则组详情", "summary": "获取 WAF 规则详情",
"parameters": [ "parameters": [
{ {
"type": "integer", "type": "integer",
"description": "规则组 ID", "description": "规则 ID",
"name": "id", "name": "id",
"in": "path", "in": "path",
"required": true "required": true
@@ -10405,7 +10396,7 @@ const docTemplate = `{
], ],
"responses": { "responses": {
"200": { "200": {
"description": "规则组详情", "description": "规则详情",
"schema": { "schema": {
"allOf": [ "allOf": [
{ {
@@ -10415,7 +10406,7 @@ const docTemplate = `{
"type": "object", "type": "object",
"properties": { "properties": {
"data": { "data": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
} }
} }
@@ -10435,7 +10426,7 @@ const docTemplate = `{
} }
}, },
"404": { "404": {
"description": "记录不存在", "description": "无权限或不存在",
"schema": { "schema": {
"$ref": "#/definitions/response.Any" "$ref": "#/definitions/response.Any"
} }
@@ -10456,18 +10447,17 @@ const docTemplate = `{
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "按 ID 删除 WAF 规则组,需要管理员权限",
"produces": [ "produces": [
"application/json" "application/json"
], ],
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "删除 WAF 规则组", "summary": "删除 WAF 规则",
"parameters": [ "parameters": [
{ {
"type": "integer", "type": "integer",
"description": "规则组 ID", "description": "规则 ID",
"name": "id", "name": "id",
"in": "path", "in": "path",
"required": true "required": true
@@ -10493,7 +10483,7 @@ const docTemplate = `{
} }
}, },
"404": { "404": {
"description": "记录不存在", "description": "无权限或不存在",
"schema": { "schema": {
"$ref": "#/definitions/response.Any" "$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": { "post": {
"security": [ "security": [
{ {
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "替换 WAF 规则组关联的代理站点列表,需要管理员权限",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -10524,28 +10513,28 @@ const docTemplate = `{
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "替换规则组站点绑定", "summary": "保存 WAF 规则图",
"parameters": [ "parameters": [
{ {
"type": "integer", "type": "integer",
"description": "规则组 ID", "description": "规则 ID",
"name": "id", "name": "id",
"in": "path", "in": "path",
"required": true "required": true
}, },
{ {
"description": "站点 ID 列表", "description": "规则图和修订号",
"name": "request", "name": "request",
"in": "body", "in": "body",
"required": true, "required": true,
"schema": { "schema": {
"$ref": "#/definitions/waf.IDsRequest" "$ref": "#/definitions/waf.SaveRuleGraphInput"
} }
} }
], ],
"responses": { "responses": {
"200": { "200": {
"description": "更新后的规则组", "description": "保存成功",
"schema": { "schema": {
"allOf": [ "allOf": [
{ {
@@ -10555,7 +10544,94 @@ const docTemplate = `{
"type": "object", "type": "object",
"properties": { "properties": {
"data": { "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": { "404": {
"description": "记录不存在", "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": "记录不存在",
"schema": { "schema": {
"$ref": "#/definitions/response.Any" "$ref": "#/definitions/response.Any"
} }
@@ -18531,6 +18525,14 @@ const docTemplate = `{
} }
} }
}, },
"waf.CreateRuleInput": {
"type": "object",
"properties": {
"name": {
"type": "string"
}
}
},
"waf.IDsRequest": { "waf.IDsRequest": {
"type": "object", "type": "object",
"properties": { "properties": {
@@ -18716,139 +18718,97 @@ const docTemplate = `{
} }
} }
}, },
"waf.PoWConfig": { "waf.RuleEdge": {
"type": "object", "type": "object",
"properties": { "properties": {
"algorithm": { "id": {
"type": "string" "type": "string"
}, },
"blacklist": { "source": {
"$ref": "#/definitions/waf.PoWListConfig" "type": "string"
}, },
"challenge_ttl": { "source_handle": {
"type": "integer" "type": "string"
}, },
"difficulty": { "target": {
"type": "integer" "type": "string"
},
"session_ttl": {
"type": "integer"
},
"whitelist": {
"$ref": "#/definitions/waf.PoWListConfig"
} }
} }
}, },
"waf.PoWListConfig": { "waf.RuleGraph": {
"type": "object", "type": "object",
"properties": { "properties": {
"ip_cidrs": { "edges": {
"type": "array", "type": "array",
"items": { "items": {
"type": "string" "$ref": "#/definitions/waf.RuleEdge"
} }
}, },
"ips": { "nodes": {
"type": "array", "type": "array",
"items": { "items": {
"type": "string" "$ref": "#/definitions/waf.RuleNode"
} }
}, },
"path_regexes": { "schema_version": {
"type": "array", "type": "integer"
"items": {
"type": "string"
}
},
"paths": {
"type": "array",
"items": {
"type": "string"
}
},
"user_agents": {
"type": "array",
"items": {
"type": "string"
}
} }
} }
}, },
"waf.RuleGroupInput": { "waf.RuleNode": {
"type": "object", "type": "object",
"properties": { "properties": {
"block_response_body": { "config": {
"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": {
"type": "array", "type": "array",
"items": { "items": {
"type": "integer" "type": "integer"
} }
}, },
"ip_whitelist": { "id": {
"type": "array",
"items": {
"type": "string"
}
},
"ip_whitelist_group_ids": {
"type": "array",
"items": {
"type": "integer"
}
},
"name": {
"type": "string" "type": "string"
}, },
"pow_config": { "label": {
"type": "array", "type": "string"
"items": {
"type": "integer"
}
}, },
"pow_enabled": { "position": {
"type": "boolean" "$ref": "#/definitions/waf.RulePosition"
}, },
"region_blacklist": { "type": {
"type": "array", "$ref": "#/definitions/waf.RuleNodeType"
"items": {
"type": "string"
}
},
"region_whitelist": {
"type": "array",
"items": {
"type": "string"
}
} }
} }
}, },
"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", "type": "object",
"properties": { "properties": {
"applied_site_count": { "applied_site_count": {
@@ -18860,86 +18820,43 @@ const docTemplate = `{
"type": "integer" "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": { "created_at": {
"type": "string" "type": "string"
}, },
"enabled": { "enabled": {
"type": "boolean" "type": "boolean"
}, },
"graph": {
"$ref": "#/definitions/waf.RuleGraph"
},
"id": { "id": {
"type": "integer" "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": { "is_global": {
"type": "boolean" "type": "boolean"
}, },
"name": { "name": {
"type": "string" "type": "string"
}, },
"pow_config": { "revision": {
"$ref": "#/definitions/waf.PoWConfig" "type": "integer"
},
"pow_enabled": {
"type": "boolean"
},
"region_blacklist": {
"type": "array",
"items": {
"type": "string"
}
},
"region_whitelist": {
"type": "array",
"items": {
"type": "string"
}
}, },
"updated_at": { "updated_at": {
"type": "string" "type": "string"
} }
} }
}, },
"waf.SaveRuleGraphInput": {
"type": "object",
"properties": {
"graph": {
"$ref": "#/definitions/waf.RuleGraph"
},
"revision": {
"type": "integer"
}
}
},
"waf.SiteRuleGroupsView": { "waf.SiteRuleGroupsView": {
"type": "object", "type": "object",
"properties": { "properties": {
@@ -18952,11 +18869,11 @@ const docTemplate = `{
"applied_rule_groups": { "applied_rule_groups": {
"type": "array", "type": "array",
"items": { "items": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
}, },
"global_rule_group": { "global_rule_group": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
}, },
"route_id": { "route_id": {
"type": "integer" "type": "integer"
@@ -18964,11 +18881,22 @@ const docTemplate = `{
"rule_groups": { "rule_groups": {
"type": "array", "type": "array",
"items": { "items": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
} }
} }
}, },
"waf.UpdateRuleMetaInput": {
"type": "object",
"properties": {
"enabled": {
"type": "boolean"
},
"name": {
"type": "string"
}
}
},
"zone.DomainInput": { "zone.DomainInput": {
"type": "object", "type": "object",
"properties": { "properties": {
+15 -166
View File
@@ -1,176 +1,25 @@
# WAF 安全防护使用 # 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)。 架构、图校验和失败回滚细节见 [WAF 可编排规则设计](../design/waf-orchestration-design.md)。
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 组」中,并让规则组同时引用该手动组与自动组。
+7 -1
View File
@@ -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。 **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 ## Global Constraints
- 每张图恰好一个 `start` 和一个 `allow`;`block` 可多个;图必须无环、无悬空、无不可达节点,所有路径必须抵达 `allow` 或 `block`。 - 每张图恰好一个 `start` 和一个 `allow`;`block` 可多个;图必须无环、无悬空、无不可达节点,所有路径必须抵达 `allow` 或 `block`。
@@ -374,7 +380,7 @@ Commit: `feat(agent): execute waf graphs from worker memory`
**Interfaces:** **Interfaces:**
- Produces: `waf_ip_groups.json.checksum`;Lua `ip_groups.current()` 返回 Worker 本地对象;协调刷新间隔固定 5 秒。 - 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: 写失败测试** - [ ] **Step 1: 写失败测试**
+4
View File
@@ -286,6 +286,8 @@ Server 的所有核心基础配置定义在 `config.yaml` 中,且均支持环
| `OPENFLARE_MMDB_PATH` | WAF GeoIP mmdb 路径,可覆盖 `agent.json` | 空 | | `OPENFLARE_MMDB_PATH` | WAF GeoIP mmdb 路径,可覆盖 `agent.json` | 空 |
| `OPENFLARE_MMDB_UPDATE_INTERVAL` | WAF GeoIP mmdb 更新间隔,可覆盖 `agent.json` | 空 | | `OPENFLARE_MMDB_UPDATE_INTERVAL` | WAF GeoIP mmdb 更新间隔,可覆盖 `agent.json` | 空 |
| `OPENFLARE_MMDB_DOWNLOAD_URL` | 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` | | `runtime_config_dir` | Agent 运行时配置写入目录,如 `pow_config.json` | 否 | `data_dir/etc/openflare` |
| `pages_dir` | Pages 静态部署包解压与当前部署目录 | 否 | `data_dir/var/lib/openflare/pages` | | `pages_dir` | Pages 静态部署包解压与当前部署目录 | 否 | `data_dir/var/lib/openflare/pages` |
| `mmdb_path` | WAF GeoIP mmdb 文件路径 | 否 | `data_dir/etc/openflare/GeoLite2-Country.mmdb` | | `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_update_interval` | WAF GeoIP mmdb 更新间隔 | 否 | `86400000` 毫秒 (24h) |
| `mmdb_download_url` | WAF GeoIP mmdb 下载地址 | 否 | 内置 GeoLite2 Country 下载地址 | | `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_buffer_path` | 观测补报缓冲文件路径 | 否 | `data_dir/var/lib/openflare/observability-buffer.json` |
| `observability_replay_minutes` | 自动补传最近观测窗口分钟数 | 否 | `15` | | `observability_replay_minutes` | 自动补传最近观测窗口分钟数 | 否 | `15` |
| `state_path` | Agent 本地状态文件路径 | 否 | `data_dir/var/lib/openflare/agent-state.json` | | `state_path` | Agent 本地状态文件路径 | 否 | `data_dir/var/lib/openflare/agent-state.json` |
+204 -276
View File
@@ -10242,17 +10242,16 @@
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "返回全部 WAF 规则组,需要管理员权限",
"produces": [ "produces": [
"application/json" "application/json"
], ],
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "列出 WAF 规则组", "summary": "列出 WAF 规则",
"responses": { "responses": {
"200": { "200": {
"description": "规则组列表", "description": "规则列表",
"schema": { "schema": {
"allOf": [ "allOf": [
{ {
@@ -10264,7 +10263,7 @@
"data": { "data": {
"type": "array", "type": "array",
"items": { "items": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
} }
} }
@@ -10272,12 +10271,6 @@
] ]
} }
}, },
"400": {
"description": "参数错误",
"schema": {
"$ref": "#/definitions/response.Any"
}
},
"401": { "401": {
"description": "未登录", "description": "未登录",
"schema": { "schema": {
@@ -10304,7 +10297,6 @@
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "创建新的 WAF 规则组,需要管理员权限",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -10314,21 +10306,21 @@
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "创建 WAF 规则组", "summary": "创建 WAF 规则",
"parameters": [ "parameters": [
{ {
"description": "规则组参数", "description": "规则名称",
"name": "request", "name": "request",
"in": "body", "in": "body",
"required": true, "required": true,
"schema": { "schema": {
"$ref": "#/definitions/waf.RuleGroupInput" "$ref": "#/definitions/waf.CreateRuleInput"
} }
} }
], ],
"responses": { "responses": {
"200": { "200": {
"description": "创建成功的规则组", "description": "创建成功",
"schema": { "schema": {
"allOf": [ "allOf": [
{ {
@@ -10338,7 +10330,7 @@
"type": "object", "type": "object",
"properties": { "properties": {
"data": { "data": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
} }
} }
@@ -10379,18 +10371,17 @@
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "按 ID 返回 WAF 规则组详情,需要管理员权限",
"produces": [ "produces": [
"application/json" "application/json"
], ],
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "获取 WAF 规则组详情", "summary": "获取 WAF 规则详情",
"parameters": [ "parameters": [
{ {
"type": "integer", "type": "integer",
"description": "规则组 ID", "description": "规则 ID",
"name": "id", "name": "id",
"in": "path", "in": "path",
"required": true "required": true
@@ -10398,7 +10389,7 @@
], ],
"responses": { "responses": {
"200": { "200": {
"description": "规则组详情", "description": "规则详情",
"schema": { "schema": {
"allOf": [ "allOf": [
{ {
@@ -10408,7 +10399,7 @@
"type": "object", "type": "object",
"properties": { "properties": {
"data": { "data": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
} }
} }
@@ -10428,7 +10419,7 @@
} }
}, },
"404": { "404": {
"description": "记录不存在", "description": "无权限或不存在",
"schema": { "schema": {
"$ref": "#/definitions/response.Any" "$ref": "#/definitions/response.Any"
} }
@@ -10449,18 +10440,17 @@
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "按 ID 删除 WAF 规则组,需要管理员权限",
"produces": [ "produces": [
"application/json" "application/json"
], ],
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "删除 WAF 规则组", "summary": "删除 WAF 规则",
"parameters": [ "parameters": [
{ {
"type": "integer", "type": "integer",
"description": "规则组 ID", "description": "规则 ID",
"name": "id", "name": "id",
"in": "path", "in": "path",
"required": true "required": true
@@ -10486,7 +10476,7 @@
} }
}, },
"404": { "404": {
"description": "记录不存在", "description": "无权限或不存在",
"schema": { "schema": {
"$ref": "#/definitions/response.Any" "$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": { "post": {
"security": [ "security": [
{ {
"SessionCookie": [] "SessionCookie": []
} }
], ],
"description": "替换 WAF 规则组关联的代理站点列表,需要管理员权限",
"consumes": [ "consumes": [
"application/json" "application/json"
], ],
@@ -10517,28 +10506,28 @@
"tags": [ "tags": [
"openflare-waf" "openflare-waf"
], ],
"summary": "替换规则组站点绑定", "summary": "保存 WAF 规则图",
"parameters": [ "parameters": [
{ {
"type": "integer", "type": "integer",
"description": "规则组 ID", "description": "规则 ID",
"name": "id", "name": "id",
"in": "path", "in": "path",
"required": true "required": true
}, },
{ {
"description": "站点 ID 列表", "description": "规则图和修订号",
"name": "request", "name": "request",
"in": "body", "in": "body",
"required": true, "required": true,
"schema": { "schema": {
"$ref": "#/definitions/waf.IDsRequest" "$ref": "#/definitions/waf.SaveRuleGraphInput"
} }
} }
], ],
"responses": { "responses": {
"200": { "200": {
"description": "更新后的规则组", "description": "保存成功",
"schema": { "schema": {
"allOf": [ "allOf": [
{ {
@@ -10548,7 +10537,94 @@
"type": "object", "type": "object",
"properties": { "properties": {
"data": { "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": { "404": {
"description": "记录不存在", "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": "记录不存在",
"schema": { "schema": {
"$ref": "#/definitions/response.Any" "$ref": "#/definitions/response.Any"
} }
@@ -18524,6 +18518,14 @@
} }
} }
}, },
"waf.CreateRuleInput": {
"type": "object",
"properties": {
"name": {
"type": "string"
}
}
},
"waf.IDsRequest": { "waf.IDsRequest": {
"type": "object", "type": "object",
"properties": { "properties": {
@@ -18709,139 +18711,97 @@
} }
} }
}, },
"waf.PoWConfig": { "waf.RuleEdge": {
"type": "object", "type": "object",
"properties": { "properties": {
"algorithm": { "id": {
"type": "string" "type": "string"
}, },
"blacklist": { "source": {
"$ref": "#/definitions/waf.PoWListConfig" "type": "string"
}, },
"challenge_ttl": { "source_handle": {
"type": "integer" "type": "string"
}, },
"difficulty": { "target": {
"type": "integer" "type": "string"
},
"session_ttl": {
"type": "integer"
},
"whitelist": {
"$ref": "#/definitions/waf.PoWListConfig"
} }
} }
}, },
"waf.PoWListConfig": { "waf.RuleGraph": {
"type": "object", "type": "object",
"properties": { "properties": {
"ip_cidrs": { "edges": {
"type": "array", "type": "array",
"items": { "items": {
"type": "string" "$ref": "#/definitions/waf.RuleEdge"
} }
}, },
"ips": { "nodes": {
"type": "array", "type": "array",
"items": { "items": {
"type": "string" "$ref": "#/definitions/waf.RuleNode"
} }
}, },
"path_regexes": { "schema_version": {
"type": "array", "type": "integer"
"items": {
"type": "string"
}
},
"paths": {
"type": "array",
"items": {
"type": "string"
}
},
"user_agents": {
"type": "array",
"items": {
"type": "string"
}
} }
} }
}, },
"waf.RuleGroupInput": { "waf.RuleNode": {
"type": "object", "type": "object",
"properties": { "properties": {
"block_response_body": { "config": {
"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": {
"type": "array", "type": "array",
"items": { "items": {
"type": "integer" "type": "integer"
} }
}, },
"ip_whitelist": { "id": {
"type": "array",
"items": {
"type": "string"
}
},
"ip_whitelist_group_ids": {
"type": "array",
"items": {
"type": "integer"
}
},
"name": {
"type": "string" "type": "string"
}, },
"pow_config": { "label": {
"type": "array", "type": "string"
"items": {
"type": "integer"
}
}, },
"pow_enabled": { "position": {
"type": "boolean" "$ref": "#/definitions/waf.RulePosition"
}, },
"region_blacklist": { "type": {
"type": "array", "$ref": "#/definitions/waf.RuleNodeType"
"items": {
"type": "string"
}
},
"region_whitelist": {
"type": "array",
"items": {
"type": "string"
}
} }
} }
}, },
"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", "type": "object",
"properties": { "properties": {
"applied_site_count": { "applied_site_count": {
@@ -18853,86 +18813,43 @@
"type": "integer" "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": { "created_at": {
"type": "string" "type": "string"
}, },
"enabled": { "enabled": {
"type": "boolean" "type": "boolean"
}, },
"graph": {
"$ref": "#/definitions/waf.RuleGraph"
},
"id": { "id": {
"type": "integer" "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": { "is_global": {
"type": "boolean" "type": "boolean"
}, },
"name": { "name": {
"type": "string" "type": "string"
}, },
"pow_config": { "revision": {
"$ref": "#/definitions/waf.PoWConfig" "type": "integer"
},
"pow_enabled": {
"type": "boolean"
},
"region_blacklist": {
"type": "array",
"items": {
"type": "string"
}
},
"region_whitelist": {
"type": "array",
"items": {
"type": "string"
}
}, },
"updated_at": { "updated_at": {
"type": "string" "type": "string"
} }
} }
}, },
"waf.SaveRuleGraphInput": {
"type": "object",
"properties": {
"graph": {
"$ref": "#/definitions/waf.RuleGraph"
},
"revision": {
"type": "integer"
}
}
},
"waf.SiteRuleGroupsView": { "waf.SiteRuleGroupsView": {
"type": "object", "type": "object",
"properties": { "properties": {
@@ -18945,11 +18862,11 @@
"applied_rule_groups": { "applied_rule_groups": {
"type": "array", "type": "array",
"items": { "items": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
}, },
"global_rule_group": { "global_rule_group": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
}, },
"route_id": { "route_id": {
"type": "integer" "type": "integer"
@@ -18957,11 +18874,22 @@
"rule_groups": { "rule_groups": {
"type": "array", "type": "array",
"items": { "items": {
"$ref": "#/definitions/waf.RuleGroupView" "$ref": "#/definitions/waf.RuleView"
} }
} }
} }
}, },
"waf.UpdateRuleMetaInput": {
"type": "object",
"properties": {
"enabled": {
"type": "boolean"
},
"name": {
"type": "string"
}
}
},
"zone.DomainInput": { "zone.DomainInput": {
"type": "object", "type": "object",
"properties": { "properties": {
+157 -205
View File
@@ -3641,6 +3641,11 @@ definitions:
website: website:
type: string type: string
type: object type: object
waf.CreateRuleInput:
properties:
name:
type: string
type: object
waf.IDsRequest: waf.IDsRequest:
properties: properties:
ids: ids:
@@ -3762,94 +3767,69 @@ definitions:
updated_at: updated_at:
type: string type: string
type: object type: object
waf.PoWConfig: waf.RuleEdge:
properties: properties:
algorithm: id:
type: string type: string
blacklist: source:
$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:
type: string 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 type: integer
country_blacklist: type: object
items: waf.RuleNode:
type: string properties:
type: array config:
country_whitelist:
items:
type: string
type: array
enabled:
type: boolean
ip_blacklist:
items:
type: string
type: array
ip_blacklist_group_ids:
items: items:
type: integer type: integer
type: array type: array
ip_whitelist: id:
items:
type: string
type: array
ip_whitelist_group_ids:
items:
type: integer
type: array
name:
type: string type: string
pow_config: label:
items: type: string
type: integer position:
type: array $ref: '#/definitions/waf.RulePosition'
pow_enabled: type:
type: boolean $ref: '#/definitions/waf.RuleNodeType'
region_blacklist:
items:
type: string
type: array
region_whitelist:
items:
type: string
type: array
type: object 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: properties:
applied_site_count: applied_site_count:
type: integer type: integer
@@ -3857,59 +3837,30 @@ definitions:
items: items:
type: integer type: integer
type: array 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: created_at:
type: string type: string
enabled: enabled:
type: boolean type: boolean
graph:
$ref: '#/definitions/waf.RuleGraph'
id: id:
type: integer 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: is_global:
type: boolean type: boolean
name: name:
type: string type: string
pow_config: revision:
$ref: '#/definitions/waf.PoWConfig' type: integer
pow_enabled:
type: boolean
region_blacklist:
items:
type: string
type: array
region_whitelist:
items:
type: string
type: array
updated_at: updated_at:
type: string type: string
type: object type: object
waf.SaveRuleGraphInput:
properties:
graph:
$ref: '#/definitions/waf.RuleGraph'
revision:
type: integer
type: object
waf.SiteRuleGroupsView: waf.SiteRuleGroupsView:
properties: properties:
applied_ids: applied_ids:
@@ -3918,17 +3869,24 @@ definitions:
type: array type: array
applied_rule_groups: applied_rule_groups:
items: items:
$ref: '#/definitions/waf.RuleGroupView' $ref: '#/definitions/waf.RuleView'
type: array type: array
global_rule_group: global_rule_group:
$ref: '#/definitions/waf.RuleGroupView' $ref: '#/definitions/waf.RuleView'
route_id: route_id:
type: integer type: integer
rule_groups: rule_groups:
items: items:
$ref: '#/definitions/waf.RuleGroupView' $ref: '#/definitions/waf.RuleView'
type: array type: array
type: object type: object
waf.UpdateRuleMetaInput:
properties:
enabled:
type: boolean
name:
type: string
type: object
zone.DomainInput: zone.DomainInput:
properties: properties:
cert_id: cert_id:
@@ -10168,25 +10126,20 @@ paths:
- openflare-waf - openflare-waf
/api/v1/d/waf/rule-groups: /api/v1/d/waf/rule-groups:
get: get:
description: 返回全部 WAF 规则组,需要管理员权限
produces: produces:
- application/json - application/json
responses: responses:
"200": "200":
description: 规则组列表 description: 规则列表
schema: schema:
allOf: allOf:
- $ref: '#/definitions/response.Any' - $ref: '#/definitions/response.Any'
- properties: - properties:
data: data:
items: items:
$ref: '#/definitions/waf.RuleGroupView' $ref: '#/definitions/waf.RuleView'
type: array type: array
type: object type: object
"400":
description: 参数错误
schema:
$ref: '#/definitions/response.Any'
"401": "401":
description: 未登录 description: 未登录
schema: schema:
@@ -10201,31 +10154,30 @@ paths:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
security: security:
- SessionCookie: [] - SessionCookie: []
summary: 列出 WAF 规则组 summary: 列出 WAF 规则
tags: tags:
- openflare-waf - openflare-waf
post: post:
consumes: consumes:
- application/json - application/json
description: 创建新的 WAF 规则组,需要管理员权限
parameters: parameters:
- description: 规则组参数 - description: 规则名称
in: body in: body
name: request name: request
required: true required: true
schema: schema:
$ref: '#/definitions/waf.RuleGroupInput' $ref: '#/definitions/waf.CreateRuleInput'
produces: produces:
- application/json - application/json
responses: responses:
"200": "200":
description: 创建成功的规则组 description: 创建成功
schema: schema:
allOf: allOf:
- $ref: '#/definitions/response.Any' - $ref: '#/definitions/response.Any'
- properties: - properties:
data: data:
$ref: '#/definitions/waf.RuleGroupView' $ref: '#/definitions/waf.RuleView'
type: object type: object
"400": "400":
description: 参数错误 description: 参数错误
@@ -10245,14 +10197,13 @@ paths:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
security: security:
- SessionCookie: [] - SessionCookie: []
summary: 创建 WAF 规则组 summary: 创建 WAF 规则
tags: tags:
- openflare-waf - openflare-waf
/api/v1/d/waf/rule-groups/{id}: /api/v1/d/waf/rule-groups/{id}:
get: get:
description: 按 ID 返回 WAF 规则组详情,需要管理员权限
parameters: parameters:
- description: 规则组 ID - description: 规则 ID
in: path in: path
name: id name: id
required: true required: true
@@ -10261,13 +10212,13 @@ paths:
- application/json - application/json
responses: responses:
"200": "200":
description: 规则组详情 description: 规则详情
schema: schema:
allOf: allOf:
- $ref: '#/definitions/response.Any' - $ref: '#/definitions/response.Any'
- properties: - properties:
data: data:
$ref: '#/definitions/waf.RuleGroupView' $ref: '#/definitions/waf.RuleView'
type: object type: object
"400": "400":
description: 参数错误 description: 参数错误
@@ -10278,7 +10229,7 @@ paths:
schema: schema:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
"404": "404":
description: 记录不存在 description: 无权限或不存在
schema: schema:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
"500": "500":
@@ -10287,14 +10238,13 @@ paths:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
security: security:
- SessionCookie: [] - SessionCookie: []
summary: 获取 WAF 规则组详情 summary: 获取 WAF 规则详情
tags: tags:
- openflare-waf - openflare-waf
/api/v1/d/waf/rule-groups/{id}/delete: /api/v1/d/waf/rule-groups/{id}/delete:
post: post:
description: 按 ID 删除 WAF 规则组,需要管理员权限
parameters: parameters:
- description: 规则组 ID - description: 规则 ID
in: path in: path
name: id name: id
required: true required: true
@@ -10315,7 +10265,7 @@ paths:
schema: schema:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
"404": "404":
description: 记录不存在 description: 无权限或不存在
schema: schema:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
"500": "500":
@@ -10324,37 +10274,89 @@ paths:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
security: security:
- SessionCookie: [] - SessionCookie: []
summary: 删除 WAF 规则组 summary: 删除 WAF 规则
tags: tags:
- openflare-waf - openflare-waf
/api/v1/d/waf/rule-groups/{id}/sites: /api/v1/d/waf/rule-groups/{id}/graph:
post: post:
consumes: consumes:
- application/json - application/json
description: 替换 WAF 规则组关联的代理站点列表,需要管理员权限
parameters: parameters:
- description: 规则组 ID - description: 规则 ID
in: path in: path
name: id name: id
required: true required: true
type: integer type: integer
- description: 站点 ID 列表 - description: 规则图和修订号
in: body in: body
name: request name: request
required: true required: true
schema: schema:
$ref: '#/definitions/waf.IDsRequest' $ref: '#/definitions/waf.SaveRuleGraphInput'
produces: produces:
- application/json - application/json
responses: responses:
"200": "200":
description: 更新后的规则组 description: 保存成功
schema: schema:
allOf: allOf:
- $ref: '#/definitions/response.Any' - $ref: '#/definitions/response.Any'
- properties: - properties:
data: 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 type: object
"400": "400":
description: 参数错误 description: 参数错误
@@ -10365,7 +10367,7 @@ paths:
schema: schema:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
"404": "404":
description: 记录不存在 description: 无权限或不存在
schema: schema:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
"500": "500":
@@ -10374,57 +10376,7 @@ paths:
$ref: '#/definitions/response.Any' $ref: '#/definitions/response.Any'
security: security:
- SessionCookie: [] - SessionCookie: []
summary: 替换规则组站点绑定 summary: 更新 WAF 规则元数据
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 规则组
tags: tags:
- openflare-waf - openflare-waf
/api/v1/d/waf/sites/{route_id}/rule-groups: /api/v1/d/waf/sites/{route_id}/rule-groups:
@@ -1,6 +1,6 @@
'use client'; '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 {Badge} from '@/components/ui/badge';
import {Button} from '@/components/ui/button'; import {Button} from '@/components/ui/button';
@@ -20,14 +20,12 @@ interface RuleGroupsTableProps {
groups: WAFRule[]; groups: WAFRule[];
onEdit: (group: WAFRule) => void; onEdit: (group: WAFRule) => void;
onDelete: (group: WAFRule) => void; onDelete: (group: WAFRule) => void;
onBindSites: (group: WAFRule) => void;
} }
export function RuleGroupsTable({ export function RuleGroupsTable({
groups, groups,
onEdit, onEdit,
onDelete, onDelete,
onBindSites,
}: RuleGroupsTableProps) { }: RuleGroupsTableProps) {
return ( return (
<Table> <Table>
@@ -85,12 +83,6 @@ export function RuleGroupsTable({
<Pencil /> <Pencil />
编排 编排
</DropdownMenuItem> </DropdownMenuItem>
{!group.is_global ? (
<DropdownMenuItem onClick={() => onBindSites(group)}>
<Users />
绑定网站
</DropdownMenuItem>
) : null}
</DropdownMenuGroup> </DropdownMenuGroup>
{!group.is_global ? ( {!group.is_global ? (
<> <>
@@ -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<number[]>([]);
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 (
<Sheet open={open} onOpenChange={onOpenChange}>
<SheetContent side="right" className="w-full sm:max-w-lg overflow-y-auto">
<SheetHeader>
<SheetTitle>{group ? `绑定 ${group.name}` : '绑定规则组'}</SheetTitle>
<SheetDescription>
选择这个自定义规则组要叠加到哪些网站。
</SheetDescription>
</SheetHeader>
<div className="space-y-4 px-4 pb-4">
<div className="flex items-center gap-2 rounded-md border px-3 py-2">
<Search className="size-4 text-muted-foreground" />
<Input
value={keyword}
placeholder="搜索网站或域名"
className="border-0 shadow-none focus-visible:ring-0"
onChange={(event) => setKeyword(event.target.value)}
/>
<Button type="button" variant="ghost" size="sm" onClick={selectFiltered}>
全选当前
</Button>
</div>
<div className="space-y-2">
{filteredRoutes.map((route) => (
<button
key={route.id}
type="button"
onClick={() => toggleID(route.id)}
className={cn(
'flex w-full items-center gap-3 rounded-md border px-3 py-2 text-left transition',
selectedSet.has(route.id) && 'border-primary bg-muted/50',
)}
>
<span
className={cn(
'flex size-5 items-center justify-center rounded border',
selectedSet.has(route.id) && 'border-primary bg-primary text-primary-foreground',
)}
>
{selectedSet.has(route.id) ? <Check className="size-3" /> : null}
</span>
<span className="min-w-0 flex-1">
<span className="block truncate text-sm font-medium">
{route.site_name}
</span>
<span className="block truncate text-xs text-muted-foreground">
{(route.zone_domains ?? []).map((item) => item.domain).join(', ') ||
'未绑定域名'}
</span>
</span>
</button>
))}
</div>
</div>
<SheetFooter className="px-4">
<Button type="button" variant="outline" onClick={() => onOpenChange(false)}>
取消
</Button>
<Button
type="button"
disabled={!group || pending}
onClick={() => onSave(selectedIDs)}
>
{pending ? '保存中...' : '保存应用范围'}
</Button>
</SheetFooter>
</SheetContent>
</Sheet>
);
}
+3 -38
View File
@@ -22,33 +22,25 @@ import {EmptyStateWithBorder} from '@/components/layout/empty';
import {ErrorInline} from '@/components/layout/error'; import {ErrorInline} from '@/components/layout/error';
import {LoadingStateWithBorder} from '@/components/layout/loading'; import {LoadingStateWithBorder} from '@/components/layout/loading';
import type {WAFRule} from '@/lib/services/openflare'; 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 {CreateRuleDialog} from './components/create-rule-dialog';
import {getErrorMessage} from './components/helpers'; import {getErrorMessage} from './components/helpers';
import {RuleGroupsTable} from './components/rule-groups-table'; import {RuleGroupsTable} from './components/rule-groups-table';
import {SiteBindingSheet} from './components/site-binding-sheet';
const ruleGroupsQueryKey = ['openflare', 'waf', 'rule-groups']; const ruleGroupsQueryKey = ['openflare', 'waf', 'rule-groups'];
const routesQueryKey = ['openflare', 'proxy-routes'];
export default function WafPage() { export default function WafPage() {
const router = useRouter(); const router = useRouter();
const queryClient = useQueryClient(); const queryClient = useQueryClient();
const [createOpen, setCreateOpen] = useState(false); const [createOpen, setCreateOpen] = useState(false);
const [deleteTarget, setDeleteTarget] = useState<WAFRule | null>(null); const [deleteTarget, setDeleteTarget] = useState<WAFRule | null>(null);
const [bindingGroup, setBindingGroup] = useState<WAFRule | null>(null);
const groupsQuery = useQuery({ const groupsQuery = useQuery({
queryKey: ruleGroupsQueryKey, queryKey: ruleGroupsQueryKey,
queryFn: () => WafService.listRuleGroups(), queryFn: () => WafService.listRuleGroups(),
}); });
const routesQuery = useQuery({
queryKey: routesQueryKey,
queryFn: () => ProxyRouteService.list(),
});
const invalidate = async () => { const invalidate = async () => {
await Promise.all([ await Promise.all([
queryClient.invalidateQueries({ queryKey: ruleGroupsQueryKey }), 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 = () => { const handleRefresh = () => {
void queryClient.invalidateQueries({ queryKey: ruleGroupsQueryKey }); void queryClient.invalidateQueries({ queryKey: ruleGroupsQueryKey });
}; };
const groups = groupsQuery.data ?? []; const groups = groupsQuery.data ?? [];
const loading = groupsQuery.isLoading || routesQuery.isLoading; const loading = groupsQuery.isLoading;
const error = groupsQuery.error ?? routesQuery.error ?? null; const error = groupsQuery.error ?? null;
return ( return (
<div className="w-full py-6 px-1 space-y-6"> <div className="w-full py-6 px-1 space-y-6">
@@ -162,7 +141,6 @@ export default function WafPage() {
groups={groups} groups={groups}
onEdit={(rule) => router.push(`/waf/rules/editor?id=${rule.id}`)} onEdit={(rule) => router.push(`/waf/rules/editor?id=${rule.id}`)}
onDelete={setDeleteTarget} onDelete={setDeleteTarget}
onBindSites={setBindingGroup}
/> />
)} )}
</CardContent> </CardContent>
@@ -177,19 +155,6 @@ export default function WafPage() {
}} }}
/> />
<SiteBindingSheet
group={bindingGroup}
routes={routesQuery.data ?? []}
open={Boolean(bindingGroup)}
pending={bindMutation.isPending}
onOpenChange={(open) => !open && setBindingGroup(null)}
onSave={(ids) => {
if (bindingGroup) {
bindMutation.mutate({ id: bindingGroup.id, ids });
}
}}
/>
<AlertDialog <AlertDialog
open={Boolean(deleteTarget)} open={Boolean(deleteTarget)}
onOpenChange={(open) => !open && setDeleteTarget(null)} onOpenChange={(open) => !open && setDeleteTarget(null)}
@@ -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});
});
@@ -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<Record<WAFRuleNode['type'], string[]>> = {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;
}
@@ -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);
});
@@ -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<Record<WAFRuleNode['type'], string[]>> = {
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<string>();
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<string>();
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<string, string[]>): Set<string> {
const seen = new Set<string>();
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)};
}
@@ -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<WAFRuleNode['type'], 'ip_match' | 'geo_match' | 'pow' | 'block'>;
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 <div className="flex items-center gap-2">{items.map(({type, label, icon: Icon}) => <Button key={type} variant="outline" size="sm" onClick={() => onAdd(type)}><Icon data-icon="inline-start" />{label}</Button>)}</div>;
}
@@ -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(<NodeProperties node={node} ipGroups={[group]} onChange={onChange}/>);
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(<NodeProperties node={node} ipGroups={[]} onChange={vi.fn()}/>);
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(<NodeProperties node={node} ipGroups={[]} onChange={onChange}/>);
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']})}));
});
@@ -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 <aside className="w-80 shrink-0 border-l bg-card"><ScrollArea className="h-full"><div className="flex flex-col gap-5 p-5"><div className="flex items-center gap-2"><Settings2 className="size-5 text-primary"/><div><h2 className="text-sm font-semibold">节点属性</h2><p className="text-xs text-muted-foreground">配置当前处理单元</p></div></div><Separator />{!node ? <p className="text-sm text-muted-foreground">选择画布中的节点以查看配置。</p> : <PropertyFields node={node} ipGroups={ipGroups} onChange={onChange}/>}</div></ScrollArea></aside>;
}
function PropertyFields({node, ipGroups, onChange}: {node: WAFRuleNode; ipGroups: WAFIPGroup[]; onChange: (node: WAFRuleNode) => void}) {
if (node.type === 'start' || node.type === 'allow') return <p className="text-sm text-muted-foreground">系统节点无需配置。</p>;
if (node.type === 'ip_match') return <FieldGroup><CsvField id={`${node.id}-ips`} label="IP 地址" value={node.config.ips} onChange={(ips) => onChange({...node, config: {...node.config, ips}})}/><CsvField id={`${node.id}-cidrs`} label="CIDR 网段" value={node.config.cidrs} onChange={(cidrs) => onChange({...node, config: {...node.config, cidrs}})}/><MultiSelect id={`${node.id}-groups`} label="IP 组" options={ipGroups.map((group) => ({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)}})}/></FieldGroup>;
if (node.type === 'geo_match') return <FieldGroup><MultiSelect id={`${node.id}-countries`} label="国家代码" options={countries} value={node.config.countries} creatablePattern={/^[A-Z]{2}$/} onChange={(countries) => onChange({...node, config: {...node.config, countries}})}/><MultiSelect id={`${node.id}-regions`} label="地区代码" options={regions} value={node.config.regions} creatablePattern={/^[A-Z]{2}-[A-Z0-9]{1,3}$/} onChange={(regions) => onChange({...node, config: {...node.config, regions}})}/></FieldGroup>;
if (node.type === 'pow') return <FieldGroup><Field><FieldLabel htmlFor={`${node.id}-algorithm`}>算法</FieldLabel><Select value={node.config.algorithm} onValueChange={(algorithm: 'fast' | 'slow') => onChange({...node, config: {...node.config, algorithm}})}><SelectTrigger id={`${node.id}-algorithm`} className="w-full"><SelectValue/></SelectTrigger><SelectContent><SelectGroup><SelectItem value="fast">快速</SelectItem><SelectItem value="slow">稳健</SelectItem></SelectGroup></SelectContent></Select></Field>{(['difficulty', 'session_ttl', 'challenge_ttl'] as const).map((key) => <NumberField id={`${node.id}-${key}`} key={key} min={{difficulty: 1, session_ttl: 60, challenge_ttl: 30}[key]} max={key === 'difficulty' ? 16 : undefined} label={{difficulty: '难度', session_ttl: '会话 TTL(秒)', challenge_ttl: '挑战 TTL(秒)'}[key]} value={node.config[key]} onChange={(value) => onChange({...node, config: {...node.config, [key]: value}})}/>)}</FieldGroup>;
return <FieldGroup><NumberField id={`${node.id}-status`} min={400} max={599} label="HTTP 状态码" value={node.config.status_code} onChange={(status_code) => onChange({...node, config: {...node.config, status_code}})}/><Field><FieldLabel htmlFor={`${node.id}-body`}>HTML 响应体</FieldLabel><Textarea id={`${node.id}-body`} rows={9} value={node.config.response_body} onChange={(event) => onChange({...node, config: {...node.config, response_body: event.target.value}})}/><FieldDescription>{new TextEncoder().encode(node.config.response_body).length} / 16384 字节</FieldDescription></Field></FieldGroup>;
}
function CsvField({id, label, value, onChange}: {id: string; label: string; value: string[]; onChange: (value: string[]) => void}) { return <Field><FieldLabel htmlFor={id}>{label}</FieldLabel><Textarea id={id} value={value.join('\n')} onChange={(event) => onChange(event.target.value.split(/[\n,]/).map((item) => item.trim()).filter(Boolean))}/><FieldDescription>每行一个值</FieldDescription></Field>; }
function NumberField({id, label, value, min, max, onChange}: {id: string; label: string; value: number; min?: number; max?: number; onChange: (value: number) => void}) { return <Field><FieldLabel htmlFor={id}>{label}</FieldLabel><Input id={id} min={min} max={max} type="number" value={value} onChange={(event) => onChange(Number(event.target.value))}/></Field>; }
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 <Field><FieldLabel htmlFor={id}>{label}</FieldLabel><Popover><PopoverTrigger asChild><Button id={id} variant="outline" className="w-full justify-start">{value.length ? `已选择 ${value.length} 项` : '请选择'}</Button></PopoverTrigger><PopoverContent align="start" className="flex max-h-64 flex-col gap-2 overflow-y-auto">{creatablePattern && <div className="flex gap-2"><Input aria-label={`新建${label}`} placeholder="输入代码并添加" value={draft} onChange={(event) => setDraft(event.target.value)}/><Button size="sm" disabled={!canCreate} onClick={() => { onChange([...value, normalized]); setDraft(''); }}>添加代码</Button></div>}{visible.length === 0 ? <p className="text-sm text-muted-foreground">暂无可选项</p> : visible.map((option) => <label key={option.value} className="flex cursor-pointer items-center gap-2 text-sm"><Checkbox checked={value.includes(option.value)} onCheckedChange={(checked) => onChange(checked ? [...value, option.value] : value.filter((item) => item !== option.value))}/><span>{option.label}</span><span className="ml-auto font-mono text-xs text-muted-foreground">{option.value}</span></label>)}</PopoverContent></Popover></Field>;
}
@@ -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<ReactFlowInstance<Node<RuleFlowNodeData>, Edge> | null>(null);
const nodes = useMemo<Node<RuleFlowNodeData>[]>(() => 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<Edge[]>(() => 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 <section className="relative min-w-0 flex-1 bg-muted/20"><div className="absolute left-4 top-4 z-10 rounded-lg border bg-background/95 p-2 shadow-sm backdrop-blur"><NodeLibrary onAdd={addNode}/></div><ReactFlow nodes={nodes} edges={edges} nodeTypes={nodeTypes} onInit={(value) => { 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']}><Background gap={20} size={1}/><MiniMap pannable zoomable/><Controls/></ReactFlow></section>;
}
function toRuleEdge(edge: Edge): WAFRuleEdge { return {id: edge.id, source: edge.source, source_handle: edge.sourceHandle ?? '', target: edge.target}; }
@@ -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<string, unknown> { 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<Record<WAFRuleNode['type'], string[]>> = {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 (
<div className={cn('min-w-44 rounded-lg border bg-card shadow-sm transition-shadow', selected && 'ring-2 ring-ring', issues > 0 && 'border-destructive')}>
{rule.type !== 'start' && <Handle type="target" position={Position.Left} />}
<div className="flex items-center gap-3 px-4 py-3">
<Icon className="size-5 text-primary" />
<div className="flex min-w-0 flex-1 flex-col gap-0.5">
<span className="text-sm font-medium">{label}</span>
<span className="font-mono text-[10px] text-muted-foreground">{rule.id}</span>
</div>
{issues > 0 && <Badge variant="destructive">{issues}</Badge>}
</div>
{(outputHandles[rule.type] ?? []).map((handle, index, all) => (
<Handle key={handle} id={handle} type="source" position={Position.Right} style={{top: `${((index + 1) / (all.length + 1)) * 100}%`}}>
<span className="absolute right-3 -translate-y-1/2 text-[9px] text-muted-foreground">{handle}</span>
</Handle>
))}
</div>
);
}
@@ -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(<><UnsavedChanges dirty/><a href="/waf">WAF</a></>);
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(<><UnsavedChanges dirty={false}/><a href="/waf">WAF</a></>);
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(<UnsavedChanges dirty/>);
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(<UnsavedChanges dirty/>);
window.dispatchEvent(new PopStateEvent('popstate', {state: {legacy: true}}));
expect(window.confirm).toHaveBeenCalledOnce();
expect(push).toHaveBeenCalledWith(expect.objectContaining({__wafEditorIndex: 4}), '', '/waf/rules/editor?id=9');
});
@@ -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;
}
@@ -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<typeof import('@/lib/services')>();
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}) => <div><button onClick={() => onGraphChange({...graph, nodes: graph.nodes.map((node) => node.id === 'start' ? {...node, position: {x: 10, y: 0}} : node)})}>修改画布</button>{focusTarget && <span>focus:{focusTarget.kind}:{focusTarget.id}</span>}</div>}));
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(<QueryClientProvider client={client}><WAFRuleEditorPage/></QueryClientProvider>);
}
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<WAFRule>((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<WAFRule>((_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());
});
@@ -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 <Suspense fallback={<EditorSkeleton/>}><EditorContent/></Suspense>;
}
function EditorContent() {
const router = useRouter();
const searchParams = useSearchParams();
const queryClient = useQueryClient();
const id = Number(searchParams.get('id'));
const [graph, setGraph] = useState<WAFRuleGraph>();
const [revision, setRevision] = useState(0);
const [selectedId, setSelectedId] = useState<string>();
const [selectedEdgeId, setSelectedEdgeId] = useState<string>();
const [focusTarget, setFocusTarget] = useState<GraphErrorTarget>();
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<WAFRule>(ruleQueryKey);
queryClient.setQueryData<WAFRule>(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 <div className="w-full px-1 py-6"><p className="text-sm text-destructive">缺少有效的规则 ID。</p></div>;
if (ruleQuery.isError) return <div className="flex w-full flex-col items-start gap-3 px-1 py-6"><p className="text-sm text-destructive">规则加载失败,请重试。</p><Button variant="outline" onClick={() => void ruleQuery.refetch()}>重新加载</Button></div>;
if (ruleQuery.isLoading || !graph || !ruleQuery.data) return <EditorSkeleton/>;
return (
<div className="flex h-[calc(100dvh-4rem)] w-full flex-col px-1 py-6">
<UnsavedChanges dirty={dirty}/>
<header className="mb-4 flex items-center justify-between gap-4">
<div className="flex min-w-0 items-center gap-2"><GitBranch className="size-5 text-primary"/><h1 className="text-2xl font-semibold tracking-tight">{ruleQuery.data.name}</h1><Badge variant={ruleQuery.data.enabled ? 'default' : 'secondary'}>{ruleQuery.data.enabled ? '已启用' : '已停用'}</Badge><Badge variant={issues.length === 0 ? 'outline' : 'destructive'}>{issues.length === 0 ? '图校验通过' : `${issues.length} 个问题`}</Badge>{dirty && <Badge variant="secondary">未保存</Badge>}</div>
<div className="flex shrink-0 items-center gap-2">
<div className="flex items-center gap-2">
<Switch
id="rule-enabled"
checked={ruleQuery.data.enabled}
disabled={dirty || issues.length > 0 || saveMutation.isPending || metaMutation.isPending}
onCheckedChange={(enabled) => metaMutation.mutate(enabled)}
/>
<Label htmlFor="rule-enabled">启用规则</Label>
</div>
<Button variant="outline" onClick={leave}><ArrowLeft data-icon="inline-start"/>返回</Button>
{conflict && <Button variant="outline" onClick={() => { setDirty(false); setConflict(false); void ruleQuery.refetch(); }}>重新加载</Button>}
<Button disabled={!dirty || issues.length > 0 || saveMutation.isPending} onClick={() => saveMutation.mutate()}><Save data-icon="inline-start"/>{saveMutation.isPending ? '保存中...' : '保存'}</Button>
</div>
</header>
<div className="flex min-h-0 flex-1 overflow-hidden rounded-xl border bg-background shadow-sm"><RuleFlowCanvas graph={graph} issues={issues} selectedId={selectedId} selectedEdgeId={selectedEdgeId} focusTarget={focusTarget} onGraphChange={changeGraph} onSelect={setSelectedId} onSelectEdge={setSelectedEdgeId}/><NodeProperties node={selected} ipGroups={ipGroupsQuery.data ?? []} onChange={changeNode}/></div>
</div>
);
}
function EditorSkeleton() { return <div className="flex w-full flex-col gap-4 px-1 py-6"><div className="flex items-center gap-2"><Skeleton className="size-5"/><Skeleton className="h-8 w-64"/></div><Skeleton className="h-[70dvh] w-full"/></div>; }
+1
View File
@@ -76,6 +76,7 @@ export type {
WAFRuleGraph, WAFRuleGraph,
WAFRuleNode, WAFRuleNode,
WAFSaveRuleGraphPayload, WAFSaveRuleGraphPayload,
WAFUpdateRuleMetaPayload,
WAFRuleGroup, WAFRuleGroup,
WAFRuleGroupPayload, WAFRuleGroupPayload,
WAFSiteRuleGroups, WAFSiteRuleGroups,
+5
View File
@@ -773,6 +773,11 @@ export interface WAFSaveRuleGraphPayload {
graph: WAFRuleGraph; graph: WAFRuleGraph;
} }
export interface WAFUpdateRuleMetaPayload {
name: string;
enabled: boolean;
}
export interface WAFRuleGroupPayload { export interface WAFRuleGroupPayload {
name: string; name: string;
enabled: boolean; enabled: boolean;
@@ -9,6 +9,7 @@ import type {
WAFRule, WAFRule,
WAFSaveRuleGraphPayload, WAFSaveRuleGraphPayload,
WAFSiteRuleGroups, WAFSiteRuleGroups,
WAFUpdateRuleMetaPayload,
} from './types'; } from './types';
export class WafService extends OpenFlareBaseService { export class WafService extends OpenFlareBaseService {
@@ -33,12 +34,15 @@ export class WafService extends OpenFlareBaseService {
return this.post<WAFRule>(`/rule-groups/${id}/graph`, payload); return this.post<WAFRule>(`/rule-groups/${id}/graph`, payload);
} }
static async deleteRuleGroup(id: number): Promise<void> { static async updateRuleMeta(
return this.post<void>(`/rule-groups/${id}/delete`); id: number,
payload: WAFUpdateRuleMetaPayload,
): Promise<WAFRule> {
return this.post<WAFRule>(`/rule-groups/${id}/meta`, payload);
} }
static async updateRuleGroupSites(id: number, ids: number[]): Promise<WAFRule> { static async deleteRuleGroup(id: number): Promise<void> {
return this.post<WAFRule>(`/rule-groups/${id}/sites`, { ids }); return this.post<void>(`/rule-groups/${id}/delete`);
} }
static async listSiteRuleGroups(routeId: number): Promise<WAFSiteRuleGroups> { static async listSiteRuleGroups(routeId: number): Promise<WAFSiteRuleGroups> {
+22 -8
View File
@@ -6,7 +6,7 @@ import {beforeEach, describe, expect, it, vi} from 'vitest';
import WafPage from '@/app/(main)/waf/page'; import WafPage from '@/app/(main)/waf/page';
import apiClient from '@/lib/services/core/api-client'; 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 type {WAFSiteRuleGroups} from '@/lib/services/openflare';
import {WafService as DirectWafService} from '@/lib/services/openflare/waf.service'; import {WafService as DirectWafService} from '@/lib/services/openflare/waf.service';
@@ -32,11 +32,6 @@ vi.mock('@/lib/services/openflare', async (importOriginal) => {
listIPGroups: vi.fn(), listIPGroups: vi.fn(),
createRule: vi.fn(), createRule: vi.fn(),
deleteRuleGroup: vi.fn(), deleteRuleGroup: vi.fn(),
updateRuleGroupSites: vi.fn(),
},
ProxyRouteService: {
...actual.ProxyRouteService,
list: vi.fn(),
}, },
}; };
}); });
@@ -109,6 +104,27 @@ describe('WafService rule graph API', () => {
undefined, 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', () => { describe('WAF rule creation flow', () => {
@@ -118,8 +134,6 @@ describe('WAF rule creation flow', () => {
vi.mocked(WafService.listRuleGroups).mockResolvedValue([]); vi.mocked(WafService.listRuleGroups).mockResolvedValue([]);
vi.mocked(WafService.listIPGroups).mockReset(); vi.mocked(WafService.listIPGroups).mockReset();
vi.mocked(WafService.listIPGroups).mockResolvedValue([]); vi.mocked(WafService.listIPGroups).mockResolvedValue([]);
vi.mocked(ProxyRouteService.list).mockReset();
vi.mocked(ProxyRouteService.list).mockResolvedValue([]);
vi.mocked(WafService.createRule).mockReset(); vi.mocked(WafService.createRule).mockReset();
vi.mocked(WafService.createRule).mockResolvedValue(ruleSummary); vi.mocked(WafService.createRule).mockResolvedValue(ruleSummary);
}); });
+15
View File
@@ -24,6 +24,7 @@ const (
defaultRuntimeConfigDirRelativePath = "etc/openflare" defaultRuntimeConfigDirRelativePath = "etc/openflare"
defaultPagesDirRelativePath = "var/lib/openflare/pages" defaultPagesDirRelativePath = "var/lib/openflare/pages"
defaultMMDBRelativePath = "etc/openflare/GeoLite2-Country.mmdb" defaultMMDBRelativePath = "etc/openflare/GeoLite2-Country.mmdb"
defaultCityMMDBRelativePath = "etc/openflare/GeoLite2-City.mmdb"
defaultAccessLogRelativePath = "var/log/openflare/access.log" defaultAccessLogRelativePath = "var/log/openflare/access.log"
defaultStateRelativePath = "var/lib/openflare/agent-state.json" defaultStateRelativePath = "var/lib/openflare/agent-state.json"
defaultObservabilityBufferRelativePath = "var/lib/openflare/observability-buffer.json" defaultObservabilityBufferRelativePath = "var/lib/openflare/observability-buffer.json"
@@ -31,6 +32,7 @@ const (
defaultObservabilityReplayMinutes = 15 defaultObservabilityReplayMinutes = 15
defaultMMDBUpdateInterval = 24 * time.Hour defaultMMDBUpdateInterval = 24 * time.Hour
defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb" 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 defaultHeartbeatInterval = 10 * time.Second
defaultRequestTimeout = 10 * time.Second defaultRequestTimeout = 10 * time.Second
configFilePerm = 0o600 configFilePerm = 0o600
@@ -58,8 +60,10 @@ type Config struct {
RuntimeConfigDir string `json:"runtime_config_dir"` RuntimeConfigDir string `json:"runtime_config_dir"`
PagesDir string `json:"pages_dir"` PagesDir string `json:"pages_dir"`
MMDBPath string `json:"mmdb_path"` MMDBPath string `json:"mmdb_path"`
CityMMDBPath string `json:"city_mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"` MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"` MMDBDownloadURL string `json:"mmdb_download_url"`
CityMMDBDownloadURL string `json:"city_mmdb_download_url"`
OpenrestyObservabilityPort int `json:"openresty_observability_port"` OpenrestyObservabilityPort int `json:"openresty_observability_port"`
ObservabilityBufferPath string `json:"observability_buffer_path"` ObservabilityBufferPath string `json:"observability_buffer_path"`
ObservabilityReplayMinutes int `json:"observability_replay_minutes"` ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
@@ -89,8 +93,10 @@ type configFile struct {
RuntimeConfigDir string `json:"runtime_config_dir"` RuntimeConfigDir string `json:"runtime_config_dir"`
PagesDir string `json:"pages_dir"` PagesDir string `json:"pages_dir"`
MMDBPath string `json:"mmdb_path"` MMDBPath string `json:"mmdb_path"`
CityMMDBPath string `json:"city_mmdb_path"`
MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"` MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"`
MMDBDownloadURL string `json:"mmdb_download_url"` MMDBDownloadURL string `json:"mmdb_download_url"`
CityMMDBDownloadURL string `json:"city_mmdb_download_url"`
OpenrestyObservabilityPort int `json:"openresty_observability_port"` OpenrestyObservabilityPort int `json:"openresty_observability_port"`
ObservabilityBufferPath string `json:"observability_buffer_path"` ObservabilityBufferPath string `json:"observability_buffer_path"`
ObservabilityReplayMinutes int `json:"observability_replay_minutes"` ObservabilityReplayMinutes int `json:"observability_replay_minutes"`
@@ -178,6 +184,7 @@ func applyAgentPathDefaults(cfg *Config, baseDir string) {
{&cfg.RuntimeConfigDir, defaultRuntimeConfigDirRelativePath}, {&cfg.RuntimeConfigDir, defaultRuntimeConfigDirRelativePath},
{&cfg.PagesDir, defaultPagesDirRelativePath}, {&cfg.PagesDir, defaultPagesDirRelativePath},
{&cfg.MMDBPath, defaultMMDBRelativePath}, {&cfg.MMDBPath, defaultMMDBRelativePath},
{&cfg.CityMMDBPath, defaultCityMMDBRelativePath},
{&cfg.ObservabilityBufferPath, defaultObservabilityBufferRelativePath}, {&cfg.ObservabilityBufferPath, defaultObservabilityBufferRelativePath},
} }
for _, item := range pathDefaults { for _, item := range pathDefaults {
@@ -200,6 +207,9 @@ func applyAgentTimingDefaults(cfg *Config) {
if cfg.MMDBDownloadURL == "" { if cfg.MMDBDownloadURL == "" {
cfg.MMDBDownloadURL = defaultMMDBDownloadURL cfg.MMDBDownloadURL = defaultMMDBDownloadURL
} }
if cfg.CityMMDBDownloadURL == "" {
cfg.CityMMDBDownloadURL = defaultCityMMDBDownloadURL
}
if cfg.OpenrestyObservabilityPort <= 0 { if cfg.OpenrestyObservabilityPort <= 0 {
cfg.OpenrestyObservabilityPort = defaultOpenRestyObservabilityPort cfg.OpenrestyObservabilityPort = defaultOpenRestyObservabilityPort
} }
@@ -232,6 +242,7 @@ func normalizeManagedPaths(cfg *Config) {
&cfg.StatePath, &cfg.StatePath,
&cfg.ObservabilityBufferPath, &cfg.ObservabilityBufferPath,
&cfg.MMDBPath, &cfg.MMDBPath,
&cfg.CityMMDBPath,
} }
for _, p := range paths { for _, p := range paths {
if usesSlashPath(*p) { if usesSlashPath(*p) {
@@ -256,6 +267,8 @@ func hasEnvConfig() bool {
"OPENFLARE_MMDB_PATH", "OPENFLARE_MMDB_PATH",
"OPENFLARE_MMDB_UPDATE_INTERVAL", "OPENFLARE_MMDB_UPDATE_INTERVAL",
"OPENFLARE_MMDB_DOWNLOAD_URL", "OPENFLARE_MMDB_DOWNLOAD_URL",
"OPENFLARE_CITY_MMDB_PATH",
"OPENFLARE_CITY_MMDB_DOWNLOAD_URL",
} { } {
if strings.TrimSpace(os.Getenv(key)) != "" { if strings.TrimSpace(os.Getenv(key)) != "" {
return true return true
@@ -283,6 +296,8 @@ func applyEnvOverrides(cfg *Config) {
overrideString("OPENFLARE_PAGES_DIR", &cfg.PagesDir) overrideString("OPENFLARE_PAGES_DIR", &cfg.PagesDir)
overrideString("OPENFLARE_MMDB_PATH", &cfg.MMDBPath) overrideString("OPENFLARE_MMDB_PATH", &cfg.MMDBPath)
overrideString("OPENFLARE_MMDB_DOWNLOAD_URL", &cfg.MMDBDownloadURL) 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 value := strings.TrimSpace(os.Getenv("OPENFLARE_HEARTBEAT_INTERVAL")); value != "" {
if duration, err := parseDurationValue(value); err == nil { if duration, err := parseDurationValue(value); err == nil {
cfg.HeartbeatInterval = duration cfg.HeartbeatInterval = duration
+30
View File
@@ -61,6 +61,12 @@ func TestLoadDefaultsToManagedBinaryPaths(t *testing.T) {
if cfg.RuntimeConfigDir != filepath.Join(dir, "data", defaultRuntimeConfigDirRelativePath) { if cfg.RuntimeConfigDir != filepath.Join(dir, "data", defaultRuntimeConfigDirRelativePath) {
t.Fatalf("unexpected runtime config dir: %s", cfg.RuntimeConfigDir) 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 { if cfg.OpenrestyCertDir != cfg.CertDir {
t.Fatalf("unexpected openresty cert dir: %s", cfg.OpenrestyCertDir) 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_SERVER_URL", "http://new:3000")
t.Setenv("OPENFLARE_AGENT_TOKEN", "new-token") t.Setenv("OPENFLARE_AGENT_TOKEN", "new-token")
t.Setenv("OPENFLARE_OPENRESTY_PATH", "/new/openresty") 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) cfg, err := Load(configPath)
if err != nil { if err != nil {
@@ -347,6 +355,25 @@ func TestLoadEnvOverridesConfigFile(t *testing.T) {
if cfg.OpenrestyPath != "/new/openresty" { if cfg.OpenrestyPath != "/new/openresty" {
t.Fatalf("expected openresty path from env, got %s", cfg.OpenrestyPath) 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) { func TestLoadUsesMillisecondsForIntervals(t *testing.T) {
@@ -431,6 +458,9 @@ func TestSavePersistsMillisecondsAndOmitsRuntimeVersions(t *testing.T) {
if decoded["observability_replay_minutes"] != float64(defaultObservabilityReplayMinutes) { if decoded["observability_replay_minutes"] != float64(defaultObservabilityReplayMinutes) {
t.Fatalf("unexpected observability replay minutes: %#v", decoded["observability_replay_minutes"]) 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 { if _, ok := decoded["nginx_path"]; ok {
t.Fatal("legacy nginx_path should not be persisted") t.Fatal("legacy nginx_path should not be persisted")
} }
+66 -9
View File
@@ -3,6 +3,7 @@ package geoipupdate
import ( import (
"context" "context"
"errors"
"fmt" "fmt"
"io/fs" "io/fs"
"log/slog" "log/slog"
@@ -22,9 +23,12 @@ const (
// Updater periodically downloads a fresh GeoIP MMDB file and seeds the // Updater periodically downloads a fresh GeoIP MMDB file and seeds the
// initial embedded database when none is present on disk. // initial embedded database when none is present on disk.
type Updater struct { type Updater struct {
MMDBPath string MMDBPath string
DownloadURL string DownloadURL string
UpdateInterval time.Duration 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. // 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 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. // Run starts the periodic GeoIP update loop and blocks until ctx is cancelled.
func (u *Updater) Run(ctx context.Context) { func (u *Updater) Run(ctx context.Context) {
if u == nil || u.MMDBPath == "" || u.UpdateInterval <= 0 { if u == nil || u.MMDBPath == "" || u.UpdateInterval <= 0 {
return return
} }
if err := u.EnsureInitialDatabase(); err != nil { if err := u.EnsureInitialDatabases(ctx); err != nil {
slog.Warn("initialize GeoIP mmdb failed", "path", u.MMDBPath, "error", err) slog.Warn("initialize GeoIP databases failed", "country_path", u.MMDBPath, "city_path", u.CityMMDBPath, "error", err)
} }
ticker := time.NewTicker(u.UpdateInterval) ticker := time.NewTicker(u.UpdateInterval)
defer ticker.Stop() defer ticker.Stop()
@@ -67,11 +126,9 @@ func (u *Updater) Run(ctx context.Context) {
case <-ctx.Done(): case <-ctx.Done():
return return
case <-ticker.C: case <-ticker.C:
if err := geoip.DownloadMaxMindDatabase(ctx, u.MMDBPath, u.DownloadURL); err != nil { if err := u.updateDatabases(ctx); err != nil {
slog.Warn("update GeoIP mmdb failed", "path", u.MMDBPath, "error", err) slog.Warn("update GeoIP databases failed", "error", err)
continue
} }
slog.Info("GeoIP mmdb updated", "path", u.MMDBPath)
} }
} }
} }
@@ -1,8 +1,11 @@
package geoipupdate package geoipupdate
import ( import (
"context"
"errors"
"os" "os"
"path/filepath" "path/filepath"
"slices"
"testing" "testing"
) )
@@ -22,3 +25,77 @@ func TestEnsureInitialDatabaseCopiesEmbeddedMMDB(t *testing.T) {
t.Fatal("expected copied mmdb to be non-empty") 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)
}
}
+192 -11
View File
@@ -18,9 +18,12 @@ import (
"path/filepath" "path/filepath"
"regexp" "regexp"
"sort" "sort"
"strconv"
"strings" "strings"
"sync"
"time" "time"
sharedprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"github.com/Rain-kl/Wavelet/pkg/utils" "github.com/Rain-kl/Wavelet/pkg/utils"
@@ -31,11 +34,23 @@ import (
// RuntimeConfigDirPlaceholder is substituted into generated configs at apply time. // RuntimeConfigDirPlaceholder is substituted into generated configs at apply time.
const RuntimeConfigDirPlaceholder = "__OPENFLARE_RUNTIME_CONFIG_DIR__" 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. // ResolverDirectivePlaceholder is substituted into generated configs at apply time.
const ResolverDirectivePlaceholder = "__OPENFLARE_RESOLVER_DIRECTIVE__" const ResolverDirectivePlaceholder = "__OPENFLARE_RESOLVER_DIRECTIVE__"
// WAFIPGroupsConfigFileName is the runtime filename for synced WAF IP group data. // WAFIPGroupsConfigFileName is the runtime filename for synced WAF IP group data.
const WAFIPGroupsConfigFileName = "waf_ip_groups.json" 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 powConfigFileName = "pow_config.json"
const ( const (
@@ -45,6 +60,7 @@ const (
stubStatusCheckTimeout = 1500 * time.Millisecond stubStatusCheckTimeout = 1500 * time.Millisecond
nginxVersionSubmatchCount = 2 nginxVersionSubmatchCount = 2
resolverAddressCapacity = 2 resolverAddressCapacity = 2
workerInitSubmatchCount = 2
) )
// Executor controls OpenResty validation, reload, health, and lifecycle operations. // Executor controls OpenResty validation, reload, health, and lifecycle operations.
@@ -169,11 +185,15 @@ type Manager struct {
LuaDir string LuaDir string
NginxLuaDir string NginxLuaDir string
RuntimeConfigDir string RuntimeConfigDir string
MMDBPath string
CityMMDBPath string
PagesDir string PagesDir string
OpenrestyObservabilityListen string OpenrestyObservabilityListen string
OpenrestyObservabilityPort int OpenrestyObservabilityPort int
OpenrestyResolverDirective string OpenrestyResolverDirective string
Executor Executor Executor Executor
atomicFileWriter func(path string, data []byte, perm os.FileMode) error
wafIPGroupsMu sync.Mutex
} }
// ApplyStatus reports the outcome of an OpenResty configuration apply. // 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. // WAFIPGroupChecksums returns checksums for locally synced WAF IP groups.
func (m *Manager) WAFIPGroupChecksums() (map[string]string, error) { func (m *Manager) WAFIPGroupChecksums() (map[string]string, error) {
m.wafIPGroupsMu.Lock()
defer m.wafIPGroupsMu.Unlock()
config, err := m.readWAFIPGroupsRuntimeConfig() config, err := m.readWAFIPGroupsRuntimeConfig()
if err != nil { if err != nil {
return nil, err return nil, err
} }
if err = m.ensureWAFIPGroupsChecksum(); err != nil {
return nil, err
}
result := make(map[string]string, len(config.Groups)) result := make(map[string]string, len(config.Groups))
for id, group := range config.Groups { for id, group := range config.Groups {
if id != strconv.FormatUint(uint64(group.ID), 10) {
continue
}
if strings.TrimSpace(group.Checksum) != "" { if strings.TrimSpace(group.Checksum) != "" {
result[id] = 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 return result, nil
} }
// SyncWAFIPGroups writes WAF IP group definitions to the runtime config directory. // ReconcileWAFIPGroups atomically replaces the runtime snapshot with exactly the
func (m *Manager) SyncWAFIPGroups(groups []protocol.WAFIPGroup) error { // authoritative target IDs, retaining local definitions that did not change.
if m.RuntimeConfigDir == "" || len(groups) == 0 { func (m *Manager) ReconcileWAFIPGroups(targetIDs []uint, changed []protocol.WAFIPGroup) error {
if m.RuntimeConfigDir == "" {
return nil return nil
} }
m.wafIPGroupsMu.Lock()
defer m.wafIPGroupsMu.Unlock()
config, err := m.readWAFIPGroupsRuntimeConfig() config, err := m.readWAFIPGroupsRuntimeConfig()
if err != nil { if err != nil {
return err return err
} }
if config.Groups == nil { target := make(map[string]protocol.WAFIPGroup, len(targetIDs))
config.Groups = make(map[string]protocol.WAFIPGroup) targetSet := make(map[uint]struct{}, len(targetIDs))
} for _, id := range targetIDs {
for _, group := range groups { if id == 0 {
if group.ID == 0 {
continue 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 { if err != nil {
return err 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 { if err := os.MkdirAll(m.RuntimeConfigDir, nginxDirPerm); err != nil {
return err return err
} }
path := filepath.Join(m.RuntimeConfigDir, WAFIPGroupsConfigFileName) 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) 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)) slog.Info("synced waf ip groups", "path", path, "group_count", len(groups))
return nil 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) { func (m *Manager) readWAFIPGroupsRuntimeConfig() (*wafIPGroupsRuntimeConfig, error) {
config := &wafIPGroupsRuntimeConfig{Groups: map[string]protocol.WAFIPGroup{}} config := &wafIPGroupsRuntimeConfig{Groups: map[string]protocol.WAFIPGroup{}}
if m.RuntimeConfigDir == "" { if m.RuntimeConfigDir == "" {
@@ -1276,6 +1430,9 @@ func (m *Manager) renderMainConfig(content string) string {
} }
if luaDir := m.luaRuntimePath(); luaDir != "" { if luaDir := m.luaRuntimePath(); luaDir != "" {
rendered = strings.ReplaceAll(rendered, openrestyrender.LuaDirPlaceholder, 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 != "" { if listen := strings.TrimSpace(m.OpenrestyObservabilityListen); listen != "" {
rendered = strings.ReplaceAll(rendered, openrestyrender.ObservabilityListenPlaceholder, listen) rendered = strings.ReplaceAll(rendered, openrestyrender.ObservabilityListenPlaceholder, listen)
@@ -1289,6 +1446,19 @@ func (m *Manager) renderMainConfig(content string) string {
return rendered 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 { func (m *Manager) managedPowLuaFiles() []protocol.SupportFile {
files := ManagedPowLuaFiles() files := ManagedPowLuaFiles()
runtimeConfigDir := filepath.ToSlash(strings.TrimSpace(m.RuntimeConfigDir)) runtimeConfigDir := filepath.ToSlash(strings.TrimSpace(m.RuntimeConfigDir))
@@ -1301,8 +1471,19 @@ func (m *Manager) managedPowLuaFiles() []protocol.SupportFile {
func (m *Manager) managedWAFLuaFiles() []protocol.SupportFile { func (m *Manager) managedWAFLuaFiles() []protocol.SupportFile {
files := ManagedWAFLuaFiles() files := ManagedWAFLuaFiles()
runtimeConfigDir := filepath.ToSlash(strings.TrimSpace(m.RuntimeConfigDir)) 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 { for index := range files {
files[index].Content = strings.ReplaceAll(files[index].Content, RuntimeConfigDirPlaceholder, runtimeConfigDir) 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 return files
} }
+340 -14
View File
@@ -3,6 +3,7 @@ package nginx
import ( import (
"context" "context"
"errors" "errors"
"fmt"
"net" "net"
"net/http" "net/http"
"os" "os"
@@ -13,6 +14,7 @@ import (
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol" "github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
sharedprotocol "github.com/Rain-kl/Wavelet/pkg/protocol"
) )
type runCall struct { 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) { func TestManagerCheckHealthUsesStubStatusInsteadOfConfigTest(t *testing.T) {
listener, err := net.Listen("tcp", "127.0.0.1:0") listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil { 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 { if _, err := os.Stat(filepath.Join(manager.LuaDir, "pow", "check.lua")); err != nil {
t.Fatalf("failed to stat pow lua file: %v", err) 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 { if err != nil {
t.Fatalf("failed to read pow lua file: %v", err) t.Fatalf("failed to read pow lua file: %v", err)
} }
if !strings.Contains(string(data), filepath.ToSlash(manager.RuntimeConfigDir)+"/waf_config.json") { if !strings.Contains(string(data), filepath.ToSlash(manager.RuntimeConfigDir)) || !strings.Contains(string(data), `runtime_dir .. "/waf_config.json"`) {
t.Fatalf("expected pow lua to read runtime config dir, got %s", string(data)) 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) { 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") t.Fatal("expected pow runtime lua to internally execute make-challenge instead of issuing a 302 redirect")
} }
if strings.Contains(openRestyPowRuntimeLua, "ngx.redirect(") { if strings.Contains(openRestyPowRuntimeLua, "ngx.redirect(") {
@@ -699,12 +763,28 @@ func TestManagedPowLuaFilesUseInternalChallengeFlow(t *testing.T) {
} }
} }
func TestManagedWAFLuaTreatsWhitelistAsBypass(t *testing.T) { func TestManagedPowLuaFilesPreserveConfigAcrossInternalRedirect(t *testing.T) {
if !strings.Contains(openRestyWAFRuntimeLua, "if ip_matches(group.ip_whitelist, ip)") { for _, expected := range []string{
t.Fatal("expected waf runtime to bypass request when ip matches whitelist") `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") { if !strings.Contains(openRestyPowChallengeLua, `pow_config_dict:get(config_key)`) {
t.Fatal("expected waf runtime not to block requests that miss configured whitelists") 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()} 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"}, {ID: 1, Enabled: true, IPList: []string{"203.0.113.10"}, Checksum: "sum-1"},
}); err != nil { }); 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"}, {ID: 2, Enabled: true, IPList: []string{"198.51.100.10"}, Checksum: "sum-2"},
}); err != nil { }); err != nil {
t.Fatalf("SyncWAFIPGroups second delta failed: %v", err) t.Fatalf("ReconcileWAFIPGroups second delta failed: %v", err)
} }
checksums, err := manager.WAFIPGroupChecksums() 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) { func TestObservabilityListenAddress(t *testing.T) {
if got := ObservabilityListenAddress(18081); got != "127.0.0.1:18081" { if got := ObservabilityListenAddress(18081); got != "127.0.0.1:18081" {
t.Fatalf("unexpected default observability listen address: %s", got) t.Fatalf("unexpected default observability listen address: %s", got)
+87 -21
View File
@@ -13,7 +13,81 @@ var powStaticFS embed.FS
const openRestyPowRuntimeLua = `local _M = {} 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() 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 "" local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then if string.sub(source, 1, 1) == "@" then
local script_path = string.sub(source, 2) local script_path = string.sub(source, 2)
@@ -198,7 +272,7 @@ return ngx.exec("/.within.website/x/cmd/anubis/api/make-challenge")
end end
return _M return _M
` */
const openRestyPowCheckLua = `local source = debug.getinfo(1, "S").source or "" const openRestyPowCheckLua = `local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then if string.sub(source, 1, 1) == "@" then
@@ -214,8 +288,8 @@ return require("pow.runtime").check()
const openRestyPowChallengeLua = `local cjson = require "cjson.safe" 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_challenges = ngx.shared.openflare_pow_challenges
local pow_config_dict = ngx.shared.openflare_pow_config
local function generate_entropy() local function generate_entropy()
local pieces = { local pieces = {
@@ -233,28 +307,20 @@ local args = ngx.req.get_uri_args()
local host = args["host"] or ngx.var.host or "" local host = args["host"] or ngx.var.host or ""
local redir = args["redir"] or "" local redir = args["redir"] or ""
local site = ngx.var.openflare_waf_site or "" local config = ngx.ctx.openflare_pow_config
if site == "" then 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.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 return
end 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 difficulty = config.difficulty or 4
local algorithm = config.algorithm or "fast" local algorithm = config.algorithm or "fast"
local challenge_ttl = config.challenge_ttl or 300 local challenge_ttl = config.challenge_ttl or 300
+9 -270
View File
@@ -1,278 +1,16 @@
package nginx 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() //go:embed waf_runtime.lua
local cjson = require "cjson.safe" var openRestyWAFRuntimeLua string
local config_dict = ngx.shared.openflare_waf_config //go:embed waf_ip_groups.lua
var openRestyWAFIPGroupsLua string
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
`
const openRestyWAFCheckLua = `local source = debug.getinfo(1, "S").source or "" const openRestyWAFCheckLua = `local source = debug.getinfo(1, "S").source or ""
if string.sub(source, 1, 1) == "@" then if string.sub(source, 1, 1) == "@" then
@@ -290,6 +28,7 @@ return require("waf.runtime").check()
func ManagedWAFLuaFiles() []protocol.SupportFile { func ManagedWAFLuaFiles() []protocol.SupportFile {
return []protocol.SupportFile{ return []protocol.SupportFile{
{Path: "waf/runtime.lua", Content: openRestyWAFRuntimeLua}, {Path: "waf/runtime.lua", Content: openRestyWAFRuntimeLua},
{Path: "waf/ip_groups.lua", Content: openRestyWAFIPGroupsLua},
{Path: "waf/check.lua", Content: openRestyWAFCheckLua}, {Path: "waf/check.lua", Content: openRestyWAFCheckLua},
} }
} }
@@ -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)
}
}
+150
View File
@@ -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
@@ -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
+404
View File
@@ -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
@@ -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
+54 -25
View File
@@ -9,6 +9,7 @@ import (
"fmt" "fmt"
"log/slog" "log/slog"
"sort" "sort"
"strconv"
"strings" "strings"
"sync" "sync"
@@ -42,7 +43,8 @@ type NginxManager interface {
EnsureSafeFallbackRuntime(ctx context.Context, reason string) error EnsureSafeFallbackRuntime(ctx context.Context, reason string) error
CurrentChecksum() (string, error) CurrentChecksum() (string, error)
WAFIPGroupChecksums() (map[string]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 EnsureWorkerReadAccess() error
} }
@@ -132,12 +134,12 @@ func (s *Service) WAFIPGroupChecksums() (map[string]string, error) {
return s.nginxManager.WAFIPGroupChecksums() 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 { func (s *Service) ApplyWAFIPGroups(_ context.Context, groups []protocol.WAFIPGroup) error {
if len(groups) == 0 || s.nginxManager == nil { if len(groups) == 0 || s.nginxManager == nil {
return 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 { 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 { 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 { if len(ids) == 0 {
return nil if s.nginxManager == nil {
return nil
}
return s.nginxManager.ReconcileWAFIPGroups([]uint{}, nil)
} }
checksums, err := s.WAFIPGroupChecksums() checksums, err := s.WAFIPGroupChecksums()
if err != nil { if err != nil {
return err 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{ response, err := s.client.SyncWAFIPGroups(ctx, protocol.WAFIPGroupSyncRequest{
IDs: ids, IDs: ids,
Checksums: checksums, Checksums: targetChecksums,
}) })
if err != nil { if err != nil {
return err return err
} }
if response == nil || len(response.Groups) == 0 { if s.nginxManager == nil {
return 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 { type renderedActiveConfig struct {
@@ -288,7 +307,7 @@ func fromOpenRestySupportFiles(files []openrestyrender.SupportFile) []protocol.S
return result return result
} }
func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) []uint { func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) ([]uint, error) {
var content string var content string
for _, file := range supportFiles { for _, file := range supportFiles {
if file.Path == "waf_config.json" { if file.Path == "waf_config.json" {
@@ -297,28 +316,38 @@ func referencedWAFIPGroupIDs(supportFiles []protocol.SupportFile) []uint {
} }
} }
if content == "" { if content == "" {
return []uint{} return []uint{}, nil
}
var payload struct {
RuleGroups []struct {
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids"`
} `json:"rule_groups"`
} }
var payload openrestyrender.WAFDocument
if err := json.Unmarshal([]byte(content), &payload); err != nil { if err := json.Unmarshal([]byte(content), &payload); err != nil {
slog.Debug("decode waf_config.json for ip group references failed", "error", err) return nil, fmt.Errorf("decode waf_config.json for ip group references: %w", err)
return []uint{}
} }
seen := make(map[uint]struct{}) seen := make(map[uint]struct{})
for _, group := range payload.RuleGroups { for _, group := range payload.RuleGroups {
for _, id := range group.IPWhitelistGroups { for _, legacyIDs := range [][]uint{group.IPWhitelistGroups, group.IPBlacklistGroups} {
if id > 0 { for _, id := range legacyIDs {
seen[id] = struct{}{} if id > 0 {
seen[id] = struct{}{}
}
} }
} }
for _, id := range group.IPBlacklistGroups { for nodeID, node := range group.Graph.Nodes {
if id > 0 { if node.Type != "ip_match" {
seen[id] = struct{}{} 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) ids = append(ids, id)
} }
sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] }) sort.Slice(ids, func(i, j int) bool { return ids[i] < ids[j] })
return ids return ids, nil
} }
func shouldReportNoopApply(snapshot *state.Snapshot, version string, checksum string) bool { func shouldReportNoopApply(snapshot *state.Snapshot, version string, checksum string) bool {
+158 -3
View File
@@ -30,8 +30,10 @@ func testPagesSourceConfigJSON(deploymentID uint, checksum string) string {
type fakeClient struct { type fakeClient struct {
config protocol.ActiveConfigResponse config protocol.ActiveConfigResponse
reports []protocol.ApplyLogPayload reports []protocol.ApplyLogPayload
wafSyncCalls []protocol.WAFIPGroupSyncRequest
pagesPackages map[uint][]byte pagesPackages map[uint][]byte
pagesHashes map[uint]string pagesHashes map[uint]string
wafSyncResult protocol.WAFIPGroupSyncResponse
fetchCalls int fetchCalls int
hashCalls int hashCalls int
} }
@@ -47,6 +49,12 @@ type fakeManager struct {
applyMainContents []string applyMainContents []string
applyRouteContents []string applyRouteContents []string
applyFiles [][]protocol.SupportFile 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 { 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) { 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 { 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) { 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 return nil
} }
@@ -145,6 +163,143 @@ func (m *fakeManager) EnsureWorkerReadAccess() error {
return nil 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) { func TestSyncOnceSuccess(t *testing.T) {
client := &fakeClient{ client := &fakeClient{
config: protocol.ActiveConfigResponse{ config: protocol.ActiveConfigResponse{
+79 -24
View File
@@ -4,6 +4,7 @@
package agent package agent
import ( import (
"bytes"
"context" "context"
"crypto/sha256" "crypto/sha256"
"encoding/hex" "encoding/hex"
@@ -11,43 +12,32 @@ import (
"errors" "errors"
"fmt" "fmt"
"sort" "sort"
"strconv"
"strings" "strings"
"github.com/Rain-kl/Wavelet/internal/model" "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 { 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. // WAFIPGroupsForAgent builds agent-facing WAF IP group payloads for the given ids.
func WAFIPGroupsForAgent(ctx context.Context, ids []uint) ([]WAFIPGroup, error) { 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. // 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) { func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[string]string) ([]WAFIPGroup, error) {
targetIDs := uniqueUintIDs(ids) groups, err := validatedAgentWAFIPGroups(ctx, ids, true)
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)
if err != nil { if err != nil {
return nil, err return nil, err
} }
@@ -61,6 +51,45 @@ func ChangedWAFIPGroupsForAgent(ctx context.Context, ids []uint, checksums map[s
return changed, nil 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) { func buildAgentWAFIPGroups(ctx context.Context, ids []uint) ([]WAFIPGroup, error) {
ids = uniqueUintIDs(ids) ids = uniqueUintIDs(ids)
if len(ids) == 0 { if len(ids) == 0 {
@@ -142,6 +171,8 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
} }
idSet := make(map[uint]struct{}) idSet := make(map[uint]struct{})
for _, group := range snapshot.WAF.RuleGroups { 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 { for _, id := range group.IPWhitelistGroups {
if id > 0 { if id > 0 {
idSet[id] = struct{}{} idSet[id] = struct{}{}
@@ -152,6 +183,20 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
idSet[id] = struct{}{} 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)) ids := make([]uint, 0, len(idSet))
for id := range idSet { for id := range idSet {
@@ -161,6 +206,16 @@ func activeConfigWAFIPGroupIDs(ctx context.Context) ([]uint, error) {
return ids, nil 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) { func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, error) {
text := strings.TrimSpace(snapshotJSON) text := strings.TrimSpace(snapshotJSON)
if text == "" { if text == "" {
@@ -171,7 +226,7 @@ func parseActiveConfigSnapshot(snapshotJSON string) (*activeConfigSnapshot, erro
return nil, err return nil, err
} }
if snapshot.WAF.RuleGroups == nil { if snapshot.WAF.RuleGroups == nil {
snapshot.WAF.RuleGroups = []snapshotWAFRuleGroupRef{} snapshot.WAF.RuleGroups = []openrestyrender.WAFRuleGroup{}
} }
return &snapshot, nil return &snapshot, nil
} }
@@ -7,10 +7,12 @@ import (
"context" "context"
"encoding/json" "encoding/json"
"strconv" "strconv"
"strings"
"testing" "testing"
"github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/protocol"
"github.com/glebarez/sqlite" "github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -63,6 +65,115 @@ func seedActiveConfigWithWAFIPGroup(t *testing.T, ctx context.Context, ipGroupID
}).Error) }).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) { func TestChangedWAFIPGroupsForAgentReturnsChecksumDelta(t *testing.T) {
cleanup := setupWAFIPGroupTestDB(t) cleanup := setupWAFIPGroupTestDB(t)
defer cleanup() defer cleanup()
@@ -13,6 +13,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf" "github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
"github.com/Rain-kl/Wavelet/internal/db" "github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"github.com/glebarez/sqlite" "github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert" "github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
@@ -151,13 +152,7 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx) globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err) require.NoError(t, err)
customGroup := &model.OpenFlareWAFRuleGroup{ customGroup := createSnapshotRule(t, ctx, "pow-group", waf.DefaultRuleGraph())
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))
require.NoError(t, model.ReplaceOpenFlareWAFRuleGroupBindings(ctx, customGroup.ID, []uint{route.ID})) require.NoError(t, model.ReplaceOpenFlareWAFRuleGroupBindings(ctx, customGroup.ID, []uint{route.ID}))
bundle, err := buildCurrentConfigBundle(ctx, true) bundle, err := buildCurrentConfigBundle(ctx, true)
@@ -177,20 +172,24 @@ func TestBuildSnapshotWAFDocumentUsesNormalizedSiteNames(t *testing.T) {
} }
assert.True(t, found, "expected WAF binding for enabled route") assert.True(t, found, "expected WAF binding for enabled route")
var wafRuntime struct { var wafRuntime openrestyrender.WAFDocument
SiteRuleGroups map[string][]uint `json:"site_rule_groups"` foundWAFConfig := false
}
for _, file := range bundle.SupportFiles { for _, file := range bundle.SupportFiles {
if file.Path != "waf_config.json" { if file.Path != "waf_config.json" {
continue continue
} }
foundWAFConfig = true
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime)) require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
} }
require.Contains(t, wafRuntime.SiteRuleGroups, "example.com") require.True(t, foundWAFConfig, "expected rendered WAF support file")
require.Contains(t, wafRuntime.SiteRuleGroups["example.com"], customGroup.ID) require.NotEmpty(t, wafRuntime.RuleGroups)
require.Contains(t, wafRuntime.SiteRuleGroups["example.com"], globalGroup.ID) 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, `set $openflare_waf_site "example.com"`)
assert.Contains(t, bundle.RouteConfig, `require("pow.runtime").check()`)
} }
func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testing.T) { func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testing.T) {
@@ -210,34 +209,30 @@ func TestBuildCurrentConfigBundleEnablesGlobalPoWWithoutExplicitBinding(t *testi
require.NoError(t, waf.EnsureDefaultRuleGroup(ctx)) require.NoError(t, waf.EnsureDefaultRuleGroup(ctx))
globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx) globalGroup, err := model.GetGlobalOpenFlareWAFRuleGroup(ctx)
require.NoError(t, err) require.NoError(t, err)
globalGroup.PoWEnabled = true graphJSON, err := json.Marshal(snapshotPoWGraph())
globalGroup.PoWConfig = `{"difficulty":4,"algorithm":"fast","session_ttl":600,"challenge_ttl":300}` require.NoError(t, err)
require.NoError(t, model.UpdateOpenFlareWAFRuleGroup(ctx, globalGroup)) globalGroup.Graph = string(graphJSON)
require.NoError(t, db.DB(ctx).Model(globalGroup).Update("graph", globalGroup.Graph).Error)
bundle, err := buildCurrentConfigBundle(ctx, true) bundle, err := buildCurrentConfigBundle(ctx, true)
require.NoError(t, err) require.NoError(t, err)
assert.Contains(t, bundle.RouteConfig, `require("pow.runtime").check()`) var wafRuntime openrestyrender.WAFDocument
foundWAFConfig := false
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"`
}
for _, file := range bundle.SupportFiles { for _, file := range bundle.SupportFiles {
if file.Path != "waf_config.json" { if file.Path != "waf_config.json" {
continue continue
} }
foundWAFConfig = true
require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime)) require.NoError(t, json.Unmarshal([]byte(file.Content), &wafRuntime))
} }
require.Contains(t, wafRuntime.SiteRuleGroups, "pow-global.example.com") require.True(t, foundWAFConfig, "expected rendered WAF support file")
require.Contains(t, wafRuntime.SiteRuleGroups["pow-global.example.com"], globalGroup.ID)
require.NotEmpty(t, wafRuntime.RuleGroups) require.NotEmpty(t, wafRuntime.RuleGroups)
assert.True(t, wafRuntime.RuleGroups[0].PoWEnabled) assert.Equal(t, globalGroup.ID, wafRuntime.RuleGroups[0].ID)
require.NotNil(t, wafRuntime.RuleGroups[0].PoWConfig) assert.True(t, wafRuntime.RuleGroups[0].IsGlobal)
assert.Equal(t, 4, wafRuntime.RuleGroups[0].PoWConfig.Difficulty) 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)
} }
@@ -9,17 +9,21 @@ import (
"errors" "errors"
"fmt" "fmt"
"sort" "sort"
"strconv"
"strings" "strings"
oftls "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls" oftls "github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/waf" "github.com/Rain-kl/Wavelet/internal/apps/openflare/waf"
"github.com/Rain-kl/Wavelet/internal/model" "github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository" "github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/protocol"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty" openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"gorm.io/gorm"
) )
const ( const (
supportFilesPerCertificate = 2 supportFilesPerCertificate = 2
wafIPGroupChecksumHexLength = 64
// OpenResty 默认配置值 // OpenResty 默认配置值
defaultOpenRestyReturnStatus = 421 defaultOpenRestyReturnStatus = 421
@@ -66,22 +70,11 @@ type snapshotRoute struct {
} }
type snapshotWAFRuleGroup struct { type snapshotWAFRuleGroup struct {
ID uint `json:"id"` ID uint `json:"id"`
Name string `json:"name"` Name string `json:"name"`
Enabled bool `json:"enabled"` Enabled bool `json:"enabled"`
IsGlobal bool `json:"is_global"` IsGlobal bool `json:"is_global"`
BlockStatusCode int `json:"block_status_code"` Graph waf.RuntimeRuleGraph `json:"graph"`
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"`
} }
type snapshotWAFIPGroup struct { type snapshotWAFIPGroup struct {
@@ -311,35 +304,37 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
if err := waf.EnsureDefaultRuleGroup(ctx); err != nil { if err := waf.EnsureDefaultRuleGroup(ctx); err != nil {
return snapshotWAFDocument{}, err return snapshotWAFDocument{}, err
} }
views, err := waf.ListRuleGroups(ctx) groups, err := model.ListOpenFlareWAFRuleGroups(ctx)
if err != nil { if err != nil {
return snapshotWAFDocument{}, err return snapshotWAFDocument{}, err
} }
ruleGroups := make([]snapshotWAFRuleGroup, 0, len(views)) ruleGroups := make([]snapshotWAFRuleGroup, 0, len(groups))
for _, view := range views { referencedIPGroupIDs := make(map[uint]struct{})
if !view.Enabled { enabledRuleIDs := make(map[uint]struct{})
for _, group := range groups {
if !group.Enabled {
continue 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{ ruleGroups = append(ruleGroups, snapshotWAFRuleGroup{
ID: view.ID, ID: group.ID, Name: group.Name, Enabled: group.Enabled, IsGlobal: group.IsGlobal, Graph: runtimeGraph,
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),
}) })
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 { if err != nil {
return snapshotWAFDocument{}, err return snapshotWAFDocument{}, err
} }
@@ -366,16 +361,16 @@ func buildSnapshotWAFDocument(ctx context.Context, routes []*model.ProxyRoute) (
if _, ok := enabledRouteSiteNames[binding.ProxyRouteID]; !ok { if _, ok := enabledRouteSiteNames[binding.ProxyRouteID]; !ok {
continue 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)) bindings := make([]snapshotWAFBinding, 0, len(enabledRouteSiteNames))
for routeID, siteName := range 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{ bindings = append(bindings, snapshotWAFBinding{
RouteID: routeID, RouteID: routeID,
SiteName: siteName, SiteName: siteName,
RuleGroupIDs: groupIDs, RuleGroupIDs: groupIDsByRoute[routeID],
}) })
} }
sort.Slice(bindings, func(i, j int) bool { 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 return snapshotWAFDocument{RuleGroups: ruleGroups, IPGroups: ipGroups, Bindings: bindings}, nil
} }
func buildSnapshotWAFIPGroups(ctx context.Context, ruleGroups []snapshotWAFRuleGroup) ([]snapshotWAFIPGroup, error) { func validateSnapshotWAFIPGroupSize(groups []snapshotWAFIPGroup) error {
idSet := make(map[uint]struct{}) runtimeGroups := make(map[string]protocol.WAFIPGroup, len(groups))
for _, group := range ruleGroups { for _, group := range groups {
for _, id := range group.IPWhitelistGroups { ipList := group.IPList
idSet[id] = struct{}{} if !group.Enabled {
ipList = []string{}
} }
for _, id := range group.IPBlacklistGroups { runtimeGroups[strconv.FormatUint(uint64(group.ID), 10)] = protocol.WAFIPGroup{
idSet[id] = struct{}{} 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 { if len(idSet) == 0 {
return []snapshotWAFIPGroup{}, nil return []snapshotWAFIPGroup{}, nil
} }
@@ -431,9 +436,20 @@ func buildSnapshotWAFIPGroups(ctx context.Context, ruleGroups []snapshotWAFRuleG
IPList: ipList, IPList: ipList,
}) })
} }
if err = validateSnapshotWAFIPGroupSize(snapshots); err != nil {
return nil, err
}
return snapshots, nil 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) { func decodeIPList(raw string) ([]string, error) {
text := strings.TrimSpace(raw) text := strings.TrimSpace(raw)
if text == "" { if text == "" {
@@ -446,36 +462,6 @@ func decodeIPList(raw string) ([]string, error) {
return items, nil 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 { func buildOpenRestyConfigSnapshot(ctx context.Context) openRestyConfigSnapshot {
// 读取所有 OpenResty 配置,使用默认值作为降级 // 读取所有 OpenResty 配置,使用默认值作为降级
getIntConfig := func(key string, defaultVal int) int { getIntConfig := func(key string, defaultVal int) int {
@@ -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"},
}}
}
@@ -109,12 +109,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
t.Run("WAF rule group create", func(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{ rec := performJSONRequest(t, engine, http.MethodPost, apiPath("/waf/rule-groups"), map[string]any{
"name": "edge-security", "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"},
}, adminAuthHeaders(seed.Token)) }, adminAuthHeaders(seed.Token))
require.Equal(t, http.StatusOK, rec.Code) require.Equal(t, http.StatusOK, rec.Code)
@@ -124,7 +119,8 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
assert.NotZero(t, ruleGroupID) assert.NotZero(t, ruleGroupID)
assert.Equal(t, "edge-security", data["name"]) assert.Equal(t, "edge-security", data["name"])
assert.Equal(t, false, data["is_global"]) 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) { 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, t,
engine, engine,
http.MethodPost, 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{ map[string]any{
"name": "edge-security-updated", "name": "edge-security-updated", "enabled": true,
"enabled": true,
"block_status_code": 451,
}, },
adminAuthHeaders(seed.Token), adminAuthHeaders(seed.Token),
) )
@@ -187,7 +181,7 @@ func TestSecurityWAFTLSMigrationFlow(t *testing.T) {
resp := requireAPIOK(t, rec) resp := requireAPIOK(t, rec)
data := unmarshalAPIMap(t, resp.Data) data := unmarshalAPIMap(t, resp.Data)
assert.Equal(t, "edge-security-updated", data["name"]) 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) { t.Run("WAF IP group create", func(t *testing.T) {
+22 -5
View File
@@ -5,25 +5,35 @@ package waf
import "encoding/json" import "encoding/json"
// RuleGraphSchemaVersion is the current persisted rule graph schema version.
const RuleGraphSchemaVersion = 1 const RuleGraphSchemaVersion = 1
// RuleNodeType identifies the behavior of a rule graph node.
type RuleNodeType string type RuleNodeType string
const ( const (
RuleNodeStart RuleNodeType = "start" // RuleNodeStart begins graph execution.
RuleNodeAllow RuleNodeType = "allow" RuleNodeStart RuleNodeType = "start"
RuleNodeBlock RuleNodeType = "block" // RuleNodeAllow terminates execution with an allow decision.
RuleNodeIPMatch RuleNodeType = "ip_match" 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" 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 { type RuleGraph struct {
SchemaVersion int `json:"schema_version"` SchemaVersion int `json:"schema_version"`
Nodes []RuleNode `json:"nodes"` Nodes []RuleNode `json:"nodes"`
Edges []RuleEdge `json:"edges"` Edges []RuleEdge `json:"edges"`
} }
// RuleNode stores one editor node and its type-specific configuration.
type RuleNode struct { type RuleNode struct {
ID string `json:"id"` ID string `json:"id"`
Type RuleNodeType `json:"type"` Type RuleNodeType `json:"type"`
@@ -32,11 +42,13 @@ type RuleNode struct {
Config json.RawMessage `json:"config"` Config json.RawMessage `json:"config"`
} }
// RulePosition stores a node's editor canvas coordinates.
type RulePosition struct { type RulePosition struct {
X float64 `json:"x"` X float64 `json:"x"`
Y float64 `json:"y"` Y float64 `json:"y"`
} }
// RuleEdge connects one source handle to a target node.
type RuleEdge struct { type RuleEdge struct {
ID string `json:"id"` ID string `json:"id"`
Source string `json:"source"` Source string `json:"source"`
@@ -44,17 +56,20 @@ type RuleEdge struct {
Target string `json:"target"` Target string `json:"target"`
} }
// IPMatchConfig configures literal, CIDR, and managed-group IP matching.
type IPMatchConfig struct { type IPMatchConfig struct {
IPs []string `json:"ips,omitempty"` IPs []string `json:"ips,omitempty"`
CIDRs []string `json:"cidrs,omitempty"` CIDRs []string `json:"cidrs,omitempty"`
IPGroupIDs []uint `json:"ip_group_ids,omitempty"` IPGroupIDs []uint `json:"ip_group_ids,omitempty"`
} }
// GeoMatchConfig configures country and region matching.
type GeoMatchConfig struct { type GeoMatchConfig struct {
Countries []string `json:"countries,omitempty"` Countries []string `json:"countries,omitempty"`
Regions []string `json:"regions,omitempty"` Regions []string `json:"regions,omitempty"`
} }
// PoWNodeConfig configures a proof-of-work challenge node.
type PoWNodeConfig struct { type PoWNodeConfig struct {
Algorithm string `json:"algorithm"` Algorithm string `json:"algorithm"`
Difficulty int `json:"difficulty"` Difficulty int `json:"difficulty"`
@@ -62,11 +77,13 @@ type PoWNodeConfig struct {
ChallengeTTL int `json:"challenge_ttl"` ChallengeTTL int `json:"challenge_ttl"`
} }
// BlockNodeConfig configures a terminal blocking response.
type BlockNodeConfig struct { type BlockNodeConfig struct {
StatusCode int `json:"status_code"` StatusCode int `json:"status_code"`
ResponseBody string `json:"response_body,omitempty"` ResponseBody string `json:"response_body,omitempty"`
} }
// DefaultRuleGraph returns the minimal start-to-allow graph.
func DefaultRuleGraph() RuleGraph { func DefaultRuleGraph() RuleGraph {
return RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{ return RuleGraph{SchemaVersion: RuleGraphSchemaVersion, Nodes: []RuleNode{
{ID: "start", Type: RuleNodeStart, Position: RulePosition{X: 0, Y: 0}, Config: json.RawMessage(`{}`)}, {ID: "start", Type: RuleNodeStart, Position: RulePosition{X: 0, Y: 0}, Config: json.RawMessage(`{}`)},
+164 -95
View File
@@ -26,7 +26,33 @@ var (
regionCodePattern = regexp.MustCompile(`^[A-Z]{2}-[A-Z0-9]{1,3}$`) 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 { 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 { if graph.SchemaVersion != RuleGraphSchemaVersion {
return fmt.Errorf("规则图 schema_version 必须为 %d", 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 { } else if len(raw) > maxRuleGraphBytes {
return fmt.Errorf("规则图大小不能超过 256 KiB") 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, "" startCount, allowCount, startID := 0, 0, ""
for _, node := range graph.Nodes { for _, node := range graphNodes {
if strings.TrimSpace(node.ID) == "" { if strings.TrimSpace(node.ID) == "" {
return errors.New("节点 ID 不能为空") return nil, "", errors.New("节点 ID 不能为空")
} }
if _, exists := nodes[node.ID]; exists { 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 nodes[node.ID] = node
switch node.Type { switch node.Type {
@@ -60,67 +89,74 @@ func ValidateRuleGraph(ctx context.Context, graph RuleGraph, ipGroupExists func(
allowCount++ allowCount++
case RuleNodeBlock, RuleNodeIPMatch, RuleNodeGeoMatch, RuleNodePoW: case RuleNodeBlock, RuleNodeIPMatch, RuleNodeGeoMatch, RuleNodePoW:
default: 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 { if err := validateRuleNodeConfig(ctx, node, ipGroupExists); err != nil {
return err return nil, "", err
} }
} }
if startCount != 1 { if startCount != 1 {
return errors.New("规则图必须恰好包含一个开始节点") return nil, "", errors.New("规则图必须恰好包含一个开始节点")
} }
if allowCount != 1 { 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) outgoing := make(map[string][]RuleEdge)
incoming := make(map[string]int) incoming := make(map[string]int)
handleTargets := make(map[string]int) handleTargets := make(map[string]int)
for _, edge := range graph.Edges { for _, edge := range graphEdges {
if strings.TrimSpace(edge.ID) == "" { if strings.TrimSpace(edge.ID) == "" {
return errors.New("边 ID 不能为空") return nil, nil, nil, errors.New("边 ID 不能为空")
} }
if _, exists := edgeIDs[edge.ID]; exists { 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{}{} edgeIDs[edge.ID] = struct{}{}
source, ok := nodes[edge.Source] source, ok := nodes[edge.Source]
if !ok { 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 { 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) { 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 key := edge.Source + "\x00" + edge.SourceHandle
handleTargets[key]++ handleTargets[key]++
if handleTargets[key] > 1 { 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) outgoing[edge.Source] = append(outgoing[edge.Source], edge)
incoming[edge.Target]++ incoming[edge.Target]++
} }
if hasRuleGraphCycle(nodes, outgoing, incoming) { return outgoing, incoming, handleTargets, nil
return errors.New("规则图不能包含循环") }
}
for _, node := range graph.Nodes { func validateRequiredHandles(nodes []RuleNode, handleTargets map[string]int) error {
for _, node := range nodes {
for _, handle := range requiredHandles(node.Type) { for _, handle := range requiredHandles(node.Type) {
if handleTargets[node.ID+"\x00"+handle] == 0 { if handleTargets[node.ID+"\x00"+handle] == 0 {
return fmt.Errorf("节点 %s 的 %s 出口未连接", node.ID, handle) 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) reachable := walkRuleGraph(startID, outgoing)
for _, node := range graph.Nodes { for _, node := range nodes {
if !reachable[node.ID] { if !reachable[node.ID] {
return fmt.Errorf("节点 %s 无法从开始节点到达", 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 { if node.Type == RuleNodeStart && incoming[node.ID] != 0 {
return fmt.Errorf("开始节点 %s 不能有入边", node.ID) 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) return fmt.Errorf("终止节点 %s 不能有出口", node.ID)
} }
} }
if err := validateTerminalPaths(graph.Nodes, outgoing); err != nil {
return err
}
return nil return nil
} }
func validateRuleNodeConfig(ctx context.Context, node RuleNode, exists func(context.Context, uint) (bool, error)) error { func validateRuleNodeConfig(ctx context.Context, node RuleNode, exists func(context.Context, uint) (bool, error)) error {
switch node.Type { switch node.Type {
case RuleNodeStart, RuleNodeAllow: case RuleNodeStart, RuleNodeAllow:
var cfg struct{} return validateEmptyNodeConfig(node)
if err := decodeStrictConfig(node.Config, &cfg); err != nil {
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err)
}
case RuleNodeIPMatch: case RuleNodeIPMatch:
var cfg IPMatchConfig return validateIPMatchNodeConfig(ctx, node, exists)
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)
}
}
case RuleNodeGeoMatch: case RuleNodeGeoMatch:
var cfg GeoMatchConfig return validateGeoMatchNodeConfig(node)
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)
}
}
case RuleNodePoW: case RuleNodePoW:
var cfg PoWNodeConfig return validatePoWNodeConfig(node)
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)
}
case RuleNodeBlock: case RuleNodeBlock:
var cfg BlockNodeConfig return validateBlockNodeConfig(node)
if err := decodeStrictConfig(node.Config, &cfg); err != nil { }
return fmt.Errorf("节点 %s 的配置无效: %w", node.ID, err) 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 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 { func decodeStrictConfig(raw json.RawMessage, dst any) error {
trimmed := bytes.TrimSpace(raw) trimmed := bytes.TrimSpace(raw)
if bytes.Equal(trimmed, []byte("null")) { if bytes.Equal(trimmed, []byte("null")) {
+1 -1
View File
@@ -90,7 +90,7 @@ func syncOpenFlareWAFIPGroup(ctx context.Context, group *model.OpenFlareWAFIPGro
case wafIPGroupTypeAutomatic: case wafIPGroupTypeAutomatic:
return syncIPGroupAutomatic(ctx, group, now) return syncIPGroupAutomatic(ctx, group, now)
default: default:
return nil, errors.New("只有自动和订阅类型 IP 组支持同步") return nil, &RuleValidationError{Err: errors.New("只有自动和订阅类型 IP 组支持同步")}
} }
} }
File diff suppressed because it is too large Load Diff
+1 -32
View File
@@ -26,6 +26,7 @@ func setupWAFTestDB(t *testing.T) func() {
&model.OpenFlareWAFRuleGroup{}, &model.OpenFlareWAFRuleGroup{},
&model.OpenFlareWAFIPGroup{}, &model.OpenFlareWAFIPGroup{},
&model.OpenFlareWAFRuleGroupBinding{}, &model.OpenFlareWAFRuleGroupBinding{},
&model.OriginProxyRoute{},
)) ))
db.SetDB(sqliteDB) 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) { func TestPruneIPGroupExtIPs(t *testing.T) {
group := &model.OpenFlareWAFIPGroup{ 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"}]`, 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"}]`,
@@ -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
}
+10 -179
View File
@@ -12,14 +12,6 @@ import (
"github.com/gin-gonic/gin" "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) { func routeIDParam(c *gin.Context) (uint, bool) {
raw := c.Param("route_id") raw := c.Param("route_id")
if raw == "" { if raw == "" {
@@ -34,167 +26,6 @@ func routeIDParam(c *gin.Context) (uint, bool) {
return uint(id64), true 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 规则组绑定。 // GetSiteRuleGroupsHandler 获取站点的 WAF 规则组绑定。
// @Summary 获取站点 WAF 规则组 // @Summary 获取站点 WAF 规则组
// @Description 返回代理站点关联的 WAF 规则组绑定,需要管理员权限 // @Description 返回代理站点关联的 WAF 规则组绑定,需要管理员权限
@@ -215,7 +46,7 @@ func GetSiteRuleGroupsHandler(c *gin.Context) {
return return
} }
view, err := GetSiteRuleGroups(c.Request.Context(), routeID) view, err := GetSiteRuleGroups(c.Request.Context(), routeID)
if handleLogicError(c, err) { if handleRuleError(c, err) {
return return
} }
c.JSON(http.StatusOK, response.OK(view)) c.JSON(http.StatusOK, response.OK(view))
@@ -247,7 +78,7 @@ func ReplaceSiteRuleGroupsHandler(c *gin.Context) {
return return
} }
view, err := ReplaceSiteRuleGroups(c.Request.Context(), routeID, request.IDs) view, err := ReplaceSiteRuleGroups(c.Request.Context(), routeID, request.IDs)
if handleLogicError(c, err) { if handleRuleError(c, err) {
return return
} }
c.JSON(http.StatusOK, response.OK(view)) c.JSON(http.StatusOK, response.OK(view))
@@ -267,7 +98,7 @@ func ReplaceSiteRuleGroupsHandler(c *gin.Context) {
// @Router /api/v1/d/waf/ip-groups [get] // @Router /api/v1/d/waf/ip-groups [get]
func ListIPGroupsHandler(c *gin.Context) { func ListIPGroupsHandler(c *gin.Context) {
groups, err := ListIPGroups(c.Request.Context()) groups, err := ListIPGroups(c.Request.Context())
if handleLogicError(c, err) { if handleRuleError(c, err) {
return return
} }
c.JSON(http.StatusOK, response.OK(groups)) c.JSON(http.StatusOK, response.OK(groups))
@@ -293,7 +124,7 @@ func GetIPGroupHandler(c *gin.Context) {
return return
} }
group, err := GetIPGroup(c.Request.Context(), id) group, err := GetIPGroup(c.Request.Context(), id)
if handleLogicError(c, err) { if handleRuleError(c, err) {
return return
} }
c.JSON(http.StatusOK, response.OK(group)) c.JSON(http.StatusOK, response.OK(group))
@@ -319,7 +150,7 @@ func CreateIPGroupHandler(c *gin.Context) {
return return
} }
group, err := CreateIPGroup(c.Request.Context(), input) group, err := CreateIPGroup(c.Request.Context(), input)
if handleLogicError(c, err) { if handleRuleError(c, err) {
return return
} }
c.JSON(http.StatusOK, response.OK(group)) c.JSON(http.StatusOK, response.OK(group))
@@ -351,7 +182,7 @@ func UpdateIPGroupHandler(c *gin.Context) {
return return
} }
group, err := UpdateIPGroup(c.Request.Context(), id, input) group, err := UpdateIPGroup(c.Request.Context(), id, input)
if handleLogicError(c, err) { if handleRuleError(c, err) {
return return
} }
c.JSON(http.StatusOK, response.OK(group)) c.JSON(http.StatusOK, response.OK(group))
@@ -376,7 +207,7 @@ func DeleteIPGroupHandler(c *gin.Context) {
if !ok { if !ok {
return return
} }
if err := DeleteIPGroup(c.Request.Context(), id); handleLogicError(c, err) { if err := DeleteIPGroup(c.Request.Context(), id); handleRuleError(c, err) {
return return
} }
c.JSON(http.StatusOK, response.OKNil()) c.JSON(http.StatusOK, response.OKNil())
@@ -402,7 +233,7 @@ func SyncIPGroupHandler(c *gin.Context) {
return return
} }
result, err := SyncIPGroup(c.Request.Context(), id) result, err := SyncIPGroup(c.Request.Context(), id)
if handleLogicError(c, err) { if handleRuleError(c, err) {
return return
} }
c.JSON(http.StatusOK, response.OK(result)) c.JSON(http.StatusOK, response.OK(result))
@@ -428,8 +259,8 @@ func TestIPGroupAutoConfigHandler(c *gin.Context) {
return return
} }
result, err := TestIPGroupAutoConfig(c.Request.Context(), input) result, err := TestIPGroupAutoConfig(c.Request.Context(), input)
if handleLogicError(c, err) { if handleRuleError(c, err) {
return return
} }
c.JSON(http.StatusOK, response.OK(result)) c.JSON(http.StatusOK, response.OK(result))
} }
+185
View File
@@ -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
}
@@ -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
}
+186
View File
@@ -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())
}
@@ -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 '{}';
@@ -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 '{}';
+11 -35
View File
@@ -14,26 +14,14 @@ import (
// OpenFlareWAFRuleGroup stores a WAF rule group. // OpenFlareWAFRuleGroup stores a WAF rule group.
type OpenFlareWAFRuleGroup struct { type OpenFlareWAFRuleGroup struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"` ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:255;not null"` Name string `json:"name" gorm:"size:255;not null"`
Enabled bool `json:"enabled" gorm:"not null;default:true"` Enabled bool `json:"enabled" gorm:"not null;default:true"`
IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"` IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"`
BlockStatusCode int `json:"block_status_code" gorm:"not null;default:418"` Graph string `json:"graph" gorm:"type:text;not null;default:''"`
BlockResponseBody string `json:"block_response_body" gorm:"type:text;not null;default:''"` Revision uint64 `json:"revision" gorm:"not null;default:1"`
IPWhitelist string `json:"ip_whitelist" gorm:"type:text;not null;default:'[]'"` CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
IPBlacklist string `json:"ip_blacklist" gorm:"type:text;not null;default:'[]'"` UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
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"`
} }
// TableName returns the GORM table name. // TableName returns the GORM table name.
@@ -147,21 +135,9 @@ func UpdateOpenFlareWAFRuleGroup(ctx context.Context, group *OpenFlareWAFRuleGro
return err return err
} }
return conn.Model(&OpenFlareWAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ return conn.Model(&OpenFlareWAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{
"name": group.Name, "name": group.Name,
colEnabled: group.Enabled, colEnabled: group.Enabled,
"is_global": group.IsGlobal, "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,
}).Error }).Error
} }
+17 -1
View File
@@ -31,6 +31,7 @@ func wafMigrationFS(t *testing.T) fs.FS {
for _, name := range []string{ for _, name := range []string{
"202607150001_orchestrate_waf_rules.sql", "202607150001_orchestrate_waf_rules.sql",
"202607150002_reset_waf_rule_graphs.sql", "202607150002_reset_waf_rule_graphs.sql",
"202607150003_drop_legacy_waf_rule_fields.sql",
} { } {
contents, err := os.ReadFile(filepath.Join(dir, name)) contents, err := os.ReadFile(filepath.Join(dir, name))
require.NoError(t, err) require.NoError(t, err)
@@ -45,7 +46,7 @@ func TestOpenFlareWAFGraphMigrationResetsGraphsAndOrdersBindings(t *testing.T) {
sqlDB, err := conn.DB() sqlDB, err := conn.DB()
require.NoError(t, err) 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(`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_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) 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) require.JSONEq(t, defaultWAFRuleGraph, group.Graph)
assert.Equal(t, uint64(1), group.Revision) 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) require.NoError(t, conn.Exec(`INSERT INTO of_waf_rule_groups (name) VALUES ('new')`).Error)
var newGroup OpenFlareWAFRuleGroup 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, []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}) 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)
}
}
}
+6 -6
View File
@@ -21,12 +21,12 @@ func registerWAFRoutes(apiGroup *gin.RouterGroup) {
wafRoute.POST("/ip-groups/:id/delete", waf.DeleteIPGroupHandler) wafRoute.POST("/ip-groups/:id/delete", waf.DeleteIPGroupHandler)
wafRoute.POST("/ip-groups/:id/sync", waf.SyncIPGroupHandler) wafRoute.POST("/ip-groups/:id/sync", waf.SyncIPGroupHandler)
wafRoute.GET("/rule-groups", waf.ListRuleGroupsHandler) wafRoute.GET("/rule-groups", waf.ListRulesHandler)
wafRoute.GET("/rule-groups/:id", waf.GetRuleGroupHandler) wafRoute.GET("/rule-groups/:id", waf.GetRuleHandler)
wafRoute.POST("/rule-groups", waf.CreateRuleGroupHandler) wafRoute.POST("/rule-groups", waf.CreateRuleHandler)
wafRoute.POST("/rule-groups/:id/update", waf.UpdateRuleGroupHandler) wafRoute.POST("/rule-groups/:id/meta", waf.UpdateRuleMetaHandler)
wafRoute.POST("/rule-groups/:id/delete", waf.DeleteRuleGroupHandler) wafRoute.POST("/rule-groups/:id/graph", waf.SaveRuleGraphHandler)
wafRoute.POST("/rule-groups/:id/sites", waf.ReplaceRuleGroupSitesHandler) wafRoute.POST("/rule-groups/:id/delete", waf.DeleteRuleHandler)
wafRoute.GET("/sites/:route_id/rule-groups", waf.GetSiteRuleGroupsHandler) wafRoute.GET("/sites/:route_id/rule-groups", waf.GetSiteRuleGroupsHandler)
wafRoute.POST("/sites/:route_id/rule-groups", waf.ReplaceSiteRuleGroupsHandler) wafRoute.POST("/sites/:route_id/rule-groups", waf.ReplaceSiteRuleGroupsHandler)
+39
View File
@@ -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
}
@@ -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")
}
}
+10 -117
View File
@@ -112,96 +112,10 @@ func RenderRouteConfig(doc Document, certificateFiles []SupportFile) (string, er
// RenderWAFConfig serialises the WAF runtime configuration (rule groups and // RenderWAFConfig serialises the WAF runtime configuration (rule groups and
// per-site bindings) as a JSON string consumed by the OpenResty Lua runtime. // per-site bindings) as a JSON string consumed by the OpenResty Lua runtime.
func RenderWAFConfig(snapshot WAFDocument) (string, error) { func RenderWAFConfig(snapshot WAFDocument) (string, error) {
type wafRuntimeRuleGroup struct { data, err := json.Marshal(snapshot)
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})
return string(data), err 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 // ChecksumBundle returns a stable SHA-256 hex digest over the combined content
// of the main config, route config, and deduplicated support files, excluding // of the main config, route config, and deduplicated support files, excluding
// the source config JSON file itself. // the source config JSON file itself.
@@ -313,7 +227,7 @@ func renderOpenRestyLimitZoneBlock() string {
} }
func renderOpenRestyObservabilityTemplateBlock() 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 { 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) globalGroupIDs = append(globalGroupIDs, group.ID)
} }
} }
sort.Slice(globalGroupIDs, func(i, j int) bool { return globalGroupIDs[i] < globalGroupIDs[j] })
var boundGroupIDs []uint var boundGroupIDs []uint
for _, binding := range snapshot.Bindings { for _, binding := range snapshot.Bindings {
@@ -789,23 +702,20 @@ func getPoWConfigForRoute(routeID uint, snapshot WAFDocument) (bool, *PoWConfig)
activeGroupIDs := uniqueUintIDs(append(append([]uint{}, globalGroupIDs...), boundGroupIDs...)) activeGroupIDs := uniqueUintIDs(append(append([]uint{}, globalGroupIDs...), boundGroupIDs...))
for _, groupID := range activeGroupIDs { for _, groupID := range activeGroupIDs {
group := enabledGroups[groupID] group := enabledGroups[groupID]
if group.PoWEnabled { if graphContainsNodeType(group.Graph, "pow") {
config := ensurePoWConfig(true, group.PoWConfig) return true, nil
return true, config
} }
} }
return false, nil return false, nil
} }
func ensurePoWConfig(enabled bool, config *PoWConfig) *PoWConfig { func graphContainsNodeType(graph WAFRuleGraph, nodeType string) bool {
if !enabled { for _, node := range graph.Nodes {
return nil if node.Type == nodeType {
return true
}
} }
if config != nil { return false
return config
}
defaultConfig := DefaultPoWConfig()
return &defaultConfig
} }
func uniqueUintIDs(values []uint) []uint { func uniqueUintIDs(values []uint) []uint {
@@ -824,23 +734,6 @@ func uniqueUintIDs(values []uint) []uint {
return result 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 { func resolveUpstreamServerName(originURL string, originHost string) string {
parsed, err := url.Parse(originURL) parsed, err := url.Parse(originURL)
if err != nil || !strings.EqualFold(parsed.Scheme, "https") { if err != nil || !strings.EqualFold(parsed.Scheme, "https") {
+82 -35
View File
@@ -6,6 +6,16 @@ import (
"testing" "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) { func TestRenderWAFConfigIncludesAllRouteSiteNames(t *testing.T) {
doc := Document{ doc := Document{
Routes: []Route{ Routes: []Route{
@@ -15,11 +25,8 @@ func TestRenderWAFConfigIncludesAllRouteSiteNames(t *testing.T) {
WAF: WAFDocument{ WAF: WAFDocument{
RuleGroups: []WAFRuleGroup{ RuleGroups: []WAFRuleGroup{
{ {
ID: 1, ID: 1, Name: "pow-group", Enabled: true,
Name: "pow-group", Graph: WAFRuleGraph{Entry: "pow", Nodes: map[string]WAFRuleNode{"pow": {Type: "pow"}}},
Enabled: true,
PoWEnabled: true,
PoWConfig: &PoWConfig{Difficulty: 4, Algorithm: "fast", SessionTTL: 600, ChallengeTTL: 300},
}, },
}, },
Bindings: []WAFBinding{ Bindings: []WAFBinding{
@@ -34,18 +41,13 @@ func TestRenderWAFConfigIncludesAllRouteSiteNames(t *testing.T) {
t.Fatalf("RenderWAFConfig() error = %v", err) t.Fatalf("RenderWAFConfig() error = %v", err)
} }
var decoded struct { var decoded WAFDocument
SiteRuleGroups map[string][]uint `json:"site_rule_groups"`
}
if err := json.Unmarshal([]byte(wafConfig), &decoded); err != nil { if err := json.Unmarshal([]byte(wafConfig), &decoded); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err) t.Fatalf("json.Unmarshal() error = %v", err)
} }
for _, route := range doc.Routes { if len(decoded.Bindings) != 2 || decoded.Bindings[0].SiteName != "example.com" || decoded.Bindings[1].SiteName != "named-site" {
siteName := resolveRouteSiteName(route) t.Fatalf("bindings did not preserve route site names: %#v", decoded.Bindings)
if _, ok := decoded.SiteRuleGroups[siteName]; !ok {
t.Fatalf("site_rule_groups missing site %q, got %#v", siteName, decoded.SiteRuleGroups)
}
} }
routeConfig, err := RenderRouteConfig(doc, nil) 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{ doc := WAFDocument{
RuleGroups: []WAFRuleGroup{ RuleGroups: []WAFRuleGroup{
{ {
@@ -81,26 +83,15 @@ func TestRenderWAFConfigUsesDefaultPoWConfigWhenEnabledWithoutPayload(t *testing
t.Fatalf("RenderWAFConfig() error = %v", err) t.Fatalf("RenderWAFConfig() error = %v", err)
} }
var decoded struct { var decoded WAFDocument
RuleGroups []struct {
PoWEnabled bool `json:"pow_enabled"`
PoWConfig *PoWConfig `json:"pow_config"`
} `json:"rule_groups"`
}
if err := json.Unmarshal([]byte(wafConfig), &decoded); err != nil { if err := json.Unmarshal([]byte(wafConfig), &decoded); err != nil {
t.Fatalf("json.Unmarshal() error = %v", err) t.Fatalf("json.Unmarshal() error = %v", err)
} }
if len(decoded.RuleGroups) != 1 { if len(decoded.RuleGroups) != 1 {
t.Fatalf("expected 1 rule group, got %d", len(decoded.RuleGroups)) t.Fatalf("expected 1 rule group, got %d", len(decoded.RuleGroups))
} }
if !decoded.RuleGroups[0].PoWEnabled { if decoded.RuleGroups[0].PoWConfig != nil {
t.Fatal("expected pow_enabled=true") t.Fatalf("expected renderer not to synthesize legacy PoW config, got %#v", decoded.RuleGroups[0].PoWConfig)
}
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)
} }
} }
@@ -108,11 +99,8 @@ func TestGetPoWConfigForRouteUsesGlobalGroupWithoutExplicitBinding(t *testing.T)
snapshot := WAFDocument{ snapshot := WAFDocument{
RuleGroups: []WAFRuleGroup{ RuleGroups: []WAFRuleGroup{
{ {
ID: 1, ID: 1, Name: "global", Enabled: true, IsGlobal: true,
Name: "global", Graph: WAFRuleGraph{Entry: "pow", Nodes: map[string]WAFRuleNode{"pow": {Type: "pow"}}},
Enabled: true,
IsGlobal: true,
PoWEnabled: true,
}, },
}, },
Bindings: []WAFBinding{ Bindings: []WAFBinding{
@@ -124,8 +112,67 @@ func TestGetPoWConfigForRouteUsesGlobalGroupWithoutExplicitBinding(t *testing.T)
if !enabled { if !enabled {
t.Fatal("expected pow to be enabled via global rule group") t.Fatal("expected pow to be enabled via global rule group")
} }
if config == nil || config.Difficulty != 4 { if config != nil {
t.Fatalf("expected default pow config, got %#v", config) 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)
} }
} }
+34 -18
View File
@@ -1,5 +1,7 @@
package openresty package openresty
import "encoding/json"
// Placeholder constants used as sentinel values in rendered OpenResty config // Placeholder constants used as sentinel values in rendered OpenResty config
// files; the deploy process replaces them with real paths before reload. // files; the deploy process replaces them with real paths before reload.
const ( const (
@@ -177,25 +179,39 @@ type PagesDeployment struct {
LocalRoot string `json:"local_root"` LocalRoot string `json:"local_root"`
} }
// WAFRuleGroup defines a WAF rule group with IP/country/region lists, PoW // WAFRuleGraph is the compact graph executed by the OpenResty WAF runtime.
// integration, and per-group block status configuration. 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 { type WAFRuleGroup struct {
ID uint `json:"id"` ID uint `json:"id"`
Name string `json:"name"` Name string `json:"name"`
Enabled bool `json:"enabled"` Enabled bool `json:"enabled"`
IsGlobal bool `json:"is_global"` IsGlobal bool `json:"is_global"`
BlockStatusCode int `json:"block_status_code"` BlockStatusCode int `json:"block_status_code"`
BlockResponseBody string `json:"block_response_body,omitempty"` BlockResponseBody string `json:"block_response_body,omitempty"`
IPWhitelist []string `json:"ip_whitelist,omitempty"` IPWhitelist []string `json:"ip_whitelist,omitempty"`
IPBlacklist []string `json:"ip_blacklist,omitempty"` IPBlacklist []string `json:"ip_blacklist,omitempty"`
IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"` IPWhitelistGroups []uint `json:"ip_whitelist_group_ids,omitempty"`
IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"` IPBlacklistGroups []uint `json:"ip_blacklist_group_ids,omitempty"`
CountryWhitelist []string `json:"country_whitelist,omitempty"` CountryWhitelist []string `json:"country_whitelist,omitempty"`
CountryBlacklist []string `json:"country_blacklist,omitempty"` CountryBlacklist []string `json:"country_blacklist,omitempty"`
RegionWhitelist []string `json:"region_whitelist,omitempty"` RegionWhitelist []string `json:"region_whitelist,omitempty"`
RegionBlacklist []string `json:"region_blacklist,omitempty"` RegionBlacklist []string `json:"region_blacklist,omitempty"`
PoWEnabled bool `json:"pow_enabled,omitempty"` PoWEnabled bool `json:"pow_enabled,omitempty"`
PoWConfig *PoWConfig `json:"pow_config,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 // WAFIPGroup is a named, reusable list of IP addresses or CIDRs that can be