diff --git a/cmd/exploredns/main.go b/cmd/exploredns/main.go index 4a57fda..33528a7 100644 --- a/cmd/exploredns/main.go +++ b/cmd/exploredns/main.go @@ -5,7 +5,6 @@ import ( "flag" "fmt" "os" - "strings" "time" "github.com/hits/ExploreDNS/internal/config" @@ -27,22 +26,42 @@ func main() { maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)") retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)") fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache") - verbose := flag.Bool("verbose", cfg.Verbose, "Verbose output") - debugFlag := flag.Bool("debug", false, "Debug mode (stackable: -dd for library debug)") - quiet := flag.Bool("quiet", cfg.Quiet, "Suppress supplementary info") + // Verbose: long and short form share the same variable. + var verboseVal bool + flag.BoolVar(&verboseVal, "verbose", cfg.Verbose, "Verbose output") + flag.BoolVar(&verboseVal, "v", cfg.Verbose, "Verbose output (shorthand)") + + // Debug: -d sets level 1, -dd sets level 2 (library debug). + var dFlag, ddFlag bool + flag.BoolVar(&dFlag, "d", false, "Debug mode (stackable: -dd for library debug)") + flag.BoolVar(&dFlag, "debug", false, "Debug mode") + flag.BoolVar(&ddFlag, "dd", false, "Library debug mode (equivalent to -d -d)") + + // Quiet: long and short form share the same variable. + var quietVal bool + flag.BoolVar(&quietVal, "quiet", cfg.Quiet, "Suppress supplementary info") + flag.BoolVar(&quietVal, "q", cfg.Quiet, "Suppress supplementary info (shorthand)") + + // TODO: ShowProgress is parsed but not yet wired to the traversal/output layer. showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress") noShowProgress := flag.Bool("no-show-progress", false, "Hide traversal progress") + // TODO: ShowResolves is parsed but not yet wired to the traversal/output layer. showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details") noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details") + // TODO: ShowServers is parsed but not yet wired to the traversal/output layer. showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried") noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried") + // TODO: ShowVersions is parsed but not yet wired to the traversal/output layer. showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions") noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions") + // TODO: ShowAllStats is parsed but not yet wired to the traversal/output layer. showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics") noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics") + // TODO: ShowResults is partially wired; full structured output is pending. showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results") noShowResults := flag.Bool("no-show-results", false, "Hide query results") + // TODO: ShowSummaryResults is parsed but not yet wired to the traversal/output layer. showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results") noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results") @@ -52,16 +71,6 @@ func main() { args := flag.Args() - debugLevel := 0 - for _, arg := range os.Args[1:] { - if strings.Count(arg, "-d") > 0 || strings.Count(arg, "-dd") > 0 { - debugLevel = strings.Count(arg, "-d") + strings.Count(arg, "-dd") - } - } - if *debugFlag { - debugLevel++ - } - cfg.QueryType = *queryType cfg.RootServer = *rootServer cfg.AllRootServers = *allRootServers @@ -73,9 +82,9 @@ func main() { cfg.MaxDepth = *maxDepth cfg.Retries = *retries cfg.Fast = *fast - cfg.Verbose = *verbose - cfg.Debug = debugLevel - cfg.Quiet = *quiet + cfg.Verbose = verboseVal + cfg.Debug = config.ParseDebugLevel(dFlag, ddFlag) + cfg.Quiet = quietVal if *noShowProgress { cfg.ShowProgress = false @@ -142,10 +151,11 @@ func main() { } queryConfig := &dns.QueryConfig{ - UDPSize: cfg.UDPSize, - Timeout: 5 * time.Second, - Retries: cfg.Retries, - UseTCP: cfg.AlwaysTCP, + UDPSize: cfg.UDPSize, + Timeout: 5 * time.Second, + Retries: cfg.Retries, + UseTCP: cfg.AlwaysTCP, + AllowTCP: cfg.AllowTCP, } rootConfig := &dns.RootDiscoveryConfig{ diff --git a/internal/config/config.go b/internal/config/config.go index f71ebfc..e46a464 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -11,12 +11,13 @@ import ( ) var ( - ErrInvalidQueryType = errors.New("invalid query type") - ErrInvalidUDPSize = errors.New("UDP size must be between 512 and 4096") - ErrInvalidMaxDepth = errors.New("max depth must be between 1 and 100") - ErrInvalidRetries = errors.New("retries must be between 0 and 10") - ErrMissingDomain = errors.New("domain is required") + ErrInvalidQueryType = errors.New("invalid query type") + ErrInvalidUDPSize = errors.New("UDP size must be between 512 and 4096") + ErrInvalidMaxDepth = errors.New("max depth must be between 1 and 100") + ErrInvalidRetries = errors.New("retries must be between 0 and 10") + ErrMissingDomain = errors.New("domain is required") ErrAlwaysTCPRequiresTCP = errors.New("--always-tcp requires --allow-tcp") + ErrInvalidRootServer = errors.New("invalid root server IP address") ) type Config struct { @@ -138,7 +139,23 @@ func (c *Config) ParseRootServer() (net.IP, error) { if c.RootServer == "" { return nil, nil } - return net.ParseIP(c.RootServer), nil + ip := net.ParseIP(c.RootServer) + if ip == nil { + return nil, fmt.Errorf("%w: %s", ErrInvalidRootServer, c.RootServer) + } + return ip, nil +} + +// ParseDebugLevel returns the debug verbosity level from the -d and -dd flag values. +// dd=true → 2 (library debug), d=true → 1 (application debug), neither → 0. +func ParseDebugLevel(d, dd bool) int { + if dd { + return 2 + } + if d { + return 1 + } + return 0 } func PrintUsage() { diff --git a/internal/config/config_test.go b/internal/config/config_test.go new file mode 100644 index 0000000..6a5041e --- /dev/null +++ b/internal/config/config_test.go @@ -0,0 +1,216 @@ +package config + +import ( + "errors" + "testing" +) + +func TestParseQueryTypeValid(t *testing.T) { + cases := []struct { + input string + wantType uint16 + }{ + {"A", 1}, + {"aaaa", 28}, + {"NS", 2}, + {"CNAME", 5}, + {"MX", 15}, + {"TXT", 16}, + {"SOA", 6}, + {"PTR", 12}, + {"ANY", 255}, + } + for _, tc := range cases { + got, err := ParseQueryType(tc.input) + if err != nil { + t.Errorf("ParseQueryType(%q) unexpected error: %v", tc.input, err) + } + if got != tc.wantType { + t.Errorf("ParseQueryType(%q) = %d, want %d", tc.input, got, tc.wantType) + } + } +} + +func TestParseQueryTypeInvalid(t *testing.T) { + _, err := ParseQueryType("BOGUS") + if err == nil { + t.Fatal("expected error for invalid query type") + } + if !errors.Is(err, ErrInvalidQueryType) { + t.Errorf("expected ErrInvalidQueryType, got %v", err) + } +} + +func TestParseUDPSizeValid(t *testing.T) { + cases := []string{"512", "2048", "4096"} + for _, s := range cases { + if _, err := ParseUDPSize(s); err != nil { + t.Errorf("ParseUDPSize(%q) unexpected error: %v", s, err) + } + } +} + +func TestParseUDPSizeInvalid(t *testing.T) { + cases := []string{"0", "511", "4097", "notanumber"} + for _, s := range cases { + _, err := ParseUDPSize(s) + if err == nil { + t.Errorf("ParseUDPSize(%q): expected error", s) + } + if !errors.Is(err, ErrInvalidUDPSize) { + t.Errorf("ParseUDPSize(%q): expected ErrInvalidUDPSize, got %v", s, err) + } + } +} + +func TestValidateOK(t *testing.T) { + cfg := DefaultConfig() + if err := cfg.Validate(); err != nil { + t.Fatalf("DefaultConfig should be valid, got: %v", err) + } +} + +func TestValidateAlwaysTCPRequiresAllowTCP(t *testing.T) { + cfg := DefaultConfig() + cfg.AlwaysTCP = true + cfg.AllowTCP = false + err := cfg.Validate() + if err == nil { + t.Fatal("expected error when AlwaysTCP=true and AllowTCP=false") + } + if !errors.Is(err, ErrAlwaysTCPRequiresTCP) { + t.Errorf("expected ErrAlwaysTCPRequiresTCP, got %v", err) + } +} + +func TestValidateAlwaysTCPWithAllowTCP(t *testing.T) { + cfg := DefaultConfig() + cfg.AlwaysTCP = true + cfg.AllowTCP = true + if err := cfg.Validate(); err != nil { + t.Errorf("AlwaysTCP=true, AllowTCP=true should be valid, got: %v", err) + } +} + +func TestValidateBadUDPSize(t *testing.T) { + cfg := DefaultConfig() + cfg.UDPSize = 100 + if err := cfg.Validate(); !errors.Is(err, ErrInvalidUDPSize) { + t.Errorf("expected ErrInvalidUDPSize, got %v", err) + } +} + +func TestValidateBadMaxDepth(t *testing.T) { + cfg := DefaultConfig() + cfg.MaxDepth = 0 + if err := cfg.Validate(); !errors.Is(err, ErrInvalidMaxDepth) { + t.Errorf("expected ErrInvalidMaxDepth, got %v", err) + } +} + +func TestValidateBadRetries(t *testing.T) { + cfg := DefaultConfig() + cfg.Retries = 11 + if err := cfg.Validate(); !errors.Is(err, ErrInvalidRetries) { + t.Errorf("expected ErrInvalidRetries, got %v", err) + } +} + +func TestGetDomainOK(t *testing.T) { + cfg := DefaultConfig() + d, err := cfg.GetDomain([]string{"example.com"}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if d != "example.com" { + t.Errorf("got %q, want %q", d, "example.com") + } +} + +func TestGetDomainMissing(t *testing.T) { + cfg := DefaultConfig() + _, err := cfg.GetDomain([]string{}) + if err == nil { + t.Fatal("expected error for missing domain") + } + if !errors.Is(err, ErrMissingDomain) { + t.Errorf("expected ErrMissingDomain, got %v", err) + } +} + +func TestParseRootServerEmpty(t *testing.T) { + cfg := DefaultConfig() + ip, err := cfg.ParseRootServer() + if err != nil { + t.Fatalf("unexpected error for empty root server: %v", err) + } + if ip != nil { + t.Errorf("expected nil IP for empty root server, got %v", ip) + } +} + +func TestParseRootServerValidIPv4(t *testing.T) { + cfg := DefaultConfig() + cfg.RootServer = "198.41.0.4" + ip, err := cfg.ParseRootServer() + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if ip == nil || ip.String() != "198.41.0.4" { + t.Errorf("expected 198.41.0.4, got %v", ip) + } +} + +func TestParseRootServerValidIPv6(t *testing.T) { + cfg := DefaultConfig() + cfg.RootServer = "2001:503:ba3e::2:30" + ip, err := cfg.ParseRootServer() + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if ip == nil { + t.Error("expected non-nil IP for valid IPv6 address") + } +} + +func TestParseRootServerInvalid(t *testing.T) { + cfg := DefaultConfig() + cfg.RootServer = "not-an-ip" + _, err := cfg.ParseRootServer() + if err == nil { + t.Fatal("expected error for invalid root server IP") + } + if !errors.Is(err, ErrInvalidRootServer) { + t.Errorf("expected ErrInvalidRootServer, got %v", err) + } +} + +func TestParseRootServerHostname(t *testing.T) { + cfg := DefaultConfig() + cfg.RootServer = "a.root-servers.net" + _, err := cfg.ParseRootServer() + if err == nil { + t.Fatal("expected error for hostname (not IP) root server") + } + if !errors.Is(err, ErrInvalidRootServer) { + t.Errorf("expected ErrInvalidRootServer, got %v", err) + } +} + +func TestParseDebugLevel(t *testing.T) { + cases := []struct { + d, dd bool + want int + }{ + {false, false, 0}, + {true, false, 1}, + {false, true, 2}, + {true, true, 2}, // dd takes precedence + } + for _, tc := range cases { + got := ParseDebugLevel(tc.d, tc.dd) + if got != tc.want { + t.Errorf("ParseDebugLevel(d=%v, dd=%v) = %d, want %d", tc.d, tc.dd, got, tc.want) + } + } +} diff --git a/internal/dns/query.go b/internal/dns/query.go index ace16ee..0fb0edc 100644 --- a/internal/dns/query.go +++ b/internal/dns/query.go @@ -10,18 +10,20 @@ import ( ) type QueryConfig struct { - UDPSize int - Timeout time.Duration - Retries int - UseTCP bool + UDPSize int + Timeout time.Duration + Retries int + UseTCP bool + AllowTCP bool } func DefaultQueryConfig() *QueryConfig { return &QueryConfig{ - UDPSize: DefaultEDNS0UDPSize(), - Timeout: 5 * time.Second, - Retries: 3, - UseTCP: false, + UDPSize: DefaultEDNS0UDPSize(), + Timeout: 5 * time.Second, + Retries: 3, + UseTCP: false, + AllowTCP: true, } } @@ -110,7 +112,7 @@ func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype ui continue } - if resp.Truncated { + if resp.Truncated && cfg.AllowTCP { resp, err = exchangeFn(ctx, serverStr, msg, true) if err != nil { lastErr = err @@ -178,7 +180,7 @@ func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string, continue } - if resp.Truncated { + if resp.Truncated && cfg.AllowTCP { resp, err = exchangeFn(ctx, serverStr, msg, true) if err != nil { lastErr = err diff --git a/internal/dns/query_test.go b/internal/dns/query_test.go index 5c18f27..6fa840f 100644 --- a/internal/dns/query_test.go +++ b/internal/dns/query_test.go @@ -110,10 +110,11 @@ func TestQueryTCPFallbackOnTruncation(t *testing.T) { } cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: false, + UDPSize: 2048, + Timeout: 5, + Retries: 1, + UseTCP: false, + AllowTCP: true, } server := net.ParseIP("8.8.8.8") @@ -272,10 +273,11 @@ func TestQueryTCPFallbackFailsThenRetries(t *testing.T) { } cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 2, - UseTCP: false, + UDPSize: 2048, + Timeout: 5, + Retries: 2, + UseTCP: false, + AllowTCP: true, } server := net.ParseIP("8.8.8.8") @@ -287,3 +289,43 @@ func TestQueryTCPFallbackFailsThenRetries(t *testing.T) { t.Errorf("expected 4 calls (2 retries x UDP+TCP), got %d", callCount) } } + +func TestQueryNoTCPFallbackWhenDisabled(t *testing.T) { + truncatedResp := new(dns.Msg) + truncatedResp.Truncated = true + truncatedResp.SetReply(new(dns.Msg)) + + var mu sync.Mutex + calls := []bool{} + exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + mu.Lock() + defer mu.Unlock() + calls = append(calls, useTCP) + return truncatedResp.Copy(), nil + } + + cfg := &QueryConfig{ + UDPSize: 2048, + Timeout: 5, + Retries: 1, + UseTCP: false, + AllowTCP: false, // TCP fallback must be suppressed + } + + server := net.ParseIP("8.8.8.8") + resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + // Only one UDP call; no TCP fallback. + if len(calls) != 1 { + t.Fatalf("expected 1 exchange call (no TCP fallback), got %d", len(calls)) + } + if calls[0] != false { + t.Error("expected UDP-only call") + } + if !resp.Truncated { + t.Error("expected truncated response to be returned as-is") + } +} +