feat: implement CLI flags and configuration handling (HAN-382) #8
+193
-10
@@ -1,26 +1,209 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"context"
|
||||
"flag"
|
||||
"fmt"
|
||||
"os"
|
||||
"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] <domain>\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: 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")
|
||||
|
||||
flag.Usage = config.PrintUsage
|
||||
|
||||
flag.Parse()
|
||||
|
||||
if flag.NArg() == 0 {
|
||||
flag.Usage()
|
||||
args := flag.Args()
|
||||
|
||||
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 = verboseVal
|
||||
cfg.Debug = config.ParseDebugLevel(dFlag, ddFlag)
|
||||
cfg.Quiet = quietVal
|
||||
|
||||
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,
|
||||
AllowTCP: cfg.AllowTCP,
|
||||
}
|
||||
|
||||
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))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1 +1,238 @@
|
||||
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")
|
||||
ErrInvalidRootServer = errors.New("invalid root server IP address")
|
||||
)
|
||||
|
||||
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
|
||||
}
|
||||
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() {
|
||||
fmt.Fprintf(os.Stderr, "ExploreDNS - DNS reconnaissance and exploration tool\n\n")
|
||||
fmt.Fprintf(os.Stderr, "Usage:\n")
|
||||
fmt.Fprintf(os.Stderr, " exploredns [flags] <domain>\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,
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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")
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user