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 }