diff --git a/internal/cmd/check.go b/internal/cmd/check.go index 2f21930..082a810 100644 --- a/internal/cmd/check.go +++ b/internal/cmd/check.go @@ -2,7 +2,6 @@ package cmd import ( "fmt" - "net" "regexp" "strings" @@ -73,13 +72,5 @@ func isValidDomain(domain string) bool { return false } - // 尝试解析域名(不进行实际DNS查询) - _, err := net.LookupHost(domain) - if err != nil { - // 即使DNS解析失败,只要格式正确就认为是有效的 - // 因为可能是网络问题或域名确实不存在 - return true - } - return true } diff --git a/internal/cmd/csv.go b/internal/cmd/csv.go index 8314000..eac5c90 100644 --- a/internal/cmd/csv.go +++ b/internal/cmd/csv.go @@ -58,7 +58,11 @@ func (r *RootCmd) executeCSV(csvFile string) { } // 提取域名(从CERT_DOMAIN列) - domains := extractDomainsFromCSV(records) + domains, err := extractDomainsFromCSV(records) + if err != nil { + ui.PrintError(fmt.Sprintf("错误:解析CSV表头失败: %v", err)) + return + } if len(domains) == 0 { ui.PrintErrorWithDetails( "错误:未找到有效的域名", @@ -84,17 +88,33 @@ func (r *RootCmd) executeCSV(csvFile string) { } // extractDomainsFromCSV 从CSV记录中提取域名 -func extractDomainsFromCSV(records [][]string) []string { +func extractDomainsFromCSV(records [][]string) ([]string, error) { var domains []string domainSet := make(map[string]bool) // 用于去重 + if len(records) == 0 { + return nil, fmt.Errorf("CSV文件为空") + } + + certDomainIndex := -1 + for i, header := range records[0] { + header = strings.TrimPrefix(header, "\ufeff") + if strings.EqualFold(strings.TrimSpace(header), "CERT_DOMAIN") { + certDomainIndex = i + break + } + } + if certDomainIndex < 0 { + return nil, fmt.Errorf("缺少 CERT_DOMAIN 列") + } + // 跳过标题行,从第二行开始处理 for i := 1; i < len(records); i++ { - if len(records[i]) < 3 { + if len(records[i]) <= certDomainIndex { continue } - certDomain := strings.TrimSpace(records[i][2]) // CERT_DOMAIN列 + certDomain := strings.TrimSpace(records[i][certDomainIndex]) if certDomain == "" { continue } @@ -114,7 +134,7 @@ func extractDomainsFromCSV(records [][]string) []string { } } - return domains + return domains, nil } // shouldExcludeDomain 判断是否应该排除某个域名 @@ -171,5 +191,10 @@ func shouldExcludeDomain(domain string) bool { return true } + // 6. 排除证书占位文本等不符合DNS域名语法的值 + if !isValidDomain(domain) { + return true + } + return false } diff --git a/internal/cmd/csv_test.go b/internal/cmd/csv_test.go new file mode 100644 index 0000000..e840f3f --- /dev/null +++ b/internal/cmd/csv_test.go @@ -0,0 +1,65 @@ +package cmd + +import ( + "reflect" + "testing" +) + +func TestExtractDomainsFromLegacyCSV(t *testing.T) { + records := [][]string{ + {"IP", "ORIGIN", "CERT_DOMAIN", "CERT_ISSUER", "GEO_CODE"}, + {"1.2.3.4", "1.2.3.0/24", "example.com", "Let's Encrypt", "US"}, + } + + got, err := extractDomainsFromCSV(records) + if err != nil { + t.Fatal(err) + } + want := []string{"example.com"} + if !reflect.DeepEqual(got, want) { + t.Fatalf("got %v, want %v", got, want) + } +} + +func TestExtractDomainsFromCurrentCSV(t *testing.T) { + records := [][]string{ + {"IP", "ORIGIN", "TLS", "ALPN", "CURVE", "CERT_LENGTH", "CERT_SIGNATURE", "CERT_PUBLICKEY", "CERT_DOMAIN", "CERT_ISSUER", "GEO_CODE"}, + {"1.2.3.4", "1.2.3.0/24", "TLS 1.3", "h2", "X25519", "1234", "SHA256-RSA", "RSA", "example.org", "Let's Encrypt", "US"}, + {"1.2.3.5", "1.2.3.0/24", "TLS 1.3", "h2", "X25519", "1234", "SHA256-RSA", "RSA", "example.org", "Let's Encrypt", "US"}, + {"1.2.3.6", "1.2.3.0/24", "TLS 1.3", "h2", "X25519", "1234", "SHA256-RSA", "RSA", "*.wildcard.example", "Let's Encrypt", "US"}, + {"1.2.3.7", "1.2.3.0/24", "TLS 1.3", "h2", "X25519", "1234", "SHA256-RSA", "RSA", "192.0.2.1", "Let's Encrypt", "US"}, + {"1.2.3.8", "1.2.3.0/24", "TLS 1.3", "h2", "X25519", "1234", "SHA256-RSA", "RSA", "Common Name", "Let's Encrypt", "US"}, + } + + got, err := extractDomainsFromCSV(records) + if err != nil { + t.Fatal(err) + } + want := []string{"example.org"} + if !reflect.DeepEqual(got, want) { + t.Fatalf("got %v, want %v", got, want) + } +} + +func TestExtractDomainsAcceptsBOMWhitespaceAndCase(t *testing.T) { + records := [][]string{ + {"\ufeff cert_domain "}, + {"example.net"}, + } + + got, err := extractDomainsFromCSV(records) + if err != nil { + t.Fatal(err) + } + want := []string{"example.net"} + if !reflect.DeepEqual(got, want) { + t.Fatalf("got %v, want %v", got, want) + } +} + +func TestExtractDomainsRejectsMissingHeader(t *testing.T) { + _, err := extractDomainsFromCSV([][]string{{"IP", "TLS"}, {"1.2.3.4", "TLS 1.3"}}) + if err == nil { + t.Fatal("expected missing CERT_DOMAIN error") + } +}