diff --git a/api/waf_tamperrule_api.go b/api/waf_tamperrule_api.go index dadb612..19efd7c 100644 --- a/api/waf_tamperrule_api.go +++ b/api/waf_tamperrule_api.go @@ -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) diff --git a/model/request/waf_tamperrule_req.go b/model/request/waf_tamperrule_req.go index ca5685d..f54b2b8 100644 --- a/model/request/waf_tamperrule_req.go +++ b/model/request/waf_tamperrule_req.go @@ -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"` +} diff --git a/router/waf_tamperrule_router.go b/router/waf_tamperrule_router.go index fc987ea..93d1d9e 100644 --- a/router/waf_tamperrule_router.go +++ b/router/waf_tamperrule_router.go @@ -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) } diff --git a/service/waf_service/waf_tamperrule_extract.go b/service/waf_service/waf_tamperrule_extract.go new file mode 100644 index 0000000..343f408 --- /dev/null +++ b/service/waf_service/waf_tamperrule_extract.go @@ -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) + }() +} diff --git a/service/waf_service/waf_tamperrule_extract_test.go b/service/waf_service/waf_tamperrule_extract_test.go new file mode 100644 index 0000000..2e9371d --- /dev/null +++ b/service/waf_service/waf_tamperrule_extract_test.go @@ -0,0 +1,120 @@ +package waf_service + +import ( + "net/url" + "testing" +) + +func TestExtractTamperCandidates(t *testing.T) { + htmlContent := []byte(` +
+ + + + + + +
+
+about
+contact
+ext
+anchor
+js
+mail
+
+`)
+ 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))
+ }
+}
diff --git a/service/waf_service/waf_tamperrule_service.go b/service/waf_service/waf_tamperrule_service.go
index cf382b4..4df2233 100644
--- a/service/waf_service/waf_tamperrule_service.go
+++ b/service/waf_service/waf_tamperrule_service.go
@@ -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
}