From 1595c92aed029ff3b2548cca0969548423ba3b4e Mon Sep 17 00:00:00 2001 From: ZacharyZcR Date: Fri, 12 Jun 2026 19:49:07 +0800 Subject: [PATCH] =?UTF-8?q?feat:=20=E7=BB=9F=E4=B8=80=E6=9C=8D=E5=8A=A1?= =?UTF-8?q?=E7=BC=93=E5=AD=98=20+=20=E6=8C=87=E7=BA=B9=E9=A9=B1=E5=8A=A8?= =?UTF-8?q?=E6=8F=92=E4=BB=B6=E5=8C=B9=E9=85=8D?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 将 webServiceCache 扩展为通用 serviceCache,所有指纹识别结果 统一缓存,插件匹配时端口不命中则回退到服务名称匹配。 删除多余的 service_cache.go,复用已有的 ServiceInfo 体系。 补充 nil 防御、Explicit 标记、大量单元/集成/回归测试。 --- common/config_builder.go | 39 ++- common/config_builder_test.go | 118 +++++++ common/config_struct.go | 37 ++- common/flag_config.go | 8 +- common/globals.go | 6 +- common/globals_test.go | 7 + common/network.go | 19 +- common/output/stdout_writer_test.go | 35 ++ common/output/writers.go | 6 + common/output/writers_test.go | 3 + common/session_test.go | 20 ++ core/base_scan_strategy.go | 6 +- core/network_profiler.go | 8 +- core/network_profiler_test.go | 38 ++- core/port_scan.go | 5 +- core/port_scan_test.go | 42 +++ core/service_cache.go | 36 --- core/service_cache_test.go | 354 +++++++++++++++++++++ core/service_scanner_test.go | 37 ++- core/tuning_cli_integration_test.go | 82 +++++ core/web_scanner.go | 46 ++- core/web_scanner_test.go | 12 +- libs/grdp/login/screen.go | 14 +- libs/grdp/login/screen_test.go | 24 ++ plugins/init_test.go | 32 +- plugins/local/systeminfo_dc_url.go | 11 + plugins/local/systeminfo_dc_url_test.go | 24 ++ plugins/local/systeminfo_dc_windows.go | 4 +- plugins/services/cassandra.go | 16 +- plugins/services/credential_tester.go | 88 ++++- plugins/services/credential_tester_test.go | 220 ++++++++++++- plugins/services/kafka.go | 8 +- plugins/services/mongodb.go | 4 +- plugins/services/protocol_ids_test.go | 46 +++ plugins/web/webtitle.go | 28 +- plugins/web/webtitle_test.go | 21 ++ web/api/result.go | 35 +- web/api/result_test.go | 31 ++ webscan/lib/Client.go | 10 +- webscan/lib/Eval.go | 13 +- webscan/lib/client_test.go | 24 ++ webscan/lib/eval_test.go | 9 + webscan/web_scan.go | 29 +- webscan/web_scan_test.go | 20 ++ 44 files changed, 1533 insertions(+), 142 deletions(-) create mode 100644 common/output/stdout_writer_test.go delete mode 100644 core/service_cache.go create mode 100644 core/service_cache_test.go create mode 100644 core/tuning_cli_integration_test.go create mode 100644 libs/grdp/login/screen_test.go create mode 100644 plugins/local/systeminfo_dc_url.go create mode 100644 plugins/local/systeminfo_dc_url_test.go create mode 100644 plugins/services/protocol_ids_test.go create mode 100644 web/api/result_test.go create mode 100644 webscan/lib/client_test.go diff --git a/common/config_builder.go b/common/config_builder.go index d5b7210..26ce43b 100644 --- a/common/config_builder.go +++ b/common/config_builder.go @@ -4,6 +4,7 @@ import ( "encoding/hex" "fmt" "net" + "net/url" "strconv" "strings" @@ -171,6 +172,7 @@ func parseUserPassPairs(fv *FlagVars) ([]config.CredentialPair, error) { // 如果命令行同时指定了单个用户名和单个密码(不是逗号分隔的多个) if fv.Username != "" && fv.Password != "" && !strings.Contains(fv.Username, ",") && !strings.Contains(fv.Password, ",") && + fv.AddUsers == "" && fv.AddPasswords == "" && fv.UsersFile == "" && fv.PasswordsFile == "" && fv.UserPassFile == "" { pairs = append(pairs, config.CredentialPair{ Username: strings.TrimSpace(fv.Username), @@ -294,9 +296,42 @@ func normalizeURL(rawURL string) string { } lowerURL := strings.ToLower(rawURL) if !strings.HasPrefix(lowerURL, "http://") && !strings.HasPrefix(lowerURL, "https://") { - return "http://" + rawURL + return "http://" + normalizeSchemelessURLTarget(rawURL) } - return rawURL + parsed, err := url.Parse(rawURL) + if err != nil || parsed.Host == "" { + return rawURL + } + normalizedHost := normalizeURLHost(parsed.Host) + if normalizedHost == parsed.Host { + return rawURL + } + parsed.Host = normalizedHost + normalized := parsed.String() + if schemeEnd := strings.Index(rawURL, "://"); schemeEnd >= 0 { + return rawURL[:schemeEnd] + normalized[len(parsed.Scheme):] + } + return normalized +} + +func normalizeSchemelessURLTarget(rawURL string) string { + authority := rawURL + suffix := "" + if idx := strings.IndexAny(rawURL, "/?#"); idx >= 0 { + authority = rawURL[:idx] + suffix = rawURL[idx:] + } + return normalizeURLHost(authority) + suffix +} + +func normalizeURLHost(host string) string { + if strings.HasPrefix(host, "[") { + return host + } + if ip := net.ParseIP(host); ip != nil && strings.Contains(host, ":") { + return "[" + host + "]" + } + return host } // ============================================================================= diff --git a/common/config_builder_test.go b/common/config_builder_test.go index 376ac61..a90186c 100644 --- a/common/config_builder_test.go +++ b/common/config_builder_test.go @@ -3,6 +3,8 @@ package common import ( "reflect" "testing" + + fscanconfig "github.com/shadow1ng/fscan/common/config" ) func TestParsePasswordsKeepsPrimaryPasswordLiteral(t *testing.T) { @@ -49,6 +51,100 @@ func TestBuildConfigRejectsInvalidHashValue(t *testing.T) { } } +func TestBuildConfigDefaultsAreIndependentCopies(t *testing.T) { + cfg, _, err := BuildConfig(&FlagVars{Username: "custom-user"}, &HostInfo{}) + if err != nil { + t.Fatalf("BuildConfig error = %v", err) + } + + defaultSSHUsers := fscanconfig.DefaultUserDict["ssh"] + if len(defaultSSHUsers) == 1 && defaultSSHUsers[0] == "custom-user" { + t.Fatal("BuildConfig mutated DefaultUserDict") + } + + cfg.Credentials.Userdict["ssh"][0] = "mutated-user" + if fscanconfig.DefaultUserDict["ssh"][0] == "mutated-user" { + t.Fatal("Config userdict shares backing storage with DefaultUserDict") + } + + cfg.Credentials.Passwords[0] = "mutated-password" + if fscanconfig.DefaultPasswords[0] == "mutated-password" { + t.Fatal("Config passwords share backing storage with DefaultPasswords") + } + + port := 80 + cfg.PortMap[port][0] = "mutated-probe" + if fscanconfig.DefaultPortMap[port][0] == "mutated-probe" { + t.Fatal("Config port map shares backing storage with DefaultPortMap") + } + + cfg.DefaultMap[0] = "mutated-default-probe" + if fscanconfig.DefaultProbeMap[0] == "mutated-default-probe" { + t.Fatal("Config default map shares backing storage with DefaultProbeMap") + } +} + +func TestParseUserPassPairsKeepsAdditionalCredentialFlags(t *testing.T) { + tests := []struct { + name string + fv *FlagVars + }{ + { + name: "additional passwords", + fv: &FlagVars{ + Username: "root", + Password: "primary", + AddPasswords: "extra", + }, + }, + { + name: "additional users", + fv: &FlagVars{ + Username: "root", + Password: "primary", + AddUsers: "admin", + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + pairs, err := parseUserPassPairs(tt.fv) + if err != nil { + t.Fatalf("parseUserPassPairs error = %v", err) + } + if len(pairs) != 0 { + t.Fatalf("parseUserPassPairs returned exact pairs %#v; additional credential flags would be ignored", pairs) + } + }) + } +} + +func TestNewConfigDefaultsAreIndependentCopies(t *testing.T) { + cfg := NewConfig() + + cfg.Credentials.Userdict["ssh"][0] = "mutated-user" + if fscanconfig.DefaultUserDict["ssh"][0] == "mutated-user" { + t.Fatal("NewConfig userdict shares backing storage with DefaultUserDict") + } + + cfg.Credentials.Passwords[0] = "mutated-password" + if fscanconfig.DefaultPasswords[0] == "mutated-password" { + t.Fatal("NewConfig passwords share backing storage with DefaultPasswords") + } + + port := 80 + cfg.PortMap[port][0] = "mutated-probe" + if fscanconfig.DefaultPortMap[port][0] == "mutated-probe" { + t.Fatal("NewConfig port map shares backing storage with DefaultPortMap") + } + + cfg.DefaultMap[0] = "mutated-default-probe" + if fscanconfig.DefaultProbeMap[0] == "mutated-default-probe" { + t.Fatal("NewConfig default map shares backing storage with DefaultProbeMap") + } +} + func TestParseTargetsHostPortDoesNotLeaveSyntheticHost(t *testing.T) { fv := &FlagVars{Ports: "22"} info := &HostInfo{Host: "127.0.0.1:8080"} @@ -73,3 +169,25 @@ func TestNormalizeURLKeepsUppercaseScheme(t *testing.T) { t.Fatalf("normalizeURL() = %q", got) } } + +func TestNormalizeURLBracketsIPv6Literals(t *testing.T) { + tests := []struct { + name string + in string + want string + }{ + {name: "bare ipv6 without scheme", in: "2001:db8::1", want: "http://[2001:db8::1]"}, + {name: "bracketed ipv6 without scheme", in: "[2001:db8::1]", want: "http://[2001:db8::1]"}, + {name: "bare ipv6 with scheme", in: "http://2001:db8::1", want: "http://[2001:db8::1]"}, + {name: "bare ipv6 path without scheme", in: "2001:db8::1/admin", want: "http://[2001:db8::1]/admin"}, + {name: "bare ipv6 query without scheme", in: "2001:db8::1?debug=1", want: "http://[2001:db8::1]?debug=1"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := normalizeURL(tt.in); got != tt.want { + t.Fatalf("normalizeURL(%q) = %q, want %q", tt.in, got, tt.want) + } + }) + } +} diff --git a/common/config_struct.go b/common/config_struct.go index 164748c..10a1f24 100644 --- a/common/config_struct.go +++ b/common/config_struct.go @@ -145,6 +145,35 @@ type LocalExploitConfig struct { DownloadSavePath string // 下载保存路径 } +func cloneStringSlice(values []string) []string { + if values == nil { + return nil + } + return append([]string(nil), values...) +} + +func cloneStringSliceMap(values map[string][]string) map[string][]string { + if values == nil { + return nil + } + cloned := make(map[string][]string, len(values)) + for key, value := range values { + cloned[key] = cloneStringSlice(value) + } + return cloned +} + +func clonePortMap(values map[int][]string) map[int][]string { + if values == nil { + return nil + } + cloned := make(map[int][]string, len(values)) + for key, value := range values { + cloned[key] = cloneStringSlice(value) + } + return cloned +} + // NewConfig 创建带默认值的Config(后备用,正常流程使用BuildConfigFromFlags) func NewConfig() *Config { return &Config{ @@ -163,13 +192,13 @@ func NewConfig() *Config { MaxRetries: 3, // 高级功能 - 使用默认配置 - PortMap: config.DefaultPortMap, - DefaultMap: config.DefaultProbeMap, + PortMap: clonePortMap(config.DefaultPortMap), + DefaultMap: cloneStringSlice(config.DefaultProbeMap), // 分组配置 - 使用默认字典 Credentials: CredentialConfig{ - Userdict: config.DefaultUserDict, - Passwords: config.DefaultPasswords, + Userdict: cloneStringSliceMap(config.DefaultUserDict), + Passwords: cloneStringSlice(config.DefaultPasswords), UserPassPairs: nil, }, Network: NetworkConfig{ diff --git a/common/flag_config.go b/common/flag_config.go index 085966b..4f09030 100644 --- a/common/flag_config.go +++ b/common/flag_config.go @@ -164,8 +164,8 @@ func BuildConfigFromFlags(fv *FlagVars) *Config { DNSLog: fv.DNSLog, PersistenceTargetFile: fv.PersistenceTargetFile, WinPEFile: fv.WinPEFile, - PortMap: config.DefaultPortMap, - DefaultMap: config.DefaultProbeMap, + PortMap: clonePortMap(config.DefaultPortMap), + DefaultMap: cloneStringSlice(config.DefaultProbeMap), // SOCKS5代理端口 Socks5ProxyPort: fv.Socks5ProxyPort, @@ -175,8 +175,8 @@ func BuildConfigFromFlags(fv *FlagVars) *Config { Username: fv.Username, Password: fv.Password, Domain: fv.Domain, - Userdict: config.DefaultUserDict, - Passwords: config.DefaultPasswords, + Userdict: cloneStringSliceMap(config.DefaultUserDict), + Passwords: cloneStringSlice(config.DefaultPasswords), UserPassPairs: nil, // 后续解析 SSHKeyPath: fv.SSHKeyPath, }, diff --git a/common/globals.go b/common/globals.go index 9ee1929..2500645 100644 --- a/common/globals.go +++ b/common/globals.go @@ -31,7 +31,11 @@ type HostInfo struct { // Target 返回 host:port 格式字符串 func (h *HostInfo) Target() string { - return net.JoinHostPort(h.Host, strconv.Itoa(h.Port)) + host := h.Host + if strings.HasPrefix(host, "[") && strings.HasSuffix(host, "]") { + host = strings.TrimPrefix(strings.TrimSuffix(host, "]"), "[") + } + return net.JoinHostPort(host, strconv.Itoa(h.Port)) } // ============================================================================= diff --git a/common/globals_test.go b/common/globals_test.go index 7933e8b..3807e81 100644 --- a/common/globals_test.go +++ b/common/globals_test.go @@ -8,3 +8,10 @@ func TestHostInfoTargetUsesBracketedIPv6(t *testing.T) { t.Fatalf("Target() = %q, want %q", got, want) } } + +func TestHostInfoTargetDoesNotDoubleBracketIPv6(t *testing.T) { + info := &HostInfo{Host: "[2001:db8::1]", Port: 443} + if got, want := info.Target(), "[2001:db8::1]:443"; got != want { + t.Fatalf("Target() = %q, want %q", got, want) + } +} diff --git a/common/network.go b/common/network.go index f45abb8..ffef90d 100644 --- a/common/network.go +++ b/common/network.go @@ -52,16 +52,31 @@ func getGlobalDialer(timeout time.Duration) (proxy.Dialer, error) { // parseProxyURL 解析代理URL,提取地址和认证信息 func parseProxyURL(proxyURL, fallback string) (host, username, password string) { + if !strings.Contains(proxyURL, "://") { + if host, username, password, ok := parseProxyURLCandidate("http://" + proxyURL); ok { + return host, username, password + } + } + if host, username, password, ok := parseProxyURLCandidate(proxyURL); ok { + return host, username, password + } + return fallback, "", "" +} + +func parseProxyURLCandidate(proxyURL string) (host, username, password string, ok bool) { parsedURL, err := url.Parse(proxyURL) if err != nil { - return fallback, "", "" + return "", "", "", false } host = parsedURL.Host + if host == "" { + return "", "", "", false + } if parsedURL.User != nil { username = parsedURL.User.Username() password, _ = parsedURL.User.Password() } - return + return host, username, password, true } // createProxyConfig 根据全局设置创建代理配置 diff --git a/common/output/stdout_writer_test.go b/common/output/stdout_writer_test.go new file mode 100644 index 0000000..0c73531 --- /dev/null +++ b/common/output/stdout_writer_test.go @@ -0,0 +1,35 @@ +package output + +import "testing" + +func TestSplitHostPort(t *testing.T) { + tests := []struct { + name string + target string + wantHost string + wantPort int + wantOK bool + }{ + {name: "ipv4", target: "192.168.1.1:80", wantHost: "192.168.1.1", wantPort: 80, wantOK: true}, + {name: "hostname", target: "example.com:443", wantHost: "example.com", wantPort: 443, wantOK: true}, + {name: "bracketed ipv6", target: "[2001:db8::1]:8443", wantHost: "2001:db8::1", wantPort: 8443, wantOK: true}, + {name: "bare ipv6 without port", target: "2001:db8::1", wantOK: false}, + {name: "invalid port", target: "example.com:abc", wantOK: false}, + {name: "port out of range", target: "example.com:65536", wantOK: false}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + host, port, ok := splitHostPort(tt.target) + if ok != tt.wantOK { + t.Fatalf("splitHostPort(%q) ok = %v, want %v", tt.target, ok, tt.wantOK) + } + if !ok { + return + } + if host != tt.wantHost || port != tt.wantPort { + t.Fatalf("splitHostPort(%q) = (%q, %d), want (%q, %d)", tt.target, host, port, tt.wantHost, tt.wantPort) + } + }) + } +} diff --git a/common/output/writers.go b/common/output/writers.go index 328d710..973e12e 100644 --- a/common/output/writers.go +++ b/common/output/writers.go @@ -46,6 +46,12 @@ func targetWithPort(target string, port interface{}) string { return target } portText := fmt.Sprint(port) + if strings.TrimSpace(portText) == "" { + return target + } + if strings.HasPrefix(target, "[") && strings.HasSuffix(target, "]") { + target = strings.TrimPrefix(strings.TrimSuffix(target, "]"), "[") + } if strings.Count(target, ":") == 1 { return target } diff --git a/common/output/writers_test.go b/common/output/writers_test.go index 7118ebb..c9fdedf 100644 --- a/common/output/writers_test.go +++ b/common/output/writers_test.go @@ -67,7 +67,10 @@ func TestTargetWithPortIPv6(t *testing.T) { {name: "ipv4 without port", target: "192.168.1.1", port: 80, want: "192.168.1.1:80"}, {name: "ipv4 with port", target: "192.168.1.1:80", port: 443, want: "192.168.1.1:80"}, {name: "ipv6 without port", target: "2001:db8::1", port: 443, want: "[2001:db8::1]:443"}, + {name: "bracketed ipv6 without port", target: "[2001:db8::1]", port: 443, want: "[2001:db8::1]:443"}, {name: "ipv6 with port", target: "[2001:db8::1]:443", port: 80, want: "[2001:db8::1]:443"}, + {name: "empty port", target: "example.com", port: "", want: "example.com"}, + {name: "blank port", target: "example.com", port: " \t", want: "example.com"}, } for _, tt := range tests { diff --git a/common/session_test.go b/common/session_test.go index cb03d17..aa7f361 100644 --- a/common/session_test.go +++ b/common/session_test.go @@ -141,6 +141,26 @@ func TestScanSessionProxyStateComesFromConfig(t *testing.T) { } } +func TestParseProxyURLFallsBackWhenHostIsEmpty(t *testing.T) { + host, username, password := parseProxyURL("127.0.0.1:8080", "127.0.0.1:8080") + if host != "127.0.0.1:8080" { + t.Fatalf("host = %q, want fallback address", host) + } + if username != "" || password != "" { + t.Fatalf("unexpected credentials: %q/%q", username, password) + } +} + +func TestParseProxyURLExtractsAuthWithoutScheme(t *testing.T) { + host, username, password := parseProxyURL("user:pass@127.0.0.1:8080", "user:pass@127.0.0.1:8080") + if host != "127.0.0.1:8080" { + t.Fatalf("host = %q, want proxy address", host) + } + if username != "user" || password != "pass" { + t.Fatalf("credentials = %q/%q, want user/pass", username, password) + } +} + type roundTripFunc func(*http.Request) (*http.Response, error) func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) { diff --git a/core/base_scan_strategy.go b/core/base_scan_strategy.go index a1912ab..c553f4a 100644 --- a/core/base_scan_strategy.go +++ b/core/base_scan_strategy.go @@ -137,11 +137,11 @@ func (b *BaseScanStrategy) isPluginApplicableToPortWithHost(pluginName string, t } } - // 端口不匹配时,按服务识别结果匹配 + // 端口不匹配时,按指纹识别结果匹配 // 例:8881 端口上识别到 ssh 服务 → ssh 插件应该执行 if targetHost != "" && targetPort > 0 { - if svcName, ok := GetServiceName(targetHost, targetPort); ok { - if strings.EqualFold(svcName, pluginName) { + if info, ok := GetCachedServiceInfo(targetHost, targetPort); ok && info != nil { + if strings.EqualFold(info.Name, pluginName) { return true } } diff --git a/core/network_profiler.go b/core/network_profiler.go index 6accdd3..5aba343 100644 --- a/core/network_profiler.go +++ b/core/network_profiler.go @@ -2,10 +2,10 @@ package core import ( "context" - "fmt" "math" "net" "sort" + "strconv" "sync" "time" @@ -97,6 +97,10 @@ func (p *NetworkProfile) RecommendConcurrency(userThreadNum int, explicit bool) // probePorts 探测用的端口列表(高响应率的常见端口) var probePorts = []int{80, 443, 22} +func networkProbeAddress(host string, port int) string { + return net.JoinHostPort(host, strconv.Itoa(port)) +} + // ProbeNetwork 探测目标网络环境 // 从 hosts 中抽样,用低并发 TCP 连接测量 RTT 和丢包率 // 整个过程控制在数秒内完成 @@ -140,7 +144,7 @@ func ProbeNetwork(ctx context.Context, hosts []string, session *common.ScanSessi go func(h string, p int) { defer func() { <-sem; wg.Done() }() - addr := fmt.Sprintf("%s:%d", h, p) + addr := networkProbeAddress(h, p) start := time.Now() conn, err := session.DialTCP(ctx, "tcp", addr, probeTimeout) rtt := time.Since(start) diff --git a/core/network_profiler_test.go b/core/network_profiler_test.go index 252c385..3aef559 100644 --- a/core/network_profiler_test.go +++ b/core/network_profiler_test.go @@ -60,7 +60,7 @@ func TestClassifyNetwork(t *testing.T) { t.Run("公网 RTT 分布(低丢包)", func(t *testing.T) { rtts := makeDurations([]int{60, 70, 80, 90, 100, 110, 120, 150, 200, 300}) // ms - p := classifyNetwork(rtts, 0, 10) // 无丢包 + p := classifyNetwork(rtts, 0, 10) // 无丢包 if p.Env != EnvInternet { t.Errorf("env = %v, want Internet", p.Env) @@ -72,7 +72,7 @@ func TestClassifyNetwork(t *testing.T) { t.Run("高丢包归类为慢速", func(t *testing.T) { rtts := makeDurations([]int{60, 70, 80, 90, 100}) // ms, 5 responded - p := classifyNetwork(rtts, 5, 10) // 50% loss + p := classifyNetwork(rtts, 5, 10) // 50% loss if p.Env != EnvSlow { t.Errorf("env = %v, want Slow (高丢包)", p.Env) @@ -96,14 +96,14 @@ func TestClassifyNetwork(t *testing.T) { func TestRecommendConcurrency(t *testing.T) { tests := []struct { - env NetworkEnv - lossRate float64 - userT int - explicit bool - wantTMin int - wantTMax int - wantCeil int - desc string + env NetworkEnv + lossRate float64 + userT int + explicit bool + wantTMin int + wantTMax int + wantCeil int + desc string }{ {EnvLAN, 0.0, 600, false, 800, 1000, -1, "内网自动: ×1.5"}, {EnvWAN, 0.0, 600, false, 550, 650, -1, "局域网自动: ×1.0"}, @@ -156,6 +156,24 @@ func TestPickSamples(t *testing.T) { } } +func TestNetworkProbeAddressUsesJoinHostPort(t *testing.T) { + tests := []struct { + host string + port int + want string + }{ + {"127.0.0.1", 80, "127.0.0.1:80"}, + {"::1", 443, "[::1]:443"}, + {"2001:db8::1", 22, "[2001:db8::1]:22"}, + } + + for _, tt := range tests { + if got := networkProbeAddress(tt.host, tt.port); got != tt.want { + t.Fatalf("networkProbeAddress(%q, %d) = %q, want %q", tt.host, tt.port, got, tt.want) + } + } +} + // ============================================================================= // 辅助 // ============================================================================= diff --git a/core/port_scan.go b/core/port_scan.go index a816c2c..882fda6 100644 --- a/core/port_scan.go +++ b/core/port_scan.go @@ -707,8 +707,8 @@ func processServiceResult(ctx context.Context, host string, port int, addr strin return } - // 缓存服务名称,供插件按服务类型匹配(解决非标准端口问题) - MarkServiceName(host, port, serviceInfo.Name) + // 缓存指纹识别结果,供插件按服务类型匹配(解决非标准端口问题) + CacheServiceInfo(host, port, serviceInfo) // 保存并输出服务信息 details := buildServiceDetails(port, serviceInfo) @@ -716,7 +716,6 @@ func processServiceResult(ctx context.Context, host string, port int, addr strin if isWeb { details["is_web"] = true - MarkAsWebService(host, port, serviceInfo) } _ = session.SaveResult(&output.ScanResult{ diff --git a/core/port_scan_test.go b/core/port_scan_test.go index 0760338..7d8e0be 100644 --- a/core/port_scan_test.go +++ b/core/port_scan_test.go @@ -212,6 +212,48 @@ func TestFormatAddress(t *testing.T) { } } +func TestBuildWebServiceURLIPv6(t *testing.T) { + tests := []struct { + name string + addr string + serviceInfo *ServiceInfo + want string + }{ + { + name: "http default port", + addr: "[2001:db8::1]:80", + serviceInfo: &ServiceInfo{ + Name: "http", + }, + want: "http://[2001:db8::1]", + }, + { + name: "https default port", + addr: "[2001:db8::1]:443", + serviceInfo: &ServiceInfo{ + Name: "https", + }, + want: "https://[2001:db8::1]", + }, + { + name: "http non-default port", + addr: "[2001:db8::1]:8080", + serviceInfo: &ServiceInfo{ + Name: "http", + }, + want: "http://[2001:db8::1]:8080", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := buildWebServiceURL(tt.addr, tt.serviceInfo); got != tt.want { + t.Fatalf("buildWebServiceURL(%q) = %q, want %q", tt.addr, got, tt.want) + } + }) + } +} + // ============================================================================= // 排除端口逻辑测试(从EnhancedPortScan:28-32行提取) // ============================================================================= diff --git a/core/service_cache.go b/core/service_cache.go deleted file mode 100644 index 4f1054b..0000000 --- a/core/service_cache.go +++ /dev/null @@ -1,36 +0,0 @@ -package core - -import ( - "net" - "strconv" - "strings" - "sync" -) - -// 服务识别缓存:host:port → 服务名称 -// 端口扫描阶段写入,插件匹配阶段读取 -// 解决非标准端口上的服务无法匹配对应插件的问题 -var ( - serviceNameCache = make(map[string]string) - serviceCacheMu sync.RWMutex -) - -// MarkServiceName 记录端口上识别到的服务名称 -func MarkServiceName(host string, port int, serviceName string) { - if serviceName == "" || serviceName == "unknown" { - return - } - key := net.JoinHostPort(host, strconv.Itoa(port)) - serviceCacheMu.Lock() - serviceNameCache[key] = strings.ToLower(serviceName) - serviceCacheMu.Unlock() -} - -// GetServiceName 查询端口上的服务名称 -func GetServiceName(host string, port int) (string, bool) { - key := net.JoinHostPort(host, strconv.Itoa(port)) - serviceCacheMu.RLock() - name, ok := serviceNameCache[key] - serviceCacheMu.RUnlock() - return name, ok -} diff --git a/core/service_cache_test.go b/core/service_cache_test.go new file mode 100644 index 0000000..c277b2c --- /dev/null +++ b/core/service_cache_test.go @@ -0,0 +1,354 @@ +package core + +import ( + "sync" + "testing" + + "github.com/shadow1ng/fscan/plugins" +) + +// registerTestPlugins 注册测试用插件(名字和服务识别结果一致) +func registerTestPlugins(t *testing.T) { + t.Helper() + plugins.RegisterWithOptions("ssh", func() plugins.Plugin { return nil }, []int{22, 2222}, nil, true) + plugins.RegisterWithOptions("mysql", func() plugins.Plugin { return nil }, []int{3306}, nil, true) + plugins.RegisterWithOptions("ftp", func() plugins.Plugin { return nil }, []int{21}, nil, true) + plugins.RegisterWithOptions("redis", func() plugins.Plugin { return nil }, []int{6379}, nil, true) + plugins.RegisterWithOptions("postgresql", func() plugins.Plugin { return nil }, []int{5432}, nil, true) + plugins.RegisterWithOptions("telnet", func() plugins.Plugin { return nil }, []int{23}, nil, true) + plugins.RegisterWithOptions("mssql", func() plugins.Plugin { return nil }, []int{1433}, nil, true) + plugins.RegisterWithOptions("vnc", func() plugins.Plugin { return nil }, []int{5900}, nil, true) + plugins.RegisterWithOptions("webtitle", func() plugins.Plugin { return nil }, []int{}, []string{plugins.PluginTypeWeb}, true) +} + +func clearServiceCache() { + serviceCacheMutex.Lock() + serviceCache = make(map[string]*ServiceInfo) + serviceCacheMutex.Unlock() +} + +// ============================================================================= +// 单元测试:CacheServiceInfo / GetCachedServiceInfo +// ============================================================================= + +func TestCacheServiceInfo_BasicCRUD(t *testing.T) { + clearServiceCache() + + t.Run("缓存后可读取", func(t *testing.T) { + CacheServiceInfo("10.0.0.1", 22, &ServiceInfo{Name: "ssh", Version: "OpenSSH_8.9"}) + info, ok := GetCachedServiceInfo("10.0.0.1", 22) + if !ok { + t.Fatal("缓存未命中") + } + if info.Name != "ssh" || info.Version != "OpenSSH_8.9" { + t.Errorf("got Name=%q Version=%q", info.Name, info.Version) + } + }) + + t.Run("不同端口独立", func(t *testing.T) { + CacheServiceInfo("10.0.0.1", 3306, &ServiceInfo{Name: "mysql"}) + CacheServiceInfo("10.0.0.1", 5432, &ServiceInfo{Name: "postgresql"}) + i1, _ := GetCachedServiceInfo("10.0.0.1", 3306) + i2, _ := GetCachedServiceInfo("10.0.0.1", 5432) + if i1.Name != "mysql" || i2.Name != "postgresql" { + t.Errorf("端口混淆: 3306=%q 5432=%q", i1.Name, i2.Name) + } + }) + + t.Run("不同主机独立", func(t *testing.T) { + CacheServiceInfo("10.0.0.1", 22, &ServiceInfo{Name: "ssh"}) + CacheServiceInfo("10.0.0.2", 22, &ServiceInfo{Name: "telnet"}) + i1, _ := GetCachedServiceInfo("10.0.0.1", 22) + i2, _ := GetCachedServiceInfo("10.0.0.2", 22) + if i1.Name != "ssh" || i2.Name != "telnet" { + t.Errorf("主机混淆: .1=%q .2=%q", i1.Name, i2.Name) + } + }) + + t.Run("覆盖写入", func(t *testing.T) { + CacheServiceInfo("10.0.0.5", 80, &ServiceInfo{Name: "unknown"}) + CacheServiceInfo("10.0.0.5", 80, &ServiceInfo{Name: "http"}) + info, _ := GetCachedServiceInfo("10.0.0.5", 80) + if info.Name != "http" { + t.Errorf("覆盖失败: %q", info.Name) + } + }) + + t.Run("未缓存返回 false", func(t *testing.T) { + if _, ok := GetCachedServiceInfo("192.168.99.99", 12345); ok { + t.Error("应返回 false") + } + }) +} + +// ============================================================================= +// 单元测试:Web 服务过滤 +// ============================================================================= + +func TestWebServiceFiltering(t *testing.T) { + clearServiceCache() + + webNames := []string{"http", "https", "ssl", "tls", "nginx", "apache", "iis", "tomcat"} + nonWebNames := []string{"ssh", "mysql", "postgresql", "redis", "mongodb", "ftp", "smtp", "telnet", "vnc", "rdp"} + + for _, name := range webNames { + clearServiceCache() + CacheServiceInfo("10.0.0.1", 443, &ServiceInfo{Name: name}) + if !IsMarkedWebService("10.0.0.1", 443) { + t.Errorf("%q 应被识别为 Web 服务", name) + } + } + + for _, name := range nonWebNames { + clearServiceCache() + CacheServiceInfo("10.0.0.1", 9999, &ServiceInfo{Name: name}) + if IsMarkedWebService("10.0.0.1", 9999) { + t.Errorf("%q 不应被识别为 Web 服务", name) + } + } + + t.Run("GetWebServiceInfo 过滤非 Web", func(t *testing.T) { + clearServiceCache() + CacheServiceInfo("10.0.0.1", 3306, &ServiceInfo{Name: "mysql"}) + if _, ok := GetWebServiceInfo("10.0.0.1", 3306); ok { + t.Error("mysql 不应通过 GetWebServiceInfo") + } + }) + + t.Run("GetWebServiceInfo 返回 Web", func(t *testing.T) { + clearServiceCache() + CacheServiceInfo("10.0.0.1", 8080, &ServiceInfo{Name: "nginx"}) + info, ok := GetWebServiceInfo("10.0.0.1", 8080) + if !ok || info.Name != "nginx" { + t.Error("nginx 应通过 GetWebServiceInfo") + } + }) +} + +// ============================================================================= +// 集成测试:指纹驱动插件匹配 +// ============================================================================= + +func TestIntegration_FingerprintDrivenPluginMatch(t *testing.T) { + clearServiceCache() + registerTestPlugins(t) + + CacheServiceInfo("10.0.0.1", 22, &ServiceInfo{Name: "ssh"}) + CacheServiceInfo("10.0.0.1", 8881, &ServiceInfo{Name: "ssh"}) + CacheServiceInfo("10.0.0.1", 13306, &ServiceInfo{Name: "mysql"}) + CacheServiceInfo("10.0.0.1", 80, &ServiceInfo{Name: "http"}) + CacheServiceInfo("10.0.0.1", 9443, &ServiceInfo{Name: "https"}) + CacheServiceInfo("10.0.0.1", 2121, &ServiceInfo{Name: "ftp"}) + CacheServiceInfo("10.0.0.1", 6380, &ServiceInfo{Name: "redis"}) + + strategy := NewServiceScanStrategy() + + tests := []struct { + plugin, host string + port int + want bool + desc string + }{ + {"ssh", "10.0.0.1", 22, true, "SSH 标准端口"}, + {"ssh", "10.0.0.1", 8881, true, "SSH 非标准端口(指纹匹配)"}, + {"mysql", "10.0.0.1", 13306, true, "MySQL 非标准端口"}, + {"ftp", "10.0.0.1", 2121, true, "FTP 非标准端口"}, + {"redis", "10.0.0.1", 6380, true, "Redis 非标准端口"}, + {"ssh", "10.0.0.1", 13306, false, "SSH 不匹配 MySQL 端口"}, + {"mysql", "10.0.0.1", 8881, false, "MySQL 不匹配 SSH 端口"}, + {"redis", "10.0.0.1", 22, false, "Redis 不匹配 SSH 标准端口"}, + {"ssh", "10.0.0.1", 65000, false, "SSH 不匹配未识别端口"}, + {"webtitle", "10.0.0.1", 80, true, "Web 匹配 http"}, + {"webtitle", "10.0.0.1", 9443, true, "Web 匹配 https 非标准"}, + {"webtitle", "10.0.0.1", 22, false, "Web 不匹配 SSH"}, + } + + for _, tt := range tests { + t.Run(tt.desc, func(t *testing.T) { + got := strategy.isPluginApplicableToPortWithHost(tt.plugin, tt.host, tt.port) + if got != tt.want { + t.Errorf("plugin=%q port=%d: got %v, want %v", tt.plugin, tt.port, got, tt.want) + } + }) + } +} + +// ============================================================================= +// 集成测试:非标准端口完整流程 +// ============================================================================= + +func TestIntegration_NonStandardPortScanFlow(t *testing.T) { + clearServiceCache() + registerTestPlugins(t) + + host := "172.16.0.100" + CacheServiceInfo(host, 8881, &ServiceInfo{ + Name: "ssh", Version: "OpenSSH_8.2p1", + Banner: "SSH-2.0-OpenSSH_8.2p1", Extras: map[string]string{"os": "Linux"}, + }) + + strategy := NewServiceScanStrategy() + + if !strategy.isPluginApplicableToPortWithHost("ssh", host, 8881) { + t.Error("SSH 应匹配 8881") + } + if strategy.isPluginApplicableToPortWithHost("mysql", host, 8881) { + t.Error("MySQL 不应匹配 8881 上的 SSH") + } + if IsMarkedWebService(host, 8881) { + t.Error("SSH 不应标记为 Web") + } +} + +// ============================================================================= +// 集成测试:同一主机多服务 +// ============================================================================= + +func TestIntegration_MultiServiceSameHost(t *testing.T) { + clearServiceCache() + registerTestPlugins(t) + + host := "192.168.1.100" + CacheServiceInfo(host, 2222, &ServiceInfo{Name: "ssh"}) + CacheServiceInfo(host, 33060, &ServiceInfo{Name: "mysql"}) + CacheServiceInfo(host, 8080, &ServiceInfo{Name: "http"}) + CacheServiceInfo(host, 63790, &ServiceInfo{Name: "redis"}) + + strategy := NewServiceScanStrategy() + + checks := []struct { + plugin string + port int + want bool + }{ + {"ssh", 2222, true}, {"ssh", 33060, false}, {"ssh", 8080, false}, + {"mysql", 33060, true}, {"mysql", 2222, false}, + {"redis", 63790, true}, {"redis", 2222, false}, + {"webtitle", 8080, true}, {"webtitle", 2222, false}, + } + + for _, c := range checks { + got := strategy.isPluginApplicableToPortWithHost(c.plugin, host, c.port) + if got != c.want { + t.Errorf("plugin=%q port=%d: got %v, want %v", c.plugin, c.port, got, c.want) + } + } +} + +// ============================================================================= +// 边界测试 +// ============================================================================= + +func TestServiceCache_EdgeCases(t *testing.T) { + clearServiceCache() + registerTestPlugins(t) + strategy := NewServiceScanStrategy() + + t.Run("空服务名不匹配", func(t *testing.T) { + CacheServiceInfo("10.0.0.1", 9999, &ServiceInfo{Name: ""}) + if strategy.isPluginApplicableToPortWithHost("ssh", "10.0.0.1", 9999) { + t.Error("空服务名不应匹配") + } + }) + + t.Run("unknown 不匹配", func(t *testing.T) { + CacheServiceInfo("10.0.0.1", 8888, &ServiceInfo{Name: "unknown"}) + if strategy.isPluginApplicableToPortWithHost("ssh", "10.0.0.1", 8888) { + t.Error("unknown 不应匹配") + } + }) + + t.Run("大小写不敏感", func(t *testing.T) { + clearServiceCache() + CacheServiceInfo("10.0.0.1", 5555, &ServiceInfo{Name: "SSH"}) + if !strategy.isPluginApplicableToPortWithHost("ssh", "10.0.0.1", 5555) { + t.Error("SSH 大写应匹配 ssh 插件") + } + }) + + t.Run("host 为空不查缓存", func(t *testing.T) { + CacheServiceInfo("10.0.0.1", 8881, &ServiceInfo{Name: "ssh"}) + if strategy.isPluginApplicableToPortWithHost("ssh", "", 8881) { + t.Error("host 为空不应匹配") + } + }) + + t.Run("nil ServiceInfo 不 panic", func(t *testing.T) { + CacheServiceInfo("10.0.0.1", 7777, nil) + got := strategy.isPluginApplicableToPortWithHost("ssh", "10.0.0.1", 7777) + if got { + t.Error("nil ServiceInfo 不应匹配") + } + }) + + t.Run("IPv6", func(t *testing.T) { + clearServiceCache() + CacheServiceInfo("::1", 22, &ServiceInfo{Name: "ssh"}) + if _, ok := GetCachedServiceInfo("::1", 22); !ok { + t.Error("IPv6 缓存失败") + } + }) +} + +// ============================================================================= +// 并发安全 +// ============================================================================= + +func TestServiceCache_ConcurrentSafety(t *testing.T) { + clearServiceCache() + var wg sync.WaitGroup + + for i := 0; i < 100; i++ { + wg.Add(3) + go func(p int) { defer wg.Done(); CacheServiceInfo("10.0.0.1", p, &ServiceInfo{Name: "ssh"}) }(i) + go func(p int) { defer wg.Done(); GetCachedServiceInfo("10.0.0.1", p) }(i) + go func(p int) { defer wg.Done(); IsMarkedWebService("10.0.0.1", p) }(i) + } + wg.Wait() + + for i := 0; i < 100; i++ { + if _, ok := GetCachedServiceInfo("10.0.0.1", i); !ok { + t.Errorf("并发写入丢失: port=%d", i) + } + } +} + +// ============================================================================= +// 回归测试:#588 +// ============================================================================= + +func TestRegression_Issue588(t *testing.T) { + clearServiceCache() + registerTestPlugins(t) + + CacheServiceInfo("192.168.1.50", 8881, &ServiceInfo{Name: "ssh", Version: "OpenSSH_7.4"}) + strategy := NewServiceScanStrategy() + + if !strategy.isPluginApplicableToPortWithHost("ssh", "192.168.1.50", 8881) { + t.Fatal("#588: SSH 应匹配 8881") + } + for _, p := range []string{"mysql", "ftp", "redis", "postgresql", "telnet", "vnc", "mssql"} { + if strategy.isPluginApplicableToPortWithHost(p, "192.168.1.50", 8881) { + t.Errorf("#588: %q 不应匹配 8881 上的 SSH", p) + } + } +} + +// ============================================================================= +// 端口匹配优先于缓存 +// ============================================================================= + +func TestIntegration_PortMatchPrecedence(t *testing.T) { + clearServiceCache() + registerTestPlugins(t) + + CacheServiceInfo("10.0.0.1", 22, &ServiceInfo{Name: "http"}) + strategy := NewServiceScanStrategy() + + if !strategy.isPluginApplicableToPortWithHost("ssh", "10.0.0.1", 22) { + t.Error("SSH 应通过端口匹配命中 22(即使缓存是 http)") + } + if !IsMarkedWebService("10.0.0.1", 22) { + t.Error("缓存是 http,应标记为 Web") + } +} diff --git a/core/service_scanner_test.go b/core/service_scanner_test.go index 13c15c6..c1cd949 100644 --- a/core/service_scanner_test.go +++ b/core/service_scanner_test.go @@ -65,6 +65,16 @@ func TestParsePortList_BasicParsing(t *testing.T) { input: "22,80,443,3306", expected: []int{22, 80, 443, 3306}, }, + { + name: "端口范围", + input: "80-82", + expected: []int{80, 81, 82}, + }, + { + name: "端口和范围混合", + input: "22,80-81", + expected: []int{22, 80, 81}, + }, { name: "空字符串", input: "", @@ -281,7 +291,7 @@ func TestParsePortList_ProductionScenarios(t *testing.T) { t.Run("数据库端口", func(t *testing.T) { input := "3306,5432,1433,27017" - expected := []int{3306, 5432, 1433, 27017} + expected := []int{1433, 3306, 5432, 27017} result := s.parsePortList(input) if !intSlicesEqual(result, expected) { t.Errorf("应该正确解析常见数据库端口") @@ -317,6 +327,13 @@ func TestParsePortList_ProductionScenarios(t *testing.T) { t.Errorf("应该正确解析高端口号") } }) + + t.Run("端口组", func(t *testing.T) { + result := s.parsePortList("web") + if !sliceContains(result, 80) || !sliceContains(result, 443) { + t.Errorf("web端口组应该包含80和443, 实际 %v", result) + } + }) } // TestParsePortList_ReturnValue 测试返回值特性 @@ -330,14 +347,11 @@ func TestParsePortList_ReturnValue(t *testing.T) { } }) - t.Run("端口不重复-但不保证去重", func(t *testing.T) { - // 注意:当前实现不去重,如果用户输入 "22,22",会返回 [22, 22] - // 这是可以接受的,因为上层逻辑会处理重复 + t.Run("重复端口会去重", func(t *testing.T) { input := "22,22" result := s.parsePortList(input) - // 这里我们只测试解析是否正确,不测试去重 - if len(result) != 2 || result[0] != 22 || result[1] != 22 { - t.Errorf("当前实现不去重,应该返回两个22") + if len(result) != 1 || result[0] != 22 { + t.Errorf("重复端口应该去重, 实际 %v", result) } }) } @@ -355,6 +369,15 @@ func intSlicesEqual(a, b []int) bool { return true } +func sliceContains(values []int, target int) bool { + for _, value := range values { + if value == target { + return true + } + } + return false +} + // ============================================================================= // 存活检测判断测试 // ============================================================================= diff --git a/core/tuning_cli_integration_test.go b/core/tuning_cli_integration_test.go new file mode 100644 index 0000000..f1c1611 --- /dev/null +++ b/core/tuning_cli_integration_test.go @@ -0,0 +1,82 @@ +package core + +import ( + "flag" + "os" + "testing" + "time" + + "github.com/shadow1ng/fscan/common" +) + +func TestCLIExplicitDefaultTuningFlagsSurviveTuneConfig(t *testing.T) { + oldArgs := os.Args + oldFlagSet := flag.CommandLine + oldFlagVars := *common.GetFlagVars() + defer func() { + os.Args = oldArgs + flag.CommandLine = oldFlagSet + *common.GetFlagVars() = oldFlagVars + }() + + *common.GetFlagVars() = common.FlagVars{} + flag.CommandLine = flag.NewFlagSet("fscan-test", flag.ContinueOnError) + os.Args = []string{ + "fscan-test", + "-silent", + "-h", "127.0.0.1", + "-time", "3", + "-mt", "20", + "-retry", "3", + "-icmp-rate", "0.1", + "-num", "20", + } + + info := &common.HostInfo{} + if err := common.Flag(info); err != nil { + t.Fatalf("Flag error = %v", err) + } + + cfg, _, err := common.BuildConfig(common.GetFlagVars(), info) + if err != nil { + t.Fatalf("BuildConfig error = %v", err) + } + if !cfg.TimeoutExplicit || !cfg.ModuleThreadNumExplicit || + !cfg.MaxRetriesExplicit || !cfg.Network.ICMPRateExplicit || + !cfg.POC.NumExplicit { + t.Fatalf("explicit flags not propagated: timeout=%v mt=%v retry=%v icmp=%v num=%v", + cfg.TimeoutExplicit, + cfg.ModuleThreadNumExplicit, + cfg.MaxRetriesExplicit, + cfg.Network.ICMPRateExplicit, + cfg.POC.NumExplicit) + } + + ep := &EnvironmentProfile{ + Net: NetworkProfile{ + Env: EnvLAN, + RTTMedian: time.Millisecond, + RTTStddev: 200 * time.Microsecond, + LossRate: 0, + Samples: 30, + }, + System: SystemProfile{FDLimit: 65536, NumCPU: 8}, + } + ep.TuneConfig(cfg, makeTestSession(cfg)) + + if cfg.Timeout != 3*time.Second { + t.Fatalf("Timeout = %v, want explicit default 3s", cfg.Timeout) + } + if cfg.ModuleThreadNum != 20 { + t.Fatalf("ModuleThreadNum = %d, want explicit default 20", cfg.ModuleThreadNum) + } + if cfg.MaxRetries != 3 { + t.Fatalf("MaxRetries = %d, want explicit default 3", cfg.MaxRetries) + } + if cfg.Network.ICMPRate != 0.1 { + t.Fatalf("ICMPRate = %.2f, want explicit default 0.10", cfg.Network.ICMPRate) + } + if cfg.POC.Num != 20 { + t.Fatalf("POC.Num = %d, want explicit default 20", cfg.POC.Num) + } +} diff --git a/core/web_scanner.go b/core/web_scanner.go index 3522802..e48d5bb 100644 --- a/core/web_scanner.go +++ b/core/web_scanner.go @@ -204,10 +204,11 @@ func (w *WebPortDetector) tryHTTP(ctx context.Context, client *http.Client, sess // 基于服务指纹的Web服务识别 // =============================== -// Web服务缓存 - 简化的全局缓存 +// 服务识别缓存 - 存储所有识别到的服务(不仅限于 Web) +// 端口扫描阶段写入,插件匹配阶段读取 var ( - webServiceCache = make(map[string]*ServiceInfo) - webCacheMutex sync.RWMutex + serviceCache = make(map[string]*ServiceInfo) + serviceCacheMutex sync.RWMutex ) // IsWebServiceByFingerprint 基于服务指纹判断Web服务 - 保持API兼容 @@ -259,28 +260,45 @@ func IsWebServiceByFingerprint(serviceInfo *ServiceInfo) bool { return false } -// MarkAsWebService 标记Web服务 - 保持API兼容 -func MarkAsWebService(host string, port int, serviceInfo *ServiceInfo) { +// CacheServiceInfo 缓存识别到的服务信息 +func CacheServiceInfo(host string, port int, serviceInfo *ServiceInfo) { cacheKey := net.JoinHostPort(host, strconv.Itoa(port)) - webCacheMutex.Lock() - defer webCacheMutex.Unlock() + serviceCacheMutex.Lock() + defer serviceCacheMutex.Unlock() - webServiceCache[cacheKey] = serviceInfo + serviceCache[cacheKey] = serviceInfo } -// GetWebServiceInfo 获取Web服务信息 -func GetWebServiceInfo(host string, port int) (*ServiceInfo, bool) { +// MarkAsWebService 标记 Web 服务(兼容旧调用) +func MarkAsWebService(host string, port int, serviceInfo *ServiceInfo) { + CacheServiceInfo(host, port, serviceInfo) +} + +// GetCachedServiceInfo 获取缓存的服务信息 +func GetCachedServiceInfo(host string, port int) (*ServiceInfo, bool) { cacheKey := net.JoinHostPort(host, strconv.Itoa(port)) - webCacheMutex.RLock() - defer webCacheMutex.RUnlock() + serviceCacheMutex.RLock() + defer serviceCacheMutex.RUnlock() - serviceInfo, exists := webServiceCache[cacheKey] + serviceInfo, exists := serviceCache[cacheKey] return serviceInfo, exists } -// IsMarkedWebService 检查是否已标记为Web服务 +// GetWebServiceInfo 获取 Web 服务信息(兼容旧调用) +func GetWebServiceInfo(host string, port int) (*ServiceInfo, bool) { + info, exists := GetCachedServiceInfo(host, port) + if !exists { + return nil, false + } + if !IsWebServiceByFingerprint(info) { + return nil, false + } + return info, true +} + +// IsMarkedWebService 检查是否为 Web 服务 func IsMarkedWebService(host string, port int) bool { _, exists := GetWebServiceInfo(host, port) return exists diff --git a/core/web_scanner_test.go b/core/web_scanner_test.go index 54c52c9..3745269 100644 --- a/core/web_scanner_test.go +++ b/core/web_scanner_test.go @@ -440,9 +440,9 @@ func TestCreateTargetFromURL(t *testing.T) { // TestWebServiceCache 测试Web服务缓存操作 func TestWebServiceCache(t *testing.T) { // 清空缓存 - webCacheMutex.Lock() - webServiceCache = make(map[string]*ServiceInfo) - webCacheMutex.Unlock() + serviceCacheMutex.Lock() + serviceCache = make(map[string]*ServiceInfo) + serviceCacheMutex.Unlock() t.Run("存储和读取", func(t *testing.T) { serviceInfo := &ServiceInfo{ @@ -517,9 +517,9 @@ func TestWebServiceCache(t *testing.T) { // TestWebServiceCache_Concurrent 测试并发安全性 func TestWebServiceCache_Concurrent(t *testing.T) { // 清空缓存 - webCacheMutex.Lock() - webServiceCache = make(map[string]*ServiceInfo) - webCacheMutex.Unlock() + serviceCacheMutex.Lock() + serviceCache = make(map[string]*ServiceInfo) + serviceCacheMutex.Unlock() t.Run("不同key并发写入", func(t *testing.T) { var wg sync.WaitGroup diff --git a/libs/grdp/login/screen.go b/libs/grdp/login/screen.go index a994bc8..aa2c44e 100644 --- a/libs/grdp/login/screen.go +++ b/libs/grdp/login/screen.go @@ -176,8 +176,7 @@ func (g *Client) ProbeOSInfo(host, domain, user, pwd string, timeout int64, rdpP exitFlag := make(chan bool, 1) info = make(map[string]any) - targetSlice := strings.Split(g.Host, ":") - ip := targetSlice[0] + ip := rdpTargetHost(g.Host) conn, err := WrapperTcpWithTimeout("tcp", g.Host, time.Duration(timeout)*time.Second) if err != nil { return @@ -273,3 +272,14 @@ loop: glog.Debug("loop ended, elapsed time: ", time.Since(start)) return info } + +func rdpTargetHost(target string) string { + host, _, err := net.SplitHostPort(target) + if err == nil { + return host + } + if strings.Count(target, ":") == 1 { + return strings.SplitN(target, ":", 2)[0] + } + return target +} diff --git a/libs/grdp/login/screen_test.go b/libs/grdp/login/screen_test.go new file mode 100644 index 0000000..4a38460 --- /dev/null +++ b/libs/grdp/login/screen_test.go @@ -0,0 +1,24 @@ +package login + +import "testing" + +func TestRDPTargetHost(t *testing.T) { + tests := []struct { + name string + target string + want string + }{ + {name: "ipv4 with port", target: "192.168.1.1:3389", want: "192.168.1.1"}, + {name: "hostname with port", target: "rdp.example.com:3389", want: "rdp.example.com"}, + {name: "bracketed ipv6 with port", target: "[2001:db8::1]:3389", want: "2001:db8::1"}, + {name: "bare ipv6 without port", target: "2001:db8::1", want: "2001:db8::1"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := rdpTargetHost(tt.target); got != tt.want { + t.Fatalf("rdpTargetHost(%q) = %q, want %q", tt.target, got, tt.want) + } + }) + } +} diff --git a/plugins/init_test.go b/plugins/init_test.go index 0ac8fd9..f2fdf8d 100644 --- a/plugins/init_test.go +++ b/plugins/init_test.go @@ -192,9 +192,9 @@ func TestGenerateCredentials_PlaceholderReplacement(t *testing.T) { // 验证:{user} 被正确替换 expectedCombos := map[string]string{ - "root:root": "root", // {user} → root - "root:root123": "root", // {user}123 → root123 - "mysql:mysql": "mysql", // {user} → mysql + "root:root": "root", // {user} → root + "root:root123": "root", // {user}123 → root123 + "mysql:mysql": "mysql", // {user} → mysql "mysql:mysql123": "mysql", // {user}123 → mysql123 } @@ -244,7 +244,7 @@ func TestGenerateCredentials_DefaultValues(t *testing.T) { cfg.Credentials.UserPassPairs = []config.CredentialPair{} cfg.Credentials.Userdict = map[string][]string{} // 空字典 - cfg.Credentials.Passwords = []string{} // 空密码列表 + cfg.Credentials.Passwords = []string{} // 空密码列表 result := GenerateCredentials("unknown_service", cfg) @@ -327,3 +327,27 @@ func TestGenerateCredentials_EmptyUserPassPairs(t *testing.T) { t.Logf("✓ 空 UserPassPairs 正确回退到笛卡尔积") } + +func TestBuildConfigAdditionalPasswordsAreNotShadowedByExactPair(t *testing.T) { + cfg, _, err := common.BuildConfig(&common.FlagVars{ + Username: "root", + Password: "primary", + AddPasswords: "extra", + }, &common.HostInfo{}) + if err != nil { + t.Fatalf("BuildConfig error = %v", err) + } + + result := GenerateCredentials("ssh", cfg) + found := map[string]bool{} + for _, cred := range result { + found[cred.Username+":"+cred.Password] = true + } + + if !found["root:primary"] { + t.Fatal("missing primary password credential") + } + if !found["root:extra"] { + t.Fatal("additional password was shadowed by exact user/password pair") + } +} diff --git a/plugins/local/systeminfo_dc_url.go b/plugins/local/systeminfo_dc_url.go new file mode 100644 index 0000000..97acb38 --- /dev/null +++ b/plugins/local/systeminfo_dc_url.go @@ -0,0 +1,11 @@ +package local + +import ( + "fmt" + "net" + "strconv" +) + +func ldapURL(host string, port int) string { + return fmt.Sprintf("ldap://%s", net.JoinHostPort(host, strconv.Itoa(port))) +} diff --git a/plugins/local/systeminfo_dc_url_test.go b/plugins/local/systeminfo_dc_url_test.go new file mode 100644 index 0000000..089065d --- /dev/null +++ b/plugins/local/systeminfo_dc_url_test.go @@ -0,0 +1,24 @@ +package local + +import "testing" + +func TestLDAPURLUsesJoinHostPort(t *testing.T) { + tests := []struct { + name string + host string + port int + want string + }{ + {name: "hostname", host: "dc.example.local", port: 389, want: "ldap://dc.example.local:389"}, + {name: "ipv4", host: "192.168.1.10", port: 389, want: "ldap://192.168.1.10:389"}, + {name: "ipv6", host: "2001:db8::10", port: 389, want: "ldap://[2001:db8::10]:389"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := ldapURL(tt.host, tt.port); got != tt.want { + t.Fatalf("ldapURL(%q, %d) = %q, want %q", tt.host, tt.port, got, tt.want) + } + }) + } +} diff --git a/plugins/local/systeminfo_dc_windows.go b/plugins/local/systeminfo_dc_windows.go index 8dcb55b..9326cf7 100644 --- a/plugins/local/systeminfo_dc_windows.go +++ b/plugins/local/systeminfo_dc_windows.go @@ -82,10 +82,10 @@ func (p *SystemInfoPlugin) connectToDomain(domain string) (*domainInfo, error) { } defer func() { _ = client.Close() }() - conn, err := ldap.DialURL(fmt.Sprintf("ldap://%s:389", dcHost)) + conn, err := ldap.DialURL(ldapURL(dcHost, 389)) if err != nil { if ipv4, resolveErr := resolveIPv4(dcHost); resolveErr == nil { - conn, err = ldap.DialURL(fmt.Sprintf("ldap://%s:389", ipv4)) + conn, err = ldap.DialURL(ldapURL(ipv4, 389)) } if err != nil { return nil, fmt.Errorf("LDAP dial: %w", err) diff --git a/plugins/services/cassandra.go b/plugins/services/cassandra.go index fa083fa..b17e646 100644 --- a/plugins/services/cassandra.go +++ b/plugins/services/cassandra.go @@ -8,6 +8,7 @@ import ( "fmt" "io" "net" + "sync/atomic" "time" "github.com/shadow1ng/fscan/common" @@ -165,20 +166,19 @@ func (p *CassandraPlugin) doCassandraAuth(ctx context.Context, info *common.Host // ── CQL wire protocol 工具 ────────────────────────────────────── -var cqlStreamID int16 +var cqlStreamID uint32 + +func nextCQLStreamID() uint16 { + return uint16((atomic.AddUint32(&cqlStreamID, 1) - 1) & 0x7fff) +} func cqlSend(conn net.Conn, opcode byte, body []byte) error { - id := cqlStreamID - if cqlStreamID == 32767 { - cqlStreamID = 0 - } else { - cqlStreamID++ - } + id := nextCQLStreamID() // frame: [1B version|flags] [2B stream] [1B opcode] [4B length] [body] header := make([]byte, 8) header[0] = cqlVersion - binary.BigEndian.PutUint16(header[1:3], uint16(id)) + binary.BigEndian.PutUint16(header[1:3], id) header[3] = opcode binary.BigEndian.PutUint32(header[4:8], uint32(len(body))) diff --git a/plugins/services/credential_tester.go b/plugins/services/credential_tester.go index 37fed92..f377e67 100644 --- a/plugins/services/credential_tester.go +++ b/plugins/services/credential_tester.go @@ -7,6 +7,7 @@ import ( "io" "net" "sync" + "sync/atomic" "time" "github.com/shadow1ng/fscan/common" @@ -61,7 +62,11 @@ type AuthFunc func(ctx context.Context, cred Credential) *AuthResult // ErrorClassifier 错误分类函数 type ErrorClassifier func(err error) ErrorType -var authCleanupWait = 2 * time.Second +var authCleanupWaitNanos int64 = int64(2 * time.Second) + +func authCleanupWait() time.Duration { + return time.Duration(atomic.LoadInt64(&authCleanupWaitNanos)) +} // ============================================================================= // 单凭据测试(解决 goroutine 泄漏) @@ -70,9 +75,36 @@ var authCleanupWait = 2 * time.Second // TestSingleCredential 安全地测试单个凭据 // 正确处理 context 取消时的资源清理 func TestSingleCredential(ctx context.Context, cred Credential, authFn AuthFunc) *AuthResult { + if ctx == nil { + ctx = context.Background() + } + if authFn == nil { + return &AuthResult{ + Success: false, + ErrorType: ErrorTypeUnknown, + Error: fmt.Errorf("auth function is nil"), + } + } + if err := ctx.Err(); err != nil { + return &AuthResult{ + Success: false, + ErrorType: ErrorTypeNetwork, + Error: err, + } + } + resultChan := make(chan *AuthResult, 1) go func() { + defer func() { + if r := recover(); r != nil { + resultChan <- &AuthResult{ + Success: false, + ErrorType: ErrorTypeUnknown, + Error: fmt.Errorf("auth function panic: %v", r), + } + } + }() result := authFn(ctx, cred) resultChan <- result }() @@ -83,7 +115,7 @@ func TestSingleCredential(ctx context.Context, cred Credential, authFn AuthFunc) case <-ctx.Done(): // context 被取消后只做有界等待,避免 authFn 卡死时清理 goroutine 也永久泄漏。 go func() { - timer := time.NewTimer(authCleanupWait) + timer := time.NewTimer(authCleanupWait()) defer timer.Stop() select { @@ -116,15 +148,35 @@ type ConcurrentTestConfig struct { UseProxy bool // 代理模式下跳过直连 TCP 预检 } +func normalizeConcurrentTestConfig(testConfig ConcurrentTestConfig) ConcurrentTestConfig { + if testConfig.Concurrency <= 0 { + testConfig.Concurrency = 10 + } + if testConfig.MaxRetries <= 0 { + testConfig.MaxRetries = 3 + } + if testConfig.RetryDelay <= 0 { + testConfig.RetryDelay = time.Second + } + if testConfig.MaxConsecutiveNetErrors <= 0 { + testConfig.MaxConsecutiveNetErrors = 5 + } + return testConfig +} + // DefaultConcurrentTestConfig 默认配置 func DefaultConcurrentTestConfig(config *common.Config) ConcurrentTestConfig { concurrency := config.ModuleThreadNum if concurrency <= 0 { concurrency = 10 } + maxRetries := config.MaxRetries + if maxRetries <= 0 { + maxRetries = 3 + } return ConcurrentTestConfig{ Concurrency: concurrency, - MaxRetries: 3, + MaxRetries: maxRetries, RetryDelay: time.Second, MaxConsecutiveNetErrors: 5, UseProxy: config.Network.Socks5Proxy != "" || config.Network.HTTPProxy != "", @@ -147,6 +199,9 @@ func TestCredentialsConcurrently( serviceName string, testConfig ConcurrentTestConfig, ) *ScanResult { + if ctx == nil { + ctx = context.Background() + } if len(credentials) == 0 { return &ScanResult{ Success: false, @@ -154,11 +209,16 @@ func TestCredentialsConcurrently( Error: fmt.Errorf("%s", i18n.GetText("service_no_test_creds")), } } + testConfig = normalizeConcurrentTestConfig(testConfig) // TCP 预检:快速验证目标可达,避免对不可达目标浪费全部凭据尝试 // 代理模式下跳过:net.DialTimeout 直连无法到达代理后的内网目标 if testConfig.TargetAddr != "" && !testConfig.UseProxy { - preConn, err := net.DialTimeout("tcp", testConfig.TargetAddr, 3*time.Second) + dialCtx, dialCancel := context.WithTimeout(ctx, 3*time.Second) + defer dialCancel() + + var dialer net.Dialer + preConn, err := dialer.DialContext(dialCtx, "tcp", testConfig.TargetAddr) if err != nil { return &ScanResult{ Success: false, @@ -240,10 +300,6 @@ func workerTestCredentials( testConfig ConcurrentTestConfig, ) { consecutiveNetErrors := 0 - maxNetErrors := testConfig.MaxConsecutiveNetErrors - if maxNetErrors <= 0 { - maxNetErrors = 5 - } for cred := range credChan { // 检查是否应该停止 @@ -254,7 +310,7 @@ func workerTestCredentials( } // 连续网络错误达到阈值,目标可能不可达,提前退出 - if consecutiveNetErrors >= maxNetErrors { + if consecutiveNetErrors >= testConfig.MaxConsecutiveNetErrors { return } @@ -292,10 +348,18 @@ func testCredentialWithRetry( // 测试凭据 result := TestSingleCredential(ctx, cred, authFn) + if result == nil { + result = &AuthResult{ + Success: false, + ErrorType: ErrorTypeUnknown, + Error: fmt.Errorf("auth function returned nil result"), + } + } - if result.Success && result.Conn != nil { - // 成功,关闭连接并返回 - _ = result.Conn.Close() + if result.Success { + if result.Conn != nil { + _ = result.Conn.Close() + } return &ScanResult{ Type: plugins.ResultTypeCredential, Success: true, diff --git a/plugins/services/credential_tester_test.go b/plugins/services/credential_tester_test.go index 1e3511e..8827928 100644 --- a/plugins/services/credential_tester_test.go +++ b/plugins/services/credential_tester_test.go @@ -8,6 +8,8 @@ import ( "sync/atomic" "testing" "time" + + "github.com/shadow1ng/fscan/common" ) /* @@ -187,6 +189,84 @@ func TestMatchIgnoreCase(t *testing.T) { // 并发测试 // ============================================================================= +func setAuthCleanupWaitForTest(wait time.Duration) func() { + oldWait := atomic.LoadInt64(&authCleanupWaitNanos) + atomic.StoreInt64(&authCleanupWaitNanos, int64(wait)) + return func() { atomic.StoreInt64(&authCleanupWaitNanos, oldWait) } +} + +func TestDefaultConcurrentTestConfigUsesConfigRetries(t *testing.T) { + cfg := DefaultConcurrentTestConfig(&common.Config{ + ModuleThreadNum: 7, + MaxRetries: 5, + }) + + if cfg.Concurrency != 7 { + t.Fatalf("Concurrency = %d, want 7", cfg.Concurrency) + } + if cfg.MaxRetries != 5 { + t.Fatalf("MaxRetries = %d, want config MaxRetries 5", cfg.MaxRetries) + } +} + +func TestDefaultConcurrentTestConfigRetriesFallback(t *testing.T) { + cfg := DefaultConcurrentTestConfig(&common.Config{ + ModuleThreadNum: 0, + MaxRetries: 0, + }) + + if cfg.Concurrency != 10 { + t.Fatalf("Concurrency = %d, want fallback 10", cfg.Concurrency) + } + if cfg.MaxRetries != 3 { + t.Fatalf("MaxRetries = %d, want fallback 3", cfg.MaxRetries) + } +} + +func TestTestCredentialsConcurrently_ZeroValueConfigStillRuns(t *testing.T) { + var calls atomic.Int32 + authFn := func(ctx context.Context, cred Credential) *AuthResult { + calls.Add(1) + return &AuthResult{Success: true} + } + + result := TestCredentialsConcurrently(context.Background(), []Credential{{Username: "u", Password: "p"}}, authFn, "test", ConcurrentTestConfig{}) + if !result.Success { + t.Fatalf("zero-value config should still test credentials: %v", result.Error) + } + if calls.Load() != 1 { + t.Fatalf("authFn calls = %d, want 1", calls.Load()) + } +} + +func TestTestCredentialsConcurrently_PrecheckHonorsCanceledContext(t *testing.T) { + var calls atomic.Int32 + authFn := func(ctx context.Context, cred Credential) *AuthResult { + calls.Add(1) + return &AuthResult{Success: false} + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + start := time.Now() + result := TestCredentialsConcurrently(ctx, []Credential{{Username: "u", Password: "p"}}, authFn, "test", ConcurrentTestConfig{ + Concurrency: 1, + MaxRetries: 1, + TargetAddr: "203.0.113.1:65000", + }) + + if result.Success { + t.Fatal("canceled context should not return success") + } + if calls.Load() != 0 { + t.Fatalf("authFn calls = %d, want 0 when precheck context is canceled", calls.Load()) + } + if elapsed := time.Since(start); elapsed > 200*time.Millisecond { + t.Fatalf("precheck ignored canceled context, elapsed=%v", elapsed) + } +} + // mockConn 模拟连接 type mockConn struct { closed atomic.Bool @@ -336,6 +416,60 @@ func TestTestCredentialsConcurrently_ContextCancel(t *testing.T) { } } +func TestTestCredentialsConcurrently_CancelWithStuckAuthReturnsPromptly(t *testing.T) { + defer setAuthCleanupWaitForTest(20 * time.Millisecond)() + + credentials := make([]Credential, 10) + for i := range credentials { + credentials[i] = Credential{Username: "user", Password: "pass"} + } + + authStarted := make(chan struct{}, len(credentials)) + releaseAuth := make(chan struct{}) + authFn := func(ctx context.Context, cred Credential) *AuthResult { + authStarted <- struct{}{} + <-releaseAuth + return &AuthResult{Success: false, ErrorType: ErrorTypeNetwork} + } + + config := ConcurrentTestConfig{ + Concurrency: 3, + MaxRetries: 1, + RetryDelay: time.Millisecond, + } + + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan *ScanResult, 1) + go func() { + done <- TestCredentialsConcurrently(ctx, credentials, authFn, "test", config) + }() + + for i := 0; i < config.Concurrency; i++ { + select { + case <-authStarted: + case <-time.After(time.Second): + close(releaseAuth) + t.Fatalf("authFn started %d workers, want %d", i, config.Concurrency) + } + } + + start := time.Now() + cancel() + select { + case result := <-done: + close(releaseAuth) + if result.Success { + t.Fatal("context取消后不应该返回成功") + } + if elapsed := time.Since(start); elapsed > 200*time.Millisecond { + t.Fatalf("取消后返回过慢: %v", elapsed) + } + case <-time.After(time.Second): + close(releaseAuth) + t.Fatal("authFn 卡住时并发测试没有及时返回") + } +} + // ============================================================================= // 单凭据测试 // ============================================================================= @@ -361,6 +495,49 @@ func TestTestSingleCredential_Success(t *testing.T) { } } +func TestTestSingleCredential_CanceledContextSkipsAuth(t *testing.T) { + var calls atomic.Int32 + authFn := func(ctx context.Context, cred Credential) *AuthResult { + calls.Add(1) + return &AuthResult{Success: true} + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + result := TestSingleCredential(ctx, Credential{Username: "admin", Password: "admin"}, authFn) + if result.Success { + t.Fatal("canceled context should not return success") + } + if calls.Load() != 0 { + t.Fatalf("authFn calls = %d, want 0", calls.Load()) + } +} + +func TestTestSingleCredential_NilAuthFunc(t *testing.T) { + result := TestSingleCredential(context.Background(), Credential{Username: "admin", Password: "admin"}, nil) + if result.Success { + t.Fatal("nil authFn should not return success") + } + if result.Error == nil { + t.Fatal("nil authFn should return an error") + } +} + +func TestTestSingleCredential_RecoverAuthPanic(t *testing.T) { + authFn := func(ctx context.Context, cred Credential) *AuthResult { + panic("boom") + } + + result := TestSingleCredential(context.Background(), Credential{Username: "admin", Password: "admin"}, authFn) + if result.Success { + t.Fatal("panic authFn should not return success") + } + if result.Error == nil { + t.Fatal("panic authFn should return an error") + } +} + // TestTestSingleCredential_ContextCancel 测试context取消时的资源清理 func TestTestSingleCredential_ContextCancel(t *testing.T) { conn := &mockConn{} @@ -403,9 +580,7 @@ func TestTestSingleCredential_ContextCancel(t *testing.T) { } func TestTestSingleCredential_ContextCancelCleanupIsBounded(t *testing.T) { - oldWait := authCleanupWait - authCleanupWait = 20 * time.Millisecond - defer func() { authCleanupWait = oldWait }() + defer setAuthCleanupWaitForTest(20 * time.Millisecond)() authStarted := make(chan struct{}) releaseAuth := make(chan struct{}) @@ -480,6 +655,45 @@ func TestRetryLogic_NetworkErrorRetries(t *testing.T) { } } +func TestRetryLogic_SuccessWithoutConn(t *testing.T) { + var attempts atomic.Int32 + authFn := func(ctx context.Context, cred Credential) *AuthResult { + attempts.Add(1) + return &AuthResult{Success: true} + } + + result := TestCredentialsConcurrently(context.Background(), []Credential{{Username: "admin", Password: "admin"}}, authFn, "test", ConcurrentTestConfig{ + Concurrency: 1, + MaxRetries: 3, + }) + if !result.Success { + t.Fatalf("success result without Conn should be accepted: %v", result.Error) + } + if attempts.Load() != 1 { + t.Fatalf("attempts = %d, want 1", attempts.Load()) + } +} + +func TestRetryLogic_NilAuthResultDoesNotPanic(t *testing.T) { + var attempts atomic.Int32 + authFn := func(ctx context.Context, cred Credential) *AuthResult { + attempts.Add(1) + return nil + } + + result := TestCredentialsConcurrently(context.Background(), []Credential{{Username: "admin", Password: "admin"}}, authFn, "test", ConcurrentTestConfig{ + Concurrency: 1, + MaxRetries: 2, + RetryDelay: time.Millisecond, + }) + if result.Success { + t.Fatal("nil auth result should not return success") + } + if attempts.Load() != 2 { + t.Fatalf("attempts = %d, want 2", attempts.Load()) + } +} + // TestRetryLogic_AuthErrorNoRetry 认证错误不应该重试 func TestRetryLogic_AuthErrorNoRetry(t *testing.T) { var attempts atomic.Int32 diff --git a/plugins/services/kafka.go b/plugins/services/kafka.go index 8137e2f..edc85ab 100644 --- a/plugins/services/kafka.go +++ b/plugins/services/kafka.go @@ -8,6 +8,7 @@ import ( "fmt" "io" "net" + "sync/atomic" "time" "github.com/shadow1ng/fscan/common" @@ -154,9 +155,12 @@ func (p *KafkaPlugin) doKafkaAuth(ctx context.Context, info *common.HostInfo, cr var kafkaCorrelationID int32 +func nextKafkaCorrelationID() int32 { + return atomic.AddInt32(&kafkaCorrelationID, 1) - 1 +} + func kafkaSend(conn net.Conn, apiKey, apiVersion int16, body []byte) error { - corrID := kafkaCorrelationID - kafkaCorrelationID++ + corrID := nextKafkaCorrelationID() // 请求格式: [4B len] [2B api_key] [2B api_version] [4B corr_id] [2B client_id_len] [client_id] [body] clientID := "fscan" diff --git a/plugins/services/mongodb.go b/plugins/services/mongodb.go index 0911a25..a7f794e 100644 --- a/plugins/services/mongodb.go +++ b/plugins/services/mongodb.go @@ -11,6 +11,7 @@ import ( "io" "net" "strings" + "sync/atomic" "time" "github.com/shadow1ng/fscan/common" @@ -156,8 +157,7 @@ const ( var mongoRequestID uint32 func nextRequestID() uint32 { - mongoRequestID++ - return mongoRequestID + return atomic.AddUint32(&mongoRequestID, 1) } // buildMongoCommand 构建 MongoDB 命令的 OP_MSG body (最小 BSON 实现) diff --git a/plugins/services/protocol_ids_test.go b/plugins/services/protocol_ids_test.go new file mode 100644 index 0000000..f344c76 --- /dev/null +++ b/plugins/services/protocol_ids_test.go @@ -0,0 +1,46 @@ +package services + +import ( + "sync" + "testing" +) + +func TestProtocolIDsAreConcurrentSafe(t *testing.T) { + const workers = 64 + const perWorker = 64 + + tests := []struct { + name string + next func() uint32 + }{ + {"mongodb", nextRequestID}, + {"kafka", func() uint32 { return uint32(nextKafkaCorrelationID()) }}, + {"cassandra", func() uint32 { return uint32(nextCQLStreamID()) }}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var wg sync.WaitGroup + values := make(chan uint32, workers*perWorker) + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + for j := 0; j < perWorker; j++ { + values <- tt.next() + } + }() + } + wg.Wait() + close(values) + + seen := make(map[uint32]struct{}, workers*perWorker) + for value := range values { + if _, ok := seen[value]; ok { + t.Fatalf("duplicate protocol id %d", value) + } + seen[value] = struct{}{} + } + }) + } +} diff --git a/plugins/web/webtitle.go b/plugins/web/webtitle.go index 214b400..0903715 100644 --- a/plugins/web/webtitle.go +++ b/plugins/web/webtitle.go @@ -6,9 +6,11 @@ import ( "context" "fmt" "io" + "net" "net/http" "net/url" "regexp" + "strconv" "strings" "time" "unicode/utf8" @@ -131,7 +133,7 @@ func (p *WebTitlePlugin) getWebTitle(ctx context.Context, info *common.HostInfo, isGM = true urlScheme = "https" // 国密连接仍使用 https URL 格式 } - baseURL := fmt.Sprintf("%s://%s:%d", urlScheme, info.Host, info.Port) + baseURL := webTitleURL(urlScheme, info.Host, info.Port) // 选择对应的 HTTP 客户端 clientNR, clientR := lib.ClientNoRedirect, lib.Client @@ -142,11 +144,11 @@ func (p *WebTitlePlugin) getWebTitle(ctx context.Context, info *common.HostInfo, // 构建显示用URL(隐藏标准端口) var displayURL string if isGM && info.Port == 443 { - displayURL = fmt.Sprintf("%s://%s", protocol, info.Host) + displayURL = webTitleDisplayURL(protocol, info.Host, info.Port, true) } else if (protocol == "https" && info.Port == 443) || (protocol == "http" && info.Port == 80) { - displayURL = fmt.Sprintf("%s://%s", protocol, info.Host) + displayURL = webTitleDisplayURL(protocol, info.Host, info.Port, true) } else { - displayURL = fmt.Sprintf("%s://%s:%d", protocol, info.Host, info.Port) + displayURL = webTitleDisplayURL(protocol, info.Host, info.Port, false) } req, err := http.NewRequestWithContext(ctx, "GET", baseURL, nil) @@ -221,6 +223,24 @@ func (p *WebTitlePlugin) getWebTitle(ctx context.Context, info *common.HostInfo, return title, statusCode, contentLen, server, fingerprints, displayURL, nil } +func webTitleURL(scheme, host string, port int) string { + return (&url.URL{Scheme: scheme, Host: net.JoinHostPort(host, strconv.Itoa(port))}).String() +} + +func webTitleDisplayURL(scheme, host string, port int, omitPort bool) string { + if omitPort { + return (&url.URL{Scheme: scheme, Host: urlHost(host)}).String() + } + return webTitleURL(scheme, host, port) +} + +func urlHost(host string) string { + if strings.Contains(host, ":") && !strings.HasPrefix(host, "[") { + return "[" + host + "]" + } + return host +} + // resolveRedirectURL 解析重定向URL,处理相对路径 func (p *WebTitlePlugin) resolveRedirectURL(baseURL, location string) string { // 如果是绝对URL,直接返回 diff --git a/plugins/web/webtitle_test.go b/plugins/web/webtitle_test.go index cef33dd..f92b730 100644 --- a/plugins/web/webtitle_test.go +++ b/plugins/web/webtitle_test.go @@ -35,3 +35,24 @@ func TestFetchFaviconHashHonorsContext(t *testing.T) { t.Fatalf("fetchFaviconHash returned hashes for canceled context: %#v", hashes) } } + +func TestWebTitleURLUsesJoinHostPort(t *testing.T) { + tests := []struct { + name string + got string + want string + }{ + {"ipv4", webTitleURL("http", "127.0.0.1", 8080), "http://127.0.0.1:8080"}, + {"ipv6", webTitleURL("http", "::1", 8080), "http://[::1]:8080"}, + {"ipv6 display with port", webTitleDisplayURL("https", "2001:db8::1", 8443, false), "https://[2001:db8::1]:8443"}, + {"ipv6 display omit port", webTitleDisplayURL("https", "2001:db8::1", 443, true), "https://[2001:db8::1]"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.got != tt.want { + t.Fatalf("got %q, want %q", tt.got, tt.want) + } + }) + } +} diff --git a/web/api/result.go b/web/api/result.go index 002111d..a28d1ff 100644 --- a/web/api/result.go +++ b/web/api/result.go @@ -7,9 +7,11 @@ import ( "encoding/csv" "encoding/json" "fmt" + "net" "net/http" "os" "path/filepath" + "strconv" "strings" "sync" "time" @@ -393,18 +395,39 @@ func (h *ResultHandler) Export(w http.ResponseWriter, r *http.Request) { // extractPort 从 "ip:port" 中提取端口 func extractPort(target string) string { - if idx := strings.LastIndex(target, ":"); idx != -1 { - return target[idx+1:] + _, port, ok := splitTargetHostPort(target) + if !ok { + return "" } - return "" + return port } // extractHost 从 "ip:port" 中提取主机 func extractHost(target string) string { - if idx := strings.LastIndex(target, ":"); idx != -1 { - return target[:idx] + host, _, ok := splitTargetHostPort(target) + if !ok { + return target } - return target + return host +} + +func splitTargetHostPort(target string) (string, string, bool) { + host, port, err := net.SplitHostPort(target) + if err != nil { + if strings.Count(target, ":") != 1 { + return "", "", false + } + parts := strings.SplitN(target, ":", 2) + host, port = parts[0], parts[1] + } + if host == "" || port == "" { + return "", "", false + } + portNum, err := strconv.Atoi(port) + if err != nil || portNum < 1 || portNum > 65535 { + return "", "", false + } + return host, port, true } // extractServiceInfo 从 details 中提取服务信息 diff --git a/web/api/result_test.go b/web/api/result_test.go new file mode 100644 index 0000000..d0f3935 --- /dev/null +++ b/web/api/result_test.go @@ -0,0 +1,31 @@ +//go:build web + +package api + +import "testing" + +func TestExtractHostPortIPv6(t *testing.T) { + tests := []struct { + name string + target string + wantHost string + wantPort string + }{ + {name: "ipv4", target: "192.168.1.1:80", wantHost: "192.168.1.1", wantPort: "80"}, + {name: "hostname", target: "example.com:443", wantHost: "example.com", wantPort: "443"}, + {name: "bracketed ipv6", target: "[2001:db8::1]:8443", wantHost: "2001:db8::1", wantPort: "8443"}, + {name: "bare ipv6 without port", target: "2001:db8::1", wantHost: "2001:db8::1", wantPort: ""}, + {name: "invalid port", target: "example.com:abc", wantHost: "example.com:abc", wantPort: ""}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := extractHost(tt.target); got != tt.wantHost { + t.Fatalf("extractHost(%q) = %q, want %q", tt.target, got, tt.wantHost) + } + if got := extractPort(tt.target); got != tt.wantPort { + t.Fatalf("extractPort(%q) = %q, want %q", tt.target, got, tt.wantPort) + } + }) + } +} diff --git a/webscan/lib/Client.go b/webscan/lib/Client.go index 4953d6a..2b90c65 100644 --- a/webscan/lib/Client.go +++ b/webscan/lib/Client.go @@ -9,6 +9,7 @@ import ( "net/http" "net/url" "os" + "strconv" "strings" "time" @@ -106,7 +107,7 @@ func configureHTTPProxy(tr *http.Transport, legacyProxy string, networkConfig *c } else if httpProxyURL == ProxyShortcutSocks5 { httpProxyURL = ProxySocks5URL } else if !strings.Contains(httpProxyURL, "://") { - httpProxyURL = "http://127.0.0.1:" + httpProxyURL + httpProxyURL = normalizeHTTPProxyURL(httpProxyURL) } // 验证代理类型 @@ -127,6 +128,13 @@ func configureHTTPProxy(tr *http.Transport, legacyProxy string, networkConfig *c return nil } +func normalizeHTTPProxyURL(proxyURL string) string { + if _, err := strconv.Atoi(proxyURL); err == nil { + return "http://127.0.0.1:" + proxyURL + } + return "http://" + proxyURL +} + // InitHTTPClient 创建HTTP客户端 func InitHTTPClient(ThreadsNum int, DownProxy string, Timeout time.Duration, maxRedirects int, networkConfig *common.NetworkConfig) error { // 配置基础连接参数 diff --git a/webscan/lib/Eval.go b/webscan/lib/Eval.go index 35efa43..54cb487 100644 --- a/webscan/lib/Eval.go +++ b/webscan/lib/Eval.go @@ -7,6 +7,7 @@ import ( "fmt" "io" "math/rand" //nolint:gosec // G404: math/rand用于生成测试数据,非加密用途 + "net" "net/http" "net/url" "strconv" @@ -176,7 +177,7 @@ func URLTypeToString(u *UrlType) string { builder.WriteString("//") } if host := u.Host; host != "" { - builder.WriteString(host) + builder.WriteString(urlTypeHost(host)) } } @@ -525,6 +526,16 @@ func ParseURL(u *url.URL) *UrlType { } } +func urlTypeHost(host string) string { + if strings.HasPrefix(host, "[") { + return host + } + if ip := net.ParseIP(host); ip != nil && strings.Contains(host, ":") { + return "[" + host + "]" + } + return host +} + // ParseRequest 将标准 HTTP 请求转换为自定义请求对象 func ParseRequest(oReq *http.Request) (*Request, error) { req := &Request{ diff --git a/webscan/lib/client_test.go b/webscan/lib/client_test.go new file mode 100644 index 0000000..0b675fb --- /dev/null +++ b/webscan/lib/client_test.go @@ -0,0 +1,24 @@ +package lib + +import "testing" + +func TestNormalizeHTTPProxyURL(t *testing.T) { + tests := []struct { + name string + in string + want string + }{ + {name: "port shortcut", in: "8080", want: "http://127.0.0.1:8080"}, + {name: "ipv4 host port", in: "127.0.0.1:8080", want: "http://127.0.0.1:8080"}, + {name: "hostname port", in: "proxy.local:8080", want: "http://proxy.local:8080"}, + {name: "bracketed ipv6 port", in: "[::1]:8080", want: "http://[::1]:8080"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := normalizeHTTPProxyURL(tt.in); got != tt.want { + t.Fatalf("normalizeHTTPProxyURL(%q) = %q, want %q", tt.in, got, tt.want) + } + }) + } +} diff --git a/webscan/lib/eval_test.go b/webscan/lib/eval_test.go index 9eb97ae..f626539 100644 --- a/webscan/lib/eval_test.go +++ b/webscan/lib/eval_test.go @@ -639,6 +639,15 @@ func TestURLTypeToString(t *testing.T) { }, expected: "http://example.com/test", }, + { + name: "IPv6 host", + url: &UrlType{ + Scheme: "http", + Host: "2001:db8::1", + Path: "/test", + }, + expected: "http://[2001:db8::1]/test", + }, { name: "仅路径", url: &UrlType{ diff --git a/webscan/web_scan.go b/webscan/web_scan.go index eb0cabc..5328e3a 100644 --- a/webscan/web_scan.go +++ b/webscan/web_scan.go @@ -107,7 +107,7 @@ func buildTargetURL(info *common.HostInfo) (string, error) { if info.URL == "" { info.URL = protocolHTTP + net.JoinHostPort(info.Host, fmt.Sprint(info.Port)) } else if !hasProtocolPrefix(info.URL) { - info.URL = protocolHTTP + info.URL + info.URL = protocolHTTP + normalizeSchemelessWebTarget(info.URL) } // 解析URL以提取基础部分 @@ -115,6 +115,7 @@ func buildTargetURL(info *common.HostInfo) (string, error) { if err != nil { return "", fmt.Errorf("%w: %w", ErrInvalidURL, err) } + parsedURL.Host = normalizeWebURLHost(parsedURL.Host) return fmt.Sprintf("%s://%s", parsedURL.Scheme, parsedURL.Host), nil } @@ -125,6 +126,32 @@ func hasProtocolPrefix(urlStr string) bool { return strings.HasPrefix(urlStr, protocolHTTP) || strings.HasPrefix(urlStr, protocolHTTPS) } +func normalizeSchemelessWebTarget(rawURL string) string { + authority := rawURL + suffix := "" + if idx := strings.IndexAny(rawURL, "/?#"); idx >= 0 { + authority = rawURL[:idx] + suffix = rawURL[idx:] + } + if strings.HasPrefix(authority, "[") { + return authority + suffix + } + if ip := net.ParseIP(authority); ip != nil && strings.Contains(authority, ":") { + return "[" + authority + "]" + suffix + } + return rawURL +} + +func normalizeWebURLHost(host string) string { + if strings.HasPrefix(host, "[") { + return host + } + if ip := net.ParseIP(host); ip != nil && strings.Contains(host, ":") { + return "[" + host + "]" + } + return host +} + // scanByFingerprints 根据指纹执行POC func scanByFingerprints(ctx context.Context, target string, fingerprints []string, cfg *common.Config, session *common.ScanSession) { for _, fingerprint := range fingerprints { diff --git a/webscan/web_scan_test.go b/webscan/web_scan_test.go index a258457..584d4ed 100644 --- a/webscan/web_scan_test.go +++ b/webscan/web_scan_test.go @@ -134,6 +134,26 @@ func TestBuildTargetURL(t *testing.T) { expected: "http://[2001:db8::1]:443", expectError: false, }, + { + name: "bare ipv6 url without protocol gets brackets", + hostInfo: &common.HostInfo{ + Host: "2001:db8::1", + Port: 80, + URL: "2001:db8::1/admin", + }, + expected: "http://[2001:db8::1]", + expectError: false, + }, + { + name: "bare ipv6 url with protocol gets brackets", + hostInfo: &common.HostInfo{ + Host: "2001:db8::1", + Port: 80, + URL: "http://2001:db8::1/admin", + }, + expected: "http://[2001:db8::1]", + expectError: false, + }, } for _, tt := range tests {