Fix short flags, AllowTCP wiring, invalid root server error, add config tests
CI / test (pull_request) Failing after 2m14s

- 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 <github@multica.ai>
This commit is contained in:
Gary Hansen
2026-06-08 01:29:46 +10:00
co-authored by Copilot multica-agent
parent d66dcbe067
commit 93c5ca7bba
5 changed files with 332 additions and 45 deletions
+31 -21
View File
@@ -5,7 +5,6 @@ import (
"flag" "flag"
"fmt" "fmt"
"os" "os"
"strings"
"time" "time"
"github.com/hits/ExploreDNS/internal/config" "github.com/hits/ExploreDNS/internal/config"
@@ -27,22 +26,42 @@ func main() {
maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)") maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)")
retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)") retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)")
fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache") 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") showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress")
noShowProgress := flag.Bool("no-show-progress", false, "Hide 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") showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details")
noShowResolves := flag.Bool("no-show-resolves", false, "Hide 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") showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried")
noShowServers := flag.Bool("no-show-servers", false, "Hide 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") showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions")
noShowVersions := flag.Bool("no-show-versions", false, "Hide 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") showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics")
noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide 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") showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results")
noShowResults := flag.Bool("no-show-results", false, "Hide 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") showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results")
noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results") noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results")
@@ -52,16 +71,6 @@ func main() {
args := flag.Args() 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.QueryType = *queryType
cfg.RootServer = *rootServer cfg.RootServer = *rootServer
cfg.AllRootServers = *allRootServers cfg.AllRootServers = *allRootServers
@@ -73,9 +82,9 @@ func main() {
cfg.MaxDepth = *maxDepth cfg.MaxDepth = *maxDepth
cfg.Retries = *retries cfg.Retries = *retries
cfg.Fast = *fast cfg.Fast = *fast
cfg.Verbose = *verbose cfg.Verbose = verboseVal
cfg.Debug = debugLevel cfg.Debug = config.ParseDebugLevel(dFlag, ddFlag)
cfg.Quiet = *quiet cfg.Quiet = quietVal
if *noShowProgress { if *noShowProgress {
cfg.ShowProgress = false cfg.ShowProgress = false
@@ -142,10 +151,11 @@ func main() {
} }
queryConfig := &dns.QueryConfig{ queryConfig := &dns.QueryConfig{
UDPSize: cfg.UDPSize, UDPSize: cfg.UDPSize,
Timeout: 5 * time.Second, Timeout: 5 * time.Second,
Retries: cfg.Retries, Retries: cfg.Retries,
UseTCP: cfg.AlwaysTCP, UseTCP: cfg.AlwaysTCP,
AllowTCP: cfg.AllowTCP,
} }
rootConfig := &dns.RootDiscoveryConfig{ rootConfig := &dns.RootDiscoveryConfig{
+23 -6
View File
@@ -11,12 +11,13 @@ import (
) )
var ( var (
ErrInvalidQueryType = errors.New("invalid query type") ErrInvalidQueryType = errors.New("invalid query type")
ErrInvalidUDPSize = errors.New("UDP size must be between 512 and 4096") ErrInvalidUDPSize = errors.New("UDP size must be between 512 and 4096")
ErrInvalidMaxDepth = errors.New("max depth must be between 1 and 100") ErrInvalidMaxDepth = errors.New("max depth must be between 1 and 100")
ErrInvalidRetries = errors.New("retries must be between 0 and 10") ErrInvalidRetries = errors.New("retries must be between 0 and 10")
ErrMissingDomain = errors.New("domain is required") ErrMissingDomain = errors.New("domain is required")
ErrAlwaysTCPRequiresTCP = errors.New("--always-tcp requires --allow-tcp") ErrAlwaysTCPRequiresTCP = errors.New("--always-tcp requires --allow-tcp")
ErrInvalidRootServer = errors.New("invalid root server IP address")
) )
type Config struct { type Config struct {
@@ -138,7 +139,23 @@ func (c *Config) ParseRootServer() (net.IP, error) {
if c.RootServer == "" { if c.RootServer == "" {
return nil, nil 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() { func PrintUsage() {
+216
View File
@@ -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)
}
}
}
+12 -10
View File
@@ -10,18 +10,20 @@ import (
) )
type QueryConfig struct { type QueryConfig struct {
UDPSize int UDPSize int
Timeout time.Duration Timeout time.Duration
Retries int Retries int
UseTCP bool UseTCP bool
AllowTCP bool
} }
func DefaultQueryConfig() *QueryConfig { func DefaultQueryConfig() *QueryConfig {
return &QueryConfig{ return &QueryConfig{
UDPSize: DefaultEDNS0UDPSize(), UDPSize: DefaultEDNS0UDPSize(),
Timeout: 5 * time.Second, Timeout: 5 * time.Second,
Retries: 3, Retries: 3,
UseTCP: false, UseTCP: false,
AllowTCP: true,
} }
} }
@@ -110,7 +112,7 @@ func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype ui
continue continue
} }
if resp.Truncated { if resp.Truncated && cfg.AllowTCP {
resp, err = exchangeFn(ctx, serverStr, msg, true) resp, err = exchangeFn(ctx, serverStr, msg, true)
if err != nil { if err != nil {
lastErr = err lastErr = err
@@ -178,7 +180,7 @@ func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string,
continue continue
} }
if resp.Truncated { if resp.Truncated && cfg.AllowTCP {
resp, err = exchangeFn(ctx, serverStr, msg, true) resp, err = exchangeFn(ctx, serverStr, msg, true)
if err != nil { if err != nil {
lastErr = err lastErr = err
+50 -8
View File
@@ -110,10 +110,11 @@ func TestQueryTCPFallbackOnTruncation(t *testing.T) {
} }
cfg := &QueryConfig{ cfg := &QueryConfig{
UDPSize: 2048, UDPSize: 2048,
Timeout: 5, Timeout: 5,
Retries: 1, Retries: 1,
UseTCP: false, UseTCP: false,
AllowTCP: true,
} }
server := net.ParseIP("8.8.8.8") server := net.ParseIP("8.8.8.8")
@@ -272,10 +273,11 @@ func TestQueryTCPFallbackFailsThenRetries(t *testing.T) {
} }
cfg := &QueryConfig{ cfg := &QueryConfig{
UDPSize: 2048, UDPSize: 2048,
Timeout: 5, Timeout: 5,
Retries: 2, Retries: 2,
UseTCP: false, UseTCP: false,
AllowTCP: true,
} }
server := net.ParseIP("8.8.8.8") 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) 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")
}
}