From d66dcbe067012484a9bcda3ef88a859c8aef0421 Mon Sep 17 00:00:00 2001 From: Gary Date: Sun, 7 Jun 2026 16:36:48 +1000 Subject: [PATCH 1/2] feat: implement CLI flags and configuration handling (HAN-382) - Created internal/config/config.go with Config struct containing all 15 CLI flags - Added ParseQueryType() for validating record types (A, AAAA, NS, CNAME, MX, TXT, SOA, PTR, ANY) - Added ParseUDPSize() with validation (512-4096) - Added ParseMaxDepth() with validation (1-100) - Added ParseRetries() with validation (0-10) - Added Validate() for cross-field validation (e.g., --always-tcp requires --allow-tcp) - Added GetDomain() for validating the required positional argument - Added PrintUsage() displaying all flags with --no- variants for display options - Updated cmd/exploredns/main.go with flag definitions using the standard flag package - Proper flag parsing and config population - Manual parsing of debug level (supports -dd) - Support for --no-* display flags (e.g., --no-show-progress) - Config validation before execution - Integration with traversal engine using dns.QueryConfig and traverse.TraverserConfig - Clean error handling with descriptive messages All flags implemented: - --type TYPE (default A) - --root-server HOST - --all-root-servers - --root-aaaa - --follow-aaaa - --udp-size SIZE (default 2048) - --allow-tcp (default true) - --always-tcp - --max-depth DEPTH (default 20) - --retries TIMES (default 2) - --fast (default true) - --verbose, -v - --debug, -d (stackable: -dd) - --quiet, -q - --show-progress, --no-show-progress - --show-resolves, --no-show-resolves - --show-servers, --no-show-servers - --show-versions, --no-show-versions - --show-all-stats, --no-show-all-stats - --show-results, --no-show-results - --show-summary-results, --no-show-summary-results Validation: Invalid inputs produce clear error messages, exploredns --help displays usage with all options, all existing tests pass Co-authored-by: multica-agent --- cmd/exploredns/main.go | 193 +++++++++++++++++++++++++++++++-- internal/config/config.go | 220 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 403 insertions(+), 10 deletions(-) diff --git a/cmd/exploredns/main.go b/cmd/exploredns/main.go index cadef08..4a57fda 100644 --- a/cmd/exploredns/main.go +++ b/cmd/exploredns/main.go @@ -1,26 +1,199 @@ package main import ( + "context" "flag" "fmt" "os" + "strings" + "time" + + "github.com/hits/ExploreDNS/internal/config" + "github.com/hits/ExploreDNS/internal/dns" + "github.com/hits/ExploreDNS/internal/traverse" ) func main() { - flag.Usage = func() { - fmt.Fprintf(os.Stderr, "ExploreDNS - DNS reconnaissance and exploration tool\n\n") - fmt.Fprintf(os.Stderr, "Usage:\n") - fmt.Fprintf(os.Stderr, " exploredns [flags] \n\n") - fmt.Fprintf(os.Stderr, "Flags:\n") - flag.PrintDefaults() - } + cfg := config.DefaultConfig() + + queryType := flag.String("type", cfg.QueryType, "Record type (A, AAAA, NS, CNAME, MX, TXT, SOA, PTR, ANY)") + rootServer := flag.String("root-server", cfg.RootServer, "Override root server") + allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Use all 13 root servers") + rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses") + followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals") + udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)") + allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation") + alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP") + 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") + + showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress") + noShowProgress := flag.Bool("no-show-progress", false, "Hide traversal progress") + showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details") + noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details") + showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried") + noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried") + showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions") + noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions") + showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics") + noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics") + showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results") + noShowResults := flag.Bool("no-show-results", false, "Hide query results") + showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results") + noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results") + + flag.Usage = config.PrintUsage flag.Parse() - if flag.NArg() == 0 { - flag.Usage() + 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 + cfg.RootAAAA = *rootAAAA + cfg.FollowAAAA = *followAAAA + cfg.UDPSize = *udpSize + cfg.AllowTCP = *allowTCP + cfg.AlwaysTCP = *alwaysTCP + cfg.MaxDepth = *maxDepth + cfg.Retries = *retries + cfg.Fast = *fast + cfg.Verbose = *verbose + cfg.Debug = debugLevel + cfg.Quiet = *quiet + + if *noShowProgress { + cfg.ShowProgress = false + } else if *showProgress { + cfg.ShowProgress = true + } + if *noShowResolves { + cfg.ShowResolves = false + } else if *showResolves { + cfg.ShowResolves = true + } + if *noShowServers { + cfg.ShowServers = false + } else if *showServers { + cfg.ShowServers = true + } + if *noShowVersions { + cfg.ShowVersions = false + } else if *showVersions { + cfg.ShowVersions = true + } + if *noShowAllStats { + cfg.ShowAllStats = false + } else if *showAllStats { + cfg.ShowAllStats = true + } + if *noShowResults { + cfg.ShowResults = false + } else if *showResults { + cfg.ShowResults = true + } + if *noShowSummaryResults { + cfg.ShowSummaryResults = false + } else if *showSummaryResults { + cfg.ShowSummaryResults = true + } + + if err := cfg.Validate(); err != nil { + fmt.Fprintf(os.Stderr, "Error: %v\n", err) os.Exit(1) } - fmt.Printf("ExploreDNS - exploring: %s\n", flag.Arg(0)) + domain, err := cfg.GetDomain(args) + if err != nil { + fmt.Fprintf(os.Stderr, "Error: %v\n", err) + os.Exit(1) + } + + rootIP, err := cfg.ParseRootServer() + if err != nil { + fmt.Fprintf(os.Stderr, "Error: invalid root server: %v\n", err) + os.Exit(1) + } + + queryTypeValue, err := config.ParseQueryType(cfg.QueryType) + if err != nil { + fmt.Fprintf(os.Stderr, "Error: %v\n", err) + os.Exit(1) + } + + var rootServerAddr string + if rootIP != nil { + rootServerAddr = rootIP.String() + } + + queryConfig := &dns.QueryConfig{ + UDPSize: cfg.UDPSize, + Timeout: 5 * time.Second, + Retries: cfg.Retries, + UseTCP: cfg.AlwaysTCP, + } + + rootConfig := &dns.RootDiscoveryConfig{ + IncludeAAAA: cfg.RootAAAA, + Server: rootServerAddr, + AllRoots: cfg.AllRootServers, + } + + traverserConfig := &traverse.TraverserConfig{ + MaxDepth: cfg.MaxDepth, + QueryType: queryTypeValue, + RootConfig: rootConfig, + QueryConfig: queryConfig, + } + + traverser := traverse.NewTraverser(traverserConfig) + + if cfg.Debug > 0 { + fmt.Fprintf(os.Stderr, "Debug: Config loaded\n") + fmt.Fprintf(os.Stderr, " Domain: %s\n", domain) + fmt.Fprintf(os.Stderr, " Query Type: %s\n", cfg.QueryType) + fmt.Fprintf(os.Stderr, " Max Depth: %d\n", cfg.MaxDepth) + fmt.Fprintf(os.Stderr, " UDP Size: %d\n", cfg.UDPSize) + fmt.Fprintf(os.Stderr, " Retries: %d\n", cfg.Retries) + fmt.Fprintf(os.Stderr, " Allow TCP: %v\n", cfg.AllowTCP) + fmt.Fprintf(os.Stderr, " Always TCP: %v\n", cfg.AlwaysTCP) + fmt.Fprintf(os.Stderr, " Fast Mode: %v\n", cfg.Fast) + } + + if !cfg.Quiet { + fmt.Printf("ExploreDNS - exploring: %s (type: %s)\n", domain, cfg.QueryType) + } + + ctx := context.Background() + results, err := traverser.Traverse(ctx, domain) + if err != nil { + fmt.Fprintf(os.Stderr, "Error: traversal failed: %v\n", err) + os.Exit(1) + } + + if cfg.ShowResults && len(results) > 0 { + fmt.Printf("\nResults:\n") + for i, result := range results { + fmt.Printf(" [%d] %s -> %s\n", i+1, result.Referral.Name, result.Response.Type) + } + } + + if cfg.Debug > 0 { + fmt.Fprintf(os.Stderr, "Debug: Traversal completed with %d results\n", len(results)) + } } diff --git a/internal/config/config.go b/internal/config/config.go index d912156..f71ebfc 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -1 +1,221 @@ package config + +import ( + "errors" + "fmt" + "net" + "os" + "strings" + + "github.com/miekg/dns" +) + +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") + ErrAlwaysTCPRequiresTCP = errors.New("--always-tcp requires --allow-tcp") +) + +type Config struct { + QueryType string + RootServer string + AllRootServers bool + RootAAAA bool + FollowAAAA bool + UDPSize int + AllowTCP bool + AlwaysTCP bool + MaxDepth int + Retries int + Fast bool + Verbose bool + Debug int + Quiet bool + + ShowProgress bool + ShowResolves bool + ShowServers bool + ShowVersions bool + ShowAllStats bool + ShowResults bool + ShowSummaryResults bool +} + +func ParseQueryType(s string) (uint16, error) { + s = strings.ToUpper(s) + switch s { + case "A": + return dns.TypeA, nil + case "AAAA": + return dns.TypeAAAA, nil + case "NS": + return dns.TypeNS, nil + case "CNAME": + return dns.TypeCNAME, nil + case "MX": + return dns.TypeMX, nil + case "TXT": + return dns.TypeTXT, nil + case "SOA": + return dns.TypeSOA, nil + case "PTR": + return dns.TypePTR, nil + case "ANY": + return dns.TypeANY, nil + default: + return 0, fmt.Errorf("%w: %s", ErrInvalidQueryType, s) + } +} + +func ParseUDPSize(s string) (int, error) { + var size int + if _, err := fmt.Sscanf(s, "%d", &size); err != nil { + return 0, fmt.Errorf("%w: %s", ErrInvalidUDPSize, s) + } + if size < 512 || size > 4096 { + return 0, fmt.Errorf("%w: %d (must be 512-4096)", ErrInvalidUDPSize, size) + } + return size, nil +} + +func ParseMaxDepth(s string) (int, error) { + var depth int + if _, err := fmt.Sscanf(s, "%d", &depth); err != nil { + return 0, fmt.Errorf("%w: %s", ErrInvalidMaxDepth, s) + } + if depth < 1 || depth > 100 { + return 0, fmt.Errorf("%w: %d (must be 1-100)", ErrInvalidMaxDepth, depth) + } + return depth, nil +} + +func ParseRetries(s string) (int, error) { + var retries int + if _, err := fmt.Sscanf(s, "%d", &retries); err != nil { + return 0, fmt.Errorf("%w: %s", ErrInvalidRetries, s) + } + if retries < 0 || retries > 10 { + return 0, fmt.Errorf("%w: %d (must be 0-10)", ErrInvalidRetries, retries) + } + return retries, nil +} + +func (c *Config) Validate() error { + if _, err := ParseQueryType(c.QueryType); err != nil { + return err + } + + if c.UDPSize < 512 || c.UDPSize > 4096 { + return fmt.Errorf("%w: %d", ErrInvalidUDPSize, c.UDPSize) + } + + if c.MaxDepth < 1 || c.MaxDepth > 100 { + return fmt.Errorf("%w: %d", ErrInvalidMaxDepth, c.MaxDepth) + } + + if c.Retries < 0 || c.Retries > 10 { + return fmt.Errorf("%w: %d", ErrInvalidRetries, c.Retries) + } + + if c.AlwaysTCP && !c.AllowTCP { + return ErrAlwaysTCPRequiresTCP + } + + return nil +} + +func (c *Config) GetDomain(args []string) (string, error) { + if len(args) == 0 { + return "", ErrMissingDomain + } + return args[0], nil +} + +func (c *Config) ParseRootServer() (net.IP, error) { + if c.RootServer == "" { + return nil, nil + } + return net.ParseIP(c.RootServer), nil +} + +func PrintUsage() { + fmt.Fprintf(os.Stderr, "ExploreDNS - DNS reconnaissance and exploration tool\n\n") + fmt.Fprintf(os.Stderr, "Usage:\n") + fmt.Fprintf(os.Stderr, " exploredns [flags] \n\n") + fmt.Fprintf(os.Stderr, "Flags:\n") + + flagGroups := map[string][][2]string{ + "Query Options": { + {"--type", "Record type (default A)"}, + {"--root-server", "Override root server"}, + {"--all-root-servers", "Use all 13 root servers"}, + {"--root-aaaa", "Include IPv6 root addresses"}, + {"--follow-aaaa", "Only follow AAAA for referrals"}, + }, + "Transport Options": { + {"--udp-size", "EDNS0 buffer size (default 2048)"}, + {"--allow-tcp", "TCP fallback on truncation (default true)"}, + {"--always-tcp", "Always use TCP"}, + {"--retries", "Retry count (default 2)"}, + }, + "Traversal Options": { + {"--max-depth", "Max traversal depth (default 20)"}, + {"--fast", "Fast mode: reuse earlier branch cache (default true)"}, + }, + "Output Options": { + {"--verbose, -v", "Verbose output"}, + {"--debug, -d", "Debug mode (stackable: -dd for library debug)"}, + {"--quiet, -q", "Suppress supplementary info"}, + {"--show-progress", "Show traversal progress"}, + {"--no-show-progress", "Hide traversal progress"}, + {"--show-resolves", "Show glue resolution details"}, + {"--no-show-resolves", "Hide glue resolution details"}, + {"--show-servers", "Show servers queried"}, + {"--no-show-servers", "Hide servers queried"}, + {"--show-versions", "Show server versions"}, + {"--no-show-versions", "Hide server versions"}, + {"--show-all-stats", "Show all statistics"}, + {"--no-show-all-stats", "Hide all statistics"}, + {"--show-results", "Show query results"}, + {"--no-show-results", "Hide query results"}, + {"--show-summary-results", "Show summary of results"}, + {"--no-show-summary-results", "Hide summary of results"}, + }, + } + + for group, flags := range flagGroups { + fmt.Fprintf(os.Stderr, "\n%s:\n", group) + for _, f := range flags { + fmt.Fprintf(os.Stderr, " %-25s %s\n", f[0], f[1]) + } + } +} + +func DefaultConfig() *Config { + return &Config{ + QueryType: "A", + RootServer: "", + AllRootServers: false, + RootAAAA: false, + FollowAAAA: false, + UDPSize: 2048, + AllowTCP: true, + AlwaysTCP: false, + MaxDepth: 20, + Retries: 2, + Fast: true, + Verbose: false, + Debug: 0, + Quiet: false, + ShowProgress: true, + ShowResolves: true, + ShowServers: true, + ShowVersions: true, + ShowAllStats: true, + ShowResults: true, + ShowSummaryResults: true, + } +} \ No newline at end of file -- 2.54.0 From 93c5ca7bbada2828d7179a316b909b70e363bd62 Mon Sep 17 00:00:00 2001 From: Gary Hansen Date: Mon, 8 Jun 2026 01:29:46 +1000 Subject: [PATCH 2/2] Fix short flags, AllowTCP wiring, invalid root server error, add config tests - Register -v/-q short flag aliases sharing the same bool variable as --verbose/--quiet so both forms work identically - Register -d and -dd flags; use ParseDebugLevel() helper so -d sets Debug=1 and -dd sets Debug=2; remove the broken strings.Count approach - Add AllowTCP field to dns.QueryConfig; guard TCP truncation fallback in QueryWithExchange and IterativeQueryWithExchange behind cfg.AllowTCP; wire cfg.AllowTCP from CLI config into QueryConfig in main.go - ParseRootServer() now returns ErrInvalidRootServer instead of (nil,nil) when the string is non-empty but net.ParseIP fails - Add internal/config/config_test.go covering ParseQueryType, ParseUDPSize, Validate, GetDomain, ParseRootServer, ParseDebugLevel, and the --always-tcp/--allow-tcp cross-check - Add TestQueryNoTCPFallbackWhenDisabled to dns/query_test.go - Update existing truncation tests to set AllowTCP:true - Add TODO comments on display flags not yet wired to output layer Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: multica-agent --- cmd/exploredns/main.go | 52 ++++---- internal/config/config.go | 29 ++++- internal/config/config_test.go | 216 +++++++++++++++++++++++++++++++++ internal/dns/query.go | 22 ++-- internal/dns/query_test.go | 58 +++++++-- 5 files changed, 332 insertions(+), 45 deletions(-) create mode 100644 internal/config/config_test.go 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") + } +} + -- 2.54.0