From a11549979334a0b118479aec5e0eaf4da65adfdb Mon Sep 17 00:00:00 2001 From: ZacharyZcR Date: Sat, 13 Jun 2026 18:46:14 +0800 Subject: [PATCH] =?UTF-8?q?fix+perf:=20=E4=BF=AE=E5=A4=8D10=E4=B8=AAbug=20?= =?UTF-8?q?&=2010=E9=A1=B9=E6=80=A7=E8=83=BD=E4=BC=98=E5=8C=96?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Bug修复: - clustersend CEL结果判断从字符串比较改为类型断言 - Nuclei DSL matcher安全降级为false避免误报 - clusterpoc发现漏洞后返回true修正语义 - reverseCheck加10s超时防止ceye API阻塞 - doSearch/bmatches正则编译结果缓存到sync.Map - evalset CEL求值失败时存空字符串而非原始表达式 - CEL wait()函数加nil Reverse指针检查防panic - MongoDB readMongoMsg应用timeout参数设置读超时 - TXTWriter.Close确保Sync失败后仍调用file.Close 性能优化: - 指纹regex缓存从RWMutex+map改为sync.Map消除锁竞争 - CaseInsensitive指纹词加载时预小写化避免匹配时分配 - 版本提取FindAllStringSubmatch限制返回数量 - i18n.Tr用strconv.Itoa替代Sprintf减少分配 - POC加载用atomic.Bool+DCLP消除热路径锁 - 结果缓冲map预分配容量减少rehash - HTTP连接池参数随并发数动态调整 - getRuleHash去除反射+Headers排序保证确定性dedup - POC并发加载用channel替代Mutex收集结果 --- common/i18n/i18n.go | 5 +- common/output/buffer.go | 8 +-- common/output/writers.go | 12 +++-- plugins/services/mongodb.go | 3 ++ webscan/fingerprint/enhanced.go | 81 ++++++++++++++-------------- webscan/fingerprint/enhanced_test.go | 2 +- webscan/lib/Client.go | 22 ++++++-- webscan/lib/Eval.go | 8 ++- webscan/lib/eval_misc.go | 2 +- webscan/lib/eval_string.go | 16 ++++-- webscan/lib/poc_adapter.go | 4 +- webscan/lib/poc_executor.go | 57 ++++++++++++++------ webscan/web_scan.go | 60 ++++++++++----------- 13 files changed, 169 insertions(+), 111 deletions(-) diff --git a/common/i18n/i18n.go b/common/i18n/i18n.go index e1a0437..580dec1 100644 --- a/common/i18n/i18n.go +++ b/common/i18n/i18n.go @@ -2,6 +2,7 @@ package i18n import ( "fmt" + "strconv" "sync" "github.com/nicksnyder/go-i18n/v2/i18n" @@ -80,9 +81,9 @@ func Tr(key string, args ...interface{}) string { loc := localizer mu.RUnlock() - data := make(map[string]interface{}) + data := make(map[string]interface{}, len(args)) for i, arg := range args { - data[fmt.Sprintf("Arg%d", i+1)] = arg + data["Arg"+strconv.Itoa(i+1)] = arg } msg, err := loc.Localize(&i18n.LocalizeConfig{ diff --git a/common/output/buffer.go b/common/output/buffer.go index 23fa0a0..450e4c3 100644 --- a/common/output/buffer.go +++ b/common/output/buffer.go @@ -22,10 +22,10 @@ type ResultBuffer struct { // NewResultBuffer 创建新的结果缓冲 func NewResultBuffer() *ResultBuffer { return &ResultBuffer{ - seenHosts: make(map[string]struct{}), - seenPorts: make(map[string]struct{}), - seenServices: make(map[string]int), - seenVulns: make(map[string]struct{}), + seenHosts: make(map[string]struct{}, 256), + seenPorts: make(map[string]struct{}, 512), + seenServices: make(map[string]int, 128), + seenVulns: make(map[string]struct{}, 64), } } diff --git a/common/output/writers.go b/common/output/writers.go index 114e607..825226b 100644 --- a/common/output/writers.go +++ b/common/output/writers.go @@ -350,13 +350,17 @@ func (w *TXTWriter) Close() error { os.Remove(w.realtimePath) } + var firstErr error if err := w.bufWriter.Flush(); err != nil { - return err + firstErr = err } - if err := w.file.Sync(); err != nil { - return err + if err := w.file.Sync(); err != nil && firstErr == nil { + firstErr = err } - return w.file.Close() + if err := w.file.Close(); err != nil && firstErr == nil { + firstErr = err + } + return firstErr } // writeSection 写入一个分类的所有结果 diff --git a/plugins/services/mongodb.go b/plugins/services/mongodb.go index 6cad724..afc5456 100644 --- a/plugins/services/mongodb.go +++ b/plugins/services/mongodb.go @@ -328,6 +328,9 @@ func sendMongoMsg(ctx context.Context, conn io.ReadWriter, body []byte, timeout // readMongoMsg 读取 MongoDB 响应 func readMongoMsg(conn io.Reader, timeout time.Duration) ([]byte, error) { + if tc, ok := conn.(interface{ SetReadDeadline(time.Time) error }); ok { + _ = tc.SetReadDeadline(time.Now().Add(timeout)) + } // 读取 16 字节消息头 header := make([]byte, 16) if _, err := io.ReadFull(conn, header); err != nil { diff --git a/webscan/fingerprint/enhanced.go b/webscan/fingerprint/enhanced.go index c5591b2..ec376fc 100644 --- a/webscan/fingerprint/enhanced.go +++ b/webscan/fingerprint/enhanced.go @@ -46,9 +46,7 @@ type EnhancedFingerprint struct { // EnhancedFingerprintDB 增强指纹数据库 type EnhancedFingerprintDB struct { Fingerprints []*EnhancedFingerprint - // 预编译的正则表达式缓存 - regexCache map[string]*regexp.Regexp - regexCacheMu sync.RWMutex // 保护regexCache的并发访问 + regexCache sync.Map // pattern → *regexp.Regexp,无锁并发安全 } var ( @@ -63,9 +61,22 @@ func LoadEnhancedFingerprints() error { return fmt.Errorf("%s: %w", i18n.GetText("fingerprint_enhanced_parse_failed"), err) } + // 预处理:CaseInsensitive 的 matcher 预先小写化 Words,避免匹配时重复分配 + for _, fp := range fps { + for hi := range fp.HTTP { + for mi := range fp.HTTP[hi].Matchers { + m := &fp.HTTP[hi].Matchers[mi] + if m.CaseInsensitive && m.Type == "word" { + for wi, w := range m.Words { + m.Words[wi] = strings.ToLower(w) + } + } + } + } + } + enhancedDB = &EnhancedFingerprintDB{ Fingerprints: fps, - regexCache: make(map[string]*regexp.Regexp), } return nil @@ -128,7 +139,7 @@ func MatchEnhancedFingerprints(body []byte, headers string, favicon FaviconHashe } httpRule := fp.HTTP[0] for _, matcher := range httpRule.Matchers { - if matchMatcher(matcher, bodyStr, headers, favicon, enhancedDB.regexCache) { + if matchMatcher(matcher, bodyStr, headers, favicon) { resultCh <- fingerprintMatch{ Name: fp.Info.Name, Priority: calcPriority(fp, matcher.Type), @@ -202,13 +213,13 @@ func matchMatcher(matcher struct { Part string `json:"part"` CaseInsensitive bool `json:"case-insensitive"` Condition string `json:"condition"` -}, body, headers string, favicon FaviconHashes, regexCache map[string]*regexp.Regexp) bool { +}, body, headers string, favicon FaviconHashes) bool { switch matcher.Type { case "word": return matchWords(matcher, body, headers) case "regex": - return matchRegex(matcher, body, headers, regexCache) + return matchRegex(matcher, body, headers) case "favicon": return matchFavicon(matcher, favicon) default: @@ -233,21 +244,19 @@ func matchWords(matcher struct { target = headers } - // 预处理搜索词,避免循环内重复转换 - searchWords := matcher.Words + // CaseInsensitive: target 转小写一次,Words 已在加载时预处理(直接调用时也兼容未预处理的词) if matcher.CaseInsensitive { target = strings.ToLower(target) - searchWords = make([]string, len(matcher.Words)) - for i, w := range matcher.Words { - searchWords[i] = strings.ToLower(w) - } } // 默认condition为or isAnd := matcher.Condition == "and" matchCount := 0 - for _, searchWord := range searchWords { + for _, searchWord := range matcher.Words { + if matcher.CaseInsensitive { + searchWord = strings.ToLower(searchWord) + } if strings.Contains(target, searchWord) { if !isAnd { // OR条件:匹配任一即可 @@ -273,7 +282,7 @@ func matchRegex(matcher struct { Part string `json:"part"` CaseInsensitive bool `json:"case-insensitive"` Condition string `json:"condition"` -}, body, headers string, regexCache map[string]*regexp.Regexp) bool { +}, body, headers string) bool { // 确定匹配目标 target := body @@ -285,33 +294,23 @@ func matchRegex(matcher struct { isAnd := matcher.Condition == "and" for _, pattern := range matcher.Regex { - // 从缓存获取或编译正则(线程安全) + // CaseInsensitive 正则需要前缀 + cacheKey := pattern + if matcher.CaseInsensitive { + cacheKey = "(?i)" + pattern + } + + // 从 sync.Map 缓存获取或编译正则 var re *regexp.Regexp - - // 先尝试读取缓存(读锁) - enhancedDB.regexCacheMu.RLock() - re, exists := regexCache[pattern] - enhancedDB.regexCacheMu.RUnlock() - - if !exists { - // 不存在,需要编译并写入缓存(写锁) - enhancedDB.regexCacheMu.Lock() - // Double-check:可能其他goroutine已经编译了 - re, exists = regexCache[pattern] - if !exists { - var err error - if matcher.CaseInsensitive { - re, err = regexp.Compile("(?i)" + pattern) - } else { - re, err = regexp.Compile(pattern) - } - if err != nil { - enhancedDB.regexCacheMu.Unlock() - continue - } - regexCache[pattern] = re + if cached, ok := enhancedDB.regexCache.Load(cacheKey); ok { + re = cached.(*regexp.Regexp) + } else { + compiled, err := regexp.Compile(cacheKey) + if err != nil { + continue } - enhancedDB.regexCacheMu.Unlock() + actual, _ := enhancedDB.regexCache.LoadOrStore(cacheKey, compiled) + re = actual.(*regexp.Regexp) } // 确保 re 不为 nil(防止并发场景下的 nil panic) @@ -402,7 +401,7 @@ func ExtractVersions(body string, headers string) []VersionInfo { seen := make(map[string]struct{}) for _, extractor := range versionExtractors { - matches := extractor.pattern.FindAllStringSubmatch(content, -1) + matches := extractor.pattern.FindAllStringSubmatch(content, 5) for _, match := range matches { var name, version string diff --git a/webscan/fingerprint/enhanced_test.go b/webscan/fingerprint/enhanced_test.go index 17843d2..d0df42c 100644 --- a/webscan/fingerprint/enhanced_test.go +++ b/webscan/fingerprint/enhanced_test.go @@ -447,7 +447,7 @@ func TestMatchMatcher_TypeDispatch(t *testing.T) { matcher = createMatcher(tt.matcherType, nil, nil, nil, "", "", false) } - result := matchMatcher(matcher, "nginx server", "Server: nginx", FaviconHashes{MMH3: "abc123", MD5: "def456"}, nil) + result := matchMatcher(matcher, "nginx server", "Server: nginx", FaviconHashes{MMH3: "abc123", MD5: "def456"}) if result != tt.expected { t.Errorf("matchMatcher(type=%s) = %v, 期望 %v", tt.matcherType, result, tt.expected) diff --git a/webscan/lib/Client.go b/webscan/lib/Client.go index 2b90c65..014d516 100644 --- a/webscan/lib/Client.go +++ b/webscan/lib/Client.go @@ -143,12 +143,28 @@ func InitHTTPClient(ThreadsNum int, DownProxy string, Timeout time.Duration, max KeepAlive: keepAlive, } + // 连接池参数随并发数动态调整 + maxConns := ThreadsNum * 2 + if maxConns < 20 { + maxConns = 20 + } + if maxConns > 200 { + maxConns = 200 + } + idlePerHost := ThreadsNum / 2 + if idlePerHost < 5 { + idlePerHost = 5 + } + if idlePerHost > 20 { + idlePerHost = 20 + } + // 配置Transport参数 tr := &http.Transport{ DialContext: dialer.DialContext, - MaxConnsPerHost: 100, // 增加到100,避免连接池耗尽 - MaxIdleConns: 100, // 保留100个空闲连接 - MaxIdleConnsPerHost: 10, // 每主机保留10个空闲连接 + MaxConnsPerHost: maxConns, + MaxIdleConns: maxConns, + MaxIdleConnsPerHost: idlePerHost, IdleConnTimeout: keepAlive, TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS10, InsecureSkipVerify: true}, TLSHandshakeTimeout: 5 * time.Second, diff --git a/webscan/lib/Eval.go b/webscan/lib/Eval.go index f84dd0f..d751e05 100644 --- a/webscan/lib/Eval.go +++ b/webscan/lib/Eval.go @@ -3,6 +3,7 @@ package lib import ( "bytes" "compress/gzip" + "context" "errors" "fmt" "io" @@ -368,11 +369,14 @@ func reverseCheck(r *Reverse, timeout int64) bool { apiURL := fmt.Sprintf("http://api.ceye.io/v1/records?token=%s&type=dns&filter=%s", ceyeAPI, sub) - // 创建并发送请求 - req, err := http.NewRequest("GET", apiURL, nil) + // 创建并发送请求(带超时控制,避免 ceye API 无响应时阻塞) + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + req, err := http.NewRequestWithContext(ctx, "GET", apiURL, nil) if err != nil { return false } + // nil session: CEL 回调无法获取 session,回退到全局限速(反连检查请求量极低,可接受) resp, err := DoRequest(req, false, nil) if err != nil { return false diff --git a/webscan/lib/eval_misc.go b/webscan/lib/eval_misc.go index d5c1a27..c6e0656 100644 --- a/webscan/lib/eval_misc.go +++ b/webscan/lib/eval_misc.go @@ -31,7 +31,7 @@ func registerMiscImplementations() []*functions.Overload { Operator: "reverse_wait_int", Binary: func(lhs ref.Val, rhs ref.Val) ref.Val { reverse, ok := lhs.Value().(*Reverse) - if !ok { + if !ok || reverse == nil { return types.ValOrErr(lhs, "unexpected type '%v' passed to wait", lhs.Type()) } timeout, ok := rhs.(types.Int) diff --git a/webscan/lib/eval_string.go b/webscan/lib/eval_string.go index b8aa5e6..d5af036 100644 --- a/webscan/lib/eval_string.go +++ b/webscan/lib/eval_string.go @@ -70,11 +70,19 @@ func registerStringImplementations() []*functions.Overload { if !ok { return types.ValOrErr(rhs, "unexpected type '%v' passed to bmatch", rhs.Type()) } - ok, err := regexp.Match(string(v1), v2) - if err != nil { - return types.NewErr("%v", err) + pattern := string(v1) + var re *regexp.Regexp + if cached, found := regexCache.Load(pattern); found { + re = cached.(*regexp.Regexp) + } else { + compiled, err := regexp.Compile(pattern) + if err != nil { + return types.NewErr("%v", err) + } + actual, _ := regexCache.LoadOrStore(pattern, compiled) + re = actual.(*regexp.Regexp) } - return types.Bool(ok) + return types.Bool(re.Match(v2)) }, }, { diff --git a/webscan/lib/poc_adapter.go b/webscan/lib/poc_adapter.go index cb67767..efc69e2 100644 --- a/webscan/lib/poc_adapter.go +++ b/webscan/lib/poc_adapter.go @@ -307,8 +307,8 @@ func convertNucleiMatchers(matchers []NucleiMatcher, matchersCondition string) s matcherConds = append(matcherConds, nucleiRegexCondition(m.Part, pattern)) } case "dsl": - // DSL类型暂不支持,使用默认匹配 - matcherConds = append(matcherConds, "response.status == 200") + // DSL类型暂不支持,安全降级为 false 避免误报 + matcherConds = append(matcherConds, "false") } // 单个matcher内的条件组合 diff --git a/webscan/lib/poc_executor.go b/webscan/lib/poc_executor.go index bc6be89..7c38e26 100644 --- a/webscan/lib/poc_executor.go +++ b/webscan/lib/poc_executor.go @@ -3,11 +3,13 @@ package lib import ( "crypto/md5" //nolint:gosec // G501: MD5用于POC规则去重,非加密用途 "fmt" + "io" "math/rand" //nolint:gosec // G404: math/rand用于生成测试数据,非加密用途 "net/http" "net/url" "os" "regexp" + "sort" "strings" "sync" "time" @@ -91,7 +93,7 @@ func CheckMultiPoc(req *http.Request, pocs []*Poc, workers int, pocCtx *POCConte // 因为clusterpoc已在内部处理了漏洞输出 if isVulnerable && vulName != "" { // 构造漏洞详细信息 - details := make(map[string]interface{}) + details := make(map[string]interface{}, 6) details["vulnerability_type"] = task.Poc.Name details["vulnerability_name"] = vulName @@ -341,14 +343,22 @@ func executeRules(oReq *http.Request, p *Poc, variableMap map[string]interface{} return false, "", nil } +var regexCache sync.Map + // doSearch 在响应体中执行正则匹配并提取命名捕获组 func doSearch(re string, body string) map[string]string { - // 编译正则表达式 - r, err := regexp.Compile(re) - // 正则表达式编译 - if err != nil { - common.LogError(i18n.Tr("webscan_regex_compile_error", err)) - return nil + // 编译正则表达式(带缓存) + var r *regexp.Regexp + if cached, ok := regexCache.Load(re); ok { + r = cached.(*regexp.Regexp) + } else { + compiled, err := regexp.Compile(re) + if err != nil { + common.LogError(i18n.Tr("webscan_regex_compile_error", err)) + return nil + } + actual, _ := regexCache.LoadOrStore(re, compiled) + r = actual.(*regexp.Regexp) } // 执行正则匹配 @@ -537,7 +547,7 @@ func clusterpoc(oReq *http.Request, p *Poc, variableMap map[string]interface{}, if ruleIndex == len(p.Rules)-1 { // 最终规则成功,记录完整的结果并返回 recordVulnerabilityResult(targetURL, p, strMap, false, pocCtx.Session) - return false, nil + return true, nil } break paramLoop } @@ -620,11 +630,25 @@ func applyParametersToRule( return hasReplacement, replacedParams } -// getRuleHash 计算规则的MD5哈希值用于去重 +// getRuleHash 计算规则的MD5哈希值用于去重(无反射,Headers 排序保证确定性) func getRuleHash(rule *Rules) string { //nolint:gosec // G401: MD5用于规则去重,非加密用途 - ruleDigest := md5.Sum([]byte(fmt.Sprintf("%v", rule))) - return fmt.Sprintf("%x", ruleDigest) + h := md5.New() + _, _ = io.WriteString(h, rule.Method) + _, _ = io.WriteString(h, rule.Path) + _, _ = io.WriteString(h, rule.Body) + if len(rule.Headers) > 0 { + keys := make([]string, 0, len(rule.Headers)) + for k := range rule.Headers { + keys = append(keys, k) + } + sort.Strings(keys) + for _, k := range keys { + _, _ = io.WriteString(h, k) + _, _ = io.WriteString(h, rule.Headers[k]) + } + } + return fmt.Sprintf("%x", h.Sum(nil)) } // recordVulnerabilityResult 记录漏洞检测结果 @@ -830,11 +854,10 @@ func clustersend(oReq *http.Request, variableMap map[string]interface{}, req *Re } // 检查表达式执行结果 - if fmt.Sprintf("%v", out) == "false" { - return false, nil + if flag, ok := out.Value().(bool); ok { + return flag, nil } - - return true, nil + return false, nil } // cloneRules 深度复制Rules结构体 @@ -870,8 +893,8 @@ func cloneMap(tags map[string]string) map[string]string { func evalset(env *cel.Env, variableMap map[string]interface{}, k string, expression string) (string, error) { out, err := Evaluate(env, expression, variableMap) if err != nil { - variableMap[k] = expression - return expression, err + variableMap[k] = "" + return "", err } // 根据不同类型处理输出 diff --git a/webscan/web_scan.go b/webscan/web_scan.go index 22b3f37..3bc593e 100644 --- a/webscan/web_scan.go +++ b/webscan/web_scan.go @@ -13,6 +13,7 @@ import ( "strconv" "strings" "sync" + "sync/atomic" "time" "github.com/shadow1ng/fscan/common" @@ -43,7 +44,7 @@ var ( var pocsFS embed.FS var ( pocMu sync.Mutex - pocLoaded bool + pocLoaded atomic.Bool allPocs []*lib.Poc cachedPocPath string ) @@ -53,16 +54,18 @@ func WebScan(ctx context.Context, info *common.HostInfo, cfg *common.Config, ses // 初始化POC配置(用于CEL回调函数) lib.InitPOCConfig(cfg.DNSLog) - // 加载POC(互斥保护,避免并发 race) - pocMu.Lock() - if !pocLoaded { - cachedPocPath = cfg.POC.PocPath - initPocs() - if len(allPocs) > 0 { - pocLoaded = true + // 加载POC(DCLP: 快速路径无锁,慢路径互斥保护) + if !pocLoaded.Load() { + pocMu.Lock() + if !pocLoaded.Load() { + cachedPocPath = cfg.POC.PocPath + initPocs() + if len(allPocs) > 0 { + pocLoaded.Store(true) + } } + pocMu.Unlock() } - pocMu.Unlock() // 验证输入 if info == nil { @@ -316,7 +319,7 @@ func loadExternalPocs(pocPath string) { loadPocsConcurrently(pocFiles, false, pocPath) } -// loadPocsConcurrently 并发加载POC文件 +// loadPocsConcurrently 并发加载POC文件(channel 收集,无锁竞争) func loadPocsConcurrently(pocFiles []string, isEmbedded bool, pocPath string) { pocCount := len(pocFiles) if pocCount == 0 { @@ -324,48 +327,45 @@ func loadPocsConcurrently(pocFiles []string, isEmbedded bool, pocPath string) { } var wg sync.WaitGroup - var mu sync.Mutex - var successCount, failCount int - - // 使用信号量控制并发数 + results := make(chan *lib.Poc, pocCount) semaphore := make(chan struct{}, concurrencyLimit) for _, file := range pocFiles { wg.Add(1) - semaphore <- struct{}{} // 获取信号量 + semaphore <- struct{}{} go func(filename string) { defer func() { - <-semaphore // 释放信号量 + <-semaphore wg.Done() }() var poc *lib.Poc var err error - - // 根据不同的来源加载POC if isEmbedded { poc, err = lib.LoadPoc(filename, pocsFS) } else { poc, err = lib.LoadPocbyPath(filename) } - mu.Lock() - defer mu.Unlock() - - if err != nil { - failCount++ - return - } - - if poc != nil { - allPocs = append(allPocs, poc) - successCount++ + if err == nil && poc != nil { + results <- poc } }(file) } - wg.Wait() + go func() { + wg.Wait() + close(results) + }() + + var successCount int + for poc := range results { + allPocs = append(allPocs, poc) + successCount++ + } + + failCount := pocCount - successCount common.LogInfo(i18n.Tr("poc_load_complete", pocCount, successCount, failCount)) }