mirror of
https://gitee.com/samwaf/SamWaf.git
synced 2026-08-31 01:41:39 +08:00
@@ -9,6 +9,7 @@ import (
|
||||
"SamWaf/model/spec"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -127,7 +128,75 @@ func (w *WafTamperRuleApi) RelearnApi(c *gin.Context) {
|
||||
response.FailWithMessage("操作失败:"+err.Error(), c)
|
||||
} else {
|
||||
w.NotifyWaf(bean.HostCode)
|
||||
response.OkWithMessage("已标记重新学习,下次访问该URL将重新捕获基线", c)
|
||||
response.OkWithMessage("已触发重新学习,正在后端重新抓取基线(无后端的站点将在下次访问时重学)", c)
|
||||
}
|
||||
} else {
|
||||
response.FailWithMessage("解析失败", c)
|
||||
}
|
||||
}
|
||||
|
||||
// RelearnBatchApi 批量/整站重新学习基线
|
||||
func (w *WafTamperRuleApi) RelearnBatchApi(c *gin.Context) {
|
||||
var req request.WafTamperRuleRelearnBatchReq
|
||||
err := c.ShouldBindJSON(&req)
|
||||
if err == nil {
|
||||
err = wafTamperRuleService.RelearnBatchApi(req)
|
||||
if err != nil {
|
||||
response.FailWithMessage("操作失败:"+err.Error(), c)
|
||||
} else {
|
||||
w.NotifyWaf(req.HostCode)
|
||||
response.OkWithMessage("已触发重新学习,正在后端重新抓取基线(无后端的站点将在下次访问时重学)", c)
|
||||
}
|
||||
} else {
|
||||
response.FailWithMessage("解析失败", c)
|
||||
}
|
||||
}
|
||||
|
||||
// ExtractUrlsApi 抓取当前站点后端的一个页面,返回同站静态资源候选供批量添加
|
||||
func (w *WafTamperRuleApi) ExtractUrlsApi(c *gin.Context) {
|
||||
var req request.WafTamperRuleExtractReq
|
||||
err := c.ShouldBindJSON(&req)
|
||||
if err == nil {
|
||||
list, e := wafTamperRuleService.ExtractUrlsApi(req)
|
||||
if e != nil {
|
||||
response.FailWithMessage(e.Error(), c)
|
||||
} else {
|
||||
response.OkWithDetailed(gin.H{"list": list, "total": len(list)}, "提取成功", c)
|
||||
}
|
||||
} else {
|
||||
response.FailWithMessage("解析失败", c)
|
||||
}
|
||||
}
|
||||
|
||||
// AddBatchApi 批量新增受保护 URL 规则
|
||||
func (w *WafTamperRuleApi) AddBatchApi(c *gin.Context) {
|
||||
var req request.WafTamperRuleAddBatchReq
|
||||
err := c.ShouldBindJSON(&req)
|
||||
if err == nil {
|
||||
added, skipped, e := wafTamperRuleService.AddBatchApi(req)
|
||||
if e != nil {
|
||||
response.FailWithMessage("批量添加失败:"+e.Error(), c)
|
||||
} else {
|
||||
w.NotifyWaf(req.HostCode)
|
||||
response.OkWithDetailed(gin.H{"added": added, "skipped": skipped},
|
||||
fmt.Sprintf("批量添加完成:新增 %d 条,跳过 %d 条", added, skipped), c)
|
||||
}
|
||||
} else {
|
||||
response.FailWithMessage("解析失败", c)
|
||||
}
|
||||
}
|
||||
|
||||
// DelBatchApi 批量删除受保护 URL 规则
|
||||
func (w *WafTamperRuleApi) DelBatchApi(c *gin.Context) {
|
||||
var req request.WafTamperRuleDelBatchReq
|
||||
err := c.ShouldBindJSON(&req)
|
||||
if err == nil {
|
||||
cnt, e := wafTamperRuleService.DelBatchApi(req)
|
||||
if e != nil {
|
||||
response.FailWithMessage("批量删除失败:"+e.Error(), c)
|
||||
} else {
|
||||
w.NotifyWaf(req.HostCode)
|
||||
response.OkWithMessage(fmt.Sprintf("已删除 %d 条", cnt), c)
|
||||
}
|
||||
} else {
|
||||
response.FailWithMessage("解析失败", c)
|
||||
|
||||
@@ -35,5 +35,41 @@ type WafTamperRuleBaselineReq struct {
|
||||
}
|
||||
type WafTamperRuleSearchReq struct {
|
||||
HostCode string `json:"host_code" form:"host_code"`
|
||||
// 过滤(空/ nil 表示不过滤)
|
||||
Url string `json:"url" form:"url"`
|
||||
RuleName string `json:"rule_name" form:"rule_name"`
|
||||
IsEnable *int `json:"is_enable" form:"is_enable"`
|
||||
IgnoreQuery *int `json:"ignore_query" form:"ignore_query"`
|
||||
BaselineStatus *int `json:"baseline_status" form:"baseline_status"`
|
||||
// 排序(OrderKey 走白名单,OrderDir=asc/desc)
|
||||
OrderKey string `json:"order_key" form:"order_key"`
|
||||
OrderDir string `json:"order_dir" form:"order_dir"`
|
||||
request.PageInfo
|
||||
}
|
||||
|
||||
// WafTamperRuleRelearnBatchReq 批量/整站重新学习:Ids 为空表示整站全部
|
||||
type WafTamperRuleRelearnBatchReq struct {
|
||||
HostCode string `json:"host_code" form:"host_code"`
|
||||
Ids []string `json:"ids" form:"ids"`
|
||||
}
|
||||
|
||||
// WafTamperRuleExtractReq 从页面提取受保护 URL 候选(只抓当前站点后端)
|
||||
type WafTamperRuleExtractReq struct {
|
||||
HostCode string `json:"host_code" form:"host_code"`
|
||||
Domain string `json:"domain" form:"domain"` // 选定的站点域名(host 或 BindMoreHost 之一),作为抓取 Host 头与同站过滤基准
|
||||
PageUrl string `json:"page_url" form:"page_url"` // 页面地址或路径(只取 path,走本站后端抓取)
|
||||
}
|
||||
|
||||
// WafTamperRuleDelBatchReq 批量删除受保护 URL 规则(限定在 HostCode 内)
|
||||
type WafTamperRuleDelBatchReq struct {
|
||||
HostCode string `json:"host_code" form:"host_code"`
|
||||
Ids []string `json:"ids" form:"ids"`
|
||||
}
|
||||
|
||||
// WafTamperRuleAddBatchReq 批量新增受保护 URL 规则
|
||||
type WafTamperRuleAddBatchReq struct {
|
||||
HostCode string `json:"host_code" form:"host_code"`
|
||||
Urls []string `json:"urls" form:"urls"`
|
||||
IsEnable int `json:"is_enable" form:"is_enable"`
|
||||
IgnoreQuery int `json:"ignore_query" form:"ignore_query"`
|
||||
}
|
||||
|
||||
@@ -17,5 +17,9 @@ func (receiver *WafTamperRuleRouter) InitWafTamperRuleRouter(group *gin.RouterGr
|
||||
router.POST("/api/v1/wafhost/tamperrule/edit", api.ModifyApi)
|
||||
router.GET("/api/v1/wafhost/tamperrule/del", api.DelApi)
|
||||
router.GET("/api/v1/wafhost/tamperrule/relearn", api.RelearnApi)
|
||||
router.POST("/api/v1/wafhost/tamperrule/relearnbatch", api.RelearnBatchApi)
|
||||
router.POST("/api/v1/wafhost/tamperrule/extract", api.ExtractUrlsApi)
|
||||
router.POST("/api/v1/wafhost/tamperrule/addbatch", api.AddBatchApi)
|
||||
router.POST("/api/v1/wafhost/tamperrule/delbatch", api.DelBatchApi)
|
||||
router.GET("/api/v1/wafhost/tamperrule/baseline", api.GetBaselineApi)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,408 @@
|
||||
package waf_service
|
||||
|
||||
import (
|
||||
"SamWaf/enums"
|
||||
"SamWaf/global"
|
||||
"SamWaf/model"
|
||||
"SamWaf/model/request"
|
||||
"SamWaf/model/spec"
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"crypto/tls"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"golang.org/x/net/html"
|
||||
)
|
||||
|
||||
const (
|
||||
tamperExtractTimeout = 10 * time.Second // 抓取超时
|
||||
tamperExtractMaxBody = 4 * 1024 * 1024 // 抓取正文上限 4MB
|
||||
tamperExtractMaxURLs = 300 // 返回候选上限
|
||||
)
|
||||
|
||||
// DiscoveredURL 从页面提取到的受保护 URL 候选
|
||||
type DiscoveredURL struct {
|
||||
Url string `json:"url"`
|
||||
Type string `json:"type"` // js/css/html/img/other
|
||||
}
|
||||
|
||||
// ExtractUrlsApi 抓取「当前站点后端」的一个页面,解析出同站静态资源引用供批量添加。
|
||||
// 安全边界:只抓该站点已配置的后端(Remote_host:Remote_port,SamWaf 本就代理的目标),
|
||||
// 不接受任意外部地址;用户填的 page_url 只取其 path,强制走本站后端,无新增对外请求出口、无 SSRF。
|
||||
func (receiver *WafTamperRuleService) ExtractUrlsApi(req request.WafTamperRuleExtractReq) ([]DiscoveredURL, error) {
|
||||
if strings.TrimSpace(req.HostCode) == "" {
|
||||
return nil, errors.New("缺少站点标识")
|
||||
}
|
||||
var host model.Hosts
|
||||
global.GWAF_LOCAL_DB.Where("code=?", req.HostCode).First(&host)
|
||||
if host.Id == "" {
|
||||
return nil, errors.New("站点不存在")
|
||||
}
|
||||
if strings.TrimSpace(host.Remote_host) == "" {
|
||||
return nil, errors.New("该站点未配置后端地址,无法提取")
|
||||
}
|
||||
|
||||
// 只取 path,忽略用户填的 host —— 强制走本站后端,杜绝任意地址抓取
|
||||
reqPath := extractPathOnly(req.PageUrl)
|
||||
// 选定抓取域名:只能是本站 host 或 BindMoreHost 之一,作为 Host 头与同站过滤基准
|
||||
siteHosts := buildSiteHosts(host)
|
||||
domain := pickExtractDomain(req.Domain, siteHosts, host.Host)
|
||||
|
||||
// 后端抓取地址:镜像 wafworker 的 Remote_host+":"+Remote_port
|
||||
target := host.Remote_host + ":" + strconv.Itoa(host.Remote_port)
|
||||
if !strings.Contains(target, "://") {
|
||||
target = "http://" + target
|
||||
}
|
||||
backendURL, err := url.Parse(target)
|
||||
if err != nil || backendURL.Host == "" {
|
||||
return nil, errors.New("后端地址无效")
|
||||
}
|
||||
backendURL.Path = reqPath
|
||||
backendURL.RawQuery = ""
|
||||
|
||||
httpReq, err := http.NewRequest(http.MethodGet, backendURL.String(), nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
httpReq.Host = domain // 让后端按选定站点域名路由
|
||||
httpReq.Header.Set("User-Agent", "SamWaf-TamperExtractor")
|
||||
|
||||
resp, err := buildBackendClient(host).Do(httpReq)
|
||||
if err != nil {
|
||||
return nil, errors.New("抓取后端页面失败:" + err.Error())
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, tamperExtractMaxBody))
|
||||
if err != nil {
|
||||
return nil, errors.New("读取页面内容失败:" + err.Error())
|
||||
}
|
||||
|
||||
// 归一化基准用「站点域名 + 请求路径」,这样 HTML 里写成站点绝对地址的引用也能被识别为同站
|
||||
scheme := "http"
|
||||
if host.Ssl == 1 {
|
||||
scheme = "https"
|
||||
}
|
||||
siteBase := &url.URL{Scheme: scheme, Host: domain, Path: reqPath}
|
||||
return extractTamperCandidates(body, siteBase, siteHosts), nil
|
||||
}
|
||||
|
||||
// pickExtractDomain 选定抓取用域名:请求指定且属于本站(host/BindMoreHost)才用,否则回退主域名,杜绝任意 Host
|
||||
func pickExtractDomain(reqDomain string, siteHosts map[string]bool, defaultHost string) string {
|
||||
d := strings.ToLower(strings.TrimSpace(reqDomain))
|
||||
if hh, _, err := net.SplitHostPort(d); err == nil {
|
||||
d = hh
|
||||
}
|
||||
if d != "" && siteHosts[d] {
|
||||
return d
|
||||
}
|
||||
return defaultHost
|
||||
}
|
||||
|
||||
// extractPathOnly 从用户输入里只提取 path(默认 /),忽略 host/query/fragment
|
||||
func extractPathOnly(raw string) string {
|
||||
s := strings.TrimSpace(raw)
|
||||
if s == "" {
|
||||
return "/"
|
||||
}
|
||||
if u, err := url.Parse(s); err == nil && u.Path != "" {
|
||||
return ensureLeadingSlash(u.Path)
|
||||
}
|
||||
return ensureLeadingSlash(s)
|
||||
}
|
||||
|
||||
func ensureLeadingSlash(p string) string {
|
||||
// 去掉可能带的 query/fragment
|
||||
if i := strings.IndexAny(p, "?#"); i >= 0 {
|
||||
p = p[:i]
|
||||
}
|
||||
if p == "" {
|
||||
return "/"
|
||||
}
|
||||
if !strings.HasPrefix(p, "/") {
|
||||
p = "/" + p
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
// buildSiteHosts 站点自身可识别的域名集合(主域名 + 绑定多域名),用于同站过滤
|
||||
func buildSiteHosts(host model.Hosts) map[string]bool {
|
||||
set := map[string]bool{}
|
||||
add := func(h string) {
|
||||
h = strings.ToLower(strings.TrimSpace(h))
|
||||
if h == "" {
|
||||
return
|
||||
}
|
||||
if hh, _, err := net.SplitHostPort(h); err == nil {
|
||||
h = hh
|
||||
}
|
||||
set[h] = true
|
||||
}
|
||||
add(host.Host)
|
||||
for _, line := range strings.Split(host.BindMoreHost, "\n") {
|
||||
add(line)
|
||||
}
|
||||
return set
|
||||
}
|
||||
|
||||
// buildBackendClient 构造只连本站后端的 http.Client:Remote_ip 覆盖拨号、后端 TLS 校验开关、超时、禁跟随重定向
|
||||
func buildBackendClient(host model.Hosts) *http.Client {
|
||||
dialer := &net.Dialer{Timeout: 5 * time.Second}
|
||||
tr := &http.Transport{
|
||||
TLSClientConfig: &tls.Config{InsecureSkipVerify: host.InsecureSkipVerify == 1},
|
||||
DisableKeepAlives: true,
|
||||
TLSHandshakeTimeout: 5 * time.Second,
|
||||
}
|
||||
if ip := strings.TrimSpace(host.Remote_ip); ip != "" {
|
||||
port := strconv.Itoa(host.Remote_port)
|
||||
tr.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
return dialer.DialContext(ctx, network, net.JoinHostPort(ip, port))
|
||||
}
|
||||
} else {
|
||||
tr.DialContext = dialer.DialContext
|
||||
}
|
||||
return &http.Client{
|
||||
Timeout: tamperExtractTimeout,
|
||||
Transport: tr,
|
||||
// 不自动跟随重定向,避免被后端 3xx 带去其它地址
|
||||
CheckRedirect: func(r *http.Request, via []*http.Request) error {
|
||||
return http.ErrUseLastResponse
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// extractTamperCandidates 纯函数:解析 HTML,提取同站静态资源引用(去重、归一化为 path、限量)
|
||||
func extractTamperCandidates(htmlContent []byte, siteBase *url.URL, siteHosts map[string]bool) []DiscoveredURL {
|
||||
z := html.NewTokenizer(bytes.NewReader(htmlContent))
|
||||
seen := map[string]bool{}
|
||||
out := make([]DiscoveredURL, 0, 16)
|
||||
for {
|
||||
tt := z.Next()
|
||||
if tt == html.ErrorToken {
|
||||
break
|
||||
}
|
||||
if tt != html.StartTagToken && tt != html.SelfClosingTagToken {
|
||||
continue
|
||||
}
|
||||
name, hasAttr := z.TagName()
|
||||
tag := string(name)
|
||||
var attrKey string
|
||||
switch tag {
|
||||
case "script", "img":
|
||||
attrKey = "src"
|
||||
case "link", "a":
|
||||
attrKey = "href"
|
||||
default:
|
||||
continue
|
||||
}
|
||||
ref := readTagAttr(z, hasAttr, attrKey)
|
||||
p := normalizeTamperRef(ref, siteBase, siteHosts)
|
||||
if p == "" || seen[p] {
|
||||
continue
|
||||
}
|
||||
seen[p] = true
|
||||
out = append(out, DiscoveredURL{Url: p, Type: classifyTamperURL(p, tag)})
|
||||
if len(out) >= tamperExtractMaxURLs {
|
||||
break
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// readTagAttr 从 tokenizer 当前标签里取指定属性值
|
||||
func readTagAttr(z *html.Tokenizer, hasAttr bool, key string) string {
|
||||
if !hasAttr {
|
||||
return ""
|
||||
}
|
||||
for {
|
||||
k, v, more := z.TagAttr()
|
||||
if string(k) == key {
|
||||
return string(v)
|
||||
}
|
||||
if !more {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// normalizeTamperRef 把引用归一化为「本站精确 path」,非同站/非法/带通配的返回空
|
||||
func normalizeTamperRef(ref string, siteBase *url.URL, siteHosts map[string]bool) string {
|
||||
ref = strings.TrimSpace(ref)
|
||||
if ref == "" || strings.HasPrefix(ref, "#") {
|
||||
return ""
|
||||
}
|
||||
low := strings.ToLower(ref)
|
||||
for _, skip := range []string{"data:", "javascript:", "mailto:", "tel:", "blob:", "about:"} {
|
||||
if strings.HasPrefix(low, skip) {
|
||||
return ""
|
||||
}
|
||||
}
|
||||
u, err := url.Parse(ref)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
resolved := siteBase.ResolveReference(u)
|
||||
if h := strings.ToLower(resolved.Hostname()); h != "" && !siteHosts[h] {
|
||||
return "" // 第三方资源,无法保护
|
||||
}
|
||||
p := resolved.Path
|
||||
if p == "" || !strings.HasPrefix(p, "/") || strings.Contains(p, "*") {
|
||||
return ""
|
||||
}
|
||||
return p
|
||||
}
|
||||
|
||||
// classifyTamperURL 按扩展名/标签给候选归类
|
||||
func classifyTamperURL(path, tag string) string {
|
||||
lp := strings.ToLower(path)
|
||||
switch {
|
||||
case strings.HasSuffix(lp, ".js"), strings.HasSuffix(lp, ".mjs"):
|
||||
return "js"
|
||||
case strings.HasSuffix(lp, ".css"):
|
||||
return "css"
|
||||
case strings.HasSuffix(lp, ".html"), strings.HasSuffix(lp, ".htm"):
|
||||
return "html"
|
||||
}
|
||||
for _, ext := range []string{".png", ".jpg", ".jpeg", ".gif", ".svg", ".ico", ".webp", ".bmp"} {
|
||||
if strings.HasSuffix(lp, ext) {
|
||||
return "img"
|
||||
}
|
||||
}
|
||||
switch tag {
|
||||
case "script":
|
||||
return "js"
|
||||
case "link":
|
||||
return "css"
|
||||
case "img":
|
||||
return "img"
|
||||
case "a":
|
||||
return "html"
|
||||
}
|
||||
return "other"
|
||||
}
|
||||
|
||||
// sha256HexSvc 计算内容 sha256 十六进制(service 层本地实现,避免依赖 wafenginecore)
|
||||
func sha256HexSvc(b []byte) string {
|
||||
h := sha256.Sum256(b)
|
||||
return hex.EncodeToString(h[:])
|
||||
}
|
||||
|
||||
// pushTamperReload 重新读取该站点规则并推送引擎热重载(与 api.NotifyWaf 等价)
|
||||
func pushTamperReload(hostCode string) {
|
||||
var list []model.TamperRule
|
||||
global.GWAF_LOCAL_DB.Where("host_code = ?", hostCode).Find(&list)
|
||||
global.GWAF_CHAN_MSG <- spec.ChanCommonHost{
|
||||
HostCode: hostCode,
|
||||
Type: enums.ChanTypeTamperRule,
|
||||
Content: list,
|
||||
}
|
||||
}
|
||||
|
||||
// fetchAndCaptureBaseline 后端自请求该规则 URL、抓取正文并写入基线(即时重新学习)。
|
||||
// 只抓本站后端;不广告 br/zstd,Go 自动解 gzip,正文即“解压后内容”,与引擎哈希基准一致。
|
||||
func fetchAndCaptureBaseline(host model.Hosts, rule model.TamperRule, cfg model.TamperConfig) {
|
||||
now := time.Now().Format("2006-01-02 15:04:05")
|
||||
setFail := func(msg string) {
|
||||
global.GWAF_LOCAL_DB.Model(&model.TamperRule{}).Where("id=?", rule.Id).Updates(map[string]interface{}{
|
||||
"BaselineStatus": 2, "BaselineMsg": msg, "LastLearnTime": now,
|
||||
})
|
||||
}
|
||||
target := host.Remote_host + ":" + strconv.Itoa(host.Remote_port)
|
||||
if !strings.Contains(target, "://") {
|
||||
target = "http://" + target
|
||||
}
|
||||
bu, err := url.Parse(target)
|
||||
if err != nil || bu.Host == "" {
|
||||
setFail("后端地址无效,重新学习自抓失败")
|
||||
return
|
||||
}
|
||||
bu.Path = rule.Url
|
||||
bu.RawQuery = ""
|
||||
httpReq, err := http.NewRequest(http.MethodGet, bu.String(), nil)
|
||||
if err != nil {
|
||||
setFail("构造请求失败:" + err.Error())
|
||||
return
|
||||
}
|
||||
httpReq.Host = host.Host
|
||||
httpReq.Header.Set("User-Agent", "SamWaf-TamperRelearn")
|
||||
|
||||
resp, err := buildBackendClient(host).Do(httpReq)
|
||||
if err != nil {
|
||||
setFail("重新学习自抓失败:" + err.Error())
|
||||
return
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
setFail(fmt.Sprintf("重新学习自抓返回状态码 %d,未学习", resp.StatusCode))
|
||||
return
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, tamperExtractMaxBody))
|
||||
if err != nil || len(body) == 0 {
|
||||
setFail("重新学习读取正文失败或为空")
|
||||
return
|
||||
}
|
||||
maxKB := cfg.MaxSizeKB
|
||||
if maxKB <= 0 {
|
||||
maxKB = 1024
|
||||
}
|
||||
if len(body) > maxKB*1024 {
|
||||
setFail(fmt.Sprintf("正文 %d 字节超过上限 %d KB,未学习", len(body), maxKB))
|
||||
return
|
||||
}
|
||||
global.GWAF_LOCAL_DB.Model(&model.TamperRule{}).Where("id=?", rule.Id).Updates(map[string]interface{}{
|
||||
"BaselineHash": sha256HexSvc(body),
|
||||
"BaselineContent": body,
|
||||
"ContentType": resp.Header.Get("Content-Type"),
|
||||
"StatusCode": resp.StatusCode,
|
||||
"ContentSize": len(body),
|
||||
"BaselineStatus": 1,
|
||||
"BaselineMsg": "重新学习已抓取",
|
||||
"LastLearnTime": now,
|
||||
})
|
||||
}
|
||||
|
||||
// backgroundRecapture 后台对指定规则(ids 空=该站点全部启用规则)自请求重新抓取基线,完成后热重载。
|
||||
// 只抓本站后端、限并发、异步不阻塞接口;无后端(纯静态站点)则跳过,保持惰性下次访问再学。
|
||||
func (receiver *WafTamperRuleService) backgroundRecapture(hostCode string, ids []string) {
|
||||
go func() {
|
||||
var host model.Hosts
|
||||
global.GWAF_LOCAL_DB.Where("code=?", hostCode).First(&host)
|
||||
if host.Id == "" || strings.TrimSpace(host.Remote_host) == "" {
|
||||
return
|
||||
}
|
||||
var rules []model.TamperRule
|
||||
db := global.GWAF_LOCAL_DB.Where("host_code=?", hostCode)
|
||||
if len(ids) > 0 {
|
||||
db = db.Where("id in ?", ids)
|
||||
}
|
||||
db.Find(&rules)
|
||||
cfg := model.ParseTamperConfig(host.TamperJSON)
|
||||
sem := make(chan struct{}, 4)
|
||||
var wg sync.WaitGroup
|
||||
for i := range rules {
|
||||
if rules[i].IsEnable != 1 {
|
||||
continue
|
||||
}
|
||||
wg.Add(1)
|
||||
sem <- struct{}{}
|
||||
go func(rule model.TamperRule) {
|
||||
defer wg.Done()
|
||||
defer func() { <-sem }()
|
||||
fetchAndCaptureBaseline(host, rule, cfg)
|
||||
}(rules[i])
|
||||
}
|
||||
wg.Wait()
|
||||
pushTamperReload(hostCode)
|
||||
}()
|
||||
}
|
||||
@@ -0,0 +1,120 @@
|
||||
package waf_service
|
||||
|
||||
import (
|
||||
"net/url"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestExtractTamperCandidates(t *testing.T) {
|
||||
htmlContent := []byte(`
|
||||
<html><head>
|
||||
<link rel="stylesheet" href="/css/app.css">
|
||||
<link rel="icon" href="favicon.ico">
|
||||
<script src="/js/app.js?v=123"></script>
|
||||
<script src="https://example.com/js/vendor.js"></script>
|
||||
<script src="https://cdn.other.com/lib.js"></script>
|
||||
</head><body>
|
||||
<img src="/img/logo.png">
|
||||
<img src="images/banner.jpg">
|
||||
<a href="/about.html">about</a>
|
||||
<a href="https://example.com/contact.html">contact</a>
|
||||
<a href="https://third.com/x.html">ext</a>
|
||||
<a href="#top">anchor</a>
|
||||
<a href="javascript:void(0)">js</a>
|
||||
<a href="mailto:a@b.com">mail</a>
|
||||
<script src="/js/app.js"></script>
|
||||
</body></html>`)
|
||||
siteBase, _ := url.Parse("https://example.com/")
|
||||
siteHosts := map[string]bool{"example.com": true}
|
||||
got := extractTamperCandidates(htmlContent, siteBase, siteHosts)
|
||||
|
||||
types := map[string]string{}
|
||||
for _, d := range got {
|
||||
types[d.Url] = d.Type
|
||||
}
|
||||
|
||||
// 同站资源应提取,且类型正确
|
||||
wantKeep := map[string]string{
|
||||
"/css/app.css": "css",
|
||||
"/favicon.ico": "img",
|
||||
"/js/app.js": "js", // 带 ?v=123 去参后与末尾无参项去重为一条
|
||||
"/js/vendor.js": "js", // 站点绝对地址
|
||||
"/img/logo.png": "img",
|
||||
"/images/banner.jpg": "img",
|
||||
"/about.html": "html",
|
||||
"/contact.html": "html",
|
||||
}
|
||||
for u, ty := range wantKeep {
|
||||
if types[u] != ty {
|
||||
t.Errorf("应提取 %s(type=%s),实际 type=%q", u, ty, types[u])
|
||||
}
|
||||
}
|
||||
|
||||
// 第三方资源应丢弃(不同 host)
|
||||
if _, ok := types["/lib.js"]; ok {
|
||||
t.Errorf("第三方 cdn.other.com/lib.js 不应被提取")
|
||||
}
|
||||
if _, ok := types["/x.html"]; ok {
|
||||
t.Errorf("第三方 third.com/x.html 不应被提取")
|
||||
}
|
||||
|
||||
// 去重 + data/js/mailto/anchor 均无有效 path → 总数应为 8
|
||||
if len(got) != 8 {
|
||||
t.Errorf("应提取 8 条同站候选,实际 %d 条: %+v", len(got), got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractPathOnly(t *testing.T) {
|
||||
cases := []struct {
|
||||
in string
|
||||
want string
|
||||
}{
|
||||
{"", "/"},
|
||||
{"/", "/"},
|
||||
{"/index.html", "/index.html"},
|
||||
{"index.html", "/index.html"},
|
||||
{"https://example.com/a/b.html?x=1#f", "/a/b.html"},
|
||||
{"http://1.2.3.4:8080/app.js", "/app.js"},
|
||||
{"/p?q=1", "/p"},
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got := extractPathOnly(c.in); got != c.want {
|
||||
t.Errorf("extractPathOnly(%q)=%q, 期望 %q", c.in, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestPickExtractDomain(t *testing.T) {
|
||||
siteHosts := map[string]bool{"example.com": true, "www.example.com": true}
|
||||
cases := []struct {
|
||||
name string
|
||||
req string
|
||||
want string
|
||||
}{
|
||||
{"属于本站直接用", "www.example.com", "www.example.com"},
|
||||
{"大写归一化匹配", "WWW.Example.com", "www.example.com"},
|
||||
{"带端口去端口匹配", "example.com:8443", "example.com"},
|
||||
{"非本站回退主域名", "evil.com", "example.com"},
|
||||
{"空回退主域名", "", "example.com"},
|
||||
}
|
||||
for _, c := range cases {
|
||||
if got := pickExtractDomain(c.req, siteHosts, "example.com"); got != c.want {
|
||||
t.Errorf("%s: pickExtractDomain(%q)=%q, 期望 %q", c.name, c.req, got, c.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestSha256HexSvc(t *testing.T) {
|
||||
a := sha256HexSvc([]byte("hello"))
|
||||
b := sha256HexSvc([]byte("hello"))
|
||||
c := sha256HexSvc([]byte("world"))
|
||||
if a != b {
|
||||
t.Errorf("相同内容哈希应一致")
|
||||
}
|
||||
if a == c {
|
||||
t.Errorf("不同内容哈希应不同")
|
||||
}
|
||||
if len(a) != 64 {
|
||||
t.Errorf("sha256 十六进制应为64字符,实际 %d", len(a))
|
||||
}
|
||||
}
|
||||
@@ -98,7 +98,7 @@ func (receiver *WafTamperRuleService) ModifyApi(req request.WafTamperRuleEditReq
|
||||
return err
|
||||
}
|
||||
|
||||
// RelearnApi 触发重新学习:清空基线状态,下次访问该 URL 时重新捕获
|
||||
// RelearnApi 触发重新学习:清空基线状态并即时后端自抓重建(无后端则保持惰性,下次访问再学)
|
||||
func (receiver *WafTamperRuleService) RelearnApi(req request.WafTamperRuleRelearnReq) error {
|
||||
beanMap := map[string]interface{}{
|
||||
"BaselineStatus": 0,
|
||||
@@ -108,7 +108,96 @@ func (receiver *WafTamperRuleService) RelearnApi(req request.WafTamperRuleRelear
|
||||
"BaselineMsg": "已标记重新学习",
|
||||
"UPDATE_TIME": customtype.JsonTime(time.Now()),
|
||||
}
|
||||
return global.GWAF_LOCAL_DB.Model(model.TamperRule{}).Where("id = ?", req.Id).Updates(beanMap).Error
|
||||
if err := global.GWAF_LOCAL_DB.Model(model.TamperRule{}).Where("id = ?", req.Id).Updates(beanMap).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
// 即时触发后端自抓重建基线
|
||||
var rule model.TamperRule
|
||||
global.GWAF_LOCAL_DB.Omit("baseline_content").Where("id=?", req.Id).Find(&rule)
|
||||
if rule.HostCode != "" {
|
||||
receiver.backgroundRecapture(rule.HostCode, []string{req.Id})
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RelearnBatchApi 批量/整站重新学习:Ids 为空则该站点全部规则重新学习(限定在 HostCode 内,防误伤其它站点)
|
||||
func (receiver *WafTamperRuleService) RelearnBatchApi(req request.WafTamperRuleRelearnBatchReq) error {
|
||||
if strings.TrimSpace(req.HostCode) == "" {
|
||||
return errors.New("缺少站点标识")
|
||||
}
|
||||
beanMap := map[string]interface{}{
|
||||
"BaselineStatus": 0,
|
||||
"BaselineHash": "",
|
||||
"BaselineContent": []byte{},
|
||||
"ContentSize": 0,
|
||||
"BaselineMsg": "已标记重新学习",
|
||||
"UPDATE_TIME": customtype.JsonTime(time.Now()),
|
||||
}
|
||||
db := global.GWAF_LOCAL_DB.Model(model.TamperRule{}).Where("host_code = ?", req.HostCode)
|
||||
if len(req.Ids) > 0 {
|
||||
db = db.Where("id in ?", req.Ids)
|
||||
}
|
||||
if err := db.Updates(beanMap).Error; err != nil {
|
||||
return err
|
||||
}
|
||||
// 即时触发后端自抓重建基线(后台限并发)
|
||||
receiver.backgroundRecapture(req.HostCode, req.Ids)
|
||||
return nil
|
||||
}
|
||||
|
||||
// DelBatchApi 批量删除受保护 URL(限定在 HostCode 内),返回删除条数
|
||||
func (receiver *WafTamperRuleService) DelBatchApi(req request.WafTamperRuleDelBatchReq) (int64, error) {
|
||||
if strings.TrimSpace(req.HostCode) == "" {
|
||||
return 0, errors.New("缺少站点标识")
|
||||
}
|
||||
if len(req.Ids) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
res := global.GWAF_LOCAL_DB.Where("host_code=? and id in ?", req.HostCode, req.Ids).Delete(&model.TamperRule{})
|
||||
return res.RowsAffected, res.Error
|
||||
}
|
||||
|
||||
// AddBatchApi 批量新增受保护 URL:逐条校验 + 跳过已存在,返回新增/跳过数
|
||||
func (receiver *WafTamperRuleService) AddBatchApi(req request.WafTamperRuleAddBatchReq) (int, int, error) {
|
||||
if strings.TrimSpace(req.HostCode) == "" {
|
||||
return 0, 0, errors.New("缺少站点标识")
|
||||
}
|
||||
added, skipped := 0, 0
|
||||
for _, rawUrl := range req.Urls {
|
||||
u := strings.TrimSpace(rawUrl)
|
||||
if u == "" {
|
||||
continue
|
||||
}
|
||||
if err := validateTamperUrl(u); err != nil {
|
||||
skipped++
|
||||
continue
|
||||
}
|
||||
var total int64 = 0
|
||||
global.GWAF_LOCAL_DB.Model(&model.TamperRule{}).Where("host_code=? and url=?", req.HostCode, u).Count(&total)
|
||||
if total > 0 {
|
||||
skipped++
|
||||
continue
|
||||
}
|
||||
bean := &model.TamperRule{
|
||||
BaseOrm: baseorm.BaseOrm{
|
||||
Id: uuid.GenUUID(),
|
||||
USER_CODE: global.GWAF_USER_CODE,
|
||||
Tenant_ID: global.GWAF_TENANT_ID,
|
||||
CREATE_TIME: customtype.JsonTime(time.Now()),
|
||||
UPDATE_TIME: customtype.JsonTime(time.Now()),
|
||||
},
|
||||
HostCode: req.HostCode,
|
||||
Url: u,
|
||||
IsEnable: req.IsEnable,
|
||||
IgnoreQuery: req.IgnoreQuery,
|
||||
BaselineStatus: 0,
|
||||
}
|
||||
if err := global.GWAF_LOCAL_DB.Create(bean).Error; err != nil {
|
||||
return added, skipped, err
|
||||
}
|
||||
added++
|
||||
}
|
||||
return added, skipped, nil
|
||||
}
|
||||
|
||||
func (receiver *WafTamperRuleService) GetDetailApi(req request.WafTamperRuleDetailReq) model.TamperRule {
|
||||
@@ -131,20 +220,64 @@ func (receiver *WafTamperRuleService) GetBaselineApi(id string) model.TamperRule
|
||||
return bean
|
||||
}
|
||||
|
||||
// tamperOrderCols 允许排序的列白名单(防 SQL 注入:order by 只能取这里的值)
|
||||
var tamperOrderCols = map[string]string{
|
||||
"url": "url",
|
||||
"rule_name": "rule_name",
|
||||
"is_enable": "is_enable",
|
||||
"ignore_query": "ignore_query",
|
||||
"baseline_status": "baseline_status",
|
||||
"content_size": "content_size",
|
||||
"tamper_count": "tamper_count",
|
||||
"create_time": "create_time",
|
||||
}
|
||||
|
||||
func (receiver *WafTamperRuleService) GetListApi(req request.WafTamperRuleSearchReq) ([]model.TamperRule, int64, error) {
|
||||
var list []model.TamperRule
|
||||
var total int64 = 0
|
||||
|
||||
var whereField = ""
|
||||
var whereValues []interface{}
|
||||
// 动态过滤条件(空/ nil 不过滤)
|
||||
var conds []string
|
||||
var vals []interface{}
|
||||
if len(req.HostCode) > 0 {
|
||||
whereField = " host_code=? "
|
||||
whereValues = append(whereValues, req.HostCode)
|
||||
conds = append(conds, "host_code=?")
|
||||
vals = append(vals, req.HostCode)
|
||||
}
|
||||
if s := strings.TrimSpace(req.Url); s != "" {
|
||||
conds = append(conds, "url like ?")
|
||||
vals = append(vals, "%"+s+"%")
|
||||
}
|
||||
if s := strings.TrimSpace(req.RuleName); s != "" {
|
||||
conds = append(conds, "rule_name like ?")
|
||||
vals = append(vals, "%"+s+"%")
|
||||
}
|
||||
if req.IsEnable != nil {
|
||||
conds = append(conds, "is_enable=?")
|
||||
vals = append(vals, *req.IsEnable)
|
||||
}
|
||||
if req.IgnoreQuery != nil {
|
||||
conds = append(conds, "ignore_query=?")
|
||||
vals = append(vals, *req.IgnoreQuery)
|
||||
}
|
||||
if req.BaselineStatus != nil {
|
||||
conds = append(conds, "baseline_status=?")
|
||||
vals = append(vals, *req.BaselineStatus)
|
||||
}
|
||||
whereStr := strings.Join(conds, " and ")
|
||||
|
||||
// 排序:列走白名单,方向仅 asc/desc,默认按创建时间倒序
|
||||
order := "create_time desc"
|
||||
if col, ok := tamperOrderCols[req.OrderKey]; ok {
|
||||
dir := "asc"
|
||||
if strings.EqualFold(req.OrderDir, "desc") {
|
||||
dir = "desc"
|
||||
}
|
||||
order = col + " " + dir
|
||||
}
|
||||
|
||||
// Omit baseline_content:列表绝不携带大 blob
|
||||
global.GWAF_LOCAL_DB.Model(&model.TamperRule{}).Omit("baseline_content").Where(whereField, whereValues...).Limit(req.PageSize).Offset(req.PageSize * (req.PageIndex - 1)).Order("create_time desc").Find(&list)
|
||||
global.GWAF_LOCAL_DB.Model(&model.TamperRule{}).Where(whereField, whereValues...).Count(&total)
|
||||
global.GWAF_LOCAL_DB.Model(&model.TamperRule{}).Omit("baseline_content").Where(whereStr, vals...).Order(order).Limit(req.PageSize).Offset(req.PageSize * (req.PageIndex - 1)).Find(&list)
|
||||
global.GWAF_LOCAL_DB.Model(&model.TamperRule{}).Where(whereStr, vals...).Count(&total)
|
||||
|
||||
return list, total, nil
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user