feat: implement CLI flags and configuration handling (HAN-382) #8
+27
-17
@@ -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
|
||||||
@@ -146,6 +155,7 @@ func main() {
|
|||||||
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{
|
||||||
|
|||||||
@@ -17,6 +17,7 @@ var (
|
|||||||
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() {
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -14,6 +14,7 @@ type QueryConfig struct {
|
|||||||
Timeout time.Duration
|
Timeout time.Duration
|
||||||
Retries int
|
Retries int
|
||||||
UseTCP bool
|
UseTCP bool
|
||||||
|
AllowTCP bool
|
||||||
}
|
}
|
||||||
|
|
||||||
func DefaultQueryConfig() *QueryConfig {
|
func DefaultQueryConfig() *QueryConfig {
|
||||||
@@ -22,6 +23,7 @@ func DefaultQueryConfig() *QueryConfig {
|
|||||||
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
|
||||||
|
|||||||
@@ -114,6 +114,7 @@ func TestQueryTCPFallbackOnTruncation(t *testing.T) {
|
|||||||
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")
|
||||||
@@ -276,6 +277,7 @@ func TestQueryTCPFallbackFailsThenRetries(t *testing.T) {
|
|||||||
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")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user