mirror of
https://github.com/shadow1ng/fscan.git
synced 2026-09-22 03:10:42 +08:00
test: 补充单元测试覆盖率 29.9% → 36.6%
新建 18 个测试文件,追加 30 个已有测试文件,覆盖协议解析、 错误分类、CEL 表达式求值、YAML 反序列化、字节编码等纯函数。
This commit is contained in:
@@ -3,6 +3,7 @@ package common
|
||||
import (
|
||||
"reflect"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
fscanconfig "github.com/shadow1ng/fscan/common/config"
|
||||
)
|
||||
@@ -192,3 +193,258 @@ func TestNormalizeURLBracketsIPv6Literals(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestModuleTimeout 测试模块超时计算
|
||||
func TestModuleTimeout(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
timeout time.Duration
|
||||
want time.Duration
|
||||
}{
|
||||
{"超时大于下限", 10 * time.Second, 10 * time.Second},
|
||||
{"超时等于下限", 3 * time.Second, 3 * time.Second},
|
||||
{"超时小于下限", 1 * time.Second, 3 * time.Second},
|
||||
{"零超时", 0, 3 * time.Second},
|
||||
{"负超时", -1 * time.Second, 3 * time.Second},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cfg := NewConfig()
|
||||
cfg.Timeout = tt.timeout
|
||||
got := cfg.ModuleTimeout()
|
||||
if got != tt.want {
|
||||
t.Errorf("ModuleTimeout() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseUserPassPairsExactMatch 测试精确单用户单密码路径
|
||||
func TestParseUserPassPairsExactMatch(t *testing.T) {
|
||||
fv := &FlagVars{
|
||||
Username: "admin",
|
||||
Password: "secret",
|
||||
}
|
||||
pairs, err := parseUserPassPairs(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseUserPassPairs error = %v", err)
|
||||
}
|
||||
if len(pairs) != 1 {
|
||||
t.Fatalf("期望 1 个 pair, 实际 %d", len(pairs))
|
||||
}
|
||||
if pairs[0].Username != "admin" || pairs[0].Password != "secret" {
|
||||
t.Errorf("pair = %+v, want {admin secret}", pairs[0])
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseUserPassPairsMultiUserSkips 测试多用户时不生成精确 pair
|
||||
func TestParseUserPassPairsMultiUserSkips(t *testing.T) {
|
||||
fv := &FlagVars{
|
||||
Username: "admin,root",
|
||||
Password: "pass",
|
||||
}
|
||||
pairs, err := parseUserPassPairs(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseUserPassPairs error = %v", err)
|
||||
}
|
||||
if len(pairs) != 0 {
|
||||
t.Fatalf("多用户场景不应生成精确 pair, 实际 %d 个", len(pairs))
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseURLsEmpty 测试空输入返回空列表
|
||||
func TestParseURLsEmpty(t *testing.T) {
|
||||
fv := &FlagVars{}
|
||||
urls, err := parseURLs(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseURLs error = %v", err)
|
||||
}
|
||||
if len(urls) != 0 {
|
||||
t.Fatalf("空输入应返回空 url 列表, 实际 %v", urls)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseURLsCommaSeparated 测试逗号分隔多 URL
|
||||
func TestParseURLsCommaSeparated(t *testing.T) {
|
||||
fv := &FlagVars{
|
||||
TargetURL: "http://a.com,http://b.com,http://a.com", // 含重复
|
||||
}
|
||||
urls, err := parseURLs(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseURLs error = %v", err)
|
||||
}
|
||||
if len(urls) != 2 {
|
||||
t.Fatalf("去重后应有 2 个 url, 实际 %d: %v", len(urls), urls)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseURLsMissingFile 测试缺失文件返回错误
|
||||
func TestParseURLsMissingFile(t *testing.T) {
|
||||
fv := &FlagVars{URLsFile: "nonexistent-urls.txt"}
|
||||
_, err := parseURLs(fv)
|
||||
if err == nil {
|
||||
t.Fatal("缺失文件应返回错误")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// parseHashes
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// TestParseHashesEmpty 空输入返回空结果
|
||||
func TestParseHashesEmpty(t *testing.T) {
|
||||
fv := &FlagVars{}
|
||||
vals, bytes, err := parseHashes(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseHashes error = %v", err)
|
||||
}
|
||||
if len(vals) != 0 || len(bytes) != 0 {
|
||||
t.Fatalf("空输入应返回空结果, vals=%v bytes=%v", vals, bytes)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseHashesValidNTLM 纯 32 字符 hex hash
|
||||
func TestParseHashesValidNTLM(t *testing.T) {
|
||||
hash := "aabbccddeeff00112233445566778899"
|
||||
fv := &FlagVars{HashValue: hash}
|
||||
vals, hashBytes, err := parseHashes(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseHashes error = %v", err)
|
||||
}
|
||||
if len(vals) != 1 || vals[0] != hash {
|
||||
t.Fatalf("vals = %v, want [%s]", vals, hash)
|
||||
}
|
||||
if len(hashBytes) != 1 || len(hashBytes[0]) != 16 {
|
||||
t.Fatalf("hashBytes length wrong: %v", hashBytes)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseHashesLMNTFormat LM:NT 格式,提取 NT 部分
|
||||
func TestParseHashesLMNTFormat(t *testing.T) {
|
||||
lm := "aad3b435b51404eeaad3b435b51404ee"
|
||||
nt := "31d6cfe0d16ae931b73c59d7e0c089c0"
|
||||
fv := &FlagVars{HashValue: lm + ":" + nt}
|
||||
vals, _, err := parseHashes(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseHashes error = %v", err)
|
||||
}
|
||||
if len(vals) != 1 || vals[0] != nt {
|
||||
t.Fatalf("vals = %v, want [%s]", vals, nt)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseHashesInvalidLength hash 长度不是 32 → error
|
||||
func TestParseHashesInvalidLength(t *testing.T) {
|
||||
fv := &FlagVars{HashValue: "tooshort"}
|
||||
_, _, err := parseHashes(fv)
|
||||
if err == nil {
|
||||
t.Fatal("hash 长度不足应返回错误")
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseHashesInvalidHex 32 字符但含非 hex 字符 → error
|
||||
func TestParseHashesInvalidHex(t *testing.T) {
|
||||
fv := &FlagVars{HashValue: "zzzzzzzzzzzzzzzzzzzzzzzzzzzzzzzz"}
|
||||
_, _, err := parseHashes(fv)
|
||||
if err == nil {
|
||||
t.Fatal("非 hex 字符应返回错误")
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseHashesMissingFile hash 文件不存在 → error
|
||||
func TestParseHashesMissingFile(t *testing.T) {
|
||||
fv := &FlagVars{HashFile: "nonexistent-hashes.txt"}
|
||||
_, _, err := parseHashes(fv)
|
||||
if err == nil {
|
||||
t.Fatal("缺失 hash 文件应返回错误")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// parseUsernames
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// TestParseUsernamesEmpty 空输入返回空结果
|
||||
func TestParseUsernamesEmpty(t *testing.T) {
|
||||
fv := &FlagVars{}
|
||||
got, err := parseUsernames(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseUsernames error = %v", err)
|
||||
}
|
||||
if len(got) != 0 {
|
||||
t.Fatalf("空输入应返回空, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseUsernamesCommaSeparated 逗号分隔多用户
|
||||
func TestParseUsernamesCommaSeparated(t *testing.T) {
|
||||
fv := &FlagVars{Username: "admin, root, admin"} // 含重复和空格
|
||||
got, err := parseUsernames(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseUsernames error = %v", err)
|
||||
}
|
||||
want := []string{"admin", "root"}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("got %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseUsernamesAddUsers AddUsers 追加去重
|
||||
func TestParseUsernamesAddUsers(t *testing.T) {
|
||||
fv := &FlagVars{
|
||||
Username: "admin",
|
||||
AddUsers: "root,admin", // admin 重复
|
||||
}
|
||||
got, err := parseUsernames(fv)
|
||||
if err != nil {
|
||||
t.Fatalf("parseUsernames error = %v", err)
|
||||
}
|
||||
want := []string{"admin", "root"}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Fatalf("got %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestParseUsernamesMissingFile 缺失用户文件 → error
|
||||
func TestParseUsernamesMissingFile(t *testing.T) {
|
||||
fv := &FlagVars{UsersFile: "nonexistent-users.txt"}
|
||||
_, err := parseUsernames(fv)
|
||||
if err == nil {
|
||||
t.Fatal("缺失用户文件应返回错误")
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// cloneStringSlice
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// TestCloneStringSliceNil nil 输入返回 nil
|
||||
func TestCloneStringSliceNil(t *testing.T) {
|
||||
got := cloneStringSlice(nil)
|
||||
if got != nil {
|
||||
t.Fatalf("nil 输入应返回 nil, got %v", got)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCloneStringSliceEmpty 空切片:append 无元素结果为 nil,len 为 0
|
||||
func TestCloneStringSliceEmpty(t *testing.T) {
|
||||
got := cloneStringSlice([]string{})
|
||||
if len(got) != 0 {
|
||||
t.Fatalf("got len %d, want 0", len(got))
|
||||
}
|
||||
}
|
||||
|
||||
// TestCloneStringSliceCopiesValues 正常切片:值正确且独立
|
||||
func TestCloneStringSliceCopiesValues(t *testing.T) {
|
||||
src := []string{"a", "b", "c"}
|
||||
got := cloneStringSlice(src)
|
||||
if !reflect.DeepEqual(got, src) {
|
||||
t.Fatalf("got %v, want %v", got, src)
|
||||
}
|
||||
// 修改 clone 不影响原始
|
||||
got[0] = "mutated"
|
||||
if src[0] != "a" {
|
||||
t.Fatal("cloneStringSlice 返回的切片与源共享底层数组")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1728,3 +1728,420 @@ func TestManager_ConcurrentSave(t *testing.T) {
|
||||
t.Logf("✓ 并发保存测试通过(%d个goroutine,每个%d次,输出%d行)",
|
||||
numGoroutines, savesPerGoroutine, len(lines))
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// TXTWriter - 内部格式化函数覆盖率测试
|
||||
// =============================================================================
|
||||
|
||||
// newTestTXTWriter 创建用于单元测试的 TXTWriter(写到临时文件,调用方负责 Close)
|
||||
func newTestTXTWriter(t *testing.T) *TXTWriter {
|
||||
t.Helper()
|
||||
w, err := NewTXTWriter(filepath.Join(t.TempDir(), "unit.txt"))
|
||||
if err != nil {
|
||||
t.Fatalf("创建 TXTWriter 失败: %v", err)
|
||||
}
|
||||
return w
|
||||
}
|
||||
|
||||
// TestFormatServiceLine 覆盖 formatServiceLine 的各分支
|
||||
func TestFormatServiceLine(t *testing.T) {
|
||||
w := newTestTXTWriter(t)
|
||||
defer w.Close()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
details map[string]interface{}
|
||||
want []string // 输出中必须包含的子串
|
||||
notwant []string // 输出中不应包含的子串
|
||||
}{
|
||||
{
|
||||
name: "非web服务带service和banner",
|
||||
details: map[string]interface{}{
|
||||
"port": 22,
|
||||
"service": "ssh",
|
||||
"banner": "OpenSSH_8.0",
|
||||
},
|
||||
want: []string{"ssh", "OpenSSH_8.0"},
|
||||
notwant: []string{"http://", "https://"},
|
||||
},
|
||||
{
|
||||
name: "非web服务只有service",
|
||||
details: map[string]interface{}{
|
||||
"port": 3306,
|
||||
"service": "mysql",
|
||||
},
|
||||
want: []string{"mysql"},
|
||||
notwant: []string{"http://"},
|
||||
},
|
||||
{
|
||||
name: "非web服务无banner",
|
||||
details: map[string]interface{}{
|
||||
"port": 21,
|
||||
"service": "ftp",
|
||||
},
|
||||
want: []string{"ftp"},
|
||||
},
|
||||
{
|
||||
name: "service=http 走 web 分支",
|
||||
details: map[string]interface{}{
|
||||
"port": 80,
|
||||
"service": "http",
|
||||
"title": "Home",
|
||||
"status": 200,
|
||||
},
|
||||
want: []string{"http://", "Home"},
|
||||
notwant: []string{"ssh"},
|
||||
},
|
||||
{
|
||||
name: "service=https 走 web 分支",
|
||||
details: map[string]interface{}{
|
||||
"port": 443,
|
||||
"service": "https",
|
||||
"title": "Secure",
|
||||
"status": 200,
|
||||
},
|
||||
want: []string{"https://", "Secure"},
|
||||
},
|
||||
{
|
||||
name: "is_web=true 走 web 分支",
|
||||
details: map[string]interface{}{
|
||||
"port": 8080,
|
||||
"is_web": true,
|
||||
"title": "Dashboard",
|
||||
"status": 302,
|
||||
},
|
||||
want: []string{"http://", "Dashboard"},
|
||||
},
|
||||
{
|
||||
name: "有 status 字段触发 web 分支",
|
||||
details: map[string]interface{}{
|
||||
"port": 8080,
|
||||
"status": 200,
|
||||
},
|
||||
want: []string{"http://"},
|
||||
},
|
||||
{
|
||||
name: "有 server 字段触发 web 分支",
|
||||
details: map[string]interface{}{
|
||||
"port": 8080,
|
||||
"server": "nginx",
|
||||
},
|
||||
want: []string{"http://", "nginx"},
|
||||
},
|
||||
{
|
||||
name: "banner 含控制字符被转义",
|
||||
details: map[string]interface{}{
|
||||
"port": 9999,
|
||||
"service": "custom",
|
||||
"banner": "hello\nworld\r\n",
|
||||
},
|
||||
want: []string{"\\n", "\\r"},
|
||||
notwant: []string{"http://"},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := &ScanResult{
|
||||
Target: "192.168.1.1",
|
||||
Type: TypeService,
|
||||
Details: tt.details,
|
||||
}
|
||||
got := w.formatServiceLine(result)
|
||||
for _, s := range tt.want {
|
||||
if !strings.Contains(got, s) {
|
||||
t.Errorf("formatServiceLine() = %q,缺少 %q", got, s)
|
||||
}
|
||||
}
|
||||
for _, s := range tt.notwant {
|
||||
if strings.Contains(got, s) {
|
||||
t.Errorf("formatServiceLine() = %q,不应含 %q", got, s)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestGetFingerprints 覆盖 getFingerprints 的各类型分支
|
||||
func TestGetFingerprints(t *testing.T) {
|
||||
w := newTestTXTWriter(t)
|
||||
defer w.Close()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
details map[string]interface{}
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "nil fingerprints",
|
||||
details: map[string]interface{}{},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "[]string 非空",
|
||||
details: map[string]interface{}{"fingerprints": []string{"nginx", "php"}},
|
||||
want: "[nginx,php]",
|
||||
},
|
||||
{
|
||||
name: "[]string 空slice",
|
||||
details: map[string]interface{}{"fingerprints": []string{}},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "[]interface{} 非空",
|
||||
details: map[string]interface{}{"fingerprints": []interface{}{"wordpress", "jquery"}},
|
||||
want: "[wordpress,jquery]",
|
||||
},
|
||||
{
|
||||
name: "[]interface{} 含数字",
|
||||
details: map[string]interface{}{"fingerprints": []interface{}{"apache", 2}},
|
||||
want: "[apache,2]",
|
||||
},
|
||||
{
|
||||
name: "[]interface{} 空slice",
|
||||
details: map[string]interface{}{"fingerprints": []interface{}{}},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "不支持的类型返回空",
|
||||
details: map[string]interface{}{"fingerprints": "just-a-string"},
|
||||
want: "",
|
||||
},
|
||||
{
|
||||
name: "单个元素",
|
||||
details: map[string]interface{}{"fingerprints": []string{"tomcat"}},
|
||||
want: "[tomcat]",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := &ScanResult{Target: "1.2.3.4", Details: tt.details}
|
||||
got := w.getFingerprints(result)
|
||||
if got != tt.want {
|
||||
t.Errorf("getFingerprints() = %q,want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestFormatVulnLine 覆盖 formatVulnLine 的各分支
|
||||
func TestFormatVulnLine(t *testing.T) {
|
||||
w := newTestTXTWriter(t)
|
||||
defer w.Close()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
target string
|
||||
status string
|
||||
details map[string]interface{}
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "weak_credential 带 service",
|
||||
target: "192.168.1.1:22",
|
||||
details: map[string]interface{}{
|
||||
"type": "weak_credential",
|
||||
"service": "ssh",
|
||||
"username": "root",
|
||||
"password": "123456",
|
||||
},
|
||||
want: "192.168.1.1:22 ssh root/123456",
|
||||
},
|
||||
{
|
||||
name: "weak_credential 不带 service",
|
||||
target: "192.168.1.1:3306",
|
||||
details: map[string]interface{}{
|
||||
"type": "weak_credential",
|
||||
"username": "admin",
|
||||
"password": "pass",
|
||||
},
|
||||
want: "192.168.1.1:3306 admin/pass",
|
||||
},
|
||||
{
|
||||
name: "有 vulnerability 字段",
|
||||
target: "10.0.0.1",
|
||||
details: map[string]interface{}{
|
||||
"type": "poc",
|
||||
"vulnerability": "CVE-2024-1234",
|
||||
},
|
||||
want: "10.0.0.1 CVE-2024-1234",
|
||||
},
|
||||
{
|
||||
name: "无 vulnerability 字段回退到 status",
|
||||
target: "10.0.0.2",
|
||||
status: "VULNERABLE",
|
||||
details: map[string]interface{}{
|
||||
"type": "unknown",
|
||||
},
|
||||
want: "10.0.0.2 VULNERABLE",
|
||||
},
|
||||
{
|
||||
name: "空 details 回退到 status",
|
||||
target: "10.0.0.3",
|
||||
status: "poc_hit",
|
||||
details: map[string]interface{}{},
|
||||
want: "10.0.0.3 poc_hit",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := &ScanResult{
|
||||
Target: tt.target,
|
||||
Status: tt.status,
|
||||
Type: TypeVuln,
|
||||
Details: tt.details,
|
||||
}
|
||||
got := w.formatVulnLine(result)
|
||||
if got != tt.want {
|
||||
t.Errorf("formatVulnLine() = %q,want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestIsWebService 覆盖 isWebService 的各判断分支
|
||||
func TestIsWebService(t *testing.T) {
|
||||
w := newTestTXTWriter(t)
|
||||
defer w.Close()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
details map[string]interface{}
|
||||
want bool
|
||||
}{
|
||||
{
|
||||
name: "is_web=true",
|
||||
details: map[string]interface{}{"is_web": true},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "is_web=false 无其他标志",
|
||||
details: map[string]interface{}{"is_web": false},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "有 status 字段",
|
||||
details: map[string]interface{}{"status": 200},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "status=nil 不触发",
|
||||
details: map[string]interface{}{},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "有非空 server 字段",
|
||||
details: map[string]interface{}{"server": "nginx"},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "空 server 字段不触发",
|
||||
details: map[string]interface{}{"server": ""},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "service=http",
|
||||
details: map[string]interface{}{"service": "http"},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "service=https",
|
||||
details: map[string]interface{}{"service": "https"},
|
||||
want: true,
|
||||
},
|
||||
{
|
||||
name: "service=ssh 不是 web",
|
||||
details: map[string]interface{}{"service": "ssh"},
|
||||
want: false,
|
||||
},
|
||||
{
|
||||
name: "nil Details",
|
||||
details: nil,
|
||||
want: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := &ScanResult{Target: "1.2.3.4", Details: tt.details}
|
||||
got := w.isWebService(result)
|
||||
if got != tt.want {
|
||||
t.Errorf("isWebService() = %v,want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestWebProtocol 覆盖 webProtocol 的各判断分支
|
||||
func TestWebProtocol(t *testing.T) {
|
||||
w := newTestTXTWriter(t)
|
||||
defer w.Close()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
target string
|
||||
details map[string]interface{}
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "protocol=https 直接返回",
|
||||
target: "1.2.3.4:8443",
|
||||
details: map[string]interface{}{"protocol": "https"},
|
||||
want: "https",
|
||||
},
|
||||
{
|
||||
name: "protocol=http 直接返回",
|
||||
target: "1.2.3.4:8080",
|
||||
details: map[string]interface{}{"protocol": "http"},
|
||||
want: "http",
|
||||
},
|
||||
{
|
||||
name: "protocol=HTTPS 大小写不敏感",
|
||||
target: "1.2.3.4:443",
|
||||
details: map[string]interface{}{"protocol": "HTTPS"},
|
||||
want: "https",
|
||||
},
|
||||
{
|
||||
name: "service=https 回退",
|
||||
target: "1.2.3.4:8080",
|
||||
details: map[string]interface{}{"service": "https"},
|
||||
want: "https",
|
||||
},
|
||||
{
|
||||
name: "target 含 :443 回退 https",
|
||||
target: "example.com:443",
|
||||
details: map[string]interface{}{},
|
||||
want: "https",
|
||||
},
|
||||
{
|
||||
name: "无任何标志默认 http",
|
||||
target: "1.2.3.4:8080",
|
||||
details: map[string]interface{}{},
|
||||
want: "http",
|
||||
},
|
||||
{
|
||||
name: "service=http 默认 http",
|
||||
target: "1.2.3.4:80",
|
||||
details: map[string]interface{}{"service": "http"},
|
||||
want: "http",
|
||||
},
|
||||
{
|
||||
name: "protocol 为其他值走 service 分支",
|
||||
target: "1.2.3.4:9000",
|
||||
details: map[string]interface{}{"protocol": "tcp", "service": "https"},
|
||||
want: "https",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := &ScanResult{Target: tt.target, Details: tt.details}
|
||||
got := w.webProtocol(result, tt.target)
|
||||
if got != tt.want {
|
||||
t.Errorf("webProtocol() = %q,want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -105,6 +105,24 @@ func TestInitOutputValidationAndDefaultExtension(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestCloseOutputWithStdoutWriter(t *testing.T) {
|
||||
preserveOutputAPIGlobals(t)
|
||||
|
||||
// 初始化 silent 模式以创建 StdoutWriter
|
||||
flagVars = &FlagVars{Silent: true, DisableSave: true}
|
||||
if err := InitOutput(); err != nil {
|
||||
t.Fatalf("InitOutput silent error = %v", err)
|
||||
}
|
||||
if StdoutWriter == nil {
|
||||
t.Fatal("StdoutWriter 应在 Silent 模式下被初始化")
|
||||
}
|
||||
|
||||
// CloseOutput 应正常关闭 StdoutWriter
|
||||
if err := CloseOutput(); err != nil {
|
||||
t.Fatalf("CloseOutput with StdoutWriter error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSaveResultFacadeCallbackAndDisabledSave(t *testing.T) {
|
||||
preserveOutputAPIGlobals(t)
|
||||
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
package parsers
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"os"
|
||||
"reflect"
|
||||
"strings"
|
||||
@@ -166,3 +168,777 @@ func (s *closeTrackingSource) Close() error {
|
||||
s.closed = true
|
||||
return s.err
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// newHostSource 分支覆盖
|
||||
// =============================================================================
|
||||
|
||||
// TestNewHostSource_Shortcuts 验证 192/172/10 快捷方式展开为正确 CIDR
|
||||
func TestNewHostSource_Shortcuts(t *testing.T) {
|
||||
cases := []struct {
|
||||
input string
|
||||
wantFirst string
|
||||
}{
|
||||
{"192", "192.168.0.1"},
|
||||
{"172", "172.16.0.1"},
|
||||
{"10", "10.0.0.1"},
|
||||
}
|
||||
for _, c := range cases {
|
||||
t.Run(c.input, func(t *testing.T) {
|
||||
src, err := newHostSource(c.input)
|
||||
if err != nil {
|
||||
t.Fatalf("newHostSource(%q) error = %v", c.input, err)
|
||||
}
|
||||
defer src.Close()
|
||||
host, ok, err := src.Next()
|
||||
if err != nil || !ok {
|
||||
t.Fatalf("Next() = %q/%v/%v", host, ok, err)
|
||||
}
|
||||
if host != c.wantFirst {
|
||||
t.Errorf("first host = %q, 期望 %q", host, c.wantFirst)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewHostSource_CIDRBranch 验证含 "/" 走 CIDR 分支
|
||||
func TestNewHostSource_CIDRBranch(t *testing.T) {
|
||||
src, err := newHostSource("10.0.0.0/30")
|
||||
if err != nil {
|
||||
t.Fatalf("newHostSource CIDR error = %v", err)
|
||||
}
|
||||
defer src.Close()
|
||||
host, ok, _ := src.Next()
|
||||
if !ok || host != "10.0.0.1" {
|
||||
t.Errorf("CIDR first host = %q, 期望 10.0.0.1", host)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewHostSource_InvalidCIDR 无效 CIDR 返回错误
|
||||
func TestNewHostSource_InvalidCIDR(t *testing.T) {
|
||||
_, err := newHostSource("999.0.0.0/24")
|
||||
if err == nil {
|
||||
t.Error("无效 CIDR 应返回 error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewHostSource_RangeBranch 验证 a-b 格式走 range 分支
|
||||
func TestNewHostSource_RangeBranch(t *testing.T) {
|
||||
src, err := newHostSource("192.168.1.5-192.168.1.7")
|
||||
if err != nil {
|
||||
t.Fatalf("newHostSource range error = %v", err)
|
||||
}
|
||||
defer src.Close()
|
||||
|
||||
var got []string
|
||||
for {
|
||||
h, ok, err := src.Next()
|
||||
if err != nil {
|
||||
t.Fatalf("Next() error = %v", err)
|
||||
}
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
got = append(got, h)
|
||||
}
|
||||
want := []string{"192.168.1.5", "192.168.1.6", "192.168.1.7"}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("range hosts = %v, 期望 %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewHostSource_RangeShortTail 验证短尾写法 x.x.x.a-b
|
||||
func TestNewHostSource_RangeShortTail(t *testing.T) {
|
||||
src, err := newHostSource("10.0.0.3-5")
|
||||
if err != nil {
|
||||
t.Fatalf("newHostSource short-tail range error = %v", err)
|
||||
}
|
||||
defer src.Close()
|
||||
|
||||
var got []string
|
||||
for {
|
||||
h, ok, err := src.Next()
|
||||
if err != nil {
|
||||
t.Fatalf("Next() error = %v", err)
|
||||
}
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
got = append(got, h)
|
||||
}
|
||||
want := []string{"10.0.0.3", "10.0.0.4", "10.0.0.5"}
|
||||
if !reflect.DeepEqual(got, want) {
|
||||
t.Errorf("short-tail range = %v, 期望 %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewHostSource_SingleHost 验证普通主机名走 singleHostSource 分支
|
||||
func TestNewHostSource_SingleHost(t *testing.T) {
|
||||
src, err := newHostSource("example.com")
|
||||
if err != nil {
|
||||
t.Fatalf("newHostSource single error = %v", err)
|
||||
}
|
||||
defer src.Close()
|
||||
|
||||
host, ok, err := src.Next()
|
||||
if err != nil || !ok || host != "example.com" {
|
||||
t.Errorf("single host = %q/%v/%v, 期望 example.com/true/nil", host, ok, err)
|
||||
}
|
||||
// 第二次应该耗尽
|
||||
_, ok, _ = src.Next()
|
||||
if ok {
|
||||
t.Error("singleHostSource 第二次 Next 应返回 ok=false")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// hostMatcher.add 分支覆盖
|
||||
// =============================================================================
|
||||
|
||||
// TestHostMatcherAdd_192Shortcut 验证 add("192") 展开为 192.168.0.0/16
|
||||
func TestHostMatcherAdd_192Shortcut(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.add("192"); err != nil {
|
||||
t.Fatalf("add(192) error = %v", err)
|
||||
}
|
||||
if !m.match("192.168.1.100") {
|
||||
t.Error("192.168.1.100 应命中 192.168.0.0/16")
|
||||
}
|
||||
if m.match("10.0.0.1") {
|
||||
t.Error("10.0.0.1 不应命中")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHostMatcherAdd_172Shortcut 验证 add("172")
|
||||
func TestHostMatcherAdd_172Shortcut(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.add("172"); err != nil {
|
||||
t.Fatalf("add(172) error = %v", err)
|
||||
}
|
||||
if !m.match("172.16.0.1") {
|
||||
t.Error("172.16.0.1 应命中 172.16.0.0/12")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHostMatcherAdd_10Shortcut 验证 add("10")
|
||||
func TestHostMatcherAdd_10Shortcut(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.add("10"); err != nil {
|
||||
t.Fatalf("add(10) error = %v", err)
|
||||
}
|
||||
if !m.match("10.1.2.3") {
|
||||
t.Error("10.1.2.3 应命中 10.0.0.0/8")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHostMatcherAdd_CIDR 验证 add 处理 CIDR 字符串
|
||||
func TestHostMatcherAdd_CIDR(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.add("192.168.5.0/24"); err != nil {
|
||||
t.Fatalf("add CIDR error = %v", err)
|
||||
}
|
||||
if !m.match("192.168.5.10") {
|
||||
t.Error("192.168.5.10 应命中 /24")
|
||||
}
|
||||
if m.match("192.168.6.10") {
|
||||
t.Error("192.168.6.10 不应命中")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHostMatcherAdd_Range 验证 add 处理 a-b 范围
|
||||
func TestHostMatcherAdd_Range(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.add("10.0.0.10-10.0.0.20"); err != nil {
|
||||
t.Fatalf("add range error = %v", err)
|
||||
}
|
||||
if !m.match("10.0.0.15") {
|
||||
t.Error("10.0.0.15 应命中范围")
|
||||
}
|
||||
if m.match("10.0.0.9") || m.match("10.0.0.21") {
|
||||
t.Error("边界外不应命中")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHostMatcherAdd_ExactHost 验证 add 处理普通主机名(exact 分支)
|
||||
func TestHostMatcherAdd_ExactHost(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.add("myhost.local"); err != nil {
|
||||
t.Fatalf("add exact error = %v", err)
|
||||
}
|
||||
if !m.match("myhost.local") {
|
||||
t.Error("exact 主机名应命中")
|
||||
}
|
||||
if m.match("other.local") {
|
||||
t.Error("其他主机名不应命中")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHostMatcherAdd_MultipleComma 验证逗号分隔多个值
|
||||
func TestHostMatcherAdd_MultipleComma(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.add("host1.com, host2.com, 192.168.1.0/30"); err != nil {
|
||||
t.Fatalf("add comma-separated error = %v", err)
|
||||
}
|
||||
if !m.match("host1.com") || !m.match("host2.com") || !m.match("192.168.1.1") {
|
||||
t.Error("逗号分隔的值应全部命中")
|
||||
}
|
||||
}
|
||||
|
||||
// TestHostMatcherAdd_EmptyEntry 逗号中间空串不报错
|
||||
func TestHostMatcherAdd_EmptyEntry(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.add(",,,"); err != nil {
|
||||
t.Fatalf("全空逗号不应报错: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestHostMatcherAdd_InvalidCIDR 无效 CIDR 返回 error
|
||||
func TestHostMatcherAdd_InvalidCIDR(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.add("999.0.0.0/8"); err == nil {
|
||||
t.Error("无效 CIDR 应返回 error")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// fileHostSource.Next 分支覆盖
|
||||
// =============================================================================
|
||||
|
||||
// TestFileHostSourceNext_SkipsEmptyAndComments 验证空行和注释行被跳过
|
||||
func TestFileHostSourceNext_SkipsEmptyAndComments(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := dir + "/hosts.txt"
|
||||
content := "\n# this is a comment\n\n \n10.0.0.1\n# another comment\n10.0.0.2\n"
|
||||
if err := os.WriteFile(path, []byte(content), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile error = %v", err)
|
||||
}
|
||||
|
||||
iter, err := NewHostIterator("", path)
|
||||
if err != nil {
|
||||
t.Fatalf("NewHostIterator error = %v", err)
|
||||
}
|
||||
defer iter.Close()
|
||||
|
||||
batch, err := iter.NextBatch(context.Background(), 10)
|
||||
if err != nil {
|
||||
t.Fatalf("NextBatch error = %v", err)
|
||||
}
|
||||
want := []string{"10.0.0.1", "10.0.0.2"}
|
||||
if !reflect.DeepEqual(batch, want) {
|
||||
t.Errorf("batch = %v, 期望 %v", batch, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFileHostSourceNext_MultipleSources 验证文件中每行多个 host(逗号分隔)走 multiHostSource 分支
|
||||
func TestFileHostSourceNext_MultipleSources(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := dir + "/hosts.txt"
|
||||
// 一行两个 host,触发 multiHostSource 分支
|
||||
content := "10.0.0.1,10.0.0.2\n10.0.0.3\n"
|
||||
if err := os.WriteFile(path, []byte(content), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile error = %v", err)
|
||||
}
|
||||
|
||||
iter, err := NewHostIterator("", path)
|
||||
if err != nil {
|
||||
t.Fatalf("NewHostIterator error = %v", err)
|
||||
}
|
||||
defer iter.Close()
|
||||
|
||||
batch, err := iter.NextBatch(context.Background(), 10)
|
||||
if err != nil {
|
||||
t.Fatalf("NextBatch error = %v", err)
|
||||
}
|
||||
want := []string{"10.0.0.1", "10.0.0.2", "10.0.0.3"}
|
||||
if !reflect.DeepEqual(batch, want) {
|
||||
t.Errorf("batch = %v, 期望 %v", batch, want)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFileHostSourceNext_InvalidLineSkipped 无效行(解析失败)被跳过不报错
|
||||
func TestFileHostSourceNext_InvalidLineSkipped(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := dir + "/hosts.txt"
|
||||
// 包含无效 CIDR,应被跳过
|
||||
content := "999.0.0.0/8\n10.0.0.1\n"
|
||||
if err := os.WriteFile(path, []byte(content), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile error = %v", err)
|
||||
}
|
||||
|
||||
iter, err := NewHostIterator("", path)
|
||||
if err != nil {
|
||||
t.Fatalf("NewHostIterator error = %v", err)
|
||||
}
|
||||
defer iter.Close()
|
||||
|
||||
batch, err := iter.NextBatch(context.Background(), 10)
|
||||
if err != nil {
|
||||
t.Fatalf("NextBatch error = %v", err)
|
||||
}
|
||||
// 无效行被跳过,只返回有效行
|
||||
if len(batch) != 1 || batch[0] != "10.0.0.1" {
|
||||
t.Errorf("batch = %v, 期望 [10.0.0.1]", batch)
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// NewHostIterator 错误路径
|
||||
// =============================================================================
|
||||
|
||||
// TestNewHostIterator_InvalidFilename 不存在的文件应返回 error
|
||||
func TestNewHostIterator_InvalidFilename(t *testing.T) {
|
||||
_, err := NewHostIterator("", "/nonexistent/path/hosts.txt")
|
||||
if err == nil {
|
||||
t.Error("不存在的文件应返回 error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewHostIterator_InvalidHost host 解析失败时应返回 error(并关闭已打开的文件 source)
|
||||
func TestNewHostIterator_InvalidHost_WithFile(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := dir + "/hosts.txt"
|
||||
if err := os.WriteFile(path, []byte("10.0.0.1\n"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile: %v", err)
|
||||
}
|
||||
// 无效 CIDR 会让 newHostSources 失败
|
||||
_, err := NewHostIterator("999.0.0.0/8", path)
|
||||
if err == nil {
|
||||
t.Error("无效 host 应返回 error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewHostIterator_InvalidExclude exclude 参数无效时应返回 error
|
||||
func TestNewHostIterator_InvalidExclude(t *testing.T) {
|
||||
_, err := NewHostIterator("10.0.0.1", "", "999.0.0.0/8")
|
||||
if err == nil {
|
||||
t.Error("无效 exclude 应返回 error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewHostIterator_EmptyExcludeSkipped 空白 exclude 条目应被跳过,不报错
|
||||
func TestNewHostIterator_EmptyExcludeSkipped(t *testing.T) {
|
||||
iter, err := NewHostIterator("10.0.0.1", "", " ", "")
|
||||
if err != nil {
|
||||
t.Fatalf("空白 exclude 不应报错: %v", err)
|
||||
}
|
||||
defer iter.Close()
|
||||
host, ok, err := iter.Next()
|
||||
if err != nil || !ok || host != "10.0.0.1" {
|
||||
t.Errorf("Next() = %q/%v/%v", host, ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Close 路径
|
||||
// =============================================================================
|
||||
|
||||
// TestClose_Nil nil HostIterator Close 不 panic
|
||||
func TestClose_Nil(t *testing.T) {
|
||||
var it *HostIterator
|
||||
if err := it.Close(); err != nil {
|
||||
t.Errorf("nil Close 应返回 nil, 得到 %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestClose_WithCurrent 有 current source 时 Close 应关闭它
|
||||
func TestClose_WithCurrent(t *testing.T) {
|
||||
src := &closeTrackingSource{}
|
||||
it := &HostIterator{current: src}
|
||||
if err := it.Close(); err != nil {
|
||||
t.Errorf("Close error = %v", err)
|
||||
}
|
||||
if !src.closed {
|
||||
t.Error("current source 应被关闭")
|
||||
}
|
||||
if it.current != nil {
|
||||
t.Error("Close 后 current 应为 nil")
|
||||
}
|
||||
}
|
||||
|
||||
// TestClose_SourcesError Close 中 source 返回 error 应被记录
|
||||
func TestClose_SourcesError(t *testing.T) {
|
||||
errSrc := &closeTrackingSource{err: errors.New("close error")}
|
||||
it := &HostIterator{sources: []hostSource{errSrc}}
|
||||
err := it.Close()
|
||||
if err == nil {
|
||||
t.Error("source Close 失败时应返回 error")
|
||||
}
|
||||
if !errSrc.closed {
|
||||
t.Error("出错的 source 也应被调用 Close")
|
||||
}
|
||||
}
|
||||
|
||||
// TestClose_CurrentErrorThenSources current Close 报错,后续 source Close 成功,返回 current 的 error
|
||||
func TestClose_CurrentErrorThenSources(t *testing.T) {
|
||||
currentSrc := &closeTrackingSource{err: errors.New("current close error")}
|
||||
otherSrc := &closeTrackingSource{}
|
||||
it := &HostIterator{
|
||||
current: currentSrc,
|
||||
sources: []hostSource{otherSrc},
|
||||
}
|
||||
err := it.Close()
|
||||
if err == nil {
|
||||
t.Error("应返回 current 的 error")
|
||||
}
|
||||
if !currentSrc.closed || !otherSrc.closed {
|
||||
t.Error("两个 source 都应被关闭")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// Next 错误路径
|
||||
// =============================================================================
|
||||
|
||||
// errorSource 让 Next() 返回 error
|
||||
type errorSource struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (s *errorSource) Next() (string, bool, error) { return "", false, s.err }
|
||||
func (s *errorSource) Close() error { return nil }
|
||||
|
||||
// errorOnCloseSource Next 返回 ok=false,Close 返回 error
|
||||
type errorOnCloseSource struct {
|
||||
err error
|
||||
}
|
||||
|
||||
func (s *errorOnCloseSource) Next() (string, bool, error) { return "", false, nil }
|
||||
func (s *errorOnCloseSource) Close() error { return s.err }
|
||||
|
||||
// TestNext_SourceNextError source.Next() 返回 error 时 iter.Next 应透传
|
||||
func TestNext_SourceNextError(t *testing.T) {
|
||||
it := &HostIterator{
|
||||
sources: []hostSource{&errorSource{err: errors.New("next error")}},
|
||||
}
|
||||
_, _, err := it.Next()
|
||||
if err == nil {
|
||||
t.Error("source Next error 应透传")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNext_SourceCloseError 源耗尽时 Close 报错应透传
|
||||
func TestNext_SourceCloseError(t *testing.T) {
|
||||
it := &HostIterator{
|
||||
sources: []hostSource{&errorOnCloseSource{err: errors.New("close error")}},
|
||||
}
|
||||
_, _, err := it.Next()
|
||||
if err == nil {
|
||||
t.Error("source 耗尽时 Close error 应透传")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// NextBatch 边界条件
|
||||
// =============================================================================
|
||||
|
||||
// TestNextBatch_ZeroSize size=0 应使用 DefaultHostBatchSize(实际受源数量限制)
|
||||
func TestNextBatch_ZeroSize(t *testing.T) {
|
||||
iter, err := NewHostIterator("10.0.0.1", "")
|
||||
if err != nil {
|
||||
t.Fatalf("NewHostIterator: %v", err)
|
||||
}
|
||||
defer iter.Close()
|
||||
|
||||
// size=0 触发默认 DefaultHostBatchSize 分支,源只有一个 host
|
||||
batch, err := iter.NextBatch(context.Background(), 0)
|
||||
if err != nil {
|
||||
t.Fatalf("NextBatch(0) error = %v", err)
|
||||
}
|
||||
if len(batch) != 1 || batch[0] != "10.0.0.1" {
|
||||
t.Errorf("batch = %v, 期望 [10.0.0.1]", batch)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNextBatch_NegativeSize size<0 也应使用默认值
|
||||
func TestNextBatch_NegativeSize(t *testing.T) {
|
||||
iter, err := NewHostIterator("10.0.0.2", "")
|
||||
if err != nil {
|
||||
t.Fatalf("NewHostIterator: %v", err)
|
||||
}
|
||||
defer iter.Close()
|
||||
|
||||
batch, err := iter.NextBatch(context.Background(), -1)
|
||||
if err != nil {
|
||||
t.Fatalf("NextBatch(-1) error = %v", err)
|
||||
}
|
||||
if len(batch) != 1 || batch[0] != "10.0.0.2" {
|
||||
t.Errorf("batch = %v, 期望 [10.0.0.2]", batch)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNextBatch_ContextCancelled context 取消应立即返回
|
||||
func TestNextBatch_ContextCancelled(t *testing.T) {
|
||||
iter, err := NewHostIterator("10.0.0.0/8", "")
|
||||
if err != nil {
|
||||
t.Fatalf("NewHostIterator: %v", err)
|
||||
}
|
||||
defer iter.Close()
|
||||
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel() // 立即取消
|
||||
|
||||
_, err = iter.NextBatch(ctx, 100)
|
||||
if err == nil {
|
||||
t.Error("已取消的 context 应返回 error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNextBatch_DeduplicatesHosts 重复 host 只保留一个
|
||||
func TestNextBatch_DeduplicatesHosts(t *testing.T) {
|
||||
// 两个相同的单 host source
|
||||
it := &HostIterator{
|
||||
sources: []hostSource{
|
||||
&singleHostSource{host: "10.0.0.1"},
|
||||
&singleHostSource{host: "10.0.0.1"},
|
||||
},
|
||||
}
|
||||
batch, err := it.NextBatch(context.Background(), 10)
|
||||
if err != nil {
|
||||
t.Fatalf("NextBatch error = %v", err)
|
||||
}
|
||||
if len(batch) != 1 || batch[0] != "10.0.0.1" {
|
||||
t.Errorf("batch = %v, 期望去重为 [10.0.0.1]", batch)
|
||||
}
|
||||
}
|
||||
|
||||
// TestNextBatch_NextError Next 报错时应透传
|
||||
func TestNextBatch_NextError(t *testing.T) {
|
||||
it := &HostIterator{
|
||||
sources: []hostSource{&errorSource{err: errors.New("iter error")}},
|
||||
}
|
||||
_, err := it.NextBatch(context.Background(), 10)
|
||||
if err == nil {
|
||||
t.Error("Next error 应透传到 NextBatch")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// newRangeHostSource 错误路径
|
||||
// =============================================================================
|
||||
|
||||
// TestNewRangeHostSource_TooManyDashes 超过一个 "-" 应报错(实际按首个切分:a-b-c 被 Split 成 3 段)
|
||||
func TestNewRangeHostSource_TooManyDashes(t *testing.T) {
|
||||
// "a-b-c" Split by "-" 得到 3 段,len != 2,应报错
|
||||
_, err := newRangeHostSource("10.0.0.1-10.0.0.5-extra")
|
||||
if err == nil {
|
||||
t.Error("三段格式应报错")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewRangeHostSource_InvalidStartIP 起始 IP 无效
|
||||
func TestNewRangeHostSource_InvalidStartIP(t *testing.T) {
|
||||
_, err := newRangeHostSource("notanip-10.0.0.5")
|
||||
if err == nil {
|
||||
t.Error("无效起始 IP 应报错")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewRangeHostSource_InvalidShortTailNonNumeric 短尾不是数字应报错
|
||||
func TestNewRangeHostSource_InvalidShortTailNonNumeric(t *testing.T) {
|
||||
// 尾部 "xyz" 不是数字
|
||||
_, err := newRangeHostSource("10.0.0.1-xyz")
|
||||
if err == nil {
|
||||
t.Error("非数字短尾应报错")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewRangeHostSource_InvalidShortTailOver255 短尾超过 255 应报错
|
||||
func TestNewRangeHostSource_InvalidShortTailOver255(t *testing.T) {
|
||||
_, err := newRangeHostSource("10.0.0.1-300")
|
||||
if err == nil {
|
||||
t.Error("短尾 >255 应报错")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewRangeHostSource_StartGTEnd 起始 > 结束应报错
|
||||
func TestNewRangeHostSource_StartGTEnd(t *testing.T) {
|
||||
_, err := newRangeHostSource("10.0.0.200-10.0.0.100")
|
||||
if err == nil {
|
||||
t.Error("start > end 应报错")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewRangeHostSource_InvalidFullEndIP 完整结束 IP 无效(如 "10.0.0.999")
|
||||
func TestNewRangeHostSource_InvalidFullEndIP(t *testing.T) {
|
||||
// end IP 包含 "." 但无效
|
||||
_, err := newRangeHostSource("10.0.0.1-10.0.0.999")
|
||||
if err == nil {
|
||||
t.Error("无效结束 IP 应报错")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewRangeHostSource_ShortTailStartGTEnd 短尾导致 start > end 应报错
|
||||
func TestNewRangeHostSource_ShortTailStartGTEnd(t *testing.T) {
|
||||
_, err := newRangeHostSource("10.0.0.200-100")
|
||||
if err == nil {
|
||||
t.Error("短尾结果 start > end 应报错")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// hostMatcher.addRange 错误路径
|
||||
// =============================================================================
|
||||
|
||||
// TestAddRange_InvalidRange addRange 传入无效范围应报错
|
||||
func TestAddRange_InvalidRange(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.addRange("notvalid-range"); err == nil {
|
||||
t.Error("无效 range 应返回 error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestAddRange_ValidRange addRange 正常路径
|
||||
func TestAddRange_ValidRange(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
if err := m.addRange("10.0.0.10-10.0.0.20"); err != nil {
|
||||
t.Fatalf("addRange error = %v", err)
|
||||
}
|
||||
if !m.match("10.0.0.10") || !m.match("10.0.0.20") {
|
||||
t.Error("addRange 边界值应命中")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// hostMatcher.add 错误路径(shortcut 分支中 addCIDR 失败)
|
||||
// =============================================================================
|
||||
|
||||
// TestHostMatcherAdd_InvalidRange add 的 range 格式无效
|
||||
func TestHostMatcherAdd_InvalidRange(t *testing.T) {
|
||||
m := newHostMatcher()
|
||||
// 构造一个 looksLikeIPRange 通过但 newRangeHostSource 失败的字符串
|
||||
// "10.0.0.200-10.0.0.100" start>end 会报错
|
||||
if err := m.add("10.0.0.200-10.0.0.100"); err == nil {
|
||||
t.Error("无效 range (start>end) 应返回 error")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// newCIDRHostSource IPv6 路径
|
||||
// =============================================================================
|
||||
|
||||
// TestNewCIDRHostSource_IPv6Rejected IPv6 CIDR 应报错
|
||||
func TestNewCIDRHostSource_IPv6Rejected(t *testing.T) {
|
||||
_, err := newCIDRHostSource("2001:db8::/32")
|
||||
if err == nil {
|
||||
t.Error("IPv6 CIDR 应被拒绝")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// fileHostSource.Close 路径
|
||||
// =============================================================================
|
||||
|
||||
// TestFileHostSource_CloseWithCurrent fileHostSource.Close 时 current != nil 分支
|
||||
func TestFileHostSource_CloseWithCurrent(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := dir + "/hosts.txt"
|
||||
// 写入一个 CIDR,这样 fileHostSource 会持有 current source
|
||||
if err := os.WriteFile(path, []byte("10.0.0.0/30\n"), 0o600); err != nil {
|
||||
t.Fatalf("WriteFile: %v", err)
|
||||
}
|
||||
|
||||
src, err := newFileHostSource(path)
|
||||
if err != nil {
|
||||
t.Fatalf("newFileHostSource: %v", err)
|
||||
}
|
||||
// 触发 current 被设置
|
||||
_, _, _ = src.Next()
|
||||
// 此时 current 应非 nil,Close 应正常关闭它
|
||||
if err := src.Close(); err != nil {
|
||||
t.Errorf("Close with current error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestFileHostSource_CloseNilFile file 已经为 nil 时 Close 直接返回 nil
|
||||
func TestFileHostSource_CloseNilFile(t *testing.T) {
|
||||
src := &fileHostSource{file: nil}
|
||||
if err := src.Close(); err != nil {
|
||||
t.Errorf("nil file Close error = %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// multiHostSource.Close 路径
|
||||
// =============================================================================
|
||||
|
||||
// TestMultiHostSource_CloseWithCurrent Close 时 current != nil 分支
|
||||
func TestMultiHostSource_CloseWithCurrent(t *testing.T) {
|
||||
inner := &closeTrackingSource{}
|
||||
ms := &multiHostSource{current: inner}
|
||||
if err := ms.Close(); err != nil {
|
||||
t.Errorf("Close error = %v", err)
|
||||
}
|
||||
if !inner.closed {
|
||||
t.Error("current 应被关闭")
|
||||
}
|
||||
if ms.current != nil {
|
||||
t.Error("Close 后 current 应为 nil")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// ipToUint32 IPv6 路径
|
||||
// =============================================================================
|
||||
|
||||
// TestIpToUint32_IPv6ReturnsFalse IPv6 地址应返回 false
|
||||
func TestIpToUint32_IPv6ReturnsFalse(t *testing.T) {
|
||||
ip := net.ParseIP("2001:db8::1")
|
||||
_, ok := ipToUint32(ip)
|
||||
if ok {
|
||||
t.Error("IPv6 地址应返回 ok=false")
|
||||
}
|
||||
}
|
||||
|
||||
// TestIpToUint32_NilReturnsFalse nil IP 应返回 false
|
||||
func TestIpToUint32_NilReturnsFalse(t *testing.T) {
|
||||
_, ok := ipToUint32(nil)
|
||||
if ok {
|
||||
t.Error("nil IP 应返回 ok=false")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// 剩余未覆盖路径
|
||||
// =============================================================================
|
||||
|
||||
// TestFileHostSource_CurrentNextError fileHostSource.Next 中 current.Next() 报错应透传
|
||||
func TestFileHostSource_CurrentNextError(t *testing.T) {
|
||||
src := &fileHostSource{
|
||||
current: &errorSource{err: errors.New("inner error")},
|
||||
// scanner 为 nil——不会走到 scanner 分支
|
||||
scanner: bufio.NewScanner(strings.NewReader("")),
|
||||
}
|
||||
_, _, err := src.Next()
|
||||
if err == nil {
|
||||
t.Error("current.Next() 报错应透传")
|
||||
}
|
||||
}
|
||||
|
||||
// TestMultiHostSource_InnerNextError multiHostSource.Next 中内部 source.Next() 报错应透传
|
||||
func TestMultiHostSource_InnerNextError(t *testing.T) {
|
||||
ms := &multiHostSource{
|
||||
sources: []hostSource{&errorSource{err: errors.New("inner error")}},
|
||||
}
|
||||
_, _, err := ms.Next()
|
||||
if err == nil {
|
||||
t.Error("内部 source.Next() 报错应透传到 multiHostSource.Next")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewHostSource_RangeError newHostSource range 分支中 newRangeHostSource 失败
|
||||
func TestNewHostSource_RangeError(t *testing.T) {
|
||||
// start > end,looksLikeIPRange 通过(前半部分是有效 IP),但 newRangeHostSource 返回错误
|
||||
_, err := newHostSource("10.0.0.200-10.0.0.100")
|
||||
if err == nil {
|
||||
t.Error("start>end range 应返回 error")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNewCIDRHostSource_IPv6DirectCall 直接调用 newCIDRHostSource 传入 IPv6 CIDR
|
||||
func TestNewCIDRHostSource_IPv6DirectCall(t *testing.T) {
|
||||
// IPv6 CIDR —— bits=128 != 32,触发 line 332-334
|
||||
_, err := newCIDRHostSource("::1/128")
|
||||
if err == nil {
|
||||
t.Error("IPv6 CIDR 应被 newCIDRHostSource 拒绝 (bits!=32)")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/shadow1ng/fscan/common/output"
|
||||
)
|
||||
|
||||
func TestScanSessionLogMethodsHonorSilentConfig(t *testing.T) {
|
||||
@@ -161,6 +163,85 @@ func TestParseProxyURLExtractsAuthWithoutScheme(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestScanSessionSaveResultUsesSink 测试 SaveResult 通过 ResultSink 分发
|
||||
func TestScanSessionSaveResultUsesSink(t *testing.T) {
|
||||
preserveOutputAPIGlobals(t)
|
||||
|
||||
cfg := NewConfig()
|
||||
cfg.Output.DisableSave = true
|
||||
SetGlobalConfig(cfg)
|
||||
flagVars = &FlagVars{DisableSave: true}
|
||||
_ = InitOutput()
|
||||
|
||||
var sinkGot *output.ScanResult
|
||||
session := NewScanSession(cfg, NewState(), &FlagVars{})
|
||||
session.ResultSink = func(r *output.ScanResult) error {
|
||||
sinkGot = r
|
||||
return nil
|
||||
}
|
||||
|
||||
result := &output.ScanResult{
|
||||
Type: output.TypeHost,
|
||||
Target: "10.0.0.1",
|
||||
Status: "ALIVE",
|
||||
}
|
||||
if err := session.SaveResult(result); err != nil {
|
||||
t.Fatalf("session.SaveResult error = %v", err)
|
||||
}
|
||||
if sinkGot != result {
|
||||
t.Fatalf("ResultSink 未被调用或参数不符: got %v", sinkGot)
|
||||
}
|
||||
}
|
||||
|
||||
// TestScanSessionSaveResultFallsBackToGlobal 测试无 sink 时回退到全局 SaveResult
|
||||
func TestScanSessionSaveResultFallsBackToGlobal(t *testing.T) {
|
||||
preserveOutputAPIGlobals(t)
|
||||
|
||||
cfg := NewConfig()
|
||||
cfg.Output.DisableSave = true
|
||||
SetGlobalConfig(cfg)
|
||||
flagVars = &FlagVars{DisableSave: true}
|
||||
_ = InitOutput()
|
||||
|
||||
called := false
|
||||
SetResultCallback(func(payload interface{}) {
|
||||
called = true
|
||||
})
|
||||
|
||||
session := NewScanSession(cfg, NewState(), &FlagVars{})
|
||||
// 不设置 ResultSink,应回退到全局
|
||||
|
||||
result := &output.ScanResult{
|
||||
Type: output.TypeHost,
|
||||
Target: "10.0.0.2",
|
||||
Status: "ALIVE",
|
||||
}
|
||||
if err := session.SaveResult(result); err != nil {
|
||||
t.Fatalf("session.SaveResult (fallback) error = %v", err)
|
||||
}
|
||||
if !called {
|
||||
t.Fatal("回退到全局 SaveResult 时应触发 ResultCallback")
|
||||
}
|
||||
}
|
||||
|
||||
// TestScanSessionLogMethodsEnabledByDefault 测试非 Silent 配置下 Log 方法不被屏蔽
|
||||
func TestScanSessionLogMethodsEnabledByDefault(t *testing.T) {
|
||||
cfg := NewConfig()
|
||||
cfg.Output.Silent = false
|
||||
session := NewScanSession(cfg, NewState(), &FlagVars{})
|
||||
if !session.loggingEnabled() {
|
||||
t.Fatal("非 Silent 配置下 loggingEnabled 应返回 true")
|
||||
}
|
||||
}
|
||||
|
||||
// TestNilScanSessionLoggingEnabled 测试 nil session 的 loggingEnabled
|
||||
func TestNilScanSessionLoggingEnabled(t *testing.T) {
|
||||
var session *ScanSession
|
||||
if !session.loggingEnabled() {
|
||||
t.Fatal("nil session 的 loggingEnabled 应返回 true(安全降级)")
|
||||
}
|
||||
}
|
||||
|
||||
type roundTripFunc func(*http.Request) (*http.Response, error)
|
||||
|
||||
func (f roundTripFunc) RoundTrip(req *http.Request) (*http.Response, error) {
|
||||
|
||||
@@ -215,6 +215,300 @@ func TestState_ConcurrentTaskCounters(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestState_GetOutputMutex 测试获取输出互斥锁指针
|
||||
func TestState_GetOutputMutex(t *testing.T) {
|
||||
s := NewState()
|
||||
mu := s.GetOutputMutex()
|
||||
if mu == nil {
|
||||
t.Fatal("GetOutputMutex returned nil")
|
||||
}
|
||||
// 验证返回的指针可以正常加锁
|
||||
mu.Lock()
|
||||
mu.Unlock()
|
||||
}
|
||||
|
||||
// TestState_GetICMPLimiter 测试 ICMP 限速器延迟初始化
|
||||
func TestState_GetICMPLimiter(t *testing.T) {
|
||||
s := NewState()
|
||||
|
||||
limiter := s.GetICMPLimiter(0.1)
|
||||
if limiter == nil {
|
||||
t.Fatal("GetICMPLimiter returned nil")
|
||||
}
|
||||
|
||||
// 再次调用应返回同一个实例(sync.Once 保证)
|
||||
limiter2 := s.GetICMPLimiter(0.5)
|
||||
if limiter != limiter2 {
|
||||
t.Fatal("GetICMPLimiter should return the same instance on repeated calls")
|
||||
}
|
||||
}
|
||||
|
||||
// TestState_GetICMPLimiterMinRate 测试极低速率下的 ICMP 限速器
|
||||
func TestState_GetICMPLimiterMinRate(t *testing.T) {
|
||||
s := NewState()
|
||||
// 极低速率(packetsPerSecond < 1)应被钳位到 1
|
||||
limiter := s.GetICMPLimiter(0.000001)
|
||||
if limiter == nil {
|
||||
t.Fatal("GetICMPLimiter with tiny rate returned nil")
|
||||
}
|
||||
}
|
||||
|
||||
// TestState_GetPerfStats 测试性能统计数据
|
||||
func TestState_GetPerfStats(t *testing.T) {
|
||||
s := NewState()
|
||||
|
||||
// 初始状态:全零
|
||||
stats := s.GetPerfStats()
|
||||
if stats.TotalPackets != 0 {
|
||||
t.Errorf("初始 TotalPackets 应为 0, 实际 %d", stats.TotalPackets)
|
||||
}
|
||||
if stats.SuccessRate != 0 {
|
||||
t.Errorf("初始 SuccessRate 应为 0, 实际 %f", stats.SuccessRate)
|
||||
}
|
||||
|
||||
// 增加一些计数后验证统计
|
||||
s.IncrementTCPSuccessPacketCount()
|
||||
s.IncrementTCPSuccessPacketCount()
|
||||
s.IncrementTCPFailedPacketCount()
|
||||
s.SetNum(3)
|
||||
|
||||
stats = s.GetPerfStats()
|
||||
if stats.TotalPackets != 3 {
|
||||
t.Errorf("TotalPackets 期望 3, 实际 %d", stats.TotalPackets)
|
||||
}
|
||||
if stats.TCPSuccess != 2 {
|
||||
t.Errorf("TCPSuccess 期望 2, 实际 %d", stats.TCPSuccess)
|
||||
}
|
||||
if stats.TCPFailed != 1 {
|
||||
t.Errorf("TCPFailed 期望 1, 实际 %d", stats.TCPFailed)
|
||||
}
|
||||
if stats.TargetsScanned != 3 {
|
||||
t.Errorf("TargetsScanned 期望 3, 实际 %d", stats.TargetsScanned)
|
||||
}
|
||||
// success rate = 2/3 * 100 ≈ 66.67%
|
||||
if stats.SuccessRate < 66 || stats.SuccessRate > 67 {
|
||||
t.Errorf("SuccessRate 期望约 66.67, 实际 %f", stats.SuccessRate)
|
||||
}
|
||||
}
|
||||
|
||||
// TestState_GetPerfStatsJSON 测试性能统计 JSON 序列化
|
||||
func TestState_GetPerfStatsJSON(t *testing.T) {
|
||||
s := NewState()
|
||||
s.IncrementTCPSuccessPacketCount()
|
||||
|
||||
json := s.GetPerfStatsJSON()
|
||||
if json == "" || json == "{}" {
|
||||
t.Fatalf("GetPerfStatsJSON 返回空: %q", json)
|
||||
}
|
||||
if len(json) < 10 {
|
||||
t.Fatalf("GetPerfStatsJSON 内容过短: %q", json)
|
||||
}
|
||||
// 验证包含关键字段
|
||||
for _, key := range []string{"total_packets", "tcp_success", "success_rate"} {
|
||||
if !containsStr(json, key) {
|
||||
t.Errorf("GetPerfStatsJSON 缺少字段 %q", key)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func containsStr(s, sub string) bool {
|
||||
return len(s) >= len(sub) && (s == sub || len(s) > 0 && stringContains(s, sub))
|
||||
}
|
||||
|
||||
func stringContains(s, sub string) bool {
|
||||
for i := 0; i <= len(s)-len(sub); i++ {
|
||||
if s[i:i+len(sub)] == sub {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// TestState_GetPacketLimiter 测试通用发包限速器
|
||||
func TestState_GetPacketLimiter(t *testing.T) {
|
||||
t.Run("零速率返回nil", func(t *testing.T) {
|
||||
s := NewState()
|
||||
limiter := s.GetPacketLimiter(0)
|
||||
if limiter != nil {
|
||||
t.Fatal("零速率应返回 nil limiter")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("负速率返回nil", func(t *testing.T) {
|
||||
s := NewState()
|
||||
limiter := s.GetPacketLimiter(-1)
|
||||
if limiter != nil {
|
||||
t.Fatal("负速率应返回 nil limiter")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("正速率初始化限速器", func(t *testing.T) {
|
||||
s := NewState()
|
||||
limiter := s.GetPacketLimiter(600) // 600/min = 10/s
|
||||
if limiter == nil {
|
||||
t.Fatal("正速率应返回非 nil limiter")
|
||||
}
|
||||
// 再次调用返回同一实例
|
||||
limiter2 := s.GetPacketLimiter(1200)
|
||||
if limiter != limiter2 {
|
||||
t.Fatal("GetPacketLimiter 应通过 sync.Once 复用实例")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("低速率被钳位到1pps", func(t *testing.T) {
|
||||
s := NewState()
|
||||
// 1/min < 1/s,应被钳位
|
||||
limiter := s.GetPacketLimiter(1)
|
||||
if limiter == nil {
|
||||
t.Fatal("低速率钳位后应返回非 nil limiter")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// TestState_CacheService 测试服务识别缓存
|
||||
func TestState_CacheService(t *testing.T) {
|
||||
s := NewState()
|
||||
|
||||
// 未缓存时查询返回 false
|
||||
_, ok := s.GetCachedService("192.168.1.1:80")
|
||||
if ok {
|
||||
t.Fatal("未缓存的 key 不应返回 ok=true")
|
||||
}
|
||||
|
||||
// 缓存并查询
|
||||
type fakeInfo struct{ Name string }
|
||||
info := &fakeInfo{Name: "http"}
|
||||
s.CacheService("192.168.1.1:80", info)
|
||||
|
||||
got, ok := s.GetCachedService("192.168.1.1:80")
|
||||
if !ok {
|
||||
t.Fatal("已缓存的 key 应返回 ok=true")
|
||||
}
|
||||
if got != info {
|
||||
t.Fatalf("GetCachedService 返回 %v, 期望 %v", got, info)
|
||||
}
|
||||
|
||||
// 不同 key 互不干扰
|
||||
_, ok = s.GetCachedService("192.168.1.1:443")
|
||||
if ok {
|
||||
t.Fatal("不同 key 不应命中缓存")
|
||||
}
|
||||
}
|
||||
|
||||
// =============================================================================
|
||||
// CheckAndIncrementPacketRate 测试
|
||||
// =============================================================================
|
||||
|
||||
// TestCheckAndIncrementPacketRate_ZeroLimit 速率为 0 时无限制
|
||||
func TestCheckAndIncrementPacketRate_ZeroLimit(t *testing.T) {
|
||||
s := NewState()
|
||||
for i := 0; i < 1000; i++ {
|
||||
ok, err := s.CheckAndIncrementPacketRate(0)
|
||||
if !ok || err != nil {
|
||||
t.Fatalf("零速率限制应始终允许: ok=%v err=%v", ok, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckAndIncrementPacketRate_NegativeLimit 负速率等同于无限制
|
||||
func TestCheckAndIncrementPacketRate_NegativeLimit(t *testing.T) {
|
||||
s := NewState()
|
||||
ok, err := s.CheckAndIncrementPacketRate(-1)
|
||||
if !ok || err != nil {
|
||||
t.Fatalf("负速率应允许: ok=%v err=%v", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckAndIncrementPacketRate_AllowsWhenTokensAvailable 有令牌时返回 true
|
||||
func TestCheckAndIncrementPacketRate_AllowsWhenTokensAvailable(t *testing.T) {
|
||||
s := NewState()
|
||||
// 600/min = 10/s,桶容量 20,初始满桶
|
||||
ok, err := s.CheckAndIncrementPacketRate(600)
|
||||
if !ok || err != nil {
|
||||
t.Fatalf("初始应有令牌: ok=%v err=%v", ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestCheckAndIncrementPacketRate_RateLimitedAfterExhaustion 耗尽令牌后返回 false 和 PacketLimitError
|
||||
func TestCheckAndIncrementPacketRate_RateLimitedAfterExhaustion(t *testing.T) {
|
||||
s := NewState()
|
||||
// 极低速率:1/min,桶容量为 1(钳位后 packetsPerSecond=1,capacity=2)
|
||||
// 消耗掉所有令牌后应被限速
|
||||
const limit int64 = 1
|
||||
|
||||
// 初始化限速器(第一次调用触发 sync.Once)
|
||||
s.GetPacketLimiter(limit)
|
||||
|
||||
// 消耗完所有令牌(容量 <= 2)
|
||||
for i := 0; i < 10; i++ {
|
||||
s.CheckAndIncrementPacketRate(limit) //nolint: errcheck
|
||||
}
|
||||
|
||||
// 此时令牌应已耗尽,下一次调用应被限速
|
||||
ok, err := s.CheckAndIncrementPacketRate(limit)
|
||||
if ok {
|
||||
// 桶可能还剩令牌(容量 2),多耗几次再判断
|
||||
for i := 0; i < 20; i++ {
|
||||
ok, err = s.CheckAndIncrementPacketRate(limit)
|
||||
if !ok {
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if ok {
|
||||
t.Fatal("令牌耗尽后应返回 ok=false")
|
||||
}
|
||||
if err == nil {
|
||||
t.Fatal("令牌耗尽后应返回 error")
|
||||
}
|
||||
if !isPacketLimitError(err) {
|
||||
t.Errorf("error 类型应为 PacketLimitError, 实际 %T: %v", err, err)
|
||||
}
|
||||
}
|
||||
|
||||
// isPacketLimitError 检查是否为 PacketLimitError
|
||||
func isPacketLimitError(err error) bool {
|
||||
_, ok := err.(*PacketLimitError)
|
||||
return ok
|
||||
}
|
||||
|
||||
// TestCheckAndIncrementPacketRate_ErrorUnwrapsToSentinel 验证 error 可 unwrap 到 sentinel
|
||||
func TestCheckAndIncrementPacketRate_ErrorUnwrapsToSentinel(t *testing.T) {
|
||||
s := NewState()
|
||||
const limit int64 = 1
|
||||
|
||||
// 耗尽令牌
|
||||
for i := 0; i < 50; i++ {
|
||||
s.CheckAndIncrementPacketRate(limit) //nolint: errcheck
|
||||
}
|
||||
|
||||
var lastErr error
|
||||
for i := 0; i < 10; i++ {
|
||||
ok, err := s.CheckAndIncrementPacketRate(limit)
|
||||
if !ok {
|
||||
lastErr = err
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
if lastErr == nil {
|
||||
t.Skip("未能触发限速(可能令牌桶容量较大),跳过 unwrap 测试")
|
||||
}
|
||||
|
||||
// 验证可 unwrap 到 ErrPacketRateLimited
|
||||
pErr, ok := lastErr.(*PacketLimitError)
|
||||
if !ok {
|
||||
t.Fatalf("期望 *PacketLimitError, 实际 %T", lastErr)
|
||||
}
|
||||
if pErr.Sentinel != ErrPacketRateLimited {
|
||||
t.Errorf("Sentinel = %v, 期望 ErrPacketRateLimited", pErr.Sentinel)
|
||||
}
|
||||
if pErr.Limit != limit {
|
||||
t.Errorf("Limit = %d, 期望 %d", pErr.Limit, limit)
|
||||
}
|
||||
}
|
||||
|
||||
// TestState_OutputMutex 测试输出互斥锁
|
||||
func TestState_OutputMutex(t *testing.T) {
|
||||
s := NewState()
|
||||
|
||||
Reference in New Issue
Block a user