mirror of
https://github.com/shadow1ng/fscan.git
synced 2026-09-22 03:10:42 +08:00
feat: 统一服务缓存 + 指纹驱动插件匹配
将 webServiceCache 扩展为通用 serviceCache,所有指纹识别结果 统一缓存,插件匹配时端口不命中则回退到服务名称匹配。 删除多余的 service_cache.go,复用已有的 ServiceInfo 体系。 补充 nil 防御、Explicit 标记、大量单元/集成/回归测试。
This commit is contained in:
@@ -4,6 +4,7 @@ import (
|
|||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net"
|
"net"
|
||||||
|
"net/url"
|
||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
@@ -171,6 +172,7 @@ func parseUserPassPairs(fv *FlagVars) ([]config.CredentialPair, error) {
|
|||||||
// 如果命令行同时指定了单个用户名和单个密码(不是逗号分隔的多个)
|
// 如果命令行同时指定了单个用户名和单个密码(不是逗号分隔的多个)
|
||||||
if fv.Username != "" && fv.Password != "" &&
|
if fv.Username != "" && fv.Password != "" &&
|
||||||
!strings.Contains(fv.Username, ",") && !strings.Contains(fv.Password, ",") &&
|
!strings.Contains(fv.Username, ",") && !strings.Contains(fv.Password, ",") &&
|
||||||
|
fv.AddUsers == "" && fv.AddPasswords == "" &&
|
||||||
fv.UsersFile == "" && fv.PasswordsFile == "" && fv.UserPassFile == "" {
|
fv.UsersFile == "" && fv.PasswordsFile == "" && fv.UserPassFile == "" {
|
||||||
pairs = append(pairs, config.CredentialPair{
|
pairs = append(pairs, config.CredentialPair{
|
||||||
Username: strings.TrimSpace(fv.Username),
|
Username: strings.TrimSpace(fv.Username),
|
||||||
@@ -294,9 +296,42 @@ func normalizeURL(rawURL string) string {
|
|||||||
}
|
}
|
||||||
lowerURL := strings.ToLower(rawURL)
|
lowerURL := strings.ToLower(rawURL)
|
||||||
if !strings.HasPrefix(lowerURL, "http://") && !strings.HasPrefix(lowerURL, "https://") {
|
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
|
||||||
}
|
}
|
||||||
|
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
|
|||||||
@@ -3,6 +3,8 @@ package common
|
|||||||
import (
|
import (
|
||||||
"reflect"
|
"reflect"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
|
fscanconfig "github.com/shadow1ng/fscan/common/config"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestParsePasswordsKeepsPrimaryPasswordLiteral(t *testing.T) {
|
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) {
|
func TestParseTargetsHostPortDoesNotLeaveSyntheticHost(t *testing.T) {
|
||||||
fv := &FlagVars{Ports: "22"}
|
fv := &FlagVars{Ports: "22"}
|
||||||
info := &HostInfo{Host: "127.0.0.1:8080"}
|
info := &HostInfo{Host: "127.0.0.1:8080"}
|
||||||
@@ -73,3 +169,25 @@ func TestNormalizeURLKeepsUppercaseScheme(t *testing.T) {
|
|||||||
t.Fatalf("normalizeURL() = %q", got)
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+33
-4
@@ -145,6 +145,35 @@ type LocalExploitConfig struct {
|
|||||||
DownloadSavePath string // 下载保存路径
|
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)
|
// NewConfig 创建带默认值的Config(后备用,正常流程使用BuildConfigFromFlags)
|
||||||
func NewConfig() *Config {
|
func NewConfig() *Config {
|
||||||
return &Config{
|
return &Config{
|
||||||
@@ -163,13 +192,13 @@ func NewConfig() *Config {
|
|||||||
MaxRetries: 3,
|
MaxRetries: 3,
|
||||||
|
|
||||||
// 高级功能 - 使用默认配置
|
// 高级功能 - 使用默认配置
|
||||||
PortMap: config.DefaultPortMap,
|
PortMap: clonePortMap(config.DefaultPortMap),
|
||||||
DefaultMap: config.DefaultProbeMap,
|
DefaultMap: cloneStringSlice(config.DefaultProbeMap),
|
||||||
|
|
||||||
// 分组配置 - 使用默认字典
|
// 分组配置 - 使用默认字典
|
||||||
Credentials: CredentialConfig{
|
Credentials: CredentialConfig{
|
||||||
Userdict: config.DefaultUserDict,
|
Userdict: cloneStringSliceMap(config.DefaultUserDict),
|
||||||
Passwords: config.DefaultPasswords,
|
Passwords: cloneStringSlice(config.DefaultPasswords),
|
||||||
UserPassPairs: nil,
|
UserPassPairs: nil,
|
||||||
},
|
},
|
||||||
Network: NetworkConfig{
|
Network: NetworkConfig{
|
||||||
|
|||||||
@@ -164,8 +164,8 @@ func BuildConfigFromFlags(fv *FlagVars) *Config {
|
|||||||
DNSLog: fv.DNSLog,
|
DNSLog: fv.DNSLog,
|
||||||
PersistenceTargetFile: fv.PersistenceTargetFile,
|
PersistenceTargetFile: fv.PersistenceTargetFile,
|
||||||
WinPEFile: fv.WinPEFile,
|
WinPEFile: fv.WinPEFile,
|
||||||
PortMap: config.DefaultPortMap,
|
PortMap: clonePortMap(config.DefaultPortMap),
|
||||||
DefaultMap: config.DefaultProbeMap,
|
DefaultMap: cloneStringSlice(config.DefaultProbeMap),
|
||||||
|
|
||||||
// SOCKS5代理端口
|
// SOCKS5代理端口
|
||||||
Socks5ProxyPort: fv.Socks5ProxyPort,
|
Socks5ProxyPort: fv.Socks5ProxyPort,
|
||||||
@@ -175,8 +175,8 @@ func BuildConfigFromFlags(fv *FlagVars) *Config {
|
|||||||
Username: fv.Username,
|
Username: fv.Username,
|
||||||
Password: fv.Password,
|
Password: fv.Password,
|
||||||
Domain: fv.Domain,
|
Domain: fv.Domain,
|
||||||
Userdict: config.DefaultUserDict,
|
Userdict: cloneStringSliceMap(config.DefaultUserDict),
|
||||||
Passwords: config.DefaultPasswords,
|
Passwords: cloneStringSlice(config.DefaultPasswords),
|
||||||
UserPassPairs: nil, // 后续解析
|
UserPassPairs: nil, // 后续解析
|
||||||
SSHKeyPath: fv.SSHKeyPath,
|
SSHKeyPath: fv.SSHKeyPath,
|
||||||
},
|
},
|
||||||
|
|||||||
+5
-1
@@ -31,7 +31,11 @@ type HostInfo struct {
|
|||||||
|
|
||||||
// Target 返回 host:port 格式字符串
|
// Target 返回 host:port 格式字符串
|
||||||
func (h *HostInfo) Target() string {
|
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))
|
||||||
}
|
}
|
||||||
|
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
|
|||||||
@@ -8,3 +8,10 @@ func TestHostInfoTargetUsesBracketedIPv6(t *testing.T) {
|
|||||||
t.Fatalf("Target() = %q, want %q", got, want)
|
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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+17
-2
@@ -52,16 +52,31 @@ func getGlobalDialer(timeout time.Duration) (proxy.Dialer, error) {
|
|||||||
|
|
||||||
// parseProxyURL 解析代理URL,提取地址和认证信息
|
// parseProxyURL 解析代理URL,提取地址和认证信息
|
||||||
func parseProxyURL(proxyURL, fallback string) (host, username, password string) {
|
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)
|
parsedURL, err := url.Parse(proxyURL)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fallback, "", ""
|
return "", "", "", false
|
||||||
}
|
}
|
||||||
host = parsedURL.Host
|
host = parsedURL.Host
|
||||||
|
if host == "" {
|
||||||
|
return "", "", "", false
|
||||||
|
}
|
||||||
if parsedURL.User != nil {
|
if parsedURL.User != nil {
|
||||||
username = parsedURL.User.Username()
|
username = parsedURL.User.Username()
|
||||||
password, _ = parsedURL.User.Password()
|
password, _ = parsedURL.User.Password()
|
||||||
}
|
}
|
||||||
return
|
return host, username, password, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// createProxyConfig 根据全局设置创建代理配置
|
// createProxyConfig 根据全局设置创建代理配置
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -46,6 +46,12 @@ func targetWithPort(target string, port interface{}) string {
|
|||||||
return target
|
return target
|
||||||
}
|
}
|
||||||
portText := fmt.Sprint(port)
|
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 {
|
if strings.Count(target, ":") == 1 {
|
||||||
return target
|
return target
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 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: "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: "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: "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 {
|
for _, tt := range tests {
|
||||||
|
|||||||
@@ -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:[email protected]:8080", "user:[email protected]: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)
|
type roundTripFunc func(*http.Request) (*http.Response, error)
|
||||||
|
|
||||||
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||||
|
|||||||
@@ -137,11 +137,11 @@ func (b *BaseScanStrategy) isPluginApplicableToPortWithHost(pluginName string, t
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 端口不匹配时,按服务识别结果匹配
|
// 端口不匹配时,按指纹识别结果匹配
|
||||||
// 例:8881 端口上识别到 ssh 服务 → ssh 插件应该执行
|
// 例:8881 端口上识别到 ssh 服务 → ssh 插件应该执行
|
||||||
if targetHost != "" && targetPort > 0 {
|
if targetHost != "" && targetPort > 0 {
|
||||||
if svcName, ok := GetServiceName(targetHost, targetPort); ok {
|
if info, ok := GetCachedServiceInfo(targetHost, targetPort); ok && info != nil {
|
||||||
if strings.EqualFold(svcName, pluginName) {
|
if strings.EqualFold(info.Name, pluginName) {
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,10 +2,10 @@ package core
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
|
||||||
"math"
|
"math"
|
||||||
"net"
|
"net"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strconv"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -97,6 +97,10 @@ func (p *NetworkProfile) RecommendConcurrency(userThreadNum int, explicit bool)
|
|||||||
// probePorts 探测用的端口列表(高响应率的常见端口)
|
// probePorts 探测用的端口列表(高响应率的常见端口)
|
||||||
var probePorts = []int{80, 443, 22}
|
var probePorts = []int{80, 443, 22}
|
||||||
|
|
||||||
|
func networkProbeAddress(host string, port int) string {
|
||||||
|
return net.JoinHostPort(host, strconv.Itoa(port))
|
||||||
|
}
|
||||||
|
|
||||||
// ProbeNetwork 探测目标网络环境
|
// ProbeNetwork 探测目标网络环境
|
||||||
// 从 hosts 中抽样,用低并发 TCP 连接测量 RTT 和丢包率
|
// 从 hosts 中抽样,用低并发 TCP 连接测量 RTT 和丢包率
|
||||||
// 整个过程控制在数秒内完成
|
// 整个过程控制在数秒内完成
|
||||||
@@ -140,7 +144,7 @@ func ProbeNetwork(ctx context.Context, hosts []string, session *common.ScanSessi
|
|||||||
go func(h string, p int) {
|
go func(h string, p int) {
|
||||||
defer func() { <-sem; wg.Done() }()
|
defer func() { <-sem; wg.Done() }()
|
||||||
|
|
||||||
addr := fmt.Sprintf("%s:%d", h, p)
|
addr := networkProbeAddress(h, p)
|
||||||
start := time.Now()
|
start := time.Now()
|
||||||
conn, err := session.DialTCP(ctx, "tcp", addr, probeTimeout)
|
conn, err := session.DialTCP(ctx, "tcp", addr, probeTimeout)
|
||||||
rtt := time.Since(start)
|
rtt := time.Since(start)
|
||||||
|
|||||||
@@ -60,7 +60,7 @@ func TestClassifyNetwork(t *testing.T) {
|
|||||||
|
|
||||||
t.Run("公网 RTT 分布(低丢包)", func(t *testing.T) {
|
t.Run("公网 RTT 分布(低丢包)", func(t *testing.T) {
|
||||||
rtts := makeDurations([]int{60, 70, 80, 90, 100, 110, 120, 150, 200, 300}) // ms
|
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 {
|
if p.Env != EnvInternet {
|
||||||
t.Errorf("env = %v, want Internet", p.Env)
|
t.Errorf("env = %v, want Internet", p.Env)
|
||||||
@@ -72,7 +72,7 @@ func TestClassifyNetwork(t *testing.T) {
|
|||||||
|
|
||||||
t.Run("高丢包归类为慢速", func(t *testing.T) {
|
t.Run("高丢包归类为慢速", func(t *testing.T) {
|
||||||
rtts := makeDurations([]int{60, 70, 80, 90, 100}) // ms, 5 responded
|
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 {
|
if p.Env != EnvSlow {
|
||||||
t.Errorf("env = %v, want Slow (高丢包)", p.Env)
|
t.Errorf("env = %v, want Slow (高丢包)", p.Env)
|
||||||
@@ -96,14 +96,14 @@ func TestClassifyNetwork(t *testing.T) {
|
|||||||
|
|
||||||
func TestRecommendConcurrency(t *testing.T) {
|
func TestRecommendConcurrency(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
env NetworkEnv
|
env NetworkEnv
|
||||||
lossRate float64
|
lossRate float64
|
||||||
userT int
|
userT int
|
||||||
explicit bool
|
explicit bool
|
||||||
wantTMin int
|
wantTMin int
|
||||||
wantTMax int
|
wantTMax int
|
||||||
wantCeil int
|
wantCeil int
|
||||||
desc string
|
desc string
|
||||||
}{
|
}{
|
||||||
{EnvLAN, 0.0, 600, false, 800, 1000, -1, "内网自动: ×1.5"},
|
{EnvLAN, 0.0, 600, false, 800, 1000, -1, "内网自动: ×1.5"},
|
||||||
{EnvWAN, 0.0, 600, false, 550, 650, -1, "局域网自动: ×1.0"},
|
{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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
// 辅助
|
// 辅助
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
|
|||||||
+2
-3
@@ -707,8 +707,8 @@ func processServiceResult(ctx context.Context, host string, port int, addr strin
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
// 缓存服务名称,供插件按服务类型匹配(解决非标准端口问题)
|
// 缓存指纹识别结果,供插件按服务类型匹配(解决非标准端口问题)
|
||||||
MarkServiceName(host, port, serviceInfo.Name)
|
CacheServiceInfo(host, port, serviceInfo)
|
||||||
|
|
||||||
// 保存并输出服务信息
|
// 保存并输出服务信息
|
||||||
details := buildServiceDetails(port, serviceInfo)
|
details := buildServiceDetails(port, serviceInfo)
|
||||||
@@ -716,7 +716,6 @@ func processServiceResult(ctx context.Context, host string, port int, addr strin
|
|||||||
|
|
||||||
if isWeb {
|
if isWeb {
|
||||||
details["is_web"] = true
|
details["is_web"] = true
|
||||||
MarkAsWebService(host, port, serviceInfo)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
_ = session.SaveResult(&output.ScanResult{
|
_ = session.SaveResult(&output.ScanResult{
|
||||||
|
|||||||
@@ -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行提取)
|
// 排除端口逻辑测试(从EnhancedPortScan:28-32行提取)
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
|
|||||||
@@ -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
|
|
||||||
}
|
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -65,6 +65,16 @@ func TestParsePortList_BasicParsing(t *testing.T) {
|
|||||||
input: "22,80,443,3306",
|
input: "22,80,443,3306",
|
||||||
expected: []int{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: "空字符串",
|
name: "空字符串",
|
||||||
input: "",
|
input: "",
|
||||||
@@ -281,7 +291,7 @@ func TestParsePortList_ProductionScenarios(t *testing.T) {
|
|||||||
|
|
||||||
t.Run("数据库端口", func(t *testing.T) {
|
t.Run("数据库端口", func(t *testing.T) {
|
||||||
input := "3306,5432,1433,27017"
|
input := "3306,5432,1433,27017"
|
||||||
expected := []int{3306, 5432, 1433, 27017}
|
expected := []int{1433, 3306, 5432, 27017}
|
||||||
result := s.parsePortList(input)
|
result := s.parsePortList(input)
|
||||||
if !intSlicesEqual(result, expected) {
|
if !intSlicesEqual(result, expected) {
|
||||||
t.Errorf("应该正确解析常见数据库端口")
|
t.Errorf("应该正确解析常见数据库端口")
|
||||||
@@ -317,6 +327,13 @@ func TestParsePortList_ProductionScenarios(t *testing.T) {
|
|||||||
t.Errorf("应该正确解析高端口号")
|
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 测试返回值特性
|
// TestParsePortList_ReturnValue 测试返回值特性
|
||||||
@@ -330,14 +347,11 @@ func TestParsePortList_ReturnValue(t *testing.T) {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
t.Run("端口不重复-但不保证去重", func(t *testing.T) {
|
t.Run("重复端口会去重", func(t *testing.T) {
|
||||||
// 注意:当前实现不去重,如果用户输入 "22,22",会返回 [22, 22]
|
|
||||||
// 这是可以接受的,因为上层逻辑会处理重复
|
|
||||||
input := "22,22"
|
input := "22,22"
|
||||||
result := s.parsePortList(input)
|
result := s.parsePortList(input)
|
||||||
// 这里我们只测试解析是否正确,不测试去重
|
if len(result) != 1 || result[0] != 22 {
|
||||||
if len(result) != 2 || result[0] != 22 || result[1] != 22 {
|
t.Errorf("重复端口应该去重, 实际 %v", result)
|
||||||
t.Errorf("当前实现不去重,应该返回两个22")
|
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
@@ -355,6 +369,15 @@ func intSlicesEqual(a, b []int) bool {
|
|||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func sliceContains(values []int, target int) bool {
|
||||||
|
for _, value := range values {
|
||||||
|
if value == target {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
// 存活检测判断测试
|
// 存活检测判断测试
|
||||||
// =============================================================================
|
// =============================================================================
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
+32
-14
@@ -204,10 +204,11 @@ func (w *WebPortDetector) tryHTTP(ctx context.Context, client *http.Client, sess
|
|||||||
// 基于服务指纹的Web服务识别
|
// 基于服务指纹的Web服务识别
|
||||||
// ===============================
|
// ===============================
|
||||||
|
|
||||||
// Web服务缓存 - 简化的全局缓存
|
// 服务识别缓存 - 存储所有识别到的服务(不仅限于 Web)
|
||||||
|
// 端口扫描阶段写入,插件匹配阶段读取
|
||||||
var (
|
var (
|
||||||
webServiceCache = make(map[string]*ServiceInfo)
|
serviceCache = make(map[string]*ServiceInfo)
|
||||||
webCacheMutex sync.RWMutex
|
serviceCacheMutex sync.RWMutex
|
||||||
)
|
)
|
||||||
|
|
||||||
// IsWebServiceByFingerprint 基于服务指纹判断Web服务 - 保持API兼容
|
// IsWebServiceByFingerprint 基于服务指纹判断Web服务 - 保持API兼容
|
||||||
@@ -259,28 +260,45 @@ func IsWebServiceByFingerprint(serviceInfo *ServiceInfo) bool {
|
|||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
|
|
||||||
// MarkAsWebService 标记Web服务 - 保持API兼容
|
// CacheServiceInfo 缓存识别到的服务信息
|
||||||
func MarkAsWebService(host string, port int, serviceInfo *ServiceInfo) {
|
func CacheServiceInfo(host string, port int, serviceInfo *ServiceInfo) {
|
||||||
cacheKey := net.JoinHostPort(host, strconv.Itoa(port))
|
cacheKey := net.JoinHostPort(host, strconv.Itoa(port))
|
||||||
|
|
||||||
webCacheMutex.Lock()
|
serviceCacheMutex.Lock()
|
||||||
defer webCacheMutex.Unlock()
|
defer serviceCacheMutex.Unlock()
|
||||||
|
|
||||||
webServiceCache[cacheKey] = serviceInfo
|
serviceCache[cacheKey] = serviceInfo
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetWebServiceInfo 获取Web服务信息
|
// MarkAsWebService 标记 Web 服务(兼容旧调用)
|
||||||
func GetWebServiceInfo(host string, port int) (*ServiceInfo, bool) {
|
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))
|
cacheKey := net.JoinHostPort(host, strconv.Itoa(port))
|
||||||
|
|
||||||
webCacheMutex.RLock()
|
serviceCacheMutex.RLock()
|
||||||
defer webCacheMutex.RUnlock()
|
defer serviceCacheMutex.RUnlock()
|
||||||
|
|
||||||
serviceInfo, exists := webServiceCache[cacheKey]
|
serviceInfo, exists := serviceCache[cacheKey]
|
||||||
return serviceInfo, exists
|
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 {
|
func IsMarkedWebService(host string, port int) bool {
|
||||||
_, exists := GetWebServiceInfo(host, port)
|
_, exists := GetWebServiceInfo(host, port)
|
||||||
return exists
|
return exists
|
||||||
|
|||||||
@@ -440,9 +440,9 @@ func TestCreateTargetFromURL(t *testing.T) {
|
|||||||
// TestWebServiceCache 测试Web服务缓存操作
|
// TestWebServiceCache 测试Web服务缓存操作
|
||||||
func TestWebServiceCache(t *testing.T) {
|
func TestWebServiceCache(t *testing.T) {
|
||||||
// 清空缓存
|
// 清空缓存
|
||||||
webCacheMutex.Lock()
|
serviceCacheMutex.Lock()
|
||||||
webServiceCache = make(map[string]*ServiceInfo)
|
serviceCache = make(map[string]*ServiceInfo)
|
||||||
webCacheMutex.Unlock()
|
serviceCacheMutex.Unlock()
|
||||||
|
|
||||||
t.Run("存储和读取", func(t *testing.T) {
|
t.Run("存储和读取", func(t *testing.T) {
|
||||||
serviceInfo := &ServiceInfo{
|
serviceInfo := &ServiceInfo{
|
||||||
@@ -517,9 +517,9 @@ func TestWebServiceCache(t *testing.T) {
|
|||||||
// TestWebServiceCache_Concurrent 测试并发安全性
|
// TestWebServiceCache_Concurrent 测试并发安全性
|
||||||
func TestWebServiceCache_Concurrent(t *testing.T) {
|
func TestWebServiceCache_Concurrent(t *testing.T) {
|
||||||
// 清空缓存
|
// 清空缓存
|
||||||
webCacheMutex.Lock()
|
serviceCacheMutex.Lock()
|
||||||
webServiceCache = make(map[string]*ServiceInfo)
|
serviceCache = make(map[string]*ServiceInfo)
|
||||||
webCacheMutex.Unlock()
|
serviceCacheMutex.Unlock()
|
||||||
|
|
||||||
t.Run("不同key并发写入", func(t *testing.T) {
|
t.Run("不同key并发写入", func(t *testing.T) {
|
||||||
var wg sync.WaitGroup
|
var wg sync.WaitGroup
|
||||||
|
|||||||
@@ -176,8 +176,7 @@ func (g *Client) ProbeOSInfo(host, domain, user, pwd string, timeout int64, rdpP
|
|||||||
exitFlag := make(chan bool, 1)
|
exitFlag := make(chan bool, 1)
|
||||||
info = make(map[string]any)
|
info = make(map[string]any)
|
||||||
|
|
||||||
targetSlice := strings.Split(g.Host, ":")
|
ip := rdpTargetHost(g.Host)
|
||||||
ip := targetSlice[0]
|
|
||||||
conn, err := WrapperTcpWithTimeout("tcp", g.Host, time.Duration(timeout)*time.Second)
|
conn, err := WrapperTcpWithTimeout("tcp", g.Host, time.Duration(timeout)*time.Second)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return
|
return
|
||||||
@@ -273,3 +272,14 @@ loop:
|
|||||||
glog.Debug("loop ended, elapsed time: ", time.Since(start))
|
glog.Debug("loop ended, elapsed time: ", time.Since(start))
|
||||||
return info
|
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
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+28
-4
@@ -192,9 +192,9 @@ func TestGenerateCredentials_PlaceholderReplacement(t *testing.T) {
|
|||||||
|
|
||||||
// 验证:{user} 被正确替换
|
// 验证:{user} 被正确替换
|
||||||
expectedCombos := map[string]string{
|
expectedCombos := map[string]string{
|
||||||
"root:root": "root", // {user} → root
|
"root:root": "root", // {user} → root
|
||||||
"root:root123": "root", // {user}123 → root123
|
"root:root123": "root", // {user}123 → root123
|
||||||
"mysql:mysql": "mysql", // {user} → mysql
|
"mysql:mysql": "mysql", // {user} → mysql
|
||||||
"mysql:mysql123": "mysql", // {user}123 → mysql123
|
"mysql:mysql123": "mysql", // {user}123 → mysql123
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -244,7 +244,7 @@ func TestGenerateCredentials_DefaultValues(t *testing.T) {
|
|||||||
|
|
||||||
cfg.Credentials.UserPassPairs = []config.CredentialPair{}
|
cfg.Credentials.UserPassPairs = []config.CredentialPair{}
|
||||||
cfg.Credentials.Userdict = map[string][]string{} // 空字典
|
cfg.Credentials.Userdict = map[string][]string{} // 空字典
|
||||||
cfg.Credentials.Passwords = []string{} // 空密码列表
|
cfg.Credentials.Passwords = []string{} // 空密码列表
|
||||||
|
|
||||||
result := GenerateCredentials("unknown_service", cfg)
|
result := GenerateCredentials("unknown_service", cfg)
|
||||||
|
|
||||||
@@ -327,3 +327,27 @@ func TestGenerateCredentials_EmptyUserPassPairs(t *testing.T) {
|
|||||||
|
|
||||||
t.Logf("✓ 空 UserPassPairs 正确回退到笛卡尔积")
|
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -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)))
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -82,10 +82,10 @@ func (p *SystemInfoPlugin) connectToDomain(domain string) (*domainInfo, error) {
|
|||||||
}
|
}
|
||||||
defer func() { _ = client.Close() }()
|
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 err != nil {
|
||||||
if ipv4, resolveErr := resolveIPv4(dcHost); resolveErr == 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 {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("LDAP dial: %w", err)
|
return nil, fmt.Errorf("LDAP dial: %w", err)
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/shadow1ng/fscan/common"
|
"github.com/shadow1ng/fscan/common"
|
||||||
@@ -165,20 +166,19 @@ func (p *CassandraPlugin) doCassandraAuth(ctx context.Context, info *common.Host
|
|||||||
|
|
||||||
// ── CQL wire protocol 工具 ──────────────────────────────────────
|
// ── 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 {
|
func cqlSend(conn net.Conn, opcode byte, body []byte) error {
|
||||||
id := cqlStreamID
|
id := nextCQLStreamID()
|
||||||
if cqlStreamID == 32767 {
|
|
||||||
cqlStreamID = 0
|
|
||||||
} else {
|
|
||||||
cqlStreamID++
|
|
||||||
}
|
|
||||||
|
|
||||||
// frame: [1B version|flags] [2B stream] [1B opcode] [4B length] [body]
|
// frame: [1B version|flags] [2B stream] [1B opcode] [4B length] [body]
|
||||||
header := make([]byte, 8)
|
header := make([]byte, 8)
|
||||||
header[0] = cqlVersion
|
header[0] = cqlVersion
|
||||||
binary.BigEndian.PutUint16(header[1:3], uint16(id))
|
binary.BigEndian.PutUint16(header[1:3], id)
|
||||||
header[3] = opcode
|
header[3] = opcode
|
||||||
binary.BigEndian.PutUint32(header[4:8], uint32(len(body)))
|
binary.BigEndian.PutUint32(header[4:8], uint32(len(body)))
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"sync"
|
"sync"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/shadow1ng/fscan/common"
|
"github.com/shadow1ng/fscan/common"
|
||||||
@@ -61,7 +62,11 @@ type AuthFunc func(ctx context.Context, cred Credential) *AuthResult
|
|||||||
// ErrorClassifier 错误分类函数
|
// ErrorClassifier 错误分类函数
|
||||||
type ErrorClassifier func(err error) ErrorType
|
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 泄漏)
|
// 单凭据测试(解决 goroutine 泄漏)
|
||||||
@@ -70,9 +75,36 @@ var authCleanupWait = 2 * time.Second
|
|||||||
// TestSingleCredential 安全地测试单个凭据
|
// TestSingleCredential 安全地测试单个凭据
|
||||||
// 正确处理 context 取消时的资源清理
|
// 正确处理 context 取消时的资源清理
|
||||||
func TestSingleCredential(ctx context.Context, cred Credential, authFn AuthFunc) *AuthResult {
|
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)
|
resultChan := make(chan *AuthResult, 1)
|
||||||
|
|
||||||
go func() {
|
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)
|
result := authFn(ctx, cred)
|
||||||
resultChan <- result
|
resultChan <- result
|
||||||
}()
|
}()
|
||||||
@@ -83,7 +115,7 @@ func TestSingleCredential(ctx context.Context, cred Credential, authFn AuthFunc)
|
|||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
// context 被取消后只做有界等待,避免 authFn 卡死时清理 goroutine 也永久泄漏。
|
// context 被取消后只做有界等待,避免 authFn 卡死时清理 goroutine 也永久泄漏。
|
||||||
go func() {
|
go func() {
|
||||||
timer := time.NewTimer(authCleanupWait)
|
timer := time.NewTimer(authCleanupWait())
|
||||||
defer timer.Stop()
|
defer timer.Stop()
|
||||||
|
|
||||||
select {
|
select {
|
||||||
@@ -116,15 +148,35 @@ type ConcurrentTestConfig struct {
|
|||||||
UseProxy bool // 代理模式下跳过直连 TCP 预检
|
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 默认配置
|
// DefaultConcurrentTestConfig 默认配置
|
||||||
func DefaultConcurrentTestConfig(config *common.Config) ConcurrentTestConfig {
|
func DefaultConcurrentTestConfig(config *common.Config) ConcurrentTestConfig {
|
||||||
concurrency := config.ModuleThreadNum
|
concurrency := config.ModuleThreadNum
|
||||||
if concurrency <= 0 {
|
if concurrency <= 0 {
|
||||||
concurrency = 10
|
concurrency = 10
|
||||||
}
|
}
|
||||||
|
maxRetries := config.MaxRetries
|
||||||
|
if maxRetries <= 0 {
|
||||||
|
maxRetries = 3
|
||||||
|
}
|
||||||
return ConcurrentTestConfig{
|
return ConcurrentTestConfig{
|
||||||
Concurrency: concurrency,
|
Concurrency: concurrency,
|
||||||
MaxRetries: 3,
|
MaxRetries: maxRetries,
|
||||||
RetryDelay: time.Second,
|
RetryDelay: time.Second,
|
||||||
MaxConsecutiveNetErrors: 5,
|
MaxConsecutiveNetErrors: 5,
|
||||||
UseProxy: config.Network.Socks5Proxy != "" || config.Network.HTTPProxy != "",
|
UseProxy: config.Network.Socks5Proxy != "" || config.Network.HTTPProxy != "",
|
||||||
@@ -147,6 +199,9 @@ func TestCredentialsConcurrently(
|
|||||||
serviceName string,
|
serviceName string,
|
||||||
testConfig ConcurrentTestConfig,
|
testConfig ConcurrentTestConfig,
|
||||||
) *ScanResult {
|
) *ScanResult {
|
||||||
|
if ctx == nil {
|
||||||
|
ctx = context.Background()
|
||||||
|
}
|
||||||
if len(credentials) == 0 {
|
if len(credentials) == 0 {
|
||||||
return &ScanResult{
|
return &ScanResult{
|
||||||
Success: false,
|
Success: false,
|
||||||
@@ -154,11 +209,16 @@ func TestCredentialsConcurrently(
|
|||||||
Error: fmt.Errorf("%s", i18n.GetText("service_no_test_creds")),
|
Error: fmt.Errorf("%s", i18n.GetText("service_no_test_creds")),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
testConfig = normalizeConcurrentTestConfig(testConfig)
|
||||||
|
|
||||||
// TCP 预检:快速验证目标可达,避免对不可达目标浪费全部凭据尝试
|
// TCP 预检:快速验证目标可达,避免对不可达目标浪费全部凭据尝试
|
||||||
// 代理模式下跳过:net.DialTimeout 直连无法到达代理后的内网目标
|
// 代理模式下跳过:net.DialTimeout 直连无法到达代理后的内网目标
|
||||||
if testConfig.TargetAddr != "" && !testConfig.UseProxy {
|
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 {
|
if err != nil {
|
||||||
return &ScanResult{
|
return &ScanResult{
|
||||||
Success: false,
|
Success: false,
|
||||||
@@ -240,10 +300,6 @@ func workerTestCredentials(
|
|||||||
testConfig ConcurrentTestConfig,
|
testConfig ConcurrentTestConfig,
|
||||||
) {
|
) {
|
||||||
consecutiveNetErrors := 0
|
consecutiveNetErrors := 0
|
||||||
maxNetErrors := testConfig.MaxConsecutiveNetErrors
|
|
||||||
if maxNetErrors <= 0 {
|
|
||||||
maxNetErrors = 5
|
|
||||||
}
|
|
||||||
|
|
||||||
for cred := range credChan {
|
for cred := range credChan {
|
||||||
// 检查是否应该停止
|
// 检查是否应该停止
|
||||||
@@ -254,7 +310,7 @@ func workerTestCredentials(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// 连续网络错误达到阈值,目标可能不可达,提前退出
|
// 连续网络错误达到阈值,目标可能不可达,提前退出
|
||||||
if consecutiveNetErrors >= maxNetErrors {
|
if consecutiveNetErrors >= testConfig.MaxConsecutiveNetErrors {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -292,10 +348,18 @@ func testCredentialWithRetry(
|
|||||||
|
|
||||||
// 测试凭据
|
// 测试凭据
|
||||||
result := TestSingleCredential(ctx, cred, authFn)
|
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 {
|
if result.Success {
|
||||||
// 成功,关闭连接并返回
|
if result.Conn != nil {
|
||||||
_ = result.Conn.Close()
|
_ = result.Conn.Close()
|
||||||
|
}
|
||||||
return &ScanResult{
|
return &ScanResult{
|
||||||
Type: plugins.ResultTypeCredential,
|
Type: plugins.ResultTypeCredential,
|
||||||
Success: true,
|
Success: true,
|
||||||
|
|||||||
@@ -8,6 +8,8 @@ import (
|
|||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"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 模拟连接
|
// mockConn 模拟连接
|
||||||
type mockConn struct {
|
type mockConn struct {
|
||||||
closed atomic.Bool
|
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取消时的资源清理
|
// TestTestSingleCredential_ContextCancel 测试context取消时的资源清理
|
||||||
func TestTestSingleCredential_ContextCancel(t *testing.T) {
|
func TestTestSingleCredential_ContextCancel(t *testing.T) {
|
||||||
conn := &mockConn{}
|
conn := &mockConn{}
|
||||||
@@ -403,9 +580,7 @@ func TestTestSingleCredential_ContextCancel(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestTestSingleCredential_ContextCancelCleanupIsBounded(t *testing.T) {
|
func TestTestSingleCredential_ContextCancelCleanupIsBounded(t *testing.T) {
|
||||||
oldWait := authCleanupWait
|
defer setAuthCleanupWaitForTest(20 * time.Millisecond)()
|
||||||
authCleanupWait = 20 * time.Millisecond
|
|
||||||
defer func() { authCleanupWait = oldWait }()
|
|
||||||
|
|
||||||
authStarted := make(chan struct{})
|
authStarted := make(chan struct{})
|
||||||
releaseAuth := 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 认证错误不应该重试
|
// TestRetryLogic_AuthErrorNoRetry 认证错误不应该重试
|
||||||
func TestRetryLogic_AuthErrorNoRetry(t *testing.T) {
|
func TestRetryLogic_AuthErrorNoRetry(t *testing.T) {
|
||||||
var attempts atomic.Int32
|
var attempts atomic.Int32
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/shadow1ng/fscan/common"
|
"github.com/shadow1ng/fscan/common"
|
||||||
@@ -154,9 +155,12 @@ func (p *KafkaPlugin) doKafkaAuth(ctx context.Context, info *common.HostInfo, cr
|
|||||||
|
|
||||||
var kafkaCorrelationID int32
|
var kafkaCorrelationID int32
|
||||||
|
|
||||||
|
func nextKafkaCorrelationID() int32 {
|
||||||
|
return atomic.AddInt32(&kafkaCorrelationID, 1) - 1
|
||||||
|
}
|
||||||
|
|
||||||
func kafkaSend(conn net.Conn, apiKey, apiVersion int16, body []byte) error {
|
func kafkaSend(conn net.Conn, apiKey, apiVersion int16, body []byte) error {
|
||||||
corrID := kafkaCorrelationID
|
corrID := nextKafkaCorrelationID()
|
||||||
kafkaCorrelationID++
|
|
||||||
|
|
||||||
// 请求格式: [4B len] [2B api_key] [2B api_version] [4B corr_id] [2B client_id_len] [client_id] [body]
|
// 请求格式: [4B len] [2B api_key] [2B api_version] [4B corr_id] [2B client_id_len] [client_id] [body]
|
||||||
clientID := "fscan"
|
clientID := "fscan"
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ import (
|
|||||||
"io"
|
"io"
|
||||||
"net"
|
"net"
|
||||||
"strings"
|
"strings"
|
||||||
|
"sync/atomic"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
"github.com/shadow1ng/fscan/common"
|
"github.com/shadow1ng/fscan/common"
|
||||||
@@ -156,8 +157,7 @@ const (
|
|||||||
var mongoRequestID uint32
|
var mongoRequestID uint32
|
||||||
|
|
||||||
func nextRequestID() uint32 {
|
func nextRequestID() uint32 {
|
||||||
mongoRequestID++
|
return atomic.AddUint32(&mongoRequestID, 1)
|
||||||
return mongoRequestID
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// buildMongoCommand 构建 MongoDB 命令的 OP_MSG body (最小 BSON 实现)
|
// buildMongoCommand 构建 MongoDB 命令的 OP_MSG body (最小 BSON 实现)
|
||||||
|
|||||||
@@ -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{}{}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
+24
-4
@@ -6,9 +6,11 @@ import (
|
|||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
"regexp"
|
"regexp"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
"unicode/utf8"
|
"unicode/utf8"
|
||||||
@@ -131,7 +133,7 @@ func (p *WebTitlePlugin) getWebTitle(ctx context.Context, info *common.HostInfo,
|
|||||||
isGM = true
|
isGM = true
|
||||||
urlScheme = "https" // 国密连接仍使用 https URL 格式
|
urlScheme = "https" // 国密连接仍使用 https URL 格式
|
||||||
}
|
}
|
||||||
baseURL := fmt.Sprintf("%s://%s:%d", urlScheme, info.Host, info.Port)
|
baseURL := webTitleURL(urlScheme, info.Host, info.Port)
|
||||||
|
|
||||||
// 选择对应的 HTTP 客户端
|
// 选择对应的 HTTP 客户端
|
||||||
clientNR, clientR := lib.ClientNoRedirect, lib.Client
|
clientNR, clientR := lib.ClientNoRedirect, lib.Client
|
||||||
@@ -142,11 +144,11 @@ func (p *WebTitlePlugin) getWebTitle(ctx context.Context, info *common.HostInfo,
|
|||||||
// 构建显示用URL(隐藏标准端口)
|
// 构建显示用URL(隐藏标准端口)
|
||||||
var displayURL string
|
var displayURL string
|
||||||
if isGM && info.Port == 443 {
|
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) {
|
} 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 {
|
} 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)
|
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
|
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,处理相对路径
|
// resolveRedirectURL 解析重定向URL,处理相对路径
|
||||||
func (p *WebTitlePlugin) resolveRedirectURL(baseURL, location string) string {
|
func (p *WebTitlePlugin) resolveRedirectURL(baseURL, location string) string {
|
||||||
// 如果是绝对URL,直接返回
|
// 如果是绝对URL,直接返回
|
||||||
|
|||||||
@@ -35,3 +35,24 @@ func TestFetchFaviconHashHonorsContext(t *testing.T) {
|
|||||||
t.Fatalf("fetchFaviconHash returned hashes for canceled context: %#v", hashes)
|
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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
+29
-6
@@ -7,9 +7,11 @@ import (
|
|||||||
"encoding/csv"
|
"encoding/csv"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
@@ -393,18 +395,39 @@ func (h *ResultHandler) Export(w http.ResponseWriter, r *http.Request) {
|
|||||||
|
|
||||||
// extractPort 从 "ip:port" 中提取端口
|
// extractPort 从 "ip:port" 中提取端口
|
||||||
func extractPort(target string) string {
|
func extractPort(target string) string {
|
||||||
if idx := strings.LastIndex(target, ":"); idx != -1 {
|
_, port, ok := splitTargetHostPort(target)
|
||||||
return target[idx+1:]
|
if !ok {
|
||||||
|
return ""
|
||||||
}
|
}
|
||||||
return ""
|
return port
|
||||||
}
|
}
|
||||||
|
|
||||||
// extractHost 从 "ip:port" 中提取主机
|
// extractHost 从 "ip:port" 中提取主机
|
||||||
func extractHost(target string) string {
|
func extractHost(target string) string {
|
||||||
if idx := strings.LastIndex(target, ":"); idx != -1 {
|
host, _, ok := splitTargetHostPort(target)
|
||||||
return target[:idx]
|
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 中提取服务信息
|
// extractServiceInfo 从 details 中提取服务信息
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -9,6 +9,7 @@ import (
|
|||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
"os"
|
"os"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"time"
|
"time"
|
||||||
|
|
||||||
@@ -106,7 +107,7 @@ func configureHTTPProxy(tr *http.Transport, legacyProxy string, networkConfig *c
|
|||||||
} else if httpProxyURL == ProxyShortcutSocks5 {
|
} else if httpProxyURL == ProxyShortcutSocks5 {
|
||||||
httpProxyURL = ProxySocks5URL
|
httpProxyURL = ProxySocks5URL
|
||||||
} else if !strings.Contains(httpProxyURL, "://") {
|
} 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
|
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客户端
|
// InitHTTPClient 创建HTTP客户端
|
||||||
func InitHTTPClient(ThreadsNum int, DownProxy string, Timeout time.Duration, maxRedirects int, networkConfig *common.NetworkConfig) error {
|
func InitHTTPClient(ThreadsNum int, DownProxy string, Timeout time.Duration, maxRedirects int, networkConfig *common.NetworkConfig) error {
|
||||||
// 配置基础连接参数
|
// 配置基础连接参数
|
||||||
|
|||||||
+12
-1
@@ -7,6 +7,7 @@ import (
|
|||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
"math/rand" //nolint:gosec // G404: math/rand用于生成测试数据,非加密用途
|
"math/rand" //nolint:gosec // G404: math/rand用于生成测试数据,非加密用途
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/url"
|
"net/url"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -176,7 +177,7 @@ func URLTypeToString(u *UrlType) string {
|
|||||||
builder.WriteString("//")
|
builder.WriteString("//")
|
||||||
}
|
}
|
||||||
if host := u.Host; host != "" {
|
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 请求转换为自定义请求对象
|
// ParseRequest 将标准 HTTP 请求转换为自定义请求对象
|
||||||
func ParseRequest(oReq *http.Request) (*Request, error) {
|
func ParseRequest(oReq *http.Request) (*Request, error) {
|
||||||
req := &Request{
|
req := &Request{
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -639,6 +639,15 @@ func TestURLTypeToString(t *testing.T) {
|
|||||||
},
|
},
|
||||||
expected: "http://example.com/test",
|
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: "仅路径",
|
name: "仅路径",
|
||||||
url: &UrlType{
|
url: &UrlType{
|
||||||
|
|||||||
+28
-1
@@ -107,7 +107,7 @@ func buildTargetURL(info *common.HostInfo) (string, error) {
|
|||||||
if info.URL == "" {
|
if info.URL == "" {
|
||||||
info.URL = protocolHTTP + net.JoinHostPort(info.Host, fmt.Sprint(info.Port))
|
info.URL = protocolHTTP + net.JoinHostPort(info.Host, fmt.Sprint(info.Port))
|
||||||
} else if !hasProtocolPrefix(info.URL) {
|
} else if !hasProtocolPrefix(info.URL) {
|
||||||
info.URL = protocolHTTP + info.URL
|
info.URL = protocolHTTP + normalizeSchemelessWebTarget(info.URL)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 解析URL以提取基础部分
|
// 解析URL以提取基础部分
|
||||||
@@ -115,6 +115,7 @@ func buildTargetURL(info *common.HostInfo) (string, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return "", fmt.Errorf("%w: %w", ErrInvalidURL, err)
|
return "", fmt.Errorf("%w: %w", ErrInvalidURL, err)
|
||||||
}
|
}
|
||||||
|
parsedURL.Host = normalizeWebURLHost(parsedURL.Host)
|
||||||
|
|
||||||
return fmt.Sprintf("%s://%s", parsedURL.Scheme, parsedURL.Host), nil
|
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)
|
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
|
// scanByFingerprints 根据指纹执行POC
|
||||||
func scanByFingerprints(ctx context.Context, target string, fingerprints []string, cfg *common.Config, session *common.ScanSession) {
|
func scanByFingerprints(ctx context.Context, target string, fingerprints []string, cfg *common.Config, session *common.ScanSession) {
|
||||||
for _, fingerprint := range fingerprints {
|
for _, fingerprint := range fingerprints {
|
||||||
|
|||||||
@@ -134,6 +134,26 @@ func TestBuildTargetURL(t *testing.T) {
|
|||||||
expected: "http://[2001:db8::1]:443",
|
expected: "http://[2001:db8::1]:443",
|
||||||
expectError: false,
|
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 {
|
for _, tt := range tests {
|
||||||
|
|||||||
Reference in New Issue
Block a user