diff --git a/.gitignore b/.gitignore index 9083159e..0f0d5f13 100644 --- a/.gitignore +++ b/.gitignore @@ -44,8 +44,9 @@ go.work.sum .DS_Store .codex-cache -/.gomodcache/ -*.mmdb - -*-source -*-source.* \ No newline at end of file +/.gomodcache/ +*.mmdb +!openflare_agent/internal/geoipdata/GeoLite2-Country.mmdb + +*-source +*-source.* diff --git a/README.md b/README.md index 1106c1ac..a1dc3baa 100644 --- a/README.md +++ b/README.md @@ -38,6 +38,7 @@ * 配置预览、发布、激活与历史回滚 * Agent 自动注册、心跳、同步、校验、reload 与失败回滚 * OpenResty 主配置、性能参数、缓存参数与 Lua 资源托管 +* WAF 全局/自定义规则组,支持 IP/IP 段与国家级地域黑白名单 * TLS 证书、域名资产、节点凭证与版本状态管理 * 请求聚合、访问分析、资源快照、健康事件与节点详情 @@ -173,6 +174,7 @@ curl -fsSL https://raw.githubusercontent.com/Rain-kl/OpenFlare/main/scripts/unin * 应用记录 * TLS 证书 * 域名管理 +* WAF 规则组 * 用户管理 * 设置 * 版本更新 diff --git a/docs/Guidelines.md b/docs/Guidelines.md new file mode 100644 index 00000000..3a1c7892 --- /dev/null +++ b/docs/Guidelines.md @@ -0,0 +1,188 @@ +你是一个资深 Go 后端工程师,负责维护和开发一个长期演进的 Go 应用。 + +你的目标不是“尽快写完代码”,而是产出可维护、可测试、可演进、符合 Go 生态习惯的高质量代码。禁止为了完成任务而堆砌临时代码、过度抽象、重复逻辑或破坏现有架构。 + +在任何开发前,你必须先阅读并理解现有代码结构,包括: +- 项目目录结构 +- 入口文件 +- 配置管理方式 +- 数据库/缓存/消息队列访问方式 +- HTTP/RPC/API 层设计 +- service/usecase/domain/repository 等分层方式 +- 错误处理方式 +- 日志方式 +- 测试组织方式 +- 依赖注入方式 +- 现有编码风格 + +如果你不确定某个模块的职责,先通过代码上下文推断,不要随意新建重复模块。 + +开发原则: + +1. 架构优先 +- 优先融入现有架构,而不是另起炉灶。 +- 不要随便新增 global variable、init 副作用、隐式依赖。 +- 不要把业务逻辑写进 handler/controller。 +- handler 只负责参数解析、鉴权上下文、调用 usecase/service、返回响应。 +- service/usecase 负责业务编排。 +- repository/dao 负责数据访问。 +- domain/model 负责核心业务对象和规则。 +- 基础设施代码与业务代码隔离。 + +2. Go 风格 +- 使用清晰、直接、朴素的 Go 代码。 +- 不要模仿 Java 式过度抽象。 +- interface 应该由使用方定义,而不是提供方强行定义。 +- 小接口优先。 +- 命名要准确,不使用 Manager、Helper、Util 这类含糊名称,除非确实必要。 +- 函数保持短小,单一职责。 +- 不要为了“看起来高级”引入泛型、反射、复杂设计模式。 +- 不要隐藏错误。 +- error 必须带上下文信息,必要时使用 fmt.Errorf("...: %w", err)。 +- 不要 panic,除非是程序启动阶段的不可恢复错误。 + +3. 可维护性 +- 修改前先分析影响范围。 +- 尽量最小改动,不做无关重构。 +- 不改变公开 API、数据库结构、配置格式,除非任务明确要求。 +- 如果必须改变,要说明兼容性影响和迁移方案。 +- 删除代码前确认没有调用方。 +- 避免复制粘贴已有逻辑,应抽取到合适位置,但不要过度抽象。 +- 对复杂业务逻辑添加必要注释,解释“为什么”,不要注释显而易见的“是什么”。 + +4. 测试要求 +- 新增业务逻辑必须补充单元测试。 +- 修复 bug 必须补充回归测试。 +- 测试应覆盖正常路径、异常路径、边界条件。 +- 不要为了测试方便破坏业务代码结构。 +- 外部依赖使用 mock/fake/stub 隔离。 +- 测试命名清晰,例如 TestXXX_WhenYYY_ShouldZZZ。 +- 表驱动测试优先,但不要为了表驱动牺牲可读性。 + +5. 并发与资源管理 +- goroutine 必须有退出机制。 +- 涉及 context 的地方必须正确传递 context.Context。 +- 不要随意使用 context.Background() 替代上游 context。 +- channel 必须明确关闭责任。 +- 锁的范围要小,避免死锁。 +- HTTP、数据库、文件、连接等资源必须正确关闭。 +- 注意 race condition、goroutine leak、连接泄露。 + +6. 数据库与事务 +- 数据库访问必须在 repository/dao 层。 +- 事务边界应由业务用例层控制,而不是散落在多个底层函数中。 +- 不要在循环中产生明显低效的 N+1 查询,除非数据量可控且有说明。 +- SQL 要可读、参数化,禁止拼接不可信输入。 +- schema 变更必须考虑迁移、回滚和兼容性。 + +7. API 设计 +- 请求参数必须校验。 +- 错误响应要稳定、清晰,不泄露内部敏感信息。 +- 日志中不要打印密码、token、密钥、身份证号等敏感数据。 +- 返回结构保持向后兼容。 +- HTTP 状态码要语义正确。 + +8. 日志与可观测性 +- 关键路径要有必要日志。 +- 错误日志要包含排查所需上下文,但不要泄露敏感数据。 +- 不要滥打日志。 +- 不要在库代码里直接 fmt.Println。 +- 如果项目已有 logger,要统一使用现有 logger。 + +9. 安全要求 +- 所有外部输入都不可信。 +- 不要硬编码密钥、token、密码。 +- 不要把敏感配置提交到代码。 +- 文件路径、URL、命令执行、SQL、模板渲染等位置必须注意注入风险。 +- 鉴权和权限判断必须放在明确的位置,不能依赖前端或调用方自觉。 + +10. 性能要求 +- 不要过早优化。 +- 但不能写明显低效代码。 +- 对热点路径要避免不必要的内存分配、大对象复制、重复解析。 +- 大数据量处理应考虑分页、流式处理、批量操作。 +- 如果引入缓存,必须说明一致性、过期策略和失效条件。 + +工作流程: + +每次接到开发任务,你必须按以下步骤执行: + +第一步:理解需求 +- 用自己的话简要复述需求。 +- 明确输入、输出、边界条件、异常情况。 +- 如果需求含糊,列出你的合理假设,不要直接乱写。 + +第二步:阅读现有代码 +- 找出相关模块、调用链、数据结构、接口、测试。 +- 说明当前代码是如何工作的。 +- 判断改动应该放在哪一层。 + +第三步:设计方案 +- 给出最小可行修改方案。 +- 说明为什么放在这些文件/模块中。 +- 说明是否影响已有 API、数据库、配置、测试。 +- 如果有多个方案,比较优缺点,选择更稳妥的方案。 + +第四步:编码 +- 只修改与任务相关的代码。 +- 保持现有代码风格。 +- 不引入不必要的新依赖。 +- 不制造重复逻辑。 +- 不留下 TODO、临时代码、调试代码。 + +第五步:测试 +- 补充或更新测试。 +- 说明测试覆盖了哪些场景。 +- 如果无法运行测试,要说明原因,并给出应该运行的命令。 + +第六步:交付说明 +- 总结改了什么。 +- 说明为什么这样改。 +- 说明潜在风险。 +- 给出验证方式。 +- 如果存在未完成项,必须明确列出,不要假装完成。 + +输出格式: + +你每次回复都应包含: + +1. 需求理解 +2. 现有代码分析 +3. 修改方案 +4. 具体改动 +5. 测试与验证 +6. 风险与注意事项 + +如果只是让我审查代码,则输出: +1. 问题列表 +2. 严重程度:致命 / 高 / 中 / 低 +3. 影响说明 +4. 修改建议 +5. 推荐改法示例 + +代码质量红线: + +禁止出现以下行为: +- 为了完成需求复制粘贴大段重复代码 +- 在 handler 中塞业务逻辑 +- 到处传 map[string]interface{} +- 使用全局变量绕过依赖注入 +- 随意新增 util/helper 垃圾桶包 +- 忽略 error +- catch-all 式错误处理 +- 函数超过合理长度仍继续堆逻辑 +- 修改无关代码 +- 未经说明改变已有行为 +- 无测试地修改核心逻辑 +- 引入大型依赖只为解决小问题 +- 写完代码不说明验证方式 +- 不理解现有架构就直接重构 + +当你发现现有代码已经比较混乱时: +- 不要一次性大重构。 +- 先局部止血。 +- 新代码尽量写在清晰边界内。 +- 对旧代码只做必要改动。 +- 如果需要重构,先提出分阶段计划。 + +请始终以“长期维护这个项目的人”的标准来写代码,而不是以“完成一次性任务”的标准来写代码。 \ No newline at end of file diff --git a/docs/design/architecture.md b/docs/design/architecture.md index d3e076d6..2326f60a 100644 --- a/docs/design/architecture.md +++ b/docs/design/architecture.md @@ -30,8 +30,8 @@ Origin | --- | --- | | Server | 管理端 UI、管理 API、Agent API、配置渲染、版本发布、数据存储与聚合查询 | | Agent | 注册、心跳、同步、写入文件、校验、reload、失败回滚、自更新与轻量采集 | -| OpenResty | 接收真实流量,按 OpenFlare 渲染的配置执行反向代理 | -| Frontend | 管理网站配置、源站、证书、节点、版本、用户、设置与观测页面 | +| OpenResty | 接收真实流量,按 OpenFlare 渲染的配置执行 WAF、PoW、认证与反向代理 | +| Frontend | 管理网站配置、WAF、源站、证书、节点、版本、用户、设置与观测页面 | ## Server @@ -54,6 +54,7 @@ Server 不直接 SSH 到节点,也不在线修改节点文件。它只保存 * 周期性 heartbeat,上报状态并获取激活版本摘要。 * 发现新版本后拉取配置、备份旧文件、写入新文件、校验并 reload。 * 应用失败时尝试恢复运行并回滚。 +* 维护 WAF GeoIP mmdb,启动时写入内置初始库,并按配置定期更新。 Agent 通过 `openresty_path` 指向的 OpenResty 二进制统一执行校验、reload、启动与重启;未配置时默认调用 `openresty`。Docker 部署时,Agent 镜像内置 OpenResty 二进制,仍走同一套二进制控制逻辑。 @@ -84,7 +85,7 @@ Browser -> Frontend -> /api/* -> controller -> service -> model -> database ```text Agent heartbeat -> Server 返回激活版本摘要 Agent 发现新版本 -> 拉取配置详情 -Agent 写入主配置 / 路由配置 / 证书 / Lua 资源 +Agent 写入主配置 / 路由配置 / 证书 / Lua 资源 / WAF 运行时配置 Agent 执行 OpenResty 校验与 reload Agent 上报应用结果 ``` @@ -94,11 +95,13 @@ Agent 上报应用结果 ### 反向代理流 ```text -Client -> OpenResty server block -> named upstream -> Origin +Client -> OpenResty server block -> WAF Lua -> named upstream -> Origin ``` 网站配置是反向代理聚合边界。一条网站配置可绑定多个域名,并共享站点级流量限制、反向代理和缓存配置。 +WAF 在 OpenResty `access_by_lua_file` 阶段执行。规则来自当前激活版本携带的 `waf_config.json`,全局规则组默认生效,网站可叠加自定义规则组。 + ## 核心对象 当前有效实体包括: @@ -118,6 +121,8 @@ Client -> OpenResty server block -> named upstream -> Origin * `node_metric_snapshots` * `traffic_analytics_rollups` * `node_health_events` +* `waf_rule_groups` +* `waf_rule_group_bindings` ## 关键设计决策 diff --git a/docs/design/development.md b/docs/design/development.md index ddd0aa2c..8acfe760 100644 --- a/docs/design/development.md +++ b/docs/design/development.md @@ -148,6 +148,8 @@ tests/ * `traffic_analytics_rollups` * `node_health_events` * `options` +* `waf_rule_groups` +* `waf_rule_group_bindings` 通用约束: @@ -161,6 +163,7 @@ tests/ * 上游统一使用 named `upstream` + keepalive;单上游如带 base path 或 query,应在 `proxy_pass` 上补回 URI,多上游仅允许纯 `scheme://host[:port]`。 * 流量限制、反向代理与缓存配置当前都归属站点级 `proxy_routes`。 * HTTPS 证书绑定必须通过与 `domains` 平行的 `domain_cert_ids` 逐域名保存;未绑定证书的域名不得参与 HTTPS 渲染。 +* WAF 全局规则组默认应用到所有网站,自定义规则组通过 `waf_rule_group_bindings` 绑定到网站配置;发布时必须进入完整版本快照。 * `config_versions` 必须保存完整快照与渲染结果。 * 全局同时只能有一个激活版本。 * 回滚通过重新激活旧版本实现。 @@ -237,12 +240,14 @@ Agent 必须满足: * WS 连接升级开启且连接成功时,Agent 可通过 WS 接收激活版本摘要并立即同步;WS 失败或断开必须退回 HTTP heartbeat。 * 发现新版本时先备份旧文件。 * 写入主配置、路由配置与必要证书文件。 +* 写入 WAF/PoW 运行时配置,并确保 WAF Lua 资源由 Agent 统一管理。 * 写入新配置后执行 `openresty -t -c `,再 reload;reload 发现运行时未启动时允许直接启动 OpenResty。 * 周期性运行时健康检查不得调用 `openresty -t`,避免健康探针触发 upstream 域名同步解析;应优先请求本地 `openresty_observability_port` 上的 `/openflare/stub_status`,以 HTTP `200 OK` 作为 OpenResty 主进程和 worker 正在提供服务的判断依据。 * 新配置激活失败时必须先尝试用目标配置恢复运行,再回滚到旧配置并重新拉起 OpenResty。 * 回滚后 OpenResty 恢复正常时上报警告;如果本地没有历史主配置可恢复,必须允许写入内置安全兜底配置并拉起对外只监听 `80` 端口、统一返回 `503` 的 OpenResty 运行态;兜底配置仍需保留本地 `stub_status` 健康检查入口。 * 兜底运行态不得清除失败目标的阻断状态;应用记录必须能体现目标版本失败但 fallback runtime 已启动。存在历史主配置但回滚后仍无法恢复运行时上报失败。 * 某个目标 `version + checksum` 一旦应用失败并回退,Agent 必须在本地状态中阻断该目标的重复应用。 +* Agent 维护本地 MaxMind mmdb 时,下载或刷新失败只能记录警告,不得阻断心跳、同步、配置应用或 OpenResty 健康检查。 ## 前端请求、状态与类型 diff --git a/docs/design/index.md b/docs/design/index.md index 4ccf1567..656d3947 100644 --- a/docs/design/index.md +++ b/docs/design/index.md @@ -35,6 +35,7 @@ OpenFlare 当前不定位为通用日志平台、服务网格、Kubernetes Ingre | Agent 同步 | 支持注册、心跳、同步、应用结果上报与自更新 | | OpenResty 托管 | 管理主配置模板、性能参数、缓存参数与 Lua 资源 | | HTTPS/TLS | 托管证书与域名资产,并按域名绑定证书 | +| WAF | 以全局规则组与网站自定义规则组维护 IP/IP 段、国家级地域黑白名单 | | 基础观测 | 聚合节点请求、资源快照、健康事件和访问分析 | | 节点管理 | 节点状态、令牌体系、部署与更新链路 | | 管理端前端 | 基于 Next.js 的正式管理端 | @@ -76,6 +77,8 @@ OpenFlare 当前不定位为通用日志平台、服务网格、Kubernetes Ingre * `node_metric_snapshots` * `traffic_analytics_rollups` * `node_health_events` +* `waf_rule_groups` +* `waf_rule_group_bindings` ## 网站配置约束 @@ -116,6 +119,24 @@ OpenFlare 当前不定位为通用日志平台、服务网格、Kubernetes Ingre * 未绑定证书的域名不得被自动带入 HTTPS。 * 必须将 `proxy_routes.domains` 中的全部域名一并纳入同一站点配置,避免同站点在版本快照中被拆散。 +## WAF 约束 + +WAF 以规则组为配置边界。系统固定一个全局规则组,默认应用到所有网站;网站可叠加多个自定义规则组。 + +一期支持: + +* IP / IP 段白名单与黑名单。 +* 国家级地域白名单与黑名单。 +* 规则组级拦截状态码与响应页面,默认 `418` 与空页面。 + +判定顺序: + +* 白名单是放行例外,任意启用规则组命中白名单即放行。 +* 未命中白名单时继续判断黑名单。 +* 多个黑名单命中时,全局规则组优先,其后按自定义规则组 ID 升序。 + +地域识别由 Agent 维护节点本地 MaxMind mmdb,OpenResty Lua 在请求路径中读取本地库。GeoIP 依赖不可用时只能跳过地域规则,不得影响 IP 规则与反向代理主链路。 + ## 认证源约束 `auth_sources` 是管理端第三方登录入口的配置对象,当前仅支持 `github` 与 `oidc` 两类。启用后的认证源会显示在登录页。 diff --git a/docs/design/release-model.md b/docs/design/release-model.md index ede84513..c35965fe 100644 --- a/docs/design/release-model.md +++ b/docs/design/release-model.md @@ -17,11 +17,12 @@ Server 发布时必须: 1. 读取全部启用的 `proxy_routes`。 2. 读取 Server 侧 OpenResty 主配置、性能参数、缓存参数和必要 Lua 资源。 3. 读取域名与证书绑定关系。 -4. 渲染完整 OpenResty 配置。 -5. 计算 `checksum`。 -6. 写入 `config_versions`。 -7. 切换激活版本。 -8. 让 Agent 在后续 heartbeat 中发现并应用。 +4. 读取 WAF 全局规则组、自定义规则组与网站绑定关系。 +5. 渲染完整 OpenResty 配置与 WAF 运行时配置。 +6. 计算 `checksum`。 +7. 写入 `config_versions`。 +8. 切换激活版本。 +9. 让 Agent 在后续 heartbeat 中发现并应用。 版本号格式固定为 `YYYYMMDD-NNN`。 @@ -53,7 +54,7 @@ Agent 发现新版本后会: 1. 拉取目标版本详情。 2. 备份旧文件。 -3. 写入主配置、路由配置、证书与必要 Lua 资源。 +3. 写入主配置、路由配置、证书、必要 Lua 资源与 WAF/PoW 运行时配置。 4. 执行 OpenResty 配置校验。 5. reload;如果运行时未启动,则尝试用当前配置启动 OpenResty。 6. 上报成功、警告或失败。 @@ -69,3 +70,4 @@ Agent 发现新版本后会: * Agent API 固定使用节点专属 `agent_token`,首次接入可使用 `discovery_token`。 * Server 不提供远程 shell 或任意命令执行入口。 * 配置版本必须保存完整快照、渲染结果和 `checksum`。 +* WAF 规则组和网站绑定关系必须随完整配置版本进入快照与 checksum,回滚时不得依赖当前可变 WAF 配置。 diff --git a/docs/guide/deployment.md b/docs/guide/deployment.md index f4c86c10..ccc5a430 100644 --- a/docs/guide/deployment.md +++ b/docs/guide/deployment.md @@ -43,6 +43,7 @@ Agent: | OpenResty | 本地部署需要可执行 `openresty`,或通过 `--openresty-path` 指定路径 | | Docker | 仅 Docker 部署 Agent 镜像时需要 | | 网络 | Agent 节点必须能访问 Server 地址 | +| GeoIP | WAF 地域规则使用 Agent 本地 MaxMind mmdb;Agent 内置初始库并会定期更新 | [需要确认:生产环境推荐的最低 CPU、内存与磁盘容量] @@ -225,6 +226,8 @@ export LOG_LEVEL='info' 默认情况下,Agent 在 HTTP 心跳成功后会尝试升级为 WebSocket。升级成功时,Server 发布或激活配置会立即通知 Agent;如果 WebSocket 无法建立或意外断开,Agent 会自动退回 HTTP 心跳同步。 +WAF 地域规则依赖 Agent 本地 `GeoLite2-Country.mmdb`。Agent 启动时会在 `data_dir/etc/openflare/GeoLite2-Country.mmdb` 初始化内置数据库,并按配置周期尝试更新;更新失败只记录警告,不影响配置同步与 OpenResty reload。 + ## 最小联调步骤 1. 启动 Server 并完成首次登录。 diff --git a/docs/reference/configuration.md b/docs/reference/configuration.md index cec4cb04..22ab8b57 100644 --- a/docs/reference/configuration.md +++ b/docs/reference/configuration.md @@ -148,6 +148,9 @@ OpenResty 性能参数与缓存参数继续统一保存在 `Option` 表。当前 | `OPENFLARE_HEARTBEAT_INTERVAL` | 心跳间隔,可覆盖 `agent.json` | 空 | | `OPENFLARE_REQUEST_TIMEOUT` | 请求超时,可覆盖 `agent.json` | 空 | | `OPENFLARE_OPENRESTY_OBSERVABILITY_PORT` | 本地观测端口,可覆盖 `agent.json` | 空 | +| `OPENFLARE_MMDB_PATH` | WAF GeoIP mmdb 路径,可覆盖 `agent.json` | 空 | +| `OPENFLARE_MMDB_UPDATE_INTERVAL` | WAF GeoIP mmdb 更新间隔,可覆盖 `agent.json` | 空 | +| `OPENFLARE_MMDB_DOWNLOAD_URL` | WAF GeoIP mmdb 下载地址,可覆盖 `agent.json` | 空 | ## Agent 命令行参数 @@ -178,6 +181,9 @@ OpenResty 性能参数与缓存参数继续统一保存在 `Option` 表。当前 | `lua_dir` | Lua 脚本与静态资源写入目录 | 否 | `data_dir/etc/nginx/lua` | | `openresty_lua_dir` | OpenResty 配置中读取 Lua 的目录 | 否 | 同 `lua_dir` | | `runtime_config_dir` | Agent 运行时配置写入目录,如 `pow_config.json` | 否 | `data_dir/etc/openflare` | +| `mmdb_path` | WAF GeoIP mmdb 文件路径 | 否 | `data_dir/etc/openflare/GeoLite2-Country.mmdb` | +| `mmdb_update_interval` | WAF GeoIP mmdb 更新间隔 | 否 | `86400000` 毫秒 | +| `mmdb_download_url` | WAF GeoIP mmdb 下载地址 | 否 | 内置 GeoLite2 Country 下载地址 | | `observability_buffer_path` | 观测补报缓冲文件路径 | 否 | `data_dir/var/lib/openflare/observability-buffer.json` | | `observability_replay_minutes` | 自动补传最近观测窗口分钟数 | 否 | `15` | | `state_path` | Agent 本地状态文件路径 | 否 | `data_dir/var/lib/openflare/agent-state.json` | @@ -191,6 +197,7 @@ OpenResty 性能参数与缓存参数继续统一保存在 `Option` 表。当前 * Server 运行时配置 `AgentWebsocketUpgradeEnabled` 开启时,Agent 会在 HTTP 心跳成功后尝试升级为 WebSocket;连接失败或断开后自动退回 HTTP 心跳。 * 未配置 `openresty_path` 时默认调用 `openresty`。 * Agent 周期性健康检查会请求 `http://127.0.0.1:/openflare/stub_status`,不再通过高频 `openresty -t` 判断运行时健康;配置应用、启动恢复和 reload 前校验仍会执行 `openresty -t -c `。 +* Agent 会初始化并定期更新 `mmdb_path`,供 OpenResty WAF Lua 执行国家级地域规则;更新失败只记录警告,不阻断同步或 reload。 * 如果 `agent.json` 不存在,但 `OPENFLARE_SERVER_URL` 与 Token 等环境变量足够,Agent 可以直接启动;两者同时存在时环境变量优先。 * Agent 未配置 `node_ip` 时,会优先通过 `https://realip.cc` 获取真实出口公网 IP,适配 Docker/NAT 场景;该请求失败时,才退回本机网卡探测并优先选择公网 IPv4。 * Agent 自动探测到私网 `node_ip` 时,Server 会在注册/心跳阶段优先保留 Agent 直连来源的公网地址,避免 NAT/多网卡场景误登记内网网卡地址。 diff --git a/openflare_agent/cmd/agent/main.go b/openflare_agent/cmd/agent/main.go index aa568e3f..f7cfa3ba 100644 --- a/openflare_agent/cmd/agent/main.go +++ b/openflare_agent/cmd/agent/main.go @@ -10,6 +10,7 @@ import ( "openflare-agent/internal/agent" "openflare-agent/internal/config" + "openflare-agent/internal/geoipupdate" "openflare-agent/internal/heartbeat" "openflare-agent/internal/httpclient" "openflare-agent/internal/logging" @@ -54,6 +55,7 @@ func main() { "cert_dir", cfg.CertDir, "lua_dir", cfg.LuaDir, "runtime_config_dir", cfg.RuntimeConfigDir, + "mmdb_path", cfg.MMDBPath, ) client := httpclient.New(cfg.ServerURL, cfg.InitialAuthToken(), cfg.RequestTimeout.Duration()) @@ -100,6 +102,12 @@ func main() { ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) defer stop() + geoIPUpdater := &geoipupdate.Updater{ + MMDBPath: cfg.MMDBPath, + DownloadURL: cfg.MMDBDownloadURL, + UpdateInterval: cfg.MMDBUpdateInterval.Duration(), + } + go geoIPUpdater.Run(ctx) slog.Info("agent process started") if err = runner.Run(ctx); err != nil && err != context.Canceled { diff --git a/openflare_agent/internal/config/config.go b/openflare_agent/internal/config/config.go index e4d3a186..4538186e 100644 --- a/openflare_agent/internal/config/config.go +++ b/openflare_agent/internal/config/config.go @@ -22,11 +22,14 @@ const ( defaultCertDirRelativePath = "etc/nginx/certs" defaultLuaDirRelativePath = "etc/nginx/lua" defaultRuntimeConfigDirRelativePath = "etc/openflare" + defaultMMDBRelativePath = "etc/openflare/GeoLite2-Country.mmdb" defaultAccessLogRelativePath = "var/log/openflare/access.log" defaultStateRelativePath = "var/lib/openflare/agent-state.json" defaultObservabilityBufferRelativePath = "var/lib/openflare/observability-buffer.json" defaultOpenRestyObservabilityPort = 18081 defaultObservabilityReplayMinutes = 15 + defaultMMDBUpdateInterval = 24 * time.Hour + defaultMMDBDownloadURL = "https://raw.githubusercontent.com/Loyalsoldier/geoip/release/GeoLite2-Country.mmdb" ) var ( @@ -56,6 +59,9 @@ type Config struct { LuaDir string `json:"lua_dir"` OpenrestyLuaDir string `json:"openresty_lua_dir"` RuntimeConfigDir string `json:"runtime_config_dir"` + MMDBPath string `json:"mmdb_path"` + MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"` + MMDBDownloadURL string `json:"mmdb_download_url"` OpenrestyObservabilityPort int `json:"openresty_observability_port"` ObservabilityBufferPath string `json:"observability_buffer_path"` ObservabilityReplayMinutes int `json:"observability_replay_minutes"` @@ -85,6 +91,9 @@ type configFile struct { LuaDir string `json:"lua_dir"` OpenrestyLuaDir string `json:"openresty_lua_dir"` RuntimeConfigDir string `json:"runtime_config_dir"` + MMDBPath string `json:"mmdb_path"` + MMDBUpdateInterval MillisecondDuration `json:"mmdb_update_interval"` + MMDBDownloadURL string `json:"mmdb_download_url"` OpenrestyObservabilityPort int `json:"openresty_observability_port"` ObservabilityBufferPath string `json:"observability_buffer_path"` ObservabilityReplayMinutes int `json:"observability_replay_minutes"` @@ -127,6 +136,9 @@ func Load(path string) (*Config, error) { LuaDir: file.LuaDir, OpenrestyLuaDir: file.OpenrestyLuaDir, RuntimeConfigDir: file.RuntimeConfigDir, + MMDBPath: file.MMDBPath, + MMDBUpdateInterval: file.MMDBUpdateInterval, + MMDBDownloadURL: file.MMDBDownloadURL, OpenrestyObservabilityPort: file.OpenrestyObservabilityPort, ObservabilityBufferPath: file.ObservabilityBufferPath, ObservabilityReplayMinutes: file.ObservabilityReplayMinutes, @@ -186,6 +198,15 @@ func applyDefaults(cfg *Config, baseDir string) { if cfg.RuntimeConfigDir == "" { cfg.RuntimeConfigDir = joinManagedPath(cfg.DataDir, defaultRuntimeConfigDirRelativePath) } + if cfg.MMDBPath == "" { + cfg.MMDBPath = joinManagedPath(cfg.DataDir, defaultMMDBRelativePath) + } + if cfg.MMDBUpdateInterval <= 0 { + cfg.MMDBUpdateInterval = MillisecondDuration(defaultMMDBUpdateInterval) + } + if cfg.MMDBDownloadURL == "" { + cfg.MMDBDownloadURL = defaultMMDBDownloadURL + } if cfg.OpenrestyObservabilityPort <= 0 { cfg.OpenrestyObservabilityPort = defaultOpenRestyObservabilityPort } @@ -241,6 +262,9 @@ func normalizeManagedPaths(cfg *Config) { if usesSlashPath(cfg.ObservabilityBufferPath) { cfg.ObservabilityBufferPath = filepath.ToSlash(cfg.ObservabilityBufferPath) } + if usesSlashPath(cfg.MMDBPath) { + cfg.MMDBPath = filepath.ToSlash(cfg.MMDBPath) + } } func hasEnvConfig() bool { @@ -255,6 +279,9 @@ func hasEnvConfig() bool { "OPENFLARE_HEARTBEAT_INTERVAL", "OPENFLARE_REQUEST_TIMEOUT", "OPENFLARE_OPENRESTY_OBSERVABILITY_PORT", + "OPENFLARE_MMDB_PATH", + "OPENFLARE_MMDB_UPDATE_INTERVAL", + "OPENFLARE_MMDB_DOWNLOAD_URL", } { if strings.TrimSpace(os.Getenv(key)) != "" { return true @@ -279,6 +306,8 @@ func applyEnvOverrides(cfg *Config) { overrideString("OPENFLARE_NODE_IP", &cfg.NodeIP) overrideString("OPENFLARE_DATA_DIR", &cfg.DataDir) overrideString("OPENFLARE_OPENRESTY_PATH", &cfg.OpenrestyPath) + overrideString("OPENFLARE_MMDB_PATH", &cfg.MMDBPath) + overrideString("OPENFLARE_MMDB_DOWNLOAD_URL", &cfg.MMDBDownloadURL) if value := strings.TrimSpace(os.Getenv("OPENFLARE_HEARTBEAT_INTERVAL")); value != "" { if duration, err := parseDurationValue(value); err == nil { cfg.HeartbeatInterval = duration @@ -289,6 +318,11 @@ func applyEnvOverrides(cfg *Config) { cfg.RequestTimeout = duration } } + if value := strings.TrimSpace(os.Getenv("OPENFLARE_MMDB_UPDATE_INTERVAL")); value != "" { + if duration, err := parseDurationValue(value); err == nil { + cfg.MMDBUpdateInterval = duration + } + } if value := strings.TrimSpace(os.Getenv("OPENFLARE_OPENRESTY_OBSERVABILITY_PORT")); value != "" { var port int if _, err := fmt.Sscanf(value, "%d", &port); err == nil { @@ -342,6 +376,9 @@ func validate(cfg *Config) error { if cfg.ObservabilityReplayMinutes <= 0 { return errors.New("observability_replay_minutes 必须大于 0") } + if cfg.MMDBUpdateInterval <= 0 { + return errors.New("mmdb_update_interval 必须大于 0") + } return nil } diff --git a/openflare_agent/internal/geoipdata/GeoLite2-Country.mmdb b/openflare_agent/internal/geoipdata/GeoLite2-Country.mmdb new file mode 100644 index 00000000..c7195dd1 Binary files /dev/null and b/openflare_agent/internal/geoipdata/GeoLite2-Country.mmdb differ diff --git a/openflare_agent/internal/geoipdata/data.go b/openflare_agent/internal/geoipdata/data.go new file mode 100644 index 00000000..49899533 --- /dev/null +++ b/openflare_agent/internal/geoipdata/data.go @@ -0,0 +1,8 @@ +package geoipdata + +import "embed" + +//go:embed GeoLite2-Country.mmdb +var FS embed.FS + +const DefaultMMDBName = "GeoLite2-Country.mmdb" diff --git a/openflare_agent/internal/geoipupdate/updater.go b/openflare_agent/internal/geoipupdate/updater.go new file mode 100644 index 00000000..5e7ae45c --- /dev/null +++ b/openflare_agent/internal/geoipupdate/updater.go @@ -0,0 +1,67 @@ +package geoipupdate + +import ( + "context" + "fmt" + "io/fs" + "log/slog" + "os" + "path/filepath" + "time" + + "openflare-agent/internal/geoipdata" + "openflare/utils/geoip" +) + +type Updater struct { + MMDBPath string + DownloadURL string + UpdateInterval time.Duration +} + +func (u *Updater) EnsureInitialDatabase() error { + path := filepath.Clean(u.MMDBPath) + if path == "" || path == "." { + return nil + } + if _, err := os.Stat(path); err == nil { + return nil + } else if !os.IsNotExist(err) { + return fmt.Errorf("stat mmdb file failed: %w", err) + } + data, err := fs.ReadFile(geoipdata.FS, geoipdata.DefaultMMDBName) + if err != nil { + return fmt.Errorf("read embedded mmdb failed: %w", err) + } + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + return fmt.Errorf("create mmdb directory failed: %w", err) + } + if err := os.WriteFile(path, data, 0o644); err != nil { + return fmt.Errorf("write initial mmdb failed: %w", err) + } + slog.Info("initialized GeoIP mmdb from embedded database", "path", path, "size", len(data)) + return nil +} + +func (u *Updater) Run(ctx context.Context) { + if u == nil || u.MMDBPath == "" || u.UpdateInterval <= 0 { + return + } + if err := u.EnsureInitialDatabase(); err != nil { + slog.Warn("initialize GeoIP mmdb failed", "path", u.MMDBPath, "error", err) + } + ticker := time.NewTicker(u.UpdateInterval) + defer ticker.Stop() + for { + select { + case <-ctx.Done(): + return + case <-ticker.C: + if err := geoip.DownloadMaxMindDatabase(u.MMDBPath, u.DownloadURL); err != nil { + slog.Warn("update GeoIP mmdb failed", "path", u.MMDBPath, "error", err) + continue + } + slog.Info("GeoIP mmdb updated", "path", u.MMDBPath) + } + } +} diff --git a/openflare_agent/internal/geoipupdate/updater_test.go b/openflare_agent/internal/geoipupdate/updater_test.go new file mode 100644 index 00000000..07dffda3 --- /dev/null +++ b/openflare_agent/internal/geoipupdate/updater_test.go @@ -0,0 +1,24 @@ +package geoipupdate + +import ( + "os" + "path/filepath" + "testing" +) + +func TestEnsureInitialDatabaseCopiesEmbeddedMMDB(t *testing.T) { + tempDir := t.TempDir() + path := filepath.Join(tempDir, "GeoLite2-Country.mmdb") + updater := &Updater{MMDBPath: path} + + if err := updater.EnsureInitialDatabase(); err != nil { + t.Fatalf("EnsureInitialDatabase failed: %v", err) + } + info, err := os.Stat(path) + if err != nil { + t.Fatalf("expected mmdb to exist: %v", err) + } + if info.Size() == 0 { + t.Fatal("expected copied mmdb to be non-empty") + } +} diff --git a/openflare_agent/internal/nginx/manager.go b/openflare_agent/internal/nginx/manager.go index ceb0415c..216a3285 100644 --- a/openflare_agent/internal/nginx/manager.go +++ b/openflare_agent/internal/nginx/manager.go @@ -220,6 +220,9 @@ func (m *Manager) writeTargetFiles(mainConfig string, routeConfig string, suppor if err := m.writePowConfig(supportFiles); err != nil { return err } + if err := m.writeWAFConfig(supportFiles); err != nil { + return err + } if err := m.ensureMimeTypes(); err != nil { return err } @@ -289,6 +292,7 @@ func (m *Manager) EnsureLuaAssets() error { return nil } allSupportFiles := append(ManagedObservabilityLuaFiles(), m.managedPowLuaFiles()...) + allSupportFiles = append(allSupportFiles, m.managedWAFLuaFiles()...) powStaticFiles, err := ManagedPowStaticFiles() if err != nil { return fmt.Errorf("load pow static files: %w", err) @@ -628,10 +632,30 @@ func (m *Manager) writePowConfig(supportFiles []protocol.SupportFile) error { return nil } +func (m *Manager) writeWAFConfig(supportFiles []protocol.SupportFile) error { + if m.RuntimeConfigDir == "" { + return nil + } + configPath := filepath.Join(m.RuntimeConfigDir, "waf_config.json") + for _, file := range supportFiles { + if file.Path == "waf_config.json" { + if err := os.WriteFile(configPath, []byte(file.Content), 0o644); err != nil { + return fmt.Errorf("write waf_config.json: %w", err) + } + slog.Info("wrote waf config", "path", configPath, "size", len(file.Content)) + return nil + } + } + if err := os.Remove(configPath); err != nil && !os.IsNotExist(err) { + return fmt.Errorf("remove waf_config.json: %w", err) + } + return nil +} + func (m *Manager) writeManagedCertFiles(certFiles []protocol.SupportFile) error { files := make([]managedFile, 0, len(certFiles)) for _, file := range certFiles { - if file.Path == "pow_config.json" { + if file.Path == "pow_config.json" || file.Path == "waf_config.json" { continue } targetPath, err := m.certFileTargetPath(file.Path) @@ -1014,6 +1038,15 @@ func (m *Manager) managedPowLuaFiles() []protocol.SupportFile { return files } +func (m *Manager) managedWAFLuaFiles() []protocol.SupportFile { + files := ManagedWAFLuaFiles() + runtimeConfigDir := filepath.ToSlash(strings.TrimSpace(m.RuntimeConfigDir)) + for index := range files { + files[index].Content = strings.ReplaceAll(files[index].Content, RuntimeConfigDirPlaceholder, runtimeConfigDir) + } + return files +} + func ObservabilityListenAddress(openrestyPath string, port int) string { if port <= 0 { return "" diff --git a/openflare_agent/internal/nginx/waf_assets.go b/openflare_agent/internal/nginx/waf_assets.go new file mode 100644 index 00000000..a256d9e2 --- /dev/null +++ b/openflare_agent/internal/nginx/waf_assets.go @@ -0,0 +1,214 @@ +package nginx + +import "openflare-agent/internal/protocol" + +const openRestyWAFCheckLua = `local cjson = require "cjson.safe" + +local config_dict = ngx.shared.openflare_waf_config + +local function read_file(path) + local f = io.open(path, "r") + if not f then + return nil + end + local content = f:read("*a") + f:close() + return content +end + +local function load_config() + local paths = { + "__OPENFLARE_RUNTIME_CONFIG_DIR__/waf_config.json", + "/etc/nginx/openflare-lua/waf_config.json", + "/usr/local/openresty/nginx/conf/waf_config.json" + } + for _, path in ipairs(paths) do + local content = read_file(path) + if content and content ~= "" then + local hash = ngx.md5(content) + if config_dict:get("_config_hash") == hash then + local cached = config_dict:get("_config_json") + if cached then + local decoded = cjson.decode(cached) + if decoded then + return decoded + end + end + end + local decoded = cjson.decode(content) + if decoded then + config_dict:set("_config_hash", hash, 0) + config_dict:set("_config_json", content, 0) + return decoded + end + end + end + return nil +end + +local function list_contains(items, value) + if not items 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 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 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 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.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 + return +end + +local ip = ngx.var.remote_addr or "" +local groups = active_groups(config) + +for _, group in ipairs(groups) do + if ip_matches(group.ip_whitelist, ip) then + return + end +end + +local country = nil +for _, group in ipairs(groups) do + if group.country_whitelist 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) then + return exit_with_group(group) + end +end + +for _, group in ipairs(groups) do + if group.country_blacklist 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 +` + +func ManagedWAFLuaFiles() []protocol.SupportFile { + return []protocol.SupportFile{ + {Path: "waf/check.lua", Content: openRestyWAFCheckLua}, + } +} diff --git a/openflare_server/controller/waf.go b/openflare_server/controller/waf.go new file mode 100644 index 00000000..11d9521a --- /dev/null +++ b/openflare_server/controller/waf.go @@ -0,0 +1,138 @@ +package controller + +import ( + "encoding/json" + "net/http" + "openflare/service" + "strconv" + + "github.com/gin-gonic/gin" +) + +type wafIDsRequest struct { + IDs []uint `json:"ids"` +} + +func ListWAFRuleGroups(c *gin.Context) { + groups, err := service.ListWAFRuleGroups() + if err != nil { + c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": groups}) +} + +func GetWAFRuleGroup(c *gin.Context) { + id, ok := parseUintPathParam(c, "id") + if !ok { + return + } + group, err := service.GetWAFRuleGroup(id) + if err != nil { + c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": group}) +} + +func CreateWAFRuleGroup(c *gin.Context) { + var input service.WAFRuleGroupInput + if err := json.NewDecoder(c.Request.Body).Decode(&input); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid payload"}) + return + } + group, err := service.CreateWAFRuleGroup(input) + if err != nil { + c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": group}) +} + +func UpdateWAFRuleGroup(c *gin.Context) { + id, ok := parseUintPathParam(c, "id") + if !ok { + return + } + var input service.WAFRuleGroupInput + if err := json.NewDecoder(c.Request.Body).Decode(&input); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid payload"}) + return + } + group, err := service.UpdateWAFRuleGroup(id, input) + if err != nil { + c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": group}) +} + +func DeleteWAFRuleGroup(c *gin.Context) { + id, ok := parseUintPathParam(c, "id") + if !ok { + return + } + if err := service.DeleteWAFRuleGroup(id); err != nil { + c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{"success": true, "message": ""}) +} + +func ReplaceWAFRuleGroupSites(c *gin.Context) { + id, ok := parseUintPathParam(c, "id") + if !ok { + return + } + var request wafIDsRequest + if err := json.NewDecoder(c.Request.Body).Decode(&request); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid payload"}) + return + } + group, err := service.ReplaceWAFRuleGroupSites(id, request.IDs) + if err != nil { + c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": group}) +} + +func GetWAFSiteRuleGroups(c *gin.Context) { + routeID, ok := parseUintPathParam(c, "route_id") + if !ok { + return + } + view, err := service.GetWAFSiteRuleGroups(routeID) + if err != nil { + c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": view}) +} + +func ReplaceWAFSiteRuleGroups(c *gin.Context) { + routeID, ok := parseUintPathParam(c, "route_id") + if !ok { + return + } + var request wafIDsRequest + if err := json.NewDecoder(c.Request.Body).Decode(&request); err != nil { + c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid payload"}) + return + } + view, err := service.ReplaceWAFSiteRuleGroups(routeID, request.IDs) + if err != nil { + c.JSON(http.StatusOK, gin.H{"success": false, "message": err.Error()}) + return + } + c.JSON(http.StatusOK, gin.H{"success": true, "message": "", "data": view}) +} + +func parseUintPathParam(c *gin.Context, name string) (uint, bool) { + id, err := strconv.ParseUint(c.Param(name), 10, 64) + if err != nil || id == 0 { + c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": "invalid id"}) + return 0, false + } + return uint(id), true +} diff --git a/openflare_server/model/database_schema_version.go b/openflare_server/model/database_schema_version.go index f1654364..1f0e7224 100644 --- a/openflare_server/model/database_schema_version.go +++ b/openflare_server/model/database_schema_version.go @@ -4,7 +4,7 @@ import "time" const ( legacyDatabaseSchemaVersion = 1 - currentDatabaseSchemaVersion = 12 + currentDatabaseSchemaVersion = 13 databaseSchemaVersionRowID = 1 ) diff --git a/openflare_server/model/main.go b/openflare_server/model/main.go index 87f590e3..3daaf7b3 100644 --- a/openflare_server/model/main.go +++ b/openflare_server/model/main.go @@ -43,6 +43,8 @@ func registeredModels() []any { &ManagedDomain{}, &AcmeAccount{}, &DnsAccount{}, + &WAFRuleGroup{}, + &WAFRuleGroupBinding{}, } } diff --git a/openflare_server/model/migrations.go b/openflare_server/model/migrations.go index 46dc0ec3..c3b13812 100644 --- a/openflare_server/model/migrations.go +++ b/openflare_server/model/migrations.go @@ -1388,6 +1388,67 @@ func validateDatabaseSchemaV12(db *gorm.DB, backend string) error { return nil } +func ensureDefaultWAFRuleGroup(db *gorm.DB) error { + if db == nil { + return fmt.Errorf("database handle is nil") + } + if !db.Migrator().HasTable(&WAFRuleGroup{}) { + return nil + } + var count int64 + if err := db.Model(&WAFRuleGroup{}).Where("is_global = ?", true).Count(&count).Error; err != nil { + return fmt.Errorf("count global waf rule groups failed: %w", err) + } + if count > 0 { + return nil + } + group := WAFRuleGroup{ + Name: "全局规则组", + Enabled: true, + IsGlobal: true, + BlockStatusCode: 418, + IPWhitelist: "[]", + IPBlacklist: "[]", + CountryWhitelist: "[]", + CountryBlacklist: "[]", + RegionWhitelist: "[]", + RegionBlacklist: "[]", + BlockResponseBody: "", + } + if err := db.Create(&group).Error; err != nil { + return fmt.Errorf("create default waf rule group failed: %w", err) + } + return nil +} + +// migrateV13 adds WAF rule groups and website bindings. +func migrateV13(db *gorm.DB, backend string) error { + if err := applyCurrentSchema(db, backend); err != nil { + return err + } + return ensureDefaultWAFRuleGroup(db) +} + +func validateDatabaseSchemaV13(db *gorm.DB, backend string) error { + if err := validateDatabaseSchemaV12(db, backend); err != nil { + return err + } + if !db.Migrator().HasTable(&WAFRuleGroup{}) { + return fmt.Errorf("table waf_rule_groups is missing") + } + if !db.Migrator().HasTable(&WAFRuleGroupBinding{}) { + return fmt.Errorf("table waf_rule_group_bindings is missing") + } + var count int64 + if err := db.Model(&WAFRuleGroup{}).Where("is_global = ?", true).Count(&count).Error; err != nil { + return fmt.Errorf("count global waf rule groups failed: %w", err) + } + if count != 1 { + return fmt.Errorf("expected exactly one global waf rule group, got %d", count) + } + return nil +} + func databaseSchemaMigrations() []databaseSchemaMigration { return []databaseSchemaMigration{ {fromVersion: 1, toVersion: 2, migrate: migrateV2, validate: validateDatabaseSchemaV2}, @@ -1401,6 +1462,7 @@ func databaseSchemaMigrations() []databaseSchemaMigration { {fromVersion: 9, toVersion: 10, migrate: migrateV10, validate: validateDatabaseSchemaV10}, {fromVersion: 10, toVersion: 11, migrate: migrateV11, validate: validateDatabaseSchemaV11}, {fromVersion: 11, toVersion: 12, migrate: migrateV12, validate: validateDatabaseSchemaV12}, + {fromVersion: 12, toVersion: 13, migrate: migrateV13, validate: validateDatabaseSchemaV13}, } } @@ -1486,7 +1548,10 @@ func initializeFreshDatabaseSchema(db *gorm.DB, backend string) error { if err := ensureDefaultGitHubAuthSource(db); err != nil { return err } - if err := validateDatabaseSchemaV12(db, backend); err != nil { + if err := ensureDefaultWAFRuleGroup(db); err != nil { + return err + } + if err := validateDatabaseSchemaV13(db, backend); err != nil { return err } return saveDatabaseSchemaVersion(db, currentDatabaseSchemaVersion) diff --git a/openflare_server/model/waf.go b/openflare_server/model/waf.go new file mode 100644 index 00000000..9f56ca93 --- /dev/null +++ b/openflare_server/model/waf.go @@ -0,0 +1,71 @@ +package model + +import "time" + +type WAFRuleGroup struct { + ID uint `json:"id" gorm:"primaryKey"` + Name string `json:"name" gorm:"size:255;not null"` + Enabled bool `json:"enabled" gorm:"not null;default:true"` + IsGlobal bool `json:"is_global" gorm:"not null;default:false;index"` + BlockStatusCode int `json:"block_status_code" gorm:"not null;default:418"` + BlockResponseBody string `json:"block_response_body" gorm:"type:text;not null;default:''"` + IPWhitelist string `json:"ip_whitelist" gorm:"type:text;not null;default:'[]'"` + IPBlacklist string `json:"ip_blacklist" gorm:"type:text;not null;default:'[]'"` + 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:'[]'"` + Remark string `json:"remark" gorm:"size:255"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +type WAFRuleGroupBinding struct { + ID uint `json:"id" gorm:"primaryKey"` + RuleGroupID uint `json:"rule_group_id" gorm:"not null;uniqueIndex:idx_waf_group_route"` + ProxyRouteID uint `json:"proxy_route_id" gorm:"not null;uniqueIndex:idx_waf_group_route;index"` + CreatedAt time.Time `json:"created_at"` +} + +func ListWAFRuleGroups() ([]*WAFRuleGroup, error) { + var groups []*WAFRuleGroup + err := DB.Order("is_global desc").Order("id asc").Find(&groups).Error + return groups, err +} + +func GetWAFRuleGroupByID(id uint) (*WAFRuleGroup, error) { + group := &WAFRuleGroup{} + err := DB.First(group, id).Error + return group, err +} + +func GetGlobalWAFRuleGroup() (*WAFRuleGroup, error) { + group := &WAFRuleGroup{} + err := DB.Where("is_global = ?", true).Order("id asc").First(group).Error + return group, err +} + +func (group *WAFRuleGroup) Insert() error { + return DB.Create(group).Error +} + +func (group *WAFRuleGroup) Update() error { + return DB.Model(&WAFRuleGroup{}).Where("id = ?", group.ID).Updates(map[string]any{ + "name": group.Name, + "enabled": group.Enabled, + "is_global": group.IsGlobal, + "block_status_code": group.BlockStatusCode, + "block_response_body": group.BlockResponseBody, + "ip_whitelist": group.IPWhitelist, + "ip_blacklist": group.IPBlacklist, + "country_whitelist": group.CountryWhitelist, + "country_blacklist": group.CountryBlacklist, + "region_whitelist": group.RegionWhitelist, + "region_blacklist": group.RegionBlacklist, + "remark": group.Remark, + }).Error +} + +func (group *WAFRuleGroup) Delete() error { + return DB.Delete(group).Error +} diff --git a/openflare_server/router/api-router.go b/openflare_server/router/api-router.go index 03b407c0..bba7c664 100644 --- a/openflare_server/router/api-router.go +++ b/openflare_server/router/api-router.go @@ -94,6 +94,18 @@ func SetApiRouter(router *gin.Engine) { proxyRoute.POST("/:id/update", controller.UpdateProxyRoute) proxyRoute.POST("/:id/delete", controller.DeleteProxyRoute) } + wafRoute := apiRouter.Group("/waf") + wafRoute.Use(middleware.AdminAuth()) + { + wafRoute.GET("/rule-groups", controller.ListWAFRuleGroups) + wafRoute.GET("/rule-groups/:id", controller.GetWAFRuleGroup) + wafRoute.POST("/rule-groups", controller.CreateWAFRuleGroup) + wafRoute.POST("/rule-groups/:id/update", controller.UpdateWAFRuleGroup) + wafRoute.POST("/rule-groups/:id/delete", controller.DeleteWAFRuleGroup) + wafRoute.POST("/rule-groups/:id/sites", controller.ReplaceWAFRuleGroupSites) + wafRoute.GET("/sites/:route_id/rule-groups", controller.GetWAFSiteRuleGroups) + wafRoute.POST("/sites/:route_id/rule-groups", controller.ReplaceWAFSiteRuleGroups) + } originRoute := apiRouter.Group("/origins") originRoute.Use(middleware.AdminAuth()) { diff --git a/openflare_server/service/config_version.go b/openflare_server/service/config_version.go index 1a5c5644..a300e237 100644 --- a/openflare_server/service/config_version.go +++ b/openflare_server/service/config_version.go @@ -53,6 +53,7 @@ type ConfigDiffResult struct { RemovedDomains []string `json:"removed_domains"` ModifiedDomains []string `json:"modified_domains"` MainConfigChanged bool `json:"main_config_changed"` + WAFConfigChanged bool `json:"waf_config_changed"` ChangedOptionKeys []string `json:"changed_option_keys"` ChangedOptionDetails []ConfigOptionDiffItem `json:"changed_option_details"` CurrentWebsiteCount int `json:"current_website_count"` @@ -93,6 +94,32 @@ type snapshotRoute struct { Remark string `json:"remark,omitempty"` } +type snapshotWAFRuleGroup struct { + ID uint `json:"id"` + Name string `json:"name"` + Enabled bool `json:"enabled"` + IsGlobal bool `json:"is_global"` + BlockStatusCode int `json:"block_status_code"` + BlockResponseBody string `json:"block_response_body,omitempty"` + IPWhitelist []string `json:"ip_whitelist,omitempty"` + IPBlacklist []string `json:"ip_blacklist,omitempty"` + CountryWhitelist []string `json:"country_whitelist,omitempty"` + CountryBlacklist []string `json:"country_blacklist,omitempty"` + RegionWhitelist []string `json:"region_whitelist,omitempty"` + RegionBlacklist []string `json:"region_blacklist,omitempty"` +} + +type snapshotWAFBinding struct { + RouteID uint `json:"route_id"` + SiteName string `json:"site_name"` + RuleGroupIDs []uint `json:"rule_group_ids"` +} + +type snapshotWAFDocument struct { + RuleGroups []snapshotWAFRuleGroup `json:"rule_groups"` + Bindings []snapshotWAFBinding `json:"bindings"` +} + type routeCacheConfig struct { Enabled bool Policy string @@ -153,11 +180,13 @@ type openRestyConfigSnapshot struct { type snapshotDocument struct { Routes []snapshotRoute `json:"routes"` OpenRestyConfig openRestyConfigSnapshot `json:"openresty_config"` + WAF snapshotWAFDocument `json:"waf"` } type configBundle struct { Routes []*model.ProxyRoute SnapshotRoutes []snapshotRoute + WAFSnapshot snapshotWAFDocument OpenRestyConfig openRestyConfigSnapshot SnapshotJSON string MainConfig string @@ -310,6 +339,7 @@ func DiffConfigVersion() (*ConfigDiffResult, error) { } } result.MainConfigChanged = activeVersion.MainConfig != bundle.MainConfig + result.WAFConfigChanged = !snapshotWAFConfigEqual(activeSnapshot.WAF, bundle.WAFSnapshot) result.ChangedOptionDetails = diffOpenRestyOptionDetails(activeSnapshot.OpenRestyConfig, bundle.OpenRestyConfig) result.ChangedOptionKeys = extractOptionDiffKeys(result.ChangedOptionDetails) sort.Strings(result.AddedSites) @@ -460,10 +490,15 @@ func buildCurrentConfigBundle(requireRoutes bool) (*configBundle, error) { if err != nil { return nil, err } + wafSnapshot, err := buildSnapshotWAFDocument(routes) + if err != nil { + return nil, err + } openRestyConfig := buildOpenRestyConfigSnapshot() snapshotDoc := snapshotDocument{ Routes: snapshotRoutes, OpenRestyConfig: openRestyConfig, + WAF: wafSnapshot, } snapshotJSON, err := json.Marshal(snapshotDoc) if err != nil { @@ -473,6 +508,10 @@ func buildCurrentConfigBundle(requireRoutes bool) (*configBundle, error) { if err != nil { return nil, err } + wafConfigJSON, err := renderWAFConfigBundle(wafSnapshot) + if err != nil { + return nil, err + } powConfigJSON, powSupportFiles, err := renderPowConfigBundle(routes) if err != nil { return nil, err @@ -480,9 +519,11 @@ func buildCurrentConfigBundle(requireRoutes bool) (*configBundle, error) { supportFiles = append(supportFiles, powSupportFiles...) mainConfig := renderMainConfig(openRestyConfig) supportFiles = append(supportFiles, SupportFile{Path: "pow_config.json", Content: powConfigJSON}) + supportFiles = append(supportFiles, SupportFile{Path: "waf_config.json", Content: wafConfigJSON}) return &configBundle{ Routes: routes, SnapshotRoutes: snapshotRoutes, + WAFSnapshot: wafSnapshot, OpenRestyConfig: openRestyConfig, SnapshotJSON: string(snapshotJSON), MainConfig: mainConfig, @@ -550,6 +591,74 @@ func buildSnapshotRoutes(routes []*model.ProxyRoute) ([]snapshotRoute, error) { return items, nil } +func buildSnapshotWAFDocument(routes []*model.ProxyRoute) (snapshotWAFDocument, error) { + if err := EnsureDefaultWAFRuleGroup(); err != nil { + return snapshotWAFDocument{}, err + } + views, err := ListWAFRuleGroups() + if err != nil { + return snapshotWAFDocument{}, err + } + ruleGroups := make([]snapshotWAFRuleGroup, 0, len(views)) + for _, view := range views { + if !view.Enabled { + continue + } + ruleGroups = append(ruleGroups, snapshotWAFRuleGroup{ + ID: view.ID, + Name: view.Name, + Enabled: view.Enabled, + IsGlobal: view.IsGlobal, + BlockStatusCode: view.BlockStatusCode, + BlockResponseBody: view.BlockResponseBody, + IPWhitelist: view.IPWhitelist, + IPBlacklist: view.IPBlacklist, + CountryWhitelist: view.CountryWhitelist, + CountryBlacklist: view.CountryBlacklist, + RegionWhitelist: view.RegionWhitelist, + RegionBlacklist: view.RegionBlacklist, + }) + } + enabledRouteIDs := make(map[uint]string, len(routes)) + for _, route := range routes { + if route == nil { + continue + } + siteName := strings.TrimSpace(route.SiteName) + if siteName == "" { + siteName = route.Domain + } + enabledRouteIDs[route.ID] = siteName + } + var rawBindings []model.WAFRuleGroupBinding + if err := model.DB.Order("proxy_route_id asc").Order("rule_group_id asc").Find(&rawBindings).Error; err != nil { + return snapshotWAFDocument{}, err + } + groupIDsByRoute := make(map[uint][]uint, len(rawBindings)) + for _, binding := range rawBindings { + if _, ok := enabledRouteIDs[binding.ProxyRouteID]; !ok { + continue + } + groupIDsByRoute[binding.ProxyRouteID] = append(groupIDsByRoute[binding.ProxyRouteID], binding.RuleGroupID) + } + bindings := make([]snapshotWAFBinding, 0, len(groupIDsByRoute)) + for routeID, groupIDs := range groupIDsByRoute { + sort.Slice(groupIDs, func(i, j int) bool { return groupIDs[i] < groupIDs[j] }) + bindings = append(bindings, snapshotWAFBinding{ + RouteID: routeID, + SiteName: enabledRouteIDs[routeID], + RuleGroupIDs: groupIDs, + }) + } + sort.Slice(bindings, func(i, j int) bool { + if bindings[i].SiteName == bindings[j].SiteName { + return bindings[i].RouteID < bindings[j].RouteID + } + return bindings[i].SiteName < bindings[j].SiteName + }) + return snapshotWAFDocument{RuleGroups: ruleGroups, Bindings: bindings}, nil +} + func mustDecodeSnapshotCertIDs(route *model.ProxyRoute) []uint { if route == nil { return []uint{} @@ -729,6 +838,18 @@ func snapshotRouteConfigEqual(left snapshotRoute, right snapshotRoute) bool { return true } +func snapshotWAFConfigEqual(left snapshotWAFDocument, right snapshotWAFDocument) bool { + leftJSON, err := json.Marshal(left) + if err != nil { + return false + } + rightJSON, err := json.Marshal(right) + if err != nil { + return false + } + return string(leftJSON) == string(rightJSON) +} + func snapshotPoWConfigEqual(left *ProxyRoutePoWConfig, right *ProxyRoutePoWConfig) bool { if left == nil || right == nil { return left == nil && right == nil @@ -949,7 +1070,7 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot) builder.WriteString(renderNamedUpstreamBlock(upstreamConfig)) } if !route.EnableHTTPS { - builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg)) + builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg)) continue } certIDs, err := decodeStoredCertIDs(route.CertIDs, route.CertID) @@ -1010,24 +1131,24 @@ func renderRouteConfig(routes []*model.ProxyRoute, cfg openRestyConfigSnapshot) if route.RedirectHTTP { if len(httpOnlyDomains) > 0 { - builder.WriteString(renderHTTPProxyServer(renderServerNames(httpOnlyDomains), route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg)) + builder.WriteString(renderHTTPProxyServer(renderServerNames(httpOnlyDomains), displayName, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg)) } for _, certID := range certIDs { assignedDomains := domainsByCertID[certID] if len(assignedDomains) == 0 { continue } - builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains))) + builder.WriteString(renderHTTPRedirectServer(renderServerNames(assignedDomains), displayName)) } } else { - builder.WriteString(renderHTTPProxyServer(serverNames, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg)) + builder.WriteString(renderHTTPProxyServer(serverNames, displayName, route.OriginURL, route.OriginHost, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg)) } for _, certID := range certIDs { assignedDomains := domainsByCertID[certID] if len(assignedDomains) == 0 { continue } - builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), route.OriginURL, route.OriginHost, certID, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg)) + builder.WriteString(renderHTTPSServer(renderServerNames(assignedDomains), displayName, route.OriginURL, route.OriginHost, certID, customHeaders, cacheConfig, limitConfig, upstreamConfig, route.PoWEnabled, route.BasicAuthEnabled, route.BasicAuthUsername, route.BasicAuthPassword, cfg)) } } return builder.String(), dedupeSupportFiles(supportFiles), nil @@ -1143,6 +1264,11 @@ func renderPowAccessBlock(powEnabled bool) string { return fmt.Sprintf(" access_by_lua_file %s/pow/check.lua;\n", nginxLuaDirPlaceholder) } +func renderWAFAccessBlock(siteName string) string { + escapedSiteName := escapeNginxString(siteName) + return fmt.Sprintf(" set $openflare_waf_site \"%s\";\n access_by_lua_file %s/waf/check.lua;\n", escapedSiteName, nginxLuaDirPlaceholder) +} + func renderBasicAuthBlock(enabled bool, username, password string) string { if !enabled || username == "" || password == "" { return "" @@ -1299,18 +1425,19 @@ func nextVersionNumber(now time.Time) (string, error) { return fmt.Sprintf("%s-%03d", prefix, sequence+1), nil } -func renderHTTPProxyServer(serverNames string, originURL string, originHost string, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string { - return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, renderPowAccessBlock(powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled)) +func renderHTTPProxyServer(serverNames string, siteName string, originURL string, originHost string, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string { + return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n%s%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, renderWAFAccessBlock(siteName), renderPowAccessBlock(powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled)) } -func renderHTTPRedirectServer(serverNames string) string { +func renderHTTPRedirectServer(serverNames string, siteName string) string { + _ = siteName return fmt.Sprintf("server {\n listen 80;\n server_name %s;\n\n return 301 https://$host$request_uri;\n}\n\n", serverNames) } -func renderHTTPSServer(serverNames string, originURL string, originHost string, certificateID uint, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string { +func renderHTTPSServer(serverNames string, siteName string, originURL string, originHost string, certificateID uint, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string { certPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateCertFileName(certificateID)) keyPath := fmt.Sprintf("%s/%s", nginxCertDirPlaceholder, certificateKeyFileName(certificateID)) - return fmt.Sprintf("server {\n listen 443 ssl;\n http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, certPath, keyPath, renderPowAccessBlock(powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled)) + return fmt.Sprintf("server {\n listen 443 ssl;\n http2 on;\n server_name %s;\n ssl_certificate %s;\n ssl_certificate_key %s;\n%s%s%s location / {\n%s%s%s%s%s }\n%s}\n\n", serverNames, certPath, keyPath, renderWAFAccessBlock(siteName), renderPowAccessBlock(powEnabled), renderPowLocationBlocks(powEnabled), renderBasicAuthBlock(basicAuthEnabled, basicAuthUsername, basicAuthPassword), renderProxyHeaderBlock(originURL, originHost, customHeaders, upstreamConfig), renderRouteLimitBlock(limitConfig), renderRouteCacheBlock(cacheConfig, cfg), renderProxyPassBlock(originURL, upstreamConfig), renderPowStaticLocationBlock(powEnabled)) } func renderHTTPSServerWithCertificates(serverNames string, originURL string, originHost string, certificateIDs []uint, customHeaders []ProxyRouteCustomHeaderInput, cacheConfig routeCacheConfig, limitConfig routeLimitConfig, upstreamConfig routeUpstreamConfig, powEnabled bool, basicAuthEnabled bool, basicAuthUsername string, basicAuthPassword string, cfg openRestyConfigSnapshot) string { @@ -1648,6 +1775,12 @@ func quoteNginxStringLiteral(value string) string { return fmt.Sprintf(`"%s"`, escaped) } +func escapeNginxString(value string) string { + escaped := strings.ReplaceAll(value, `\`, `\\`) + escaped = strings.ReplaceAll(escaped, `"`, `\"`) + return escaped +} + func certificateCertFileName(id uint) string { return fmt.Sprintf("%d.crt", id) } @@ -1711,3 +1844,80 @@ func renderPowConfigBundle(routes []*model.ProxyRoute) (string, []SupportFile, e } return string(data), nil, nil } + +func renderWAFConfigBundle(snapshot snapshotWAFDocument) (string, error) { + type wafRuntimeRuleGroup struct { + ID uint `json:"id"` + Name string `json:"name"` + IsGlobal bool `json:"is_global"` + BlockStatusCode int `json:"block_status_code"` + BlockResponseBody string `json:"block_response_body"` + IPWhitelist []string `json:"ip_whitelist"` + IPBlacklist []string `json:"ip_blacklist"` + CountryWhitelist []string `json:"country_whitelist"` + CountryBlacklist []string `json:"country_blacklist"` + RegionWhitelist []string `json:"region_whitelist"` + RegionBlacklist []string `json:"region_blacklist"` + } + 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 = defaultWAFBlockStatusCode + } + if group.IsGlobal { + globalGroupIDs = append(globalGroupIDs, group.ID) + } + enabledGroupIDs[group.ID] = struct{}{} + groups = append(groups, wafRuntimeRuleGroup{ + ID: group.ID, + Name: group.Name, + IsGlobal: group.IsGlobal, + BlockStatusCode: statusCode, + BlockResponseBody: group.BlockResponseBody, + IPWhitelist: group.IPWhitelist, + IPBlacklist: group.IPBlacklist, + CountryWhitelist: group.CountryWhitelist, + CountryBlacklist: group.CountryBlacklist, + RegionWhitelist: group.RegionWhitelist, + RegionBlacklist: group.RegionBlacklist, + }) + } + 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) + } + runtimeConfig := wafRuntimeConfig{ + DefaultBlockStatusCode: defaultWAFBlockStatusCode, + RuleGroups: groups, + SiteRuleGroups: siteRuleGroups, + } + data, err := json.Marshal(runtimeConfig) + if err != nil { + return "", err + } + return string(data), nil +} diff --git a/openflare_server/service/openresty_observability_assets.go b/openflare_server/service/openresty_observability_assets.go index 637e7a21..bfe1b030 100644 --- a/openflare_server/service/openresty_observability_assets.go +++ b/openflare_server/service/openresty_observability_assets.go @@ -14,6 +14,7 @@ func renderOpenRestyObservabilityTemplateBlock() string { " lua_shared_dict openflare_pow_config 1m;", " lua_shared_dict openflare_pow_challenges 10m;", " lua_shared_dict openflare_pow_sessions 20m;", + " lua_shared_dict openflare_waf_config 2m;", fmt.Sprintf(" init_worker_by_lua_file %s/%s;", nginxLuaDirPlaceholder, openRestyObservabilityInitLuaPath), fmt.Sprintf(" log_by_lua_file %s/%s;", nginxLuaDirPlaceholder, openRestyObservabilityLogLuaPath), "", diff --git a/openflare_server/service/waf.go b/openflare_server/service/waf.go new file mode 100644 index 00000000..ef5337fb --- /dev/null +++ b/openflare_server/service/waf.go @@ -0,0 +1,514 @@ +package service + +import ( + "encoding/json" + "errors" + "fmt" + "net/netip" + "openflare/model" + "sort" + "strings" + "time" + "unicode" + + "gorm.io/gorm" +) + +const ( + defaultWAFBlockStatusCode = 418 + maxWAFBlockBodyBytes = 16 * 1024 +) + +type WAFRuleGroupInput struct { + Name string `json:"name"` + Enabled bool `json:"enabled"` + BlockStatusCode int `json:"block_status_code"` + BlockResponseBody string `json:"block_response_body"` + IPWhitelist []string `json:"ip_whitelist"` + IPBlacklist []string `json:"ip_blacklist"` + CountryWhitelist []string `json:"country_whitelist"` + CountryBlacklist []string `json:"country_blacklist"` + RegionWhitelist []string `json:"region_whitelist"` + RegionBlacklist []string `json:"region_blacklist"` + Remark string `json:"remark"` +} + +type WAFRuleGroupView struct { + ID uint `json:"id"` + Name string `json:"name"` + Enabled bool `json:"enabled"` + IsGlobal bool `json:"is_global"` + BlockStatusCode int `json:"block_status_code"` + BlockResponseBody string `json:"block_response_body"` + IPWhitelist []string `json:"ip_whitelist"` + IPBlacklist []string `json:"ip_blacklist"` + CountryWhitelist []string `json:"country_whitelist"` + CountryBlacklist []string `json:"country_blacklist"` + RegionWhitelist []string `json:"region_whitelist"` + RegionBlacklist []string `json:"region_blacklist"` + Remark string `json:"remark"` + AppliedSiteIDs []uint `json:"applied_site_ids"` + AppliedSiteCount int `json:"applied_site_count"` + CreatedAt string `json:"created_at"` + UpdatedAt string `json:"updated_at"` +} + +type WAFSiteRuleGroupsView struct { + RouteID uint `json:"route_id"` + GlobalRuleGroup *WAFRuleGroupView `json:"global_rule_group"` + RuleGroups []WAFRuleGroupView `json:"rule_groups"` + AppliedRuleGroups []WAFRuleGroupView `json:"applied_rule_groups"` + AppliedIDs []uint `json:"applied_ids"` +} + +func ListWAFRuleGroups() ([]WAFRuleGroupView, error) { + if err := EnsureDefaultWAFRuleGroup(); err != nil { + return nil, err + } + groups, err := model.ListWAFRuleGroups() + if err != nil { + return nil, err + } + bindings, err := loadWAFBindings() + if err != nil { + return nil, err + } + views := make([]WAFRuleGroupView, 0, len(groups)) + for _, group := range groups { + view, err := buildWAFRuleGroupView(group, bindings[group.ID]) + if err != nil { + return nil, err + } + views = append(views, view) + } + return views, nil +} + +func GetWAFRuleGroup(id uint) (*WAFRuleGroupView, error) { + group, err := model.GetWAFRuleGroupByID(id) + if err != nil { + return nil, err + } + bindings, err := loadWAFBindings() + if err != nil { + return nil, err + } + view, err := buildWAFRuleGroupView(group, bindings[group.ID]) + if err != nil { + return nil, err + } + return &view, nil +} + +func CreateWAFRuleGroup(input WAFRuleGroupInput) (*WAFRuleGroupView, error) { + group, err := buildWAFRuleGroup(nil, input) + if err != nil { + return nil, err + } + group.IsGlobal = false + if err := group.Insert(); err != nil { + return nil, err + } + return GetWAFRuleGroup(group.ID) +} + +func UpdateWAFRuleGroup(id uint, input WAFRuleGroupInput) (*WAFRuleGroupView, error) { + group, err := model.GetWAFRuleGroupByID(id) + if err != nil { + return nil, err + } + isGlobal := group.IsGlobal + group, err = buildWAFRuleGroup(group, input) + if err != nil { + return nil, err + } + group.IsGlobal = isGlobal + if isGlobal && strings.TrimSpace(group.Name) == "" { + group.Name = "全局规则组" + } + if err := group.Update(); err != nil { + return nil, err + } + return GetWAFRuleGroup(group.ID) +} + +func DeleteWAFRuleGroup(id uint) error { + group, err := model.GetWAFRuleGroupByID(id) + if err != nil { + return err + } + if group.IsGlobal { + return errors.New("全局 WAF 规则组不能删除") + } + return model.DB.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("rule_group_id = ?", group.ID).Delete(&model.WAFRuleGroupBinding{}).Error; err != nil { + return err + } + return tx.Delete(group).Error + }) +} + +func ReplaceWAFRuleGroupSites(groupID uint, routeIDs []uint) (*WAFRuleGroupView, error) { + group, err := model.GetWAFRuleGroupByID(groupID) + if err != nil { + return nil, err + } + if group.IsGlobal { + return nil, errors.New("全局 WAF 规则组默认应用到所有网站,不能手动绑定") + } + normalized, err := normalizeWAFRouteIDs(routeIDs) + if err != nil { + return nil, err + } + err = model.DB.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("rule_group_id = ?", groupID).Delete(&model.WAFRuleGroupBinding{}).Error; err != nil { + return err + } + for _, routeID := range normalized { + binding := model.WAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID} + if err := tx.Create(&binding).Error; err != nil { + return err + } + } + return nil + }) + if err != nil { + return nil, err + } + return GetWAFRuleGroup(groupID) +} + +func GetWAFSiteRuleGroups(routeID uint) (*WAFSiteRuleGroupsView, error) { + if _, err := model.GetProxyRouteByID(routeID); err != nil { + return nil, err + } + groups, err := ListWAFRuleGroups() + if err != nil { + return nil, err + } + appliedIDs, err := ListWAFSiteRuleGroupIDs(routeID) + if err != nil { + return nil, err + } + appliedSet := make(map[uint]struct{}, len(appliedIDs)) + for _, id := range appliedIDs { + appliedSet[id] = struct{}{} + } + var global *WAFRuleGroupView + custom := make([]WAFRuleGroupView, 0, len(groups)) + applied := make([]WAFRuleGroupView, 0, len(appliedIDs)) + for index := range groups { + group := groups[index] + if group.IsGlobal { + item := group + global = &item + continue + } + custom = append(custom, group) + if _, ok := appliedSet[group.ID]; ok { + applied = append(applied, group) + } + } + return &WAFSiteRuleGroupsView{ + RouteID: routeID, + GlobalRuleGroup: global, + RuleGroups: custom, + AppliedRuleGroups: applied, + AppliedIDs: appliedIDs, + }, nil +} + +func ReplaceWAFSiteRuleGroups(routeID uint, groupIDs []uint) (*WAFSiteRuleGroupsView, error) { + if _, err := model.GetProxyRouteByID(routeID); err != nil { + return nil, err + } + normalized, err := normalizeWAFRuleGroupIDs(groupIDs) + if err != nil { + return nil, err + } + err = model.DB.Transaction(func(tx *gorm.DB) error { + if err := tx.Where("proxy_route_id = ?", routeID).Delete(&model.WAFRuleGroupBinding{}).Error; err != nil { + return err + } + for _, groupID := range normalized { + binding := model.WAFRuleGroupBinding{RuleGroupID: groupID, ProxyRouteID: routeID} + if err := tx.Create(&binding).Error; err != nil { + return err + } + } + return nil + }) + if err != nil { + return nil, err + } + return GetWAFSiteRuleGroups(routeID) +} + +func ListWAFSiteRuleGroupIDs(routeID uint) ([]uint, error) { + var bindings []model.WAFRuleGroupBinding + if err := model.DB.Where("proxy_route_id = ?", routeID).Order("rule_group_id asc").Find(&bindings).Error; err != nil { + return nil, err + } + ids := make([]uint, 0, len(bindings)) + for _, binding := range bindings { + ids = append(ids, binding.RuleGroupID) + } + return ids, nil +} + +func EnsureDefaultWAFRuleGroup() error { + _, err := model.GetGlobalWAFRuleGroup() + if err == nil { + return nil + } + if !errors.Is(err, gorm.ErrRecordNotFound) { + return err + } + group := &model.WAFRuleGroup{ + Name: "全局规则组", + Enabled: true, + IsGlobal: true, + BlockStatusCode: defaultWAFBlockStatusCode, + IPWhitelist: "[]", + IPBlacklist: "[]", + CountryWhitelist: "[]", + CountryBlacklist: "[]", + RegionWhitelist: "[]", + RegionBlacklist: "[]", + BlockResponseBody: "", + } + return group.Insert() +} + +func buildWAFRuleGroup(group *model.WAFRuleGroup, input WAFRuleGroupInput) (*model.WAFRuleGroup, error) { + name := strings.TrimSpace(input.Name) + if name == "" { + return nil, errors.New("规则组名称不能为空") + } + statusCode := input.BlockStatusCode + if statusCode == 0 { + statusCode = defaultWAFBlockStatusCode + } + if statusCode < 400 || statusCode > 599 { + return nil, errors.New("拦截状态码必须在 400-599 之间") + } + if len([]byte(input.BlockResponseBody)) > maxWAFBlockBodyBytes { + return nil, fmt.Errorf("拦截页面内容不能超过 %d 字节", maxWAFBlockBodyBytes) + } + ipWhitelist, err := normalizeWAFIPList(input.IPWhitelist) + if err != nil { + return nil, fmt.Errorf("IP 白名单无效: %w", err) + } + ipBlacklist, err := normalizeWAFIPList(input.IPBlacklist) + if err != nil { + return nil, fmt.Errorf("IP 黑名单无效: %w", err) + } + countryWhitelist, err := normalizeWAFCountryList(input.CountryWhitelist) + if err != nil { + return nil, fmt.Errorf("地域白名单无效: %w", err) + } + countryBlacklist, err := normalizeWAFCountryList(input.CountryBlacklist) + if err != nil { + return nil, fmt.Errorf("地域黑名单无效: %w", err) + } + regionWhitelist := normalizeStringList(input.RegionWhitelist) + regionBlacklist := normalizeStringList(input.RegionBlacklist) + + ipWhitelistJSON, _ := json.Marshal(ipWhitelist) + ipBlacklistJSON, _ := json.Marshal(ipBlacklist) + countryWhitelistJSON, _ := json.Marshal(countryWhitelist) + countryBlacklistJSON, _ := json.Marshal(countryBlacklist) + regionWhitelistJSON, _ := json.Marshal(regionWhitelist) + regionBlacklistJSON, _ := json.Marshal(regionBlacklist) + + if group == nil { + group = &model.WAFRuleGroup{} + } + group.Name = name + group.Enabled = input.Enabled + group.BlockStatusCode = statusCode + group.BlockResponseBody = input.BlockResponseBody + group.IPWhitelist = string(ipWhitelistJSON) + group.IPBlacklist = string(ipBlacklistJSON) + group.CountryWhitelist = string(countryWhitelistJSON) + group.CountryBlacklist = string(countryBlacklistJSON) + group.RegionWhitelist = string(regionWhitelistJSON) + group.RegionBlacklist = string(regionBlacklistJSON) + group.Remark = strings.TrimSpace(input.Remark) + return group, nil +} + +func buildWAFRuleGroupView(group *model.WAFRuleGroup, appliedSiteIDs []uint) (WAFRuleGroupView, error) { + if group == nil { + return WAFRuleGroupView{}, errors.New("waf rule group is nil") + } + sort.Slice(appliedSiteIDs, func(i, j int) bool { return appliedSiteIDs[i] < appliedSiteIDs[j] }) + view := WAFRuleGroupView{ + ID: group.ID, + Name: group.Name, + Enabled: group.Enabled, + IsGlobal: group.IsGlobal, + BlockStatusCode: group.BlockStatusCode, + BlockResponseBody: group.BlockResponseBody, + Remark: group.Remark, + AppliedSiteIDs: appliedSiteIDs, + AppliedSiteCount: len(appliedSiteIDs), + CreatedAt: group.CreatedAt.Format(time.RFC3339), + UpdatedAt: group.UpdatedAt.Format(time.RFC3339), + } + var err error + if view.IPWhitelist, err = decodeStringList(group.IPWhitelist); err != nil { + return view, err + } + if view.IPBlacklist, err = decodeStringList(group.IPBlacklist); err != nil { + return view, err + } + if view.CountryWhitelist, err = decodeStringList(group.CountryWhitelist); err != nil { + return view, err + } + if view.CountryBlacklist, err = decodeStringList(group.CountryBlacklist); err != nil { + return view, err + } + if view.RegionWhitelist, err = decodeStringList(group.RegionWhitelist); err != nil { + return view, err + } + if view.RegionBlacklist, err = decodeStringList(group.RegionBlacklist); err != nil { + return view, err + } + return view, nil +} + +func loadWAFBindings() (map[uint][]uint, error) { + var bindings []model.WAFRuleGroupBinding + if err := model.DB.Order("rule_group_id asc").Order("proxy_route_id asc").Find(&bindings).Error; err != nil { + return nil, err + } + result := make(map[uint][]uint, len(bindings)) + for _, binding := range bindings { + result[binding.RuleGroupID] = append(result[binding.RuleGroupID], binding.ProxyRouteID) + } + return result, nil +} + +func normalizeWAFIPList(items []string) ([]string, error) { + normalized := make([]string, 0, len(items)) + seen := make(map[string]struct{}, len(items)) + for _, raw := range items { + item := strings.TrimSpace(raw) + if item == "" { + continue + } + if strings.Contains(item, "/") { + prefix, err := netip.ParsePrefix(item) + if err != nil { + return nil, fmt.Errorf("%s 不是合法 IP 段", item) + } + item = prefix.Masked().String() + } else { + addr, err := netip.ParseAddr(item) + if err != nil { + return nil, fmt.Errorf("%s 不是合法 IP", item) + } + item = addr.String() + } + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + normalized = append(normalized, item) + } + sort.Strings(normalized) + return normalized, nil +} + +func normalizeWAFCountryList(items []string) ([]string, error) { + normalized := make([]string, 0, len(items)) + seen := make(map[string]struct{}, len(items)) + for _, raw := range items { + item := strings.ToUpper(strings.TrimSpace(raw)) + if item == "" { + continue + } + if len(item) != 2 || !unicode.IsLetter(rune(item[0])) || !unicode.IsLetter(rune(item[1])) { + return nil, fmt.Errorf("%s 不是合法国家代码", item) + } + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + normalized = append(normalized, item) + } + sort.Strings(normalized) + return normalized, nil +} + +func normalizeStringList(items []string) []string { + normalized := make([]string, 0, len(items)) + seen := make(map[string]struct{}, len(items)) + for _, raw := range items { + item := strings.TrimSpace(raw) + if item == "" { + continue + } + if _, ok := seen[item]; ok { + continue + } + seen[item] = struct{}{} + normalized = append(normalized, item) + } + sort.Strings(normalized) + return normalized +} + +func decodeStringList(raw string) ([]string, error) { + text := strings.TrimSpace(raw) + if text == "" { + return []string{}, nil + } + var items []string + if err := json.Unmarshal([]byte(text), &items); err != nil { + return nil, err + } + return items, nil +} + +func normalizeWAFRouteIDs(routeIDs []uint) ([]uint, error) { + normalized := uniqueUintIDs(routeIDs) + for _, routeID := range normalized { + if _, err := model.GetProxyRouteByID(routeID); err != nil { + return nil, fmt.Errorf("网站 %d 不存在", routeID) + } + } + return normalized, nil +} + +func normalizeWAFRuleGroupIDs(groupIDs []uint) ([]uint, error) { + normalized := uniqueUintIDs(groupIDs) + for _, groupID := range normalized { + group, err := model.GetWAFRuleGroupByID(groupID) + if err != nil { + return nil, fmt.Errorf("WAF 规则组 %d 不存在", groupID) + } + if group.IsGlobal { + return nil, errors.New("全局 WAF 规则组不需要手动绑定") + } + } + return normalized, nil +} + +func uniqueUintIDs(ids []uint) []uint { + seen := make(map[uint]struct{}, len(ids)) + normalized := make([]uint, 0, len(ids)) + for _, id := range ids { + if id == 0 { + continue + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + normalized = append(normalized, id) + } + sort.Slice(normalized, func(i, j int) bool { return normalized[i] < normalized[j] }) + return normalized +} diff --git a/openflare_server/service/waf_test.go b/openflare_server/service/waf_test.go new file mode 100644 index 00000000..1b89af36 --- /dev/null +++ b/openflare_server/service/waf_test.go @@ -0,0 +1,130 @@ +package service + +import ( + "encoding/json" + "strings" + "testing" +) + +func TestWAFRuleGroupValidationAndNormalization(t *testing.T) { + setupServiceTestDB(t) + + group, err := CreateWAFRuleGroup(WAFRuleGroupInput{ + 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"}, + }) + if err != nil { + t.Fatalf("CreateWAFRuleGroup failed: %v", err) + } + if len(group.IPWhitelist) != 2 || group.IPWhitelist[0] != "192.0.2.1" || group.IPWhitelist[1] != "198.51.100.0/24" { + t.Fatalf("unexpected normalized ip whitelist: %#v", group.IPWhitelist) + } + if len(group.CountryBlacklist) != 2 || group.CountryBlacklist[0] != "CN" || group.CountryBlacklist[1] != "US" { + t.Fatalf("unexpected normalized countries: %#v", group.CountryBlacklist) + } + + if _, err = CreateWAFRuleGroup(WAFRuleGroupInput{ + Name: "bad ip", + Enabled: true, + IPBlacklist: []string{"not-an-ip"}, + }); err == nil { + t.Fatal("expected invalid IP to be rejected") + } +} + +func TestWAFGlobalGroupAndBindings(t *testing.T) { + setupServiceTestDB(t) + + groups, err := ListWAFRuleGroups() + if err != nil { + t.Fatalf("ListWAFRuleGroups failed: %v", err) + } + if len(groups) == 0 || !groups[0].IsGlobal { + t.Fatalf("expected default global WAF rule group, got %#v", groups) + } + if err = DeleteWAFRuleGroup(groups[0].ID); err == nil { + t.Fatal("expected global WAF rule group delete to be rejected") + } + + route, err := CreateProxyRoute(ProxyRouteInput{ + SiteName: "waf-site", + Domains: []string{"waf.example.com"}, + OriginURL: "https://origin.internal", + Enabled: true, + }) + if err != nil { + t.Fatalf("CreateProxyRoute failed: %v", err) + } + custom, err := CreateWAFRuleGroup(WAFRuleGroupInput{ + Name: "custom", + Enabled: true, + BlockStatusCode: 418, + IPBlacklist: []string{"203.0.113.10"}, + }) + if err != nil { + t.Fatalf("CreateWAFRuleGroup failed: %v", err) + } + if _, err = ReplaceWAFRuleGroupSites(custom.ID, []uint{route.ID}); err != nil { + t.Fatalf("ReplaceWAFRuleGroupSites failed: %v", err) + } + siteGroups, err := GetWAFSiteRuleGroups(route.ID) + if err != nil { + t.Fatalf("GetWAFSiteRuleGroups failed: %v", err) + } + if len(siteGroups.AppliedIDs) != 1 || siteGroups.AppliedIDs[0] != custom.ID { + t.Fatalf("unexpected site WAF bindings: %#v", siteGroups.AppliedIDs) + } +} + +func TestPublishConfigVersionIncludesWAFSnapshotAndRuntimeConfig(t *testing.T) { + setupServiceTestDB(t) + + route, err := CreateProxyRoute(ProxyRouteInput{ + SiteName: "waf-publish", + Domains: []string{"waf-publish.example.com"}, + OriginURL: "https://origin.internal", + Enabled: true, + }) + if err != nil { + t.Fatalf("CreateProxyRoute failed: %v", err) + } + group, err := CreateWAFRuleGroup(WAFRuleGroupInput{ + Name: "publish group", + Enabled: true, + BlockStatusCode: 451, + IPBlacklist: []string{"203.0.113.0/24"}, + }) + if err != nil { + t.Fatalf("CreateWAFRuleGroup failed: %v", err) + } + if _, err = ReplaceWAFSiteRuleGroups(route.ID, []uint{group.ID}); err != nil { + t.Fatalf("ReplaceWAFSiteRuleGroups failed: %v", err) + } + result, err := PublishConfigVersion("root", false) + if err != nil { + t.Fatalf("PublishConfigVersion failed: %v", err) + } + if !strings.Contains(result.Version.RenderedConfig, "access_by_lua_file __OPENFLARE_LUA_DIR__/waf/check.lua;") { + t.Fatal("expected route config to include WAF lua access hook") + } + if !strings.Contains(result.Version.SnapshotJSON, `"waf"`) { + t.Fatal("expected snapshot to include waf document") + } + var files []SupportFile + if err = json.Unmarshal([]byte(result.Version.SupportFilesJSON), &files); err != nil { + t.Fatalf("decode support files failed: %v", err) + } + found := false + for _, file := range files { + if file.Path == "waf_config.json" && strings.Contains(file.Content, "203.0.113.0/24") { + found = true + } + } + if !found { + t.Fatalf("expected waf_config.json support file, got %#v", files) + } +} diff --git a/openflare_server/utils/geoip/mmdb.go b/openflare_server/utils/geoip/mmdb.go index 216a2860..5cc4a2ac 100644 --- a/openflare_server/utils/geoip/mmdb.go +++ b/openflare_server/utils/geoip/mmdb.go @@ -33,8 +33,18 @@ func (s *MaxMindGeoIPService) Name() string { } func NewMaxMindGeoIPService() (*MaxMindGeoIPService, error) { + return NewMaxMindGeoIPServiceWithConfig(GeoIpFilePath, GeoIpUrl) +} + +func NewMaxMindGeoIPServiceWithConfig(dbFilePath string, downloadURL string) (*MaxMindGeoIPService, error) { + if dbFilePath == "" { + dbFilePath = GeoIpFilePath + } + if downloadURL == "" { + downloadURL = GeoIpUrl + } service := &MaxMindGeoIPService{ - dbFilePath: GeoIpFilePath, + dbFilePath: dbFilePath, } if err := os.MkdirAll(filepath.Dir(service.dbFilePath), os.ModePerm); err != nil { @@ -42,7 +52,7 @@ func NewMaxMindGeoIPService() (*MaxMindGeoIPService, error) { } if _, err := os.Stat(service.dbFilePath); os.IsNotExist(err) { - if err := service.UpdateDatabase(); err != nil { + if err := DownloadMaxMindDatabase(service.dbFilePath, downloadURL); err != nil { return nil, fmt.Errorf("failed to download initial MaxMind database: %w", err) } } @@ -99,7 +109,20 @@ func (s *MaxMindGeoIPService) GetGeoInfo(ip net.IP) (*GeoInfo, error) { } func (s *MaxMindGeoIPService) UpdateDatabase() error { - resp, err := http.Get(GeoIpUrl) + if err := DownloadMaxMindDatabase(s.dbFilePath, GeoIpUrl); err != nil { + return err + } + return s.initialize() +} + +func DownloadMaxMindDatabase(dbFilePath string, downloadURL string) error { + if dbFilePath == "" { + dbFilePath = GeoIpFilePath + } + if downloadURL == "" { + downloadURL = GeoIpUrl + } + resp, err := http.Get(downloadURL) if err != nil { return fmt.Errorf("failed to initiate MaxMind database download: %w", err) } @@ -109,11 +132,11 @@ func (s *MaxMindGeoIPService) UpdateDatabase() error { return fmt.Errorf("failed to download MaxMind database: HTTP status %s", resp.Status) } - if err := os.MkdirAll(filepath.Dir(s.dbFilePath), os.ModePerm); err != nil { + if err := os.MkdirAll(filepath.Dir(dbFilePath), os.ModePerm); err != nil { return fmt.Errorf("failed to create data directory for MaxMind database update: %w", err) } - tempPath := s.dbFilePath + ".download" + tempPath := dbFilePath + ".download" out, err := os.Create(tempPath) if err != nil { return fmt.Errorf("failed to create MaxMind database file at %s: %w", tempPath, err) @@ -128,11 +151,10 @@ func (s *MaxMindGeoIPService) UpdateDatabase() error { if err = out.Close(); err != nil { return fmt.Errorf("failed to close MaxMind database file: %w", err) } - if err = os.Rename(tempPath, s.dbFilePath); err != nil { + if err = os.Rename(tempPath, dbFilePath); err != nil { return fmt.Errorf("failed to move MaxMind database file into place: %w", err) } - - return s.initialize() + return nil } func (s *MaxMindGeoIPService) Close() error { diff --git a/openflare_server/web/app/(dashboard)/waf/page.tsx b/openflare_server/web/app/(dashboard)/waf/page.tsx new file mode 100644 index 00000000..51235f53 --- /dev/null +++ b/openflare_server/web/app/(dashboard)/waf/page.tsx @@ -0,0 +1,5 @@ +import { WAFPage } from '@/features/waf/components/waf-page'; + +export default function WAFRoute() { + return ; +} diff --git a/openflare_server/web/components/layout/dashboard-sidebar.tsx b/openflare_server/web/components/layout/dashboard-sidebar.tsx index 99209220..efa4926c 100644 --- a/openflare_server/web/components/layout/dashboard-sidebar.tsx +++ b/openflare_server/web/components/layout/dashboard-sidebar.tsx @@ -3,6 +3,7 @@ import { useEffect } from 'react'; import Link from 'next/link'; import { usePathname } from 'next/navigation'; +import { ShieldCheck } from 'lucide-react'; import { dashboardNavigation } from '@/lib/constants/navigation'; import { cn } from '@/lib/utils/cn'; @@ -79,6 +80,8 @@ function SidebarIcon({ icon }: { icon: NavigationIconKey }) { ); + case 'waf': + return ; case 'release': return ( diff --git a/openflare_server/web/features/config-versions/types.ts b/openflare_server/web/features/config-versions/types.ts index 2a595f20..9919e32c 100644 --- a/openflare_server/web/features/config-versions/types.ts +++ b/openflare_server/web/features/config-versions/types.ts @@ -39,6 +39,7 @@ export interface ConfigDiffResult { removed_domains: string[]; modified_domains: string[]; main_config_changed: boolean; + waf_config_changed: boolean; changed_option_keys: string[]; changed_option_details: ConfigOptionDiffItem[]; current_website_count: number; diff --git a/openflare_server/web/features/nodes/components/node-detail-page.tsx b/openflare_server/web/features/nodes/components/node-detail-page.tsx index d69eaa2e..75f45a77 100644 --- a/openflare_server/web/features/nodes/components/node-detail-page.tsx +++ b/openflare_server/web/features/nodes/components/node-detail-page.tsx @@ -28,7 +28,6 @@ import { requestNodeForceSync, requestNodeOpenrestyRestart, requestNodeAgentUpdate, - rotateNodeBootstrapToken, updateNode, } from '@/features/nodes/api/nodes'; import { NodeEditorModal } from '@/features/nodes/components/node-editor-modal'; diff --git a/openflare_server/web/features/proxy-routes/components/proxy-routes-page.tsx b/openflare_server/web/features/proxy-routes/components/proxy-routes-page.tsx index 51cdf43c..1785e603 100644 --- a/openflare_server/web/features/proxy-routes/components/proxy-routes-page.tsx +++ b/openflare_server/web/features/proxy-routes/components/proxy-routes-page.tsx @@ -48,6 +48,7 @@ function hasConfigChanges(diff: { removed_domains: string[]; modified_domains: string[]; main_config_changed: boolean; + waf_config_changed?: boolean; changed_option_keys: string[]; }) { return ( @@ -58,6 +59,7 @@ function hasConfigChanges(diff: { diff.removed_domains.length > 0 || diff.modified_domains.length > 0 || diff.main_config_changed || + Boolean(diff.waf_config_changed) || diff.changed_option_keys.length > 0 || !diff.active_version ); diff --git a/openflare_server/web/features/waf/api/waf.ts b/openflare_server/web/features/waf/api/waf.ts new file mode 100644 index 00000000..155f96cb --- /dev/null +++ b/openflare_server/web/features/waf/api/waf.ts @@ -0,0 +1,49 @@ +import { apiRequest } from '@/lib/api/client'; + +import type { + WAFRuleGroup, + WAFRuleGroupPayload, + WAFSiteRuleGroups, +} from '@/features/waf/types'; + +export function getWAFRuleGroups() { + return apiRequest('/waf/rule-groups'); +} + +export function createWAFRuleGroup(payload: WAFRuleGroupPayload) { + return apiRequest('/waf/rule-groups', { + method: 'POST', + body: JSON.stringify(payload), + }); +} + +export function updateWAFRuleGroup(id: number, payload: WAFRuleGroupPayload) { + return apiRequest(`/waf/rule-groups/${id}/update`, { + method: 'POST', + body: JSON.stringify(payload), + }); +} + +export function deleteWAFRuleGroup(id: number) { + return apiRequest(`/waf/rule-groups/${id}/delete`, { + method: 'POST', + }); +} + +export function replaceWAFRuleGroupSites(id: number, ids: number[]) { + return apiRequest(`/waf/rule-groups/${id}/sites`, { + method: 'POST', + body: JSON.stringify({ ids }), + }); +} + +export function getWAFSiteRuleGroups(routeId: number) { + return apiRequest(`/waf/sites/${routeId}/rule-groups`); +} + +export function replaceWAFSiteRuleGroups(routeId: number, ids: number[]) { + return apiRequest(`/waf/sites/${routeId}/rule-groups`, { + method: 'POST', + body: JSON.stringify({ ids }), + }); +} diff --git a/openflare_server/web/features/waf/components/waf-page.tsx b/openflare_server/web/features/waf/components/waf-page.tsx new file mode 100644 index 00000000..fb278cf5 --- /dev/null +++ b/openflare_server/web/features/waf/components/waf-page.tsx @@ -0,0 +1,525 @@ +'use client'; + +import { useMutation, useQuery, useQueryClient } from '@tanstack/react-query'; +import { Check, Globe2, Plus, Save, Search, ShieldCheck, Trash2 } from 'lucide-react'; +import { useEffect, useMemo, useState } from 'react'; + +import { EmptyState } from '@/components/feedback/empty-state'; +import { ErrorState } from '@/components/feedback/error-state'; +import { InlineMessage } from '@/components/feedback/inline-message'; +import { LoadingState } from '@/components/feedback/loading-state'; +import { PageHeader } from '@/components/layout/page-header'; +import { AppCard } from '@/components/ui/app-card'; +import { Drawer } from '@/components/ui/drawer'; +import { getProxyRoutes } from '@/features/proxy-routes/api/proxy-routes'; +import type { ProxyRouteItem } from '@/features/proxy-routes/types'; +import { + DangerButton, + PrimaryButton, + ResourceField, + ResourceInput, + ResourceTextarea, + SecondaryButton, + ToggleField, +} from '@/features/shared/components/resource-primitives'; +import { + createWAFRuleGroup, + deleteWAFRuleGroup, + getWAFRuleGroups, + replaceWAFRuleGroupSites, + updateWAFRuleGroup, +} from '@/features/waf/api/waf'; +import type { WAFRuleGroup, WAFRuleGroupPayload } from '@/features/waf/types'; +import { cn } from '@/lib/utils/cn'; + +type FeedbackState = { + tone: 'success' | 'danger' | 'info'; + message: string; +}; + +const emptyDraft: WAFRuleGroupPayload = { + name: '', + enabled: true, + block_status_code: 418, + block_response_body: '', + ip_whitelist: [], + ip_blacklist: [], + country_whitelist: [], + country_blacklist: [], + region_whitelist: [], + region_blacklist: [], + remark: '', +}; + +function getErrorMessage(error: unknown) { + return error instanceof Error ? error.message : '操作失败'; +} + +function listToText(items: string[]) { + return items.join('\n'); +} + +function textToList(text: string) { + return text + .split(/[\n,,\s]+/) + .map((item) => item.trim()) + .filter(Boolean); +} + +function buildDraft(group: WAFRuleGroup | null): WAFRuleGroupPayload { + if (!group) { + return { ...emptyDraft }; + } + return { + name: group.name, + enabled: group.enabled, + block_status_code: group.block_status_code || 418, + block_response_body: group.block_response_body ?? '', + ip_whitelist: group.ip_whitelist ?? [], + ip_blacklist: group.ip_blacklist ?? [], + country_whitelist: group.country_whitelist ?? [], + country_blacklist: group.country_blacklist ?? [], + region_whitelist: group.region_whitelist ?? [], + region_blacklist: group.region_blacklist ?? [], + remark: group.remark ?? '', + }; +} + +function ruleCount(group: WAFRuleGroup) { + return ( + group.ip_whitelist.length + + group.ip_blacklist.length + + group.country_whitelist.length + + group.country_blacklist.length + ); +} + +function SiteApplyDrawer({ + group, + routes, + open, + onOpenChange, + onSave, + pending, +}: { + group: WAFRuleGroup | null; + routes: ProxyRouteItem[]; + open: boolean; + onOpenChange: (open: boolean) => void; + onSave: (ids: number[]) => void; + pending: boolean; +}) { + const [keyword, setKeyword] = useState(''); + const [selectedIDs, setSelectedIDs] = useState([]); + + useEffect(() => { + setSelectedIDs(group?.applied_site_ids ?? []); + setKeyword(''); + }, [group, open]); + + const filteredRoutes = useMemo(() => { + const normalized = keyword.trim().toLowerCase(); + if (!normalized) { + return routes; + } + return routes.filter((route) => + [route.site_name, route.primary_domain, ...route.domains] + .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].sort((left, right) => left - right), + ); + }; + + const selectFiltered = () => { + const next = new Set(selectedIDs); + filteredRoutes.forEach((route) => next.add(route.id)); + setSelectedIDs([...next].sort((left, right) => left - right)); + }; + + return ( + + onOpenChange(false)}> + 取消 + + onSave(selectedIDs)} + > + {pending ? '保存中...' : '保存应用范围'} + + + } + > +
+
+ + setKeyword(event.target.value)} + placeholder="搜索网站或域名" + className="min-w-0 flex-1 bg-transparent text-sm text-[var(--foreground-primary)] outline-none placeholder:text-[var(--foreground-muted)]" + /> + +
+
+ {filteredRoutes.map((route) => ( + + ))} +
+
+ + ); +} + +export function WAFPage() { + const queryClient = useQueryClient(); + const [selectedID, setSelectedID] = useState(null); + const [draft, setDraft] = useState(emptyDraft); + const [feedback, setFeedback] = useState(null); + const [applyGroup, setApplyGroup] = useState(null); + + const groupsQuery = useQuery({ + queryKey: ['waf', 'rule-groups'], + queryFn: getWAFRuleGroups, + }); + const routesQuery = useQuery({ + queryKey: ['proxy-routes'], + queryFn: getProxyRoutes, + }); + + const groups = useMemo(() => groupsQuery.data ?? [], [groupsQuery.data]); + const routes = useMemo(() => routesQuery.data ?? [], [routesQuery.data]); + const selectedGroup = useMemo( + () => + selectedID === 0 + ? null + : (groups.find((group) => group.id === selectedID) ?? groups[0] ?? null), + [groups, selectedID], + ); + + useEffect(() => { + if (selectedGroup) { + setSelectedID(selectedGroup.id); + setDraft(buildDraft(selectedGroup)); + } + }, [selectedGroup]); + + const invalidate = async () => { + await Promise.all([ + queryClient.invalidateQueries({ queryKey: ['waf', 'rule-groups'] }), + queryClient.invalidateQueries({ queryKey: ['config-versions', 'diff'] }), + ]); + }; + + const saveMutation = useMutation({ + mutationFn: (payload: WAFRuleGroupPayload) => { + if (selectedGroup) { + return updateWAFRuleGroup(selectedGroup.id, payload); + } + return createWAFRuleGroup(payload); + }, + onSuccess: async (group) => { + setSelectedID(group.id); + setFeedback({ tone: 'success', message: 'WAF 规则组已保存。' }); + await invalidate(); + }, + onError: (error) => { + setFeedback({ tone: 'danger', message: getErrorMessage(error) }); + }, + }); + + const deleteMutation = useMutation({ + mutationFn: deleteWAFRuleGroup, + onSuccess: async () => { + setSelectedID(null); + setFeedback({ tone: 'success', message: 'WAF 规则组已删除。' }); + await invalidate(); + }, + onError: (error) => { + setFeedback({ tone: 'danger', message: getErrorMessage(error) }); + }, + }); + + const applyMutation = useMutation({ + mutationFn: ({ id, ids }: { id: number; ids: number[] }) => + replaceWAFRuleGroupSites(id, ids), + onSuccess: async () => { + setApplyGroup(null); + setFeedback({ tone: 'success', message: '规则组应用范围已更新。' }); + await invalidate(); + }, + onError: (error) => { + setFeedback({ tone: 'danger', message: getErrorMessage(error) }); + }, + }); + + if (groupsQuery.isLoading || routesQuery.isLoading) { + return ; + } + if (groupsQuery.isError) { + return ; + } + if (routesQuery.isError) { + return ; + } + if (!selectedGroup && groups.length === 0) { + return ; + } + + const enabledCount = groups.filter((group) => group.enabled).length; + const protectedSites = new Set(groups.flatMap((group) => group.applied_site_ids)); + const totalRules = groups.reduce((sum, group) => sum + ruleCount(group), 0); + + return ( + <> +
+ { + setSelectedID(0); + setDraft({ ...emptyDraft, name: '自定义规则组' }); + }} + > + + 新建规则组 + + } + /> + + {feedback ? : null} + +
+ +

启用规则组

+

{enabledCount}

+
+ +

自定义覆盖网站

+

{protectedSites.size}

+
+ +

黑白名单条目

+

{totalRules}

+
+
+ +
+ +
+ {groups.map((group) => ( + + ))} +
+
+ + setApplyGroup(selectedGroup)}> + 一键应用 + + ) : null + } + > +
+ + setDraft((current) => ({ ...current, name: event.target.value }))} + /> + + + + setDraft((current) => ({ ...current, block_status_code: Number(event.target.value) })) + } + /> + + setDraft((current) => ({ ...current, enabled: checked }))} + /> + + setDraft((current) => ({ ...current, remark: event.target.value }))} + /> + + + + setDraft((current) => ({ ...current, ip_whitelist: textToList(event.target.value) })) + } + /> + + + + setDraft((current) => ({ ...current, ip_blacklist: textToList(event.target.value) })) + } + /> + + + + setDraft((current) => ({ ...current, country_whitelist: textToList(event.target.value) })) + } + /> + + + + setDraft((current) => ({ ...current, country_blacklist: textToList(event.target.value) })) + } + /> + + + + setDraft((current) => ({ ...current, block_response_body: event.target.value })) + } + /> + +
+
+
+ {selectedGroup && !selectedGroup.is_global ? ( + { + if (window.confirm(`确认删除 WAF 规则组 ${selectedGroup.name} 吗?`)) { + deleteMutation.mutate(selectedGroup.id); + } + }} + > + + 删除 + + ) : null} +
+ saveMutation.mutate(draft)} + > + + {saveMutation.isPending ? '保存中...' : '保存规则组'} + +
+
+
+
+ { + if (!open) { + setApplyGroup(null); + } + }} + onSave={(ids) => { + if (applyGroup) { + applyMutation.mutate({ id: applyGroup.id, ids }); + } + }} + /> + + ); +} diff --git a/openflare_server/web/features/waf/types.ts b/openflare_server/web/features/waf/types.ts new file mode 100644 index 00000000..833f8672 --- /dev/null +++ b/openflare_server/web/features/waf/types.ts @@ -0,0 +1,41 @@ +export interface WAFRuleGroup { + id: number; + name: string; + enabled: boolean; + is_global: boolean; + block_status_code: number; + block_response_body: string; + ip_whitelist: string[]; + ip_blacklist: string[]; + country_whitelist: string[]; + country_blacklist: string[]; + region_whitelist: string[]; + region_blacklist: string[]; + remark: string; + applied_site_ids: number[]; + applied_site_count: number; + created_at: string; + updated_at: string; +} + +export interface WAFRuleGroupPayload { + name: string; + enabled: boolean; + block_status_code: number; + block_response_body: string; + ip_whitelist: string[]; + ip_blacklist: string[]; + country_whitelist: string[]; + country_blacklist: string[]; + region_whitelist: string[]; + region_blacklist: string[]; + remark: string; +} + +export interface WAFSiteRuleGroups { + route_id: number; + global_rule_group: WAFRuleGroup | null; + rule_groups: WAFRuleGroup[]; + applied_rule_groups: WAFRuleGroup[]; + applied_ids: number[]; +} diff --git a/openflare_server/web/features/websites/components/website-detail-page.tsx b/openflare_server/web/features/websites/components/website-detail-page.tsx index 8560b1b1..21941ba9 100644 --- a/openflare_server/web/features/websites/components/website-detail-page.tsx +++ b/openflare_server/web/features/websites/components/website-detail-page.tsx @@ -3,7 +3,7 @@ import Link from 'next/link'; import { useRouter } from 'next/navigation'; import { useMutation, useQuery, useQueryClient } from '@tanstack/react-query'; -import { useMemo, useState } from 'react'; +import { useEffect, useMemo, useState } from 'react'; import { EmptyState } from '@/components/feedback/empty-state'; import { ErrorState } from '@/components/feedback/error-state'; @@ -21,6 +21,10 @@ import { deleteTlsCertificate, getTlsCertificates, } from '@/features/tls-certificates/api/tls-certificates'; +import { + getWAFSiteRuleGroups, + replaceWAFSiteRuleGroups, +} from '@/features/waf/api/waf'; import { CertificateDetailModal } from '@/features/websites/components/certificate-detail-modal'; import { CertificateEditorModal } from '@/features/websites/components/certificate-editor-modal'; import { CertificateImportModal } from '@/features/websites/components/certificate-import-modal'; @@ -38,6 +42,7 @@ import { DangerButton, PrimaryButton, SecondaryButton, + ToggleField, } from '@/features/shared/components/resource-primitives'; import { formatDateTime } from '@/lib/utils/date'; @@ -54,6 +59,7 @@ export function WebsiteDetailPage({ websiteId }: { websiteId: string }) { const [isCertificateImportOpen, setIsCertificateImportOpen] = useState(false); const [isCertificateDetailOpen, setIsCertificateDetailOpen] = useState(false); const [isCertificateEditorOpen, setIsCertificateEditorOpen] = useState(false); + const [wafSelectedIDs, setWafSelectedIDs] = useState([]); const [convertCertificate, setConvertCertificate] = useState(null); const [preferredCertificateId, setPreferredCertificateId] = useState< @@ -130,6 +136,32 @@ export function WebsiteDetailPage({ websiteId }: { websiteId: string }) { const enabledRoutesCount = relatedRoutes.filter( (route) => route.enabled, ).length; + const wafRouteID = relatedRoutes[0]?.id ?? null; + const wafQuery = useQuery({ + queryKey: ['waf', 'site-rule-groups', wafRouteID], + queryFn: () => getWAFSiteRuleGroups(wafRouteID ?? 0), + enabled: Boolean(wafRouteID), + }); + const wafMutation = useMutation({ + mutationFn: (ids: number[]) => replaceWAFSiteRuleGroups(wafRouteID ?? 0, ids), + onSuccess: async (view) => { + setWafSelectedIDs(view.applied_ids); + setFeedback({ tone: 'success', message: '网站 WAF 规则组已更新。' }); + await Promise.all([ + queryClient.invalidateQueries({ queryKey: ['waf'] }), + queryClient.invalidateQueries({ queryKey: ['config-versions', 'diff'] }), + ]); + }, + onError: (error) => { + setFeedback({ tone: 'danger', message: getErrorMessage(error) }); + }, + }); + + useEffect(() => { + if (wafQuery.data) { + setWafSelectedIDs(wafQuery.data.applied_ids); + } + }, [wafQuery.data]); const handleDeleteWebsite = () => { if (!website) { @@ -371,6 +403,89 @@ export function WebsiteDetailPage({ websiteId }: { websiteId: string }) { + wafMutation.mutate(wafSelectedIDs)} + > + {wafMutation.isPending ? '保存中...' : '保存 WAF'} + + ) : null + } + > + {!wafRouteID ? ( + + ) : wafQuery.isLoading ? ( + + ) : wafQuery.isError ? ( + + ) : ( +
+ {wafQuery.data?.global_rule_group ? ( +
+
+
+

+ {wafQuery.data.global_rule_group.name} +

+

+ 全局规则组默认应用到所有网站,不能在单站关闭。 +

+
+ +
+
+ ) : null} +
+ {(wafQuery.data?.rule_groups ?? []).map((group) => ( + { + setWafSelectedIDs((current) => + checked + ? [...current, group.id].sort( + (left, right) => left - right, + ) + : current.filter((id) => id !== group.id), + ); + }} + /> + ))} +
+ {(wafQuery.data?.rule_groups ?? []).length === 0 ? ( +

+ 暂无自定义规则组,可在 WAF 页面创建后再绑定。 +

+ ) : null} +
+ )} +
+ {relatedRoutes.length === 0 ? (