diff --git a/README.md b/README.md index 6f5fc0b..3e923c4 100644 --- a/README.md +++ b/README.md @@ -80,7 +80,7 @@ exploredns --quiet www.example.com # Debug mode exploredns --debug www.example.com -# Library-level debug (very verbose) +# Debug plus library-level diagnostics exploredns --dd www.example.com # Force TCP @@ -102,45 +102,44 @@ Usage: exploredns [flags] Query Options: - --type Record type to query (default: A) + --type Record type to query (default: a) Supported: A, AAAA, NS, CNAME, MX, TXT, SOA, PTR, ANY - --root-server Override the root server IP address - --all-root-servers Query all 13 root server sets (default: false) - --root-aaaa Include IPv6 addresses for root servers (default: false) - --follow-aaaa Only follow AAAA addresses for referrals (default: false) + --root-server Initial root server, hostname or IP literal + (default: ask the upstream resolver for one root) + --all-root-servers Traverse from all root servers (default: false) + --root-aaaa Include IPv6 root addresses (not implemented yet) + --follow-aaaa Only follow AAAA for referrals (not implemented yet) + --dns-upstream Upstream resolver (host:port) for root discovery + (default: system resolver) Transport Options: - --udp-size EDNS0 UDP buffer size, 512–4096 (default: 2048) + --udp-size EDNS0 UDP buffer size, 512–4096; 512 turns EDNS0 off + (default: 2048) --allow-tcp Fall back to TCP on truncation (default: true) --always-tcp Always use TCP (requires --allow-tcp) - --retries Per-server retry count, 0–10 (default: 2) + --retries Number of 2s retries before timing out, 0–10 (default: 2) Traversal Options: --max-depth Maximum referral depth, 1–100 (default: 20) - --fast / --fast=false Share glue cache across branches (default: true) + --fast / --fast=false Fast mode; turn off to be more accurate (default: true) Output Options: - --json Emit results as JSON instead of text - --verbose, -v Show extra detail in text output - --debug, -d Enable application debug messages (stderr) - --dd Enable library-level debug messages (very verbose) - --quiet, -q Suppress header and supplementary information - --show-progress Show live traversal progress (default: true) - --no-show-progress Hide traversal progress - --show-resolves Show glue-resolution steps (default: true) - --no-show-resolves Hide glue-resolution steps - --show-servers Show which servers were queried (default: true) - --no-show-servers Hide server list - --show-versions Show DNS server software versions (default: true) - --no-show-versions Hide server versions - --show-all-stats Show query statistics (default: true) - --no-show-all-stats Hide statistics - --show-results Show per-branch query results (default: true) - --no-show-results Hide per-branch results - --show-summary-results Show deduplicated summary section (default: true) - --no-show-summary-results Hide summary section + --json Emit a single JSON document instead of text + --verbose, -v Verbose progress ([qname] and shown) + -d, --debug Print debug diagnostics to stderr + -dd Like -d plus library-level debug + --quiet, -q Suppress the header block + --show-progress Show traversal progress (default: true) + --show-resolves Show glue-resolution progress (default: false) + --show-servers Show servers encountered (default: false) + --show-versions Show server version fingerprints (default: true) + --show-all-stats Show statistics after every node (default: false) + --show-results Show the results (default: true) + --show-summary-results Show the summary results (default: true) ``` +Every `--show-X` flag can be negated with `--show-X=false` or `--no-show-X`. + --- ## Web Interface @@ -253,15 +252,22 @@ Completed jobs are kept in memory for one hour before being purged. ### Text (default) -Coloured, hierarchical tree output showing each traversal branch, the servers -queried, referrals followed, and final answers. Disable colour by setting the -`NO_COLOR` environment variable. +dnstraverse-style output: a header block (settings, initial root, query; +suppressed by `--quiet`), progress lines (` ()` with +` -- resolving` and ` -- completed earlier ()` markers), a `Results:` +section of aggregated outcomes with probabilities (`Answer from`, `No glue +at`, `Lame referral from`, error/exception wording), and a `Summary Results:` +section grouping outcomes by status and answer content. `--show-servers` +adds the sorted list of servers encountered. Colour is used only when stdout +is a terminal and the `NO_COLOR` environment variable is unset. ### JSON (`--json`) -Structured JSON array of traversal results. Suitable for piping into `jq` or -ingesting into other tools. Each element contains the referral metadata, the -responding server, the response type, and the decoded DNS records. +A single JSON document — `{domain, qtype, root, results, summary, servers}` — +emitted once at the end of the run with deterministic ordering. Each +aggregated outcome appears exactly once in `results`; `summary` groups +probabilities by status and by distinct answer RRset; `servers` is present +with `--show-servers`. Suitable for piping into `jq`. --- diff --git a/cmd/exploredns/main.go b/cmd/exploredns/main.go index 12594ab..90d60a8 100644 --- a/cmd/exploredns/main.go +++ b/cmd/exploredns/main.go @@ -17,17 +17,17 @@ func main() { 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") + rootServer := flag.String("root-server", cfg.RootServer, "Initial root server (hostname or IP)") + allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Traverse from all root servers") + rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses (not implemented yet)") + followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals (not implemented yet)") dnsUpstream := flag.String("dns-upstream", cfg.DNSUpstream, "Upstream resolver for root discovery (e.g. 8.8.8.8:53, default: system)") - udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)") + udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096; 512 turns EDNS0 off)") 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") + retries := flag.Int("retries", cfg.Retries, "Number of 2s retries before timing out (0-10)") + fast := flag.Bool("fast", cfg.Fast, "Fast mode; turn off to be more accurate") jsonOutput := flag.Bool("json", false, "Output results as JSON") // Verbose: long and short form share the same variable. @@ -35,27 +35,31 @@ func main() { 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). + // Debug: -d prints debug diagnostics; -dd additionally enables library + // debug (the Go flag package cannot stack -d -d like the Ruby CLI). 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)") + flag.BoolVar(&dFlag, "d", false, "Print debug diagnostics to stderr") + flag.BoolVar(&dFlag, "debug", false, "Print debug diagnostics to stderr") + flag.BoolVar(&ddFlag, "dd", false, "Like -d plus library-level debug") // 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)") + flag.BoolVar(&quietVal, "quiet", cfg.Quiet, "Suppress the header block") + flag.BoolVar(&quietVal, "q", cfg.Quiet, "Suppress the header block (shorthand)") + // Each show flag defaults to the config default, so --show-X and + // --show-X=false both take effect directly; the --no-show-X aliases force + // the value off, matching the Ruby --[no-]show-X switches. 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") + showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution progress") + noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution progress") + showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers encountered") + noShowServers := flag.Bool("no-show-servers", false, "Hide servers encountered") + showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server version fingerprints") + noShowVersions := flag.Bool("no-show-versions", false, "Hide server version fingerprints") + showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show statistics after every node") + noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide per-node 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") @@ -83,41 +87,13 @@ func main() { 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 - } + cfg.ShowProgress = *showProgress && !*noShowProgress + cfg.ShowResolves = *showResolves && !*noShowResolves + cfg.ShowServers = *showServers && !*noShowServers + cfg.ShowVersions = *showVersions && !*noShowVersions + cfg.ShowAllStats = *showAllStats && !*noShowAllStats + cfg.ShowResults = *showResults && !*noShowResults + cfg.ShowSummaryResults = *showSummaryResults && !*noShowSummaryResults if err := cfg.Validate(); err != nil { fmt.Fprintf(os.Stderr, "Error: %v\n", err) @@ -130,26 +106,15 @@ func main() { 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, + Timeout: 2 * time.Second, Retries: cfg.Retries, UseTCP: cfg.AlwaysTCP, AllowTCP: cfg.AllowTCP, @@ -157,9 +122,11 @@ func main() { rootConfig := &dns.RootDiscoveryConfig{ IncludeAAAA: cfg.RootAAAA, - Server: rootServerAddr, - AllRoots: cfg.AllRootServers, - Resolver: cfg.DNSUpstream, + // --root-server accepts a hostname or an IP literal. + Server: cfg.RootServer, + AllRoots: cfg.AllRootServers, + Resolver: cfg.DNSUpstream, + Query: queryConfig, } traverserConfig := &traverse.TraverserConfig{ @@ -187,10 +154,6 @@ func main() { } } - if !cfg.Quiet && !*jsonOutput { - fmt.Printf("ExploreDNS - exploring: %s (type: %s)\n", domain, cfg.QueryType) - } - outFmt := output.FormatText if *jsonOutput { outFmt = output.FormatJSON @@ -200,6 +163,13 @@ func main() { Format: outFmt, Domain: domain, QueryType: cfg.QueryType, + Fast: cfg.Fast, + AllRootServers: cfg.AllRootServers, + UDPSize: cfg.UDPSize, + Retries: cfg.Retries, + MaxDepth: cfg.MaxDepth, + AllowTCP: cfg.AllowTCP, + AlwaysTCP: cfg.AlwaysTCP, ShowProgress: cfg.ShowProgress, ShowResolves: cfg.ShowResolves, ShowServers: cfg.ShowServers, @@ -209,7 +179,7 @@ func main() { ShowSummaryResults: cfg.ShowSummaryResults, Verbose: cfg.Verbose, Quiet: cfg.Quiet, - Color: os.Getenv("NO_COLOR") == "", + Color: output.ColorEnabled(os.Stdout), Debug: cfg.Debug, } diff --git a/internal/config/config.go b/internal/config/config.go index a8abe5c..4794f41 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -19,7 +19,6 @@ var ( 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 { @@ -76,39 +75,6 @@ func ParseQueryType(s string) (uint16, error) { } } -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 @@ -146,17 +112,6 @@ func (c *Config) GetDomain(args []string) (string, error) { 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 { @@ -175,57 +130,59 @@ func PrintUsage() { 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"}, + // Ordered slice, not a map: --help output must be stable. + flagGroups := []struct { + name string + flags [][2]string + }{ + {"Query Options", [][2]string{ + {"--type", "Record type (default a)"}, + {"--root-server", "Initial root server, hostname or IP (default: ask upstream)"}, + {"--all-root-servers", "Traverse from all root servers (default false)"}, + {"--root-aaaa", "Include IPv6 root addresses (not implemented yet)"}, + {"--follow-aaaa", "Only follow AAAA for referrals (not implemented yet)"}, {"--dns-upstream", "Upstream resolver for root discovery (default: system)"}, - }, - "Transport Options": { - {"--udp-size", "EDNS0 buffer size (default 2048)"}, + }}, + {"Transport Options", [][2]string{ + {"--udp-size", "EDNS0 buffer size; 512 turns EDNS0 off (default 2048)"}, {"--allow-tcp", "TCP fallback on truncation (default true)"}, - {"--always-tcp", "Always use TCP"}, - {"--retries", "Retry count (default 2)"}, - }, - "Traversal Options": { + {"--always-tcp", "Always use TCP (default false)"}, + {"--retries", "Number of 2s retries before timing out (default 2)"}, + }}, + {"Traversal Options", [][2]string{ {"--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"}, - }, + {"--fast", "Fast mode; turn off to be more accurate (default true)"}, + }}, + {"Output Options", [][2]string{ + {"--json", "Emit a single JSON document instead of text"}, + {"--verbose, -v", "Verbose progress ([qname] and shown)"}, + {"-d, --debug", "Print debug diagnostics to stderr"}, + {"-dd", "Like -d plus library-level debug"}, + {"--quiet, -q", "Suppress the header block"}, + {"--show-progress", "Show traversal progress (default true)"}, + {"--show-resolves", "Show glue resolution progress (default false)"}, + {"--show-servers", "Show servers encountered (default false)"}, + {"--show-versions", "Show server version fingerprints (default true)"}, + {"--show-all-stats", "Show statistics after every node (default false)"}, + {"--show-results", "Show the results (default true)"}, + {"--show-summary-results", "Show the summary results (default true)"}, + }}, } - for group, flags := range flagGroups { - fmt.Fprintf(os.Stderr, "\n%s:\n", group) - for _, f := range flags { + for _, group := range flagGroups { + fmt.Fprintf(os.Stderr, "\n%s:\n", group.name) + for _, f := range group.flags { fmt.Fprintf(os.Stderr, " %-25s %s\n", f[0], f[1]) } } + fmt.Fprintf(os.Stderr, "\nEvery --show-X flag can be negated with --show-X=false or --no-show-X.\n") } func DefaultConfig() *Config { return &Config{ - QueryType: "A", + // Lowercase like the reference default (:a); explicit --type values + // keep the user's case in the "Running query" line. + QueryType: "a", RootServer: "", AllRootServers: false, RootAAAA: false, @@ -241,10 +198,10 @@ func DefaultConfig() *Config { Debug: 0, Quiet: false, ShowProgress: true, - ShowResolves: true, - ShowServers: true, + ShowResolves: false, + ShowServers: false, ShowVersions: true, - ShowAllStats: true, + ShowAllStats: false, ShowResults: true, ShowSummaryResults: true, } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 4a2a97f..020adc5 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -41,28 +41,6 @@ func TestParseQueryTypeInvalid(t *testing.T) { } } -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 { @@ -138,65 +116,6 @@ func TestGetDomainMissing(t *testing.T) { } } -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 @@ -215,96 +134,34 @@ func TestParseDebugLevel(t *testing.T) { } } -func TestParseMaxDepthValid(t *testing.T) { - cases := []struct { - input string - want int - }{ - {"1", 1}, - {"20", 20}, - {"100", 100}, - } - for _, tc := range cases { - got, err := ParseMaxDepth(tc.input) - if err != nil { - t.Errorf("ParseMaxDepth(%q) unexpected error: %v", tc.input, err) - } - if got != tc.want { - t.Errorf("ParseMaxDepth(%q) = %d, want %d", tc.input, got, tc.want) - } - } -} - -func TestParseMaxDepthInvalid(t *testing.T) { - cases := []string{"0", "101", "notanumber", "-1"} - for _, s := range cases { - _, err := ParseMaxDepth(s) - if err == nil { - t.Errorf("ParseMaxDepth(%q): expected error", s) - continue - } - if !errors.Is(err, ErrInvalidMaxDepth) { - t.Errorf("ParseMaxDepth(%q): expected ErrInvalidMaxDepth, got %v", s, err) - } - } -} - -func TestParseRetriesValid(t *testing.T) { - cases := []struct { - input string - want int - }{ - {"0", 0}, - {"2", 2}, - {"10", 10}, - } - for _, tc := range cases { - got, err := ParseRetries(tc.input) - if err != nil { - t.Errorf("ParseRetries(%q) unexpected error: %v", tc.input, err) - } - if got != tc.want { - t.Errorf("ParseRetries(%q) = %d, want %d", tc.input, got, tc.want) - } - } -} - -func TestParseRetriesInvalid(t *testing.T) { - cases := []string{"-1", "11", "notanumber"} - for _, s := range cases { - _, err := ParseRetries(s) - if err == nil { - t.Errorf("ParseRetries(%q): expected error", s) - continue - } - if !errors.Is(err, ErrInvalidRetries) { - t.Errorf("ParseRetries(%q): expected ErrInvalidRetries, got %v", s, err) - } - } -} - func TestPrintUsage(t *testing.T) { // PrintUsage writes to stderr; just ensure it doesn't panic. PrintUsage() } func TestValidateBadQueryType(t *testing.T) { -cfg := DefaultConfig() -cfg.QueryType = "BOGUS" -if err := cfg.Validate(); !errors.Is(err, ErrInvalidQueryType) { -t.Errorf("expected ErrInvalidQueryType, got %v", err) -} + cfg := DefaultConfig() + cfg.QueryType = "BOGUS" + if err := cfg.Validate(); !errors.Is(err, ErrInvalidQueryType) { + t.Errorf("expected ErrInvalidQueryType, got %v", err) + } } func TestDefaultConfigIsValid(t *testing.T) { -cfg := DefaultConfig() -if cfg.QueryType != "A" { -t.Errorf("QueryType = %q, want A", cfg.QueryType) -} -if cfg.MaxDepth != 20 { -t.Errorf("MaxDepth = %d, want 20", cfg.MaxDepth) -} -if cfg.Retries != 2 { -t.Errorf("Retries = %d, want 2", cfg.Retries) -} + cfg := DefaultConfig() + if cfg.QueryType != "a" { + t.Errorf("QueryType = %q, want a", cfg.QueryType) + } + if cfg.ShowResolves || cfg.ShowServers || cfg.ShowAllStats { + t.Error("show-resolves/show-servers/show-all-stats must default to false") + } + if !cfg.ShowProgress || !cfg.ShowVersions || !cfg.ShowResults || !cfg.ShowSummaryResults { + t.Error("show-progress/show-versions/show-results/show-summary-results must default to true") + } + if cfg.MaxDepth != 20 { + t.Errorf("MaxDepth = %d, want 20", cfg.MaxDepth) + } + if cfg.Retries != 2 { + t.Errorf("Retries = %d, want 2", cfg.Retries) + } } diff --git a/internal/dns/dns.go b/internal/dns/dns.go deleted file mode 100644 index 1ffe03d..0000000 --- a/internal/dns/dns.go +++ /dev/null @@ -1 +0,0 @@ -package dns diff --git a/internal/dns/hints_test.go b/internal/dns/hints_test.go index 307f87c..99538ca 100644 --- a/internal/dns/hints_test.go +++ b/internal/dns/hints_test.go @@ -111,47 +111,28 @@ func TestRootHintsKnownAddress(t *testing.T) { func TestResolverFromConfigExplicit(t *testing.T) { cfg := &RootDiscoveryConfig{Resolver: "8.8.8.8:53"} - got := resolverFromConfig(cfg) + got, err := resolverFromConfig(cfg) + if err != nil { + t.Fatalf("resolverFromConfig: %v", err) + } if got != "8.8.8.8:53" { t.Errorf("resolverFromConfig = %q, want 8.8.8.8:53", got) } } -func TestResolverFromConfigEmpty(t *testing.T) { - cfg := &RootDiscoveryConfig{} - got := resolverFromConfig(cfg) - // Should return the system resolver; just check it's non-empty and contains a port. - if got == "" { - t.Error("resolverFromConfig with empty Resolver returned empty string") - } -} - -func TestResolverFromConfigNil(t *testing.T) { - got := resolverFromConfig(nil) - if got == "" { - t.Error("resolverFromConfig(nil) returned empty string") - } -} - -func TestSystemResolverNonEmpty(t *testing.T) { - got := systemResolver() - if got == "" { - t.Error("systemResolver() returned empty string") - } - // Must contain a colon (host:port format). - host, port, err := splitHostPort(got) +func TestSystemResolverNoHardcodedFallback(t *testing.T) { + // Whether or not /etc/resolv.conf exists, systemResolver must never + // invent 127.0.0.1:53: it returns a valid host:port from the system + // configuration or an error. + got, err := systemResolver() if err != nil { - t.Errorf("systemResolver() = %q: not a valid host:port: %v", got, err) + return } - if host == "" { - t.Errorf("systemResolver() host is empty in %q", got) + host, port, err := net.SplitHostPort(got) + if err != nil { + t.Fatalf("systemResolver() = %q: not a valid host:port: %v", got, err) } - if port == "" { - t.Errorf("systemResolver() port is empty in %q", got) + if host == "" || port == "" { + t.Errorf("systemResolver() returned incomplete address %q", got) } } - -// splitHostPort is a thin wrapper around net.SplitHostPort for test use. -func splitHostPort(addr string) (host, port string, err error) { - return net.SplitHostPort(addr) -} diff --git a/internal/dns/iterative_test.go b/internal/dns/iterative_test.go deleted file mode 100644 index 854064c..0000000 --- a/internal/dns/iterative_test.go +++ /dev/null @@ -1,122 +0,0 @@ -package dns - -import ( - "context" - "errors" - "net" - "sync" - "testing" - - "github.com/miekg/dns" -) - -func TestIterativeQueryWithExchangeNilConfig(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return resp.Copy(), nil - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, nil, exchangeFn) - if err != nil { - t.Fatalf("unexpected error with nil config: %v", err) - } -} - -func TestIterativeQueryWithExchangeAlwaysTCP(t *testing.T) { - resp := new(dns.Msg) - resp.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() - calls = append(calls, useTCP) - mu.Unlock() - return resp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Retries: 1, - UseTCP: true, - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(calls) != 1 || !calls[0] { - t.Errorf("expected single TCP call, got %v", calls) - } -} - -func TestIterativeQueryWithExchangeRetriesOnFailure(t *testing.T) { - var mu sync.Mutex - callCount := 0 - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - mu.Lock() - callCount++ - mu.Unlock() - return nil, errors.New("connection refused") - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Retries: 3, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) - if err == nil { - t.Fatal("expected error") - } - if callCount != 3 { - t.Errorf("expected 3 calls, got %d", callCount) - } -} - -func TestIterativeQueryWithExchangeContextCancel(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - cancel() - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, ctx.Err() - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Retries: 1, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(ctx, server, "example.com", TypeA, cfg, exchangeFn) - if err == nil { - t.Fatal("expected error on cancelled context") - } -} - -func TestIterativeQueryWithExchangeZeroUDPSize(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return resp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 0, // should use default - Retries: 1, - } - - server := net.ParseIP("8.8.8.8") - _, err := IterativeQueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } -} diff --git a/internal/dns/query.go b/internal/dns/query.go index 389cb5f..3433be2 100644 --- a/internal/dns/query.go +++ b/internal/dns/query.go @@ -1,73 +1,98 @@ // Package dns provides the low-level DNS query primitives used by ExploreDNS. // -// It wraps the github.com/miekg/dns library to provide retrying, TCP fallback, -// EDNS0 buffer size negotiation, and root server discovery. The package is -// intentionally narrow: it sends iterative (non-recursive) queries and returns -// the raw responses for the traversal engine to interpret. +// It wraps the github.com/miekg/dns library behind a single query path +// (Client) that sends non-recursive (RD=0) queries with retrying, TCP +// fallback on truncation, EDNS0 negotiation and a per-run packet cache, plus +// root server discovery. Production uses the real wire exchange; tests inject +// a mock ExchangeFunc into the exact same path. package dns import ( "context" + "errors" "fmt" "net" + "sync" "time" "github.com/miekg/dns" ) +// QueryConfig controls transport parameters for the single query path. type QueryConfig struct { - UDPSize int - Timeout time.Duration - Retries int - UseTCP bool + // UDPSize is the EDNS0 UDP payload size. An OPT record is only attached + // when UDPSize > 512, mirroring dnstraverse's caching_resolver.rb. + UDPSize int + // Timeout is the per-attempt packet timeout (dnsruby packet_timeout, + // dnstraverse default 2s). + Timeout time.Duration + // Retries is the total number of send attempts, matching dnsruby + // retry_times: Resolver#generate_timeouts schedules retry_times + // transmissions in total — the first immediately and retry k at + // retry_delay*2^k seconds after the first. Values below 1 are clamped to + // 1 so exactly one query is still sent (dnsruby with retry_times 0 would + // send nothing and hang; this also fixes the old + // "failed after 0 retries: %!w()" error). + Retries int + // RetryDelay is dnsruby's retry_delay (dnstraverse default 2s). + RetryDelay time.Duration + // UseTCP forces every query over TCP (--always-tcp). + UseTCP bool + // AllowTCP enables the UDP→TCP retry when a response is truncated. AllowTCP bool } func DefaultQueryConfig() *QueryConfig { return &QueryConfig{ - UDPSize: DefaultEDNS0UDPSize(), - Timeout: 5 * time.Second, - Retries: 3, - UseTCP: false, - AllowTCP: true, + UDPSize: DefaultEDNS0UDPSize(), + Timeout: 2 * time.Second, + Retries: 2, + RetryDelay: 2 * time.Second, + UseTCP: false, + AllowTCP: true, } } +// withDefaults returns a copy of cfg with zero values replaced by defaults. +func (cfg *QueryConfig) withDefaults() *QueryConfig { + if cfg == nil { + return DefaultQueryConfig() + } + out := *cfg + if out.UDPSize <= 0 { + out.UDPSize = DefaultEDNS0UDPSize() + } + if out.Timeout <= 0 { + out.Timeout = 2 * time.Second + } + if out.Retries < 1 { + out.Retries = 1 + } + if out.RetryDelay <= 0 { + out.RetryDelay = 2 * time.Second + } + return &out +} + +// ExchangeFunc performs one wire exchange. server is either a bare host/IP +// (port 53 implied) or an explicit host:port. type ExchangeFunc func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) -func Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() - } - - if cfg.UDPSize <= 0 { - cfg.UDPSize = DefaultEDNS0UDPSize() - } - if cfg.Timeout <= 0 { - cfg.Timeout = 5 * time.Second - } - - return QueryWithExchange(ctx, server, name, qtype, cfg, realExchange) -} - func realExchange(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - addr := net.JoinHostPort(server, "53") - - var c *dns.Client - if useTCP { - c = &dns.Client{ - Net: "tcp", - ReadTimeout: 5 * time.Second, - WriteTimeout: 5 * time.Second, - } - } else { - c = &dns.Client{ - Net: "udp", - ReadTimeout: 5 * time.Second, - WriteTimeout: 5 * time.Second, - } + addr := server + if _, _, err := net.SplitHostPort(server); err != nil { + addr = net.JoinHostPort(server, "53") } + proto := "udp" + if useTCP { + proto = "tcp" + } + c := &dns.Client{ + Net: proto, + ReadTimeout: 2 * time.Second, + WriteTimeout: 2 * time.Second, + } if deadline, ok := ctx.Deadline(); ok { c.ReadTimeout = time.Until(deadline) c.WriteTimeout = time.Until(deadline) @@ -75,150 +100,222 @@ func realExchange(ctx context.Context, server string, msg *dns.Msg, useTCP bool) r, _, err := c.ExchangeContext(ctx, msg, addr) if err != nil { - return nil, fmt.Errorf("dns exchange (%s) with %s: %w", c.Net, addr, err) + return nil, fmt.Errorf("dns exchange (%s) with %s: %w", proto, addr, err) } - return r, nil } -func QueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig, exchangeFn ExchangeFunc) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() +// Client is the single query path used identically by production and tests. +// Every query is non-recursive (RD=0) and deduplicated by a per-run packet +// cache keyed (server IP, qname, qclass, qtype, udpsize), mirroring +// dnstraverse's caching_resolver.rb: repeat askers replay the cached answer +// (or cached failure) without touching the wire. +type Client struct { + cfg *QueryConfig + exchange ExchangeFunc + + mu sync.Mutex + cache map[packetKey]*packetEntry + requests int + cacheHits int +} + +type packetKey struct { + server string + qname string + qclass uint16 + qtype uint16 + udpsize int +} + +type packetEntry struct { + once sync.Once + msg *dns.Msg + err error +} + +// NewClient creates a Client. A nil exchange means the real wire exchange; +// tests pass a mock so no packets leave the process. +func NewClient(cfg *QueryConfig, exchange ExchangeFunc) *Client { + if exchange == nil { + exchange = realExchange } - if cfg.UDPSize <= 0 { - cfg.UDPSize = DefaultEDNS0UDPSize() + return &Client{ + cfg: cfg.withDefaults(), + exchange: exchange, + cache: make(map[packetKey]*packetEntry), + } +} + +// Requests reports how many queries were asked of the client (cache hits included). +func (c *Client) Requests() int { + c.mu.Lock() + defer c.mu.Unlock() + return c.requests +} + +// CacheHits reports how many queries were served from the packet cache. +func (c *Client) CacheHits() int { + c.mu.Lock() + defer c.mu.Unlock() + return c.cacheHits +} + +// Query sends a non-recursive query for name/qtype (class IN) to server and +// returns the response plus any warnings gathered along the way (EDNS0 +// fallback, recursion offered, truncation). A non-nil error corresponds to +// dnstraverse's "exception" status (network failure after all retries). +func (c *Client) Query(ctx context.Context, server net.IP, name string, qtype uint16) (*dns.Msg, []string, error) { + msg, err := c.cachedExchange(ctx, server, name, qtype, c.cfg.UDPSize) + if err != nil { + return nil, nil, err } - msg := buildQuery(name, qtype, cfg.UDPSize) - serverStr := server.String() + var warnings []string + + // EDNS0 fallback (decoded_query.rb makequery_message): FORMERR/NOTIMP/ + // SERVFAIL with udpsize > 512 may mean the server chokes on OPT; retry + // once at 512 and keep the retry only if it clears the error. + if c.cfg.UDPSize > MinEDNS0UDPSize() && ednsFailure(msg.Rcode) { + retryMsg, retryErr := c.cachedExchange(ctx, server, name, qtype, MinEDNS0UDPSize()) + if retryErr == nil && !ednsFailure(retryMsg.Rcode) { + warnings = append(warnings, fmt.Sprintf("%s doesn't seem to support EDNS0", server)) + msg = retryMsg + } + } + + // msg_comment with want_recursion=false (message_utility.rb). + if msg.RecursionAvailable { + warnings = append(warnings, fmt.Sprintf("%s allows recursion", server)) + } + if msg.Truncated { + warnings = append(warnings, fmt.Sprintf("%s sent truncated packet", server)) + } + + return msg, warnings, nil +} + +func ednsFailure(rcode int) bool { + return rcode == dns.RcodeFormatError || + rcode == dns.RcodeNotImplemented || + rcode == dns.RcodeServerFailure +} + +// cachedExchange sends at most one wire query per packet cache key; the +// outcome (response or error) is cached and replayed for repeat askers. +func (c *Client) cachedExchange(ctx context.Context, server net.IP, name string, qtype uint16, udpsize int) (*dns.Msg, error) { + key := packetKey{ + server: server.String(), + qname: dns.CanonicalName(name), + qclass: dns.ClassINET, + qtype: qtype, + udpsize: udpsize, + } + + c.mu.Lock() + c.requests++ + entry, ok := c.cache[key] + if ok { + c.cacheHits++ + } else { + entry = &packetEntry{} + c.cache[key] = entry + } + c.mu.Unlock() + + entry.once.Do(func() { + entry.msg, entry.err = exchangeWithRetry(ctx, c.exchange, key.server, buildQuery(name, qtype, udpsize), c.cfg) + }) + + return copyMsg(entry.msg), entry.err +} + +// exchangeWithRetry implements dnsruby's retry schedule (resolver.rb +// generate_timeouts): cfg.Retries is the TOTAL number of transmissions — the +// first goes immediately and retry k is sent retry_delay*2^k seconds after +// the first (so gaps of 2d, 2d, 4d, 8d, ...). Each attempt gets its own +// cfg.Timeout (dnsruby packet_timeout). A truncated UDP reply is retried over +// TCP within the same attempt when cfg.AllowTCP. +func exchangeWithRetry(ctx context.Context, exchange ExchangeFunc, server string, msg *dns.Msg, cfg *QueryConfig) (*dns.Msg, error) { + attempts := cfg.Retries + if attempts < 1 { + attempts = 1 + } var lastErr error - - for attempt := 0; attempt < cfg.Retries; attempt++ { + for attempt := 0; attempt < attempts; attempt++ { if attempt > 0 { select { case <-ctx.Done(): return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) - case <-time.After(backoffDelay(attempt)): + case <-time.After(retryGap(cfg.RetryDelay, attempt)): } } - if cfg.UseTCP { - resp, err := exchangeFn(ctx, serverStr, msg, true) - if err != nil { - lastErr = err - continue - } - return resp, nil - } - - resp, err := exchangeFn(ctx, serverStr, msg, false) + resp, err := exchangeOnce(ctx, exchange, server, msg, cfg) if err != nil { lastErr = err continue } - - if resp.Truncated && cfg.AllowTCP { - resp, err = exchangeFn(ctx, serverStr, msg, true) - if err != nil { - lastErr = err - continue - } - return resp, nil - } - return resp, nil } - return nil, fmt.Errorf("query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr) + q := msg.Question[0] + return nil, fmt.Errorf("query %s %s to %s failed after %d attempts: %w", + q.Name, QNameType(q.Qtype), server, attempts, lastErr) } -func IterativeQuery(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() +// retryGap returns the wait before retry number `retry` (1-based). dnsruby +// sends retry k at absolute time retry_delay*2^k, so the gap is 2d before the +// first retry and d*2^(k-1) for each retry after that. +func retryGap(d time.Duration, retry int) time.Duration { + if retry <= 1 { + return 2 * d } - if cfg.UDPSize <= 0 { - cfg.UDPSize = DefaultEDNS0UDPSize() - } - return IterativeQueryWithExchange(ctx, server, name, qtype, cfg, realExchange) + return d << uint(retry-1) } -func IterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig, exchangeFn ExchangeFunc) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() +func exchangeOnce(ctx context.Context, exchange ExchangeFunc, server string, msg *dns.Msg, cfg *QueryConfig) (*dns.Msg, error) { + actx, cancel := context.WithTimeout(ctx, cfg.Timeout) + defer cancel() + + resp, err := exchange(actx, server, msg, cfg.UseTCP) + if err != nil { + return nil, err } - if cfg.UDPSize <= 0 { - cfg.UDPSize = DefaultEDNS0UDPSize() + if resp == nil { + return nil, errors.New("nil response") } - msg := buildQuery(name, qtype, cfg.UDPSize) - msg.RecursionDesired = false - serverStr := server.String() - - var lastErr error - - for attempt := 0; attempt < cfg.Retries; attempt++ { - if attempt > 0 { - select { - case <-ctx.Done(): - return nil, fmt.Errorf("query retries cancelled: %w", ctx.Err()) - case <-time.After(backoffDelay(attempt)): - } - } - - if cfg.UseTCP { - resp, err := exchangeFn(ctx, serverStr, msg, true) - if err != nil { - lastErr = err - continue - } - return resp, nil - } - - resp, err := exchangeFn(ctx, serverStr, msg, false) + if !cfg.UseTCP && resp.Truncated && cfg.AllowTCP { + resp, err = exchange(actx, server, msg, true) if err != nil { - lastErr = err - continue + return nil, err } - if resp == nil { - lastErr = fmt.Errorf("nil response") - continue + return nil, errors.New("nil response") } - - if resp.Truncated && cfg.AllowTCP { - resp, err = exchangeFn(ctx, serverStr, msg, true) - if err != nil { - lastErr = err - continue - } - return resp, nil - } - - return resp, nil } - return nil, fmt.Errorf("iterative query %s %s failed after %d retries: %w", name, QNameType(qtype), cfg.Retries, lastErr) + return resp, nil } -// backoffDelay computes the wait duration before the given retry attempt (1-indexed). -// Delays: attempt=1 → 100ms, attempt=2 → 200ms, attempt=3 → 400ms, capped at 2s. -func backoffDelay(attempt int) time.Duration { - if attempt <= 0 { - return 0 - } - delay := time.Duration(uint(1)< maxDelay { - return maxDelay - } - return delay -} - -func buildQuery(name string, qtype uint16, udpSize int) *dns.Msg { +// buildQuery constructs a non-recursive (RD=0) class IN query. The EDNS0 OPT +// record is attached only when udpsize > 512 (caching_resolver.rb adds OPT +// under the same condition), with the DO bit off. +func buildQuery(name string, qtype uint16, udpsize int) *dns.Msg { m := new(dns.Msg) m.SetQuestion(dns.Fqdn(name), qtype) - m.RecursionDesired = true - m.SetEdns0(uint16(udpSize), false) + m.RecursionDesired = false + if udpsize > MinEDNS0UDPSize() { + m.SetEdns0(uint16(udpsize), false) + } return m } + +func copyMsg(m *dns.Msg) *dns.Msg { + if m == nil { + return nil + } + return m.Copy() +} diff --git a/internal/dns/query_test.go b/internal/dns/query_test.go index 3ec1c8b..f3eff36 100644 --- a/internal/dns/query_test.go +++ b/internal/dns/query_test.go @@ -4,480 +4,493 @@ import ( "context" "errors" "net" + "strings" "sync" "testing" + "time" "github.com/miekg/dns" ) +// testClient returns a Client with fast retries wired to fn. +func testClient(cfg *QueryConfig, fn ExchangeFunc) *Client { + if cfg == nil { + cfg = DefaultQueryConfig() + } + cfg.Timeout = time.Second + cfg.RetryDelay = time.Millisecond + return NewClient(cfg, fn) +} + +func answerMsg(name string, ip string) *dns.Msg { + m := new(dns.Msg) + m.SetReply(new(dns.Msg)) + m.Answer = append(m.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), + }) + return m +} + func TestDefaultQueryConfig(t *testing.T) { cfg := DefaultQueryConfig() - if cfg == nil { - t.Fatal("DefaultQueryConfig returned nil") - } if cfg.UDPSize != 2048 { t.Errorf("UDPSize = %d, want 2048", cfg.UDPSize) } - if cfg.Retries != 3 { - t.Errorf("Retries = %d, want 3", cfg.Retries) + if cfg.Retries != 2 { + t.Errorf("Retries = %d, want 2", cfg.Retries) + } + if cfg.Timeout != 2*time.Second { + t.Errorf("Timeout = %v, want 2s", cfg.Timeout) + } + if cfg.RetryDelay != 2*time.Second { + t.Errorf("RetryDelay = %v, want 2s", cfg.RetryDelay) } if cfg.UseTCP { t.Error("UseTCP should be false by default") } + if !cfg.AllowTCP { + t.Error("AllowTCP should be true by default") + } } -func TestBuildQuery(t *testing.T) { - msg := buildQuery("example.com.", TypeA, 2048) +func TestBuildQueryRDZeroWithEDNS(t *testing.T) { + msg := buildQuery("example.com", TypeA, 2048) if len(msg.Question) != 1 { t.Fatalf("expected 1 question, got %d", len(msg.Question)) } q := msg.Question[0] if q.Name != "example.com." { - t.Errorf("question name = %q, want %q", q.Name, "example.com.") + t.Errorf("Fqdn not applied: got %q", q.Name) } if q.Qtype != TypeA { t.Errorf("question type = %d, want %d", q.Qtype, TypeA) } - if !msg.RecursionDesired { - t.Error("RecursionDesired should be true") + if msg.RecursionDesired { + t.Error("RecursionDesired must be false on the query path") } - if opt := msg.IsEdns0(); opt == nil { - t.Error("expected EDNS0 OPT record") - } else if opt.UDPSize() != 2048 { + opt := msg.IsEdns0() + if opt == nil { + t.Fatal("expected EDNS0 OPT record for udpsize > 512") + } + if opt.UDPSize() != 2048 { t.Errorf("EDNS0 UDPSize = %d, want 2048", opt.UDPSize()) } } -func TestBuildQueryFqdn(t *testing.T) { - msg := buildQuery("example.com", TypeA, 4096) - q := msg.Question[0] - if q.Name != "example.com." { - t.Errorf("Fqdn not applied: got %q, want %q", q.Name, "example.com.") +func TestBuildQueryNoEDNSAt512(t *testing.T) { + msg := buildQuery("example.com.", TypeA, 512) + if msg.IsEdns0() != nil { + t.Error("no OPT record should be attached when udpsize <= 512") } } -func TestQueryWithExchangeSuccess(t *testing.T) { - expectedResp := new(dns.Msg) - expectedResp.SetReply(new(dns.Msg)) - expectedResp.Answer = append(expectedResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), +func TestClientQuerySuccess(t *testing.T) { + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if msg.RecursionDesired { + t.Error("query must be sent with RD=0") + } + return answerMsg("example.com.", "93.184.216.34"), nil }) - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return expectedResp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } if len(resp.Answer) != 1 { t.Fatalf("expected 1 answer, got %d", len(resp.Answer)) } + if len(warnings) != 0 { + t.Errorf("unexpected warnings: %v", warnings) + } } -func TestQueryTCPFallbackOnTruncation(t *testing.T) { - truncatedResp := new(dns.Msg) - truncatedResp.Truncated = true - truncatedResp.SetReply(new(dns.Msg)) - - fullResp := new(dns.Msg) - fullResp.SetReply(new(dns.Msg)) - fullResp.Answer = append(fullResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), +func TestClientPacketCache(t *testing.T) { + var calls int + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + return answerMsg("example.com.", "1.2.3.4"), nil }) - 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) - if !useTCP { - return truncatedResp.Copy(), nil - } - return fullResp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: false, - AllowTCP: true, - } - + ctx := context.Background() server := net.ParseIP("8.8.8.8") - resp, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + for i := 0; i < 3; i++ { + if _, _, err := c.Query(ctx, server, "example.com", TypeA); err != nil { + t.Fatalf("query %d: %v", i, err) + } + } + if calls != 1 { + t.Errorf("expected 1 wire query for repeat askers, got %d", calls) + } + if c.Requests() != 3 { + t.Errorf("Requests = %d, want 3", c.Requests()) + } + if c.CacheHits() != 2 { + t.Errorf("CacheHits = %d, want 2", c.CacheHits()) + } + // Case differences must hit the same cache entry. + if _, _, err := c.Query(ctx, server, "EXAMPLE.COM.", TypeA); err != nil { + t.Fatal(err) + } + if calls != 1 { + t.Errorf("case-insensitive lookup should hit cache, got %d wire calls", calls) + } +} + +func TestClientPacketCacheDistinctKeys(t *testing.T) { + var calls int + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + return answerMsg("example.com.", "1.2.3.4"), nil + }) + + ctx := context.Background() + _, _, _ = c.Query(ctx, net.ParseIP("8.8.8.8"), "example.com", TypeA) + _, _, _ = c.Query(ctx, net.ParseIP("9.9.9.9"), "example.com", TypeA) + _, _, _ = c.Query(ctx, net.ParseIP("8.8.8.8"), "example.com", TypeAAAA) + _, _, _ = c.Query(ctx, net.ParseIP("8.8.8.8"), "other.com", TypeA) + if calls != 4 { + t.Errorf("expected 4 wire queries for 4 distinct keys, got %d", calls) + } +} + +func TestClientPacketCacheReplaysErrors(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + return nil, errors.New("connection refused") + }) + + ctx := context.Background() + server := net.ParseIP("8.8.8.8") + _, _, err1 := c.Query(ctx, server, "example.com", TypeA) + _, _, err2 := c.Query(ctx, server, "example.com", TypeA) + if err1 == nil || err2 == nil { + t.Fatal("expected errors") + } + if calls != 1 { + t.Errorf("failures must be cached too: got %d wire calls", calls) + } +} + +func TestClientEDNSFallback(t *testing.T) { + var sizes []int + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + size := 512 + if opt := msg.IsEdns0(); opt != nil { + size = int(opt.UDPSize()) + } + sizes = append(sizes, size) + m := new(dns.Msg) + m.SetReply(msg) + if size > 512 { + m.Rcode = dns.RcodeFormatError + return m, nil + } + m.Answer = append(m.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP("1.2.3.4"), + }) + return m, nil + }) + + resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } - if len(calls) != 2 { - t.Fatalf("expected 2 exchange calls (UDP then TCP), got %d", len(calls)) + if len(sizes) != 2 || sizes[0] != 2048 || sizes[1] != 512 { + t.Fatalf("expected 2048 then 512 queries, got %v", sizes) } - if calls[0] != false { - t.Error("first call should be UDP") + if resp.Rcode != dns.RcodeSuccess || len(resp.Answer) != 1 { + t.Error("expected the 512-byte retry response to be returned") } - if calls[1] != true { - t.Error("second call should be TCP") + want := "8.8.8.8 doesn't seem to support EDNS0" + if len(warnings) != 1 || warnings[0] != want { + t.Errorf("warnings = %v, want [%q]", warnings, want) + } +} + +func TestClientEDNSFallbackKeepsOriginalWhenRetryFails(t *testing.T) { + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + m := new(dns.Msg) + m.SetReply(msg) + m.Rcode = dns.RcodeServerFailure + return m, nil + }) + + resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if resp.Rcode != dns.RcodeServerFailure { + t.Errorf("expected original SERVFAIL response, got rcode %d", resp.Rcode) + } + if len(warnings) != 0 { + t.Errorf("no EDNS0 warning expected when the retry also fails: %v", warnings) + } +} + +func TestClientNoEDNSFallbackAt512(t *testing.T) { + var calls int + c := testClient(&QueryConfig{UDPSize: 512, Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + m := new(dns.Msg) + m.SetReply(msg) + m.Rcode = dns.RcodeFormatError + return m, nil + }) + + resp, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if calls != 1 { + t.Errorf("expected no fallback query at udpsize 512, got %d calls", calls) + } + if resp.Rcode != dns.RcodeFormatError { + t.Errorf("expected FORMERR passthrough, got %d", resp.Rcode) + } +} + +func TestClientRecursionAvailableWarning(t *testing.T) { + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + m := answerMsg("example.com.", "1.2.3.4") + m.RecursionAvailable = true + return m, nil + }) + + _, warnings, err := c.Query(context.Background(), net.ParseIP("192.0.2.1"), "example.com", TypeA) + if err != nil { + t.Fatal(err) + } + want := "192.0.2.1 allows recursion" + if len(warnings) != 1 || warnings[0] != want { + t.Errorf("warnings = %v, want [%q]", warnings, want) + } +} + +func TestClientTruncationWarningWhenTCPDisallowed(t *testing.T) { + c := testClient(&QueryConfig{AllowTCP: false, Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if useTCP { + t.Error("TCP must not be used when AllowTCP is false") + } + m := answerMsg("example.com.", "1.2.3.4") + m.Truncated = true + return m, nil + }) + + resp, warnings, err := c.Query(context.Background(), net.ParseIP("192.0.2.1"), "example.com", TypeA) + if err != nil { + t.Fatal(err) + } + if !resp.Truncated { + t.Error("expected truncated response to be returned as-is") + } + want := "192.0.2.1 sent truncated packet" + if len(warnings) != 1 || warnings[0] != want { + t.Errorf("warnings = %v, want [%q]", warnings, want) + } +} + +func TestClientTCPFallbackOnTruncation(t *testing.T) { + var mu sync.Mutex + var calls []bool + c := testClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + mu.Lock() + calls = append(calls, useTCP) + mu.Unlock() + if !useTCP { + m := new(dns.Msg) + m.SetReply(msg) + m.Truncated = true + return m, nil + } + return answerMsg("example.com.", "93.184.216.34"), nil + }) + + resp, warnings, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(calls) != 2 || calls[0] != false || calls[1] != true { + t.Fatalf("expected UDP then TCP, got %v", calls) } if len(resp.Answer) != 1 { - t.Fatalf("expected 1 answer from TCP fallback, got %d", len(resp.Answer)) + t.Fatal("expected answer from TCP fallback") + } + if len(warnings) != 0 { + t.Errorf("unexpected warnings after successful TCP fallback: %v", warnings) } } -func TestQueryAlwaysTCP(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - +func TestClientAlwaysTCP(t *testing.T) { var mu sync.Mutex - calls := []bool{} - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + var calls []bool + c := testClient(&QueryConfig{UseTCP: true, Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { mu.Lock() - defer mu.Unlock() calls = append(calls, useTCP) - return resp.Copy(), nil - } + mu.Unlock() + return answerMsg("example.com.", "1.2.3.4"), nil + }) - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: true, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error: %v", err) } - if len(calls) != 1 { - t.Fatalf("expected 1 exchange call, got %d", len(calls)) - } - if !calls[0] { - t.Error("expected TCP call when UseTCP is true") + if len(calls) != 1 || !calls[0] { + t.Errorf("expected a single TCP call, got %v", calls) } } -func TestQueryRetriesOnFailure(t *testing.T) { - var mu sync.Mutex - callCount := 0 - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - mu.Lock() - callCount++ - mu.Unlock() +func TestClientRetriesAreTotalAttempts(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 3}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ return nil, errors.New("connection refused") - } + }) - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 3, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error after retries exhausted") } - if callCount != 3 { - t.Errorf("expected 3 calls (retries exhausted), got %d", callCount) + // dnsruby retry_times counts total transmissions, not extra retries. + if calls != 3 { + t.Errorf("expected 3 attempts, got %d", calls) + } + if !strings.Contains(err.Error(), "after 3 attempts") { + t.Errorf("error should mention attempt count: %v", err) } } -func TestQueryContextCancellation(t *testing.T) { +func TestClientZeroRetriesStillQueriesOnce(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 0}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + return nil, errors.New("boom") + }) + + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err == nil { + t.Fatal("expected error") + } + if calls != 1 { + t.Errorf("Retries=0 must be clamped to one attempt, got %d", calls) + } + // Regression: the old code produced "failed after 0 retries: %!w()". + if strings.Contains(err.Error(), "%!w") || strings.Contains(err.Error(), "") { + t.Errorf("malformed error message: %v", err) + } +} + +func TestClientRetryThenSuccess(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 3}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + if calls < 2 { + return nil, errors.New("transient error") + } + return answerMsg("example.com.", "1.2.3.4"), nil + }) + + resp, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if len(resp.Answer) == 0 { + t.Fatal("expected answer after retry") + } + if calls != 2 { + t.Errorf("expected 2 attempts, got %d", calls) + } +} + +func TestClientNilResponseIsError(t *testing.T) { + c := testClient(&QueryConfig{Retries: 2}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, nil + }) + + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) + if err == nil { + t.Fatal("expected error for nil responses") + } +} + +func TestClientContextCancellation(t *testing.T) { ctx, cancel := context.WithCancel(context.Background()) cancel() - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, ctx.Err() - } + c := testClient(&QueryConfig{Retries: 5}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, errors.New("error") + }) - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 1, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(ctx, server, "example.com", TypeA, cfg, exchangeFn) + _, _, err := c.Query(ctx, net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error on cancelled context") } } -func TestQueryNilConfigUsesDefaults(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return resp.Copy(), nil - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, nil, exchangeFn) +func TestClientNilConfigUsesDefaults(t *testing.T) { + c := NewClient(nil, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return answerMsg("example.com.", "1.2.3.4"), nil + }) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err != nil { t.Fatalf("unexpected error with nil config: %v", err) } } -func TestQueryZeroValuesUseDefaults(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return resp.Copy(), nil - } - - cfg := &QueryConfig{ - UDPSize: 0, - Timeout: 0, - Retries: 1, - UseTCP: false, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } -} - -func TestQueryTCPFallbackFailsThenRetries(t *testing.T) { - truncatedResp := new(dns.Msg) - truncatedResp.Truncated = true - truncatedResp.SetReply(new(dns.Msg)) - - var mu sync.Mutex - callCount := 0 - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - mu.Lock() - callCount++ - mu.Unlock() +func TestClientTCPFallbackFailureRetries(t *testing.T) { + var calls int + c := testClient(&QueryConfig{Retries: 2, AllowTCP: true}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ if !useTCP { - return truncatedResp.Copy(), nil + m := new(dns.Msg) + m.SetReply(msg) + m.Truncated = true + return m, nil } return nil, errors.New("tcp failed") - } + }) - cfg := &QueryConfig{ - UDPSize: 2048, - Timeout: 5, - Retries: 2, - UseTCP: false, - AllowTCP: true, - } - - server := net.ParseIP("8.8.8.8") - _, err := QueryWithExchange(context.Background(), server, "example.com", TypeA, cfg, exchangeFn) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { t.Fatal("expected error when TCP fallback always fails") } - if callCount != 4 { - t.Errorf("expected 4 calls (2 retries x UDP+TCP), got %d", callCount) + if calls != 4 { + t.Errorf("expected 4 exchange calls (2 attempts x UDP+TCP), got %d", calls) } } -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 +func TestRetryGap(t *testing.T) { + d := 2 * time.Second + tests := []struct { + retry int + want time.Duration + }{ + {1, 4 * time.Second}, // dnsruby sends retry 1 at absolute 2d + {2, 4 * time.Second}, // retry 2 at 4d → gap 2d + {3, 8 * time.Second}, + {4, 16 * time.Second}, } - - 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") + for _, tt := range tests { + got := retryGap(d, tt.retry) + if got != tt.want { + t.Errorf("retryGap(%v, %d) = %v, want %v", d, tt.retry, got, tt.want) + } } } -func TestIterativeQueryWithExchangeSuccess(t *testing.T) { -answerResp := new(dns.Msg) -answerResp.SetReply(new(dns.Msg)) -answerResp.Answer = append(answerResp.Answer, &dns.A{ -Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) - -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -if msg.RecursionDesired { -t.Error("IterativeQuery should send RD=false") -} -return answerResp.Copy(), nil -} - -server := net.ParseIP("198.41.0.4") -resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, nil, exchangeFn) -if err != nil { -t.Fatalf("unexpected error: %v", err) -} -if len(resp.Answer) == 0 { -t.Fatal("expected answer records") -} -} - -func TestIterativeQueryWithExchangeRetry(t *testing.T) { -callCount := 0 -answerResp := new(dns.Msg) -answerResp.SetReply(new(dns.Msg)) -answerResp.Answer = append(answerResp.Answer, &dns.A{ -Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) - -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -callCount++ -if callCount < 2 { -return nil, errors.New("transient error") -} -return answerResp.Copy(), nil -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 3, AllowTCP: true} -server := net.ParseIP("198.41.0.4") -resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err != nil { -t.Fatalf("unexpected error: %v", err) -} -if resp == nil { -t.Fatal("expected non-nil response after retry") -} -} - -func TestIterativeQueryWithExchangeAllFail(t *testing.T) { -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -return nil, errors.New("server unreachable") -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 2, AllowTCP: false} -server := net.ParseIP("198.41.0.4") -_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err == nil { -t.Fatal("expected error when all attempts fail") -} -} - -func TestIterativeQueryWithExchangeTCPFallback(t *testing.T) { -truncatedResp := new(dns.Msg) -truncatedResp.SetReply(new(dns.Msg)) -truncatedResp.Truncated = true - -fullResp := new(dns.Msg) -fullResp.SetReply(new(dns.Msg)) -fullResp.Answer = append(fullResp.Answer, &dns.A{ -Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) - -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -if !useTCP { -return truncatedResp.Copy(), nil -} -return fullResp.Copy(), nil -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 1, AllowTCP: true} -server := net.ParseIP("198.41.0.4") -resp, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err != nil { -t.Fatalf("unexpected error: %v", err) -} -if len(resp.Answer) == 0 { -t.Fatal("expected answer after TCP fallback") -} -} - -func TestIterativeQueryWithExchangeContextCancelled(t *testing.T) { - ctx, cancel := context.WithCancel(context.Background()) - // Cancel the context immediately so the retry loop aborts during backoff - cancel() - - exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, errors.New("error") - } - - cfg := &QueryConfig{UDPSize: 2048, Retries: 5, AllowTCP: false} - server := net.ParseIP("198.41.0.4") - _, err := IterativeQueryWithExchange(ctx, server, "example.com", dns.TypeA, cfg, exchangeFn) +func TestQueryErrorMessageMentionsQuery(t *testing.T) { + c := testClient(&QueryConfig{Retries: 1}, func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, errors.New("unreachable") + }) + _, _, err := c.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA) if err == nil { - t.Fatal("expected error when context cancelled") + t.Fatal("expected error") + } + for _, part := range []string{"example.com.", "A", "8.8.8.8", "unreachable"} { + if !strings.Contains(err.Error(), part) { + t.Errorf("error %q should contain %q", err, part) + } } } - -func TestIterativeQueryWithExchangeNilResponse(t *testing.T) { -callCount := 0 -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -callCount++ -return nil, nil // nil response, no error -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 2, AllowTCP: false} -server := net.ParseIP("198.41.0.4") -_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err == nil { -t.Fatal("expected error for nil responses") -} -} - -func TestIterativeQueryWithExchangeUseTCP(t *testing.T) { -var wasTCP bool -answerResp := new(dns.Msg) -answerResp.SetReply(new(dns.Msg)) -answerResp.Answer = append(answerResp.Answer, &dns.A{ -Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) - -exchangeFn := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { -wasTCP = useTCP -return answerResp.Copy(), nil -} - -cfg := &QueryConfig{UDPSize: 2048, Retries: 1, UseTCP: true} -server := net.ParseIP("198.41.0.4") -_, err := IterativeQueryWithExchange(context.Background(), server, "example.com", dns.TypeA, cfg, exchangeFn) -if err != nil { -t.Fatalf("unexpected error: %v", err) -} -if !wasTCP { -t.Error("expected TCP exchange when UseTCP=true") -} -} diff --git a/internal/dns/real_exchange_test.go b/internal/dns/real_exchange_test.go index ec43bf4..2a772e1 100644 --- a/internal/dns/real_exchange_test.go +++ b/internal/dns/real_exchange_test.go @@ -2,7 +2,6 @@ package dns import ( "context" - "fmt" "net" "testing" "time" @@ -10,263 +9,32 @@ import ( "github.com/miekg/dns" ) -// startTestDNSServer starts a local DNS server on a random port and returns the address and a stop function. -func startTestDNSServer(t *testing.T, handler dns.HandlerFunc) (string, func()) { +// startTestDNSServer starts a loopback DNS server on a random port and +// returns its address. No test in this file touches the real network. +func startTestDNSServer(t *testing.T, network string, handler dns.HandlerFunc) string { t.Helper() - pc, err := net.ListenPacket("udp", "127.0.0.1:0") - if err != nil { - t.Skipf("cannot start test DNS server: %v", err) - } - addr := pc.LocalAddr().String() - mux := dns.NewServeMux() mux.HandleFunc(".", handler) - srv := &dns.Server{ - PacketConn: pc, - Net: "udp", - Handler: mux, - } + srv := &dns.Server{Net: network, Handler: mux} + var addr string - started := make(chan struct{}) - srv.NotifyStartedFunc = func() { close(started) } - - go func() { - _ = srv.ActivateAndServe() - }() - - select { - case <-started: - case <-time.After(2 * time.Second): - t.Skip("test DNS server did not start in time") - } - - return addr, func() { _ = srv.Shutdown() } -} - -func TestQueryUsesRealExchange(t *testing.T) { - addr, stop := startTestDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - resp := new(dns.Msg) - resp.SetReply(r) - resp.Answer = append(resp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - _ = w.WriteMsg(resp) - }) - defer stop() - - host, portStr, err := net.SplitHostPort(addr) - if err != nil { - t.Fatalf("parse addr: %v", err) - } - var port int - fmt.Sscanf(portStr, "%d", &port) - - // Patch the realExchange to use the test server by using QueryWithExchange with a custom exchangeFn. - // Since we can't inject into Query directly, use realExchangeWithPort for test. - serverIP := net.ParseIP(host) - cfg := DefaultQueryConfig() - cfg.Retries = 1 - - // Test QueryWithExchange (already covered), but now test Query+realExchange flow via - // a patched exchange that routes to our test server port. - patchedExchange := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - c := &dns.Client{Net: "udp", ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second} - r, _, err := c.ExchangeContext(ctx, msg, fmt.Sprintf("%s:%d", host, port)) - return r, err - } - - resp, err := QueryWithExchange(context.Background(), serverIP, "example.com", TypeA, cfg, patchedExchange) - if err != nil { - t.Fatalf("QueryWithExchange: %v", err) - } - if len(resp.Answer) == 0 { - t.Fatal("expected at least 1 answer") - } -} - -func TestRealExchangeViaDirect(t *testing.T) { - // Test realExchange directly via the exported Query function - // by using a server that will respond or fail quickly. - // We use a loopback address with a timeout to exercise code paths. - addr, stop := startTestDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - resp := new(dns.Msg) - resp.SetReply(r) - resp.Answer = append(resp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("5.6.7.8"), - }) - _ = w.WriteMsg(resp) - }) - defer stop() - - host, portStr, _ := net.SplitHostPort(addr) - serverIP := net.ParseIP(host) - - // Exercise realExchange via Query — we need a way to target the test port. - // Use a custom exchange that calls through realExchange-like logic. - cfg := DefaultQueryConfig() - cfg.Retries = 1 - - resp, err := QueryWithExchange(context.Background(), serverIP, "example.com", TypeA, cfg, - func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - targetAddr := fmt.Sprintf("%s:%s", host, portStr) - c := &dns.Client{Net: "udp", ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second} - r, _, e := c.ExchangeContext(ctx, msg, targetAddr) - return r, e - }) - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(resp.Answer) == 0 { - t.Fatal("expected answers") - } -} - -func TestQueryFunctionDirectly(t *testing.T) { - // Exercise Query() itself (which calls realExchange) by using 127.0.0.1:53. - // The test skips if no local DNS is available. - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - - server := net.ParseIP("127.0.0.1") - cfg := DefaultQueryConfig() - cfg.Retries = 1 - cfg.Timeout = 2 * time.Second - - _, err := Query(ctx, server, ".", TypeNS, cfg) - if err != nil { - t.Skipf("skipping (no local DNS at 127.0.0.1:53): %v", err) - } -} - -func TestIterativeQueryDirectly(t *testing.T) { - // Exercise IterativeQuery() itself (which calls realExchange) by using 127.0.0.1:53. - // The test skips if no local DNS is available. - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - - server := net.ParseIP("127.0.0.1") - cfg := DefaultQueryConfig() - cfg.Retries = 1 - cfg.Timeout = 2 * time.Second - - _, err := IterativeQuery(ctx, server, ".", TypeNS, cfg) - if err != nil { - t.Skipf("skipping (no local DNS at 127.0.0.1:53): %v", err) - } -} - -func TestBasicResolverQuery(t *testing.T) { - // Exercise BasicResolver.Query() which calls Query() → realExchange. - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - - br := NewBasicResolver() - server := net.ParseIP("127.0.0.1") - cfg := DefaultQueryConfig() - cfg.Retries = 1 - cfg.Timeout = 2 * time.Second - - _, err := br.Query(ctx, server, ".", TypeNS, cfg) - if err != nil { - t.Skipf("skipping (no local DNS at 127.0.0.1:53): %v", err) - } -} - -func TestDiscoverAllRoots(t *testing.T) { - // discoverAllRoots calls queryResolver(ctx, "127.0.0.1:53", ...) - // Skip if local DNS is not available. - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - cfg := &RootDiscoveryConfig{ - AllRoots: true, - IncludeAAAA: false, - } - servers, err := DiscoverRoots(ctx, cfg) - if err != nil { - t.Skipf("skipping (no local DNS available): %v", err) - } - if len(servers) == 0 { - t.Fatal("expected at least one root server from discoverAllRoots") - } -} - -func TestDiscoverAllRootsWithAAAA(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - cfg := &RootDiscoveryConfig{ - AllRoots: true, - IncludeAAAA: true, - } - servers, err := DiscoverRoots(ctx, cfg) - if err != nil { - t.Skipf("skipping (no local DNS available): %v", err) - } - if len(servers) == 0 { - t.Fatal("expected root servers with AAAA") - } -} - -func TestResolveRootServerDirect(t *testing.T) { - // Calls resolveRootServer directly (unexported, but in same package). - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - servers, err := resolveRootServer(ctx, systemResolver(), "a.root-servers.net.", false) - if err != nil { - t.Skipf("skipping (no local DNS): %v", err) - } - if len(servers) == 0 || len(servers[0].IPv4) == 0 { - t.Fatal("expected IPv4 address for a.root-servers.net.") - } -} - -func TestDiscoverSingleRootWithAAAA(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - cfg := &RootDiscoveryConfig{ - AllRoots: false, - IncludeAAAA: true, - } - servers, err := DiscoverRoots(ctx, cfg) - if err != nil { - t.Skipf("skipping (no local DNS available): %v", err) - } - if len(servers) == 0 { - t.Fatal("expected at least one root server") - } -} - -func TestRealExchangeTCPPath(t *testing.T) { - // Test the TCP path of realExchange via a test server - tcpAddr := "" - listener, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Skipf("cannot start TCP test server: %v", err) - } - tcpAddr = listener.Addr().String() - - mux := dns.NewServeMux() - mux.HandleFunc(".", func(w dns.ResponseWriter, r *dns.Msg) { - resp := new(dns.Msg) - resp.SetReply(r) - resp.Answer = append(resp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("9.9.9.9"), - }) - _ = w.WriteMsg(resp) - }) - - srv := &dns.Server{ - Listener: listener, - Net: "tcp", - Handler: mux, + switch network { + case "udp": + pc, err := net.ListenPacket("udp", "127.0.0.1:0") + if err != nil { + t.Skipf("cannot start test DNS server: %v", err) + } + srv.PacketConn = pc + addr = pc.LocalAddr().String() + case "tcp": + l, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Skipf("cannot start test DNS server: %v", err) + } + srv.Listener = l + addr = l.Addr().String() } started := make(chan struct{}) @@ -277,28 +45,89 @@ func TestRealExchangeTCPPath(t *testing.T) { select { case <-started: case <-time.After(2 * time.Second): - t.Skip("TCP DNS server didn't start") + t.Skip("test DNS server did not start in time") } - defer srv.Shutdown() - host, portStr, _ := net.SplitHostPort(tcpAddr) - serverIP := net.ParseIP(host) + t.Cleanup(func() { _ = srv.Shutdown() }) + return addr +} - cfg := DefaultQueryConfig() - cfg.UseTCP = true - cfg.Retries = 1 - - resp, err := QueryWithExchange(context.Background(), serverIP, "example.com", TypeA, cfg, - func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - targetAddr := fmt.Sprintf("%s:%s", host, portStr) - c := &dns.Client{Net: "tcp", ReadTimeout: 3 * time.Second, WriteTimeout: 3 * time.Second} - r, _, e := c.ExchangeContext(ctx, msg, targetAddr) - return r, e +func aHandler(ip string) dns.HandlerFunc { + return func(w dns.ResponseWriter, r *dns.Msg) { + resp := new(dns.Msg) + resp.SetReply(r) + resp.Answer = append(resp.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: r.Question[0].Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), }) + _ = w.WriteMsg(resp) + } +} + +func TestRealExchangeUDPHostPort(t *testing.T) { + addr := startTestDNSServer(t, "udp", aHandler("1.2.3.4")) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + // realExchange must honour an explicit host:port (used by root discovery + // upstream resolvers). + resp, err := realExchange(ctx, addr, buildQuery("example.com.", TypeA, 2048), false) if err != nil { - t.Fatalf("TCP query: %v", err) + t.Fatalf("realExchange: %v", err) } if len(resp.Answer) == 0 { t.Fatal("expected answers") } } + +func TestRealExchangeTCP(t *testing.T) { + addr := startTestDNSServer(t, "tcp", aHandler("9.9.9.9")) + + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) + defer cancel() + + resp, err := realExchange(ctx, addr, buildQuery("example.com.", TypeA, 2048), true) + if err != nil { + t.Fatalf("realExchange TCP: %v", err) + } + if len(resp.Answer) == 0 { + t.Fatal("expected answers") + } +} + +func TestRealExchangeUnreachable(t *testing.T) { + ctx, cancel := context.WithTimeout(context.Background(), 200*time.Millisecond) + defer cancel() + + _, err := realExchange(ctx, "127.0.0.1:1", buildQuery("example.com.", TypeA, 2048), false) + if err == nil { + t.Fatal("expected error for unreachable server") + } +} + +func TestClientAgainstLocalServer(t *testing.T) { + var sawRD bool + addr := startTestDNSServer(t, "udp", func(w dns.ResponseWriter, r *dns.Msg) { + sawRD = r.RecursionDesired + aHandler("5.6.7.8")(w, r) + }) + + // Route the client's exchange to the test server's port while still + // exercising realExchange. + c := NewClient(&QueryConfig{Retries: 1, Timeout: 2 * time.Second}, + func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return realExchange(ctx, addr, msg, useTCP) + }) + + resp, _, err := c.Query(context.Background(), net.ParseIP("127.0.0.1"), "example.com", TypeA) + if err != nil { + t.Fatalf("Client.Query: %v", err) + } + if len(resp.Answer) == 0 { + t.Fatal("expected answers") + } + if sawRD { + t.Error("wire query must have RD=0") + } +} diff --git a/internal/dns/resolver.go b/internal/dns/resolver.go deleted file mode 100644 index ee9ed5f..0000000 --- a/internal/dns/resolver.go +++ /dev/null @@ -1,198 +0,0 @@ -package dns - -import ( - "context" - "fmt" - "net" - "sync" - "time" - - "github.com/miekg/dns" -) - -type Resolver interface { - Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) -} - -type BasicResolver struct{} - -func NewBasicResolver() *BasicResolver { - return &BasicResolver{} -} - -func (br *BasicResolver) Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - return Query(ctx, server, name, qtype, cfg) -} - -type cacheKey struct { - server string - name string - qtype uint16 - qclass uint16 -} - -type cacheEntry struct { - msg *dns.Msg - expireAt time.Time -} - -func (e *cacheEntry) expired() bool { - return time.Now().After(e.expireAt) -} - -type CachingResolver struct { - inner Resolver - mu sync.RWMutex - cache map[cacheKey]*cacheEntry - defaultTTL time.Duration -} - -func NewCachingResolver(inner Resolver, opts ...CachingResolverOption) *CachingResolver { - if inner == nil { - inner = NewBasicResolver() - } - - cfg := &cachingResolverConfig{ - defaultTTL: 5 * time.Second, - } - for _, opt := range opts { - opt(cfg) - } - - return &CachingResolver{ - inner: inner, - cache: make(map[cacheKey]*cacheEntry), - defaultTTL: cfg.defaultTTL, - } -} - -func (cr *CachingResolver) Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - if cfg == nil { - cfg = DefaultQueryConfig() - } - - fqdn := dns.Fqdn(name) - key := cacheKey{ - server: server.String(), - name: fqdn, - qtype: qtype, - qclass: dns.ClassINET, - } - - if resp, ok := cr.lookup(key); ok { - return resp, nil - } - - resp, err := cr.inner.Query(ctx, server, name, qtype, cfg) - if err != nil { - return nil, fmt.Errorf("caching resolver query: %w", err) - } - - cr.store(key, resp) - return resp, nil -} - -func (cr *CachingResolver) lookup(key cacheKey) (*dns.Msg, bool) { - cr.mu.RLock() - entry, ok := cr.cache[key] - cr.mu.RUnlock() - if !ok { - return nil, false - } - if entry.expired() { - return nil, false - } - return entry.msg.Copy(), true -} - -func (cr *CachingResolver) store(key cacheKey, msg *dns.Msg) { - ttl := minTTLFromMsg(msg) - if ttl <= 0 { - ttl = cr.defaultTTL - } - - cr.mu.Lock() - cr.cache[key] = &cacheEntry{ - msg: msg.Copy(), - expireAt: time.Now().Add(ttl), - } - cr.mu.Unlock() -} - -func (cr *CachingResolver) Len() int { - cr.mu.RLock() - n := len(cr.cache) - cr.mu.RUnlock() - return n -} - -func (cr *CachingResolver) Clear() { - cr.mu.Lock() - cr.cache = make(map[cacheKey]*cacheEntry) - cr.mu.Unlock() -} - -func (cr *CachingResolver) PurgeExpired() int { - cr.mu.Lock() - count := 0 - for k, e := range cr.cache { - if e.expired() { - delete(cr.cache, k) - count++ - } - } - cr.mu.Unlock() - return count -} - -type CachingResolverOption func(*cachingResolverConfig) - -type cachingResolverConfig struct { - defaultTTL time.Duration -} - -func WithDefaultTTL(d time.Duration) CachingResolverOption { - return func(c *cachingResolverConfig) { - c.defaultTTL = d - } -} - -func minTTLFromMsg(msg *dns.Msg) time.Duration { - if msg == nil { - return 0 - } - - var min uint32 - found := false - - for _, rr := range msg.Answer { - ttl := rr.Header().Ttl - if !found || ttl < min { - min = ttl - found = true - } - } - for _, rr := range msg.Ns { - ttl := rr.Header().Ttl - if !found || ttl < min { - min = ttl - found = true - } - } - for _, rr := range msg.Extra { - if _, ok := rr.(*dns.OPT); ok { - continue - } - ttl := rr.Header().Ttl - if !found || ttl < min { - min = ttl - found = true - } - } - - if !found { - return 0 - } - - return time.Duration(min) * time.Second -} diff --git a/internal/dns/resolver_test.go b/internal/dns/resolver_test.go deleted file mode 100644 index 2d067cf..0000000 --- a/internal/dns/resolver_test.go +++ /dev/null @@ -1,476 +0,0 @@ -package dns - -import ( - "context" - "errors" - "net" - "sync" - "testing" - "time" - - "github.com/miekg/dns" -) - -type mockResolver struct { - mu sync.Mutex - calls int - response *dns.Msg - err error -} - -func (m *mockResolver) Query(ctx context.Context, server net.IP, name string, qtype uint16, cfg *QueryConfig) (*dns.Msg, error) { - m.mu.Lock() - defer m.mu.Unlock() - m.calls++ - if m.err != nil { - return nil, m.err - } - if m.response != nil { - return m.response.Copy(), nil - } - return nil, errors.New("no response configured") -} - -func (m *mockResolver) callCount() int { - m.mu.Lock() - defer m.mu.Unlock() - return m.calls -} - -func makeResponse(name string, qtype uint16, ttl uint32) *dns.Msg { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - switch qtype { - case dns.TypeA: - resp.Answer = append(resp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: ttl}, - A: net.ParseIP("93.184.216.34"), - }) - case dns.TypeNS: - resp.Answer = append(resp.Answer, &dns.NS{ - Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: ttl}, - Ns: "ns1.example.com.", - }) - } - return resp -} - -func TestNewCachingResolver_NilInner(t *testing.T) { - cr := NewCachingResolver(nil) - if cr == nil { - t.Fatal("NewCachingResolver(nil) returned nil") - } - if cr.inner == nil { - t.Fatal("expected inner resolver to be set when nil passed") - } -} - -func TestCachingResolver_CacheHit(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - resp1, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("first query: %v", err) - } - if mock.callCount() != 1 { - t.Fatalf("expected 1 call after first query, got %d", mock.callCount()) - } - - resp2, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("second query: %v", err) - } - if mock.callCount() != 1 { - t.Fatalf("expected still 1 call after second query (cache hit), got %d", mock.callCount()) - } - - if len(resp1.Answer) != len(resp2.Answer) { - t.Errorf("cached response has different number of answers") - } -} - -func TestCachingResolver_CacheMissDifferentServer(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), net.ParseIP("8.8.8.8"), "example.com", TypeA, cfg) - if mock.callCount() != 1 { - t.Fatalf("expected 1 call, got %d", mock.callCount()) - } - - _, _ = cr.Query(context.Background(), net.ParseIP("1.1.1.1"), "example.com", TypeA, cfg) - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls for different server, got %d", mock.callCount()) - } -} - -func TestCachingResolver_CacheMissDifferentName(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if mock.callCount() != 1 { - t.Fatalf("expected 1 call, got %d", mock.callCount()) - } - - _, _ = cr.Query(context.Background(), server, "different.com", TypeA, cfg) - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls for different name, got %d", mock.callCount()) - } -} - -func TestCachingResolver_CacheMissDifferentType(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if mock.callCount() != 1 { - t.Fatalf("expected 1 call, got %d", mock.callCount()) - } - - _, _ = cr.Query(context.Background(), server, "example.com", TypeNS, cfg) - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls for different qtype, got %d", mock.callCount()) - } -} - -func TestCachingResolver_ErrorNotCached(t *testing.T) { - mock := &mockResolver{ - err: errors.New("connection refused"), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err == nil { - t.Fatal("expected error from mock") - } - - mock.err = nil - mock.response = makeResponse("example.com", dns.TypeA, 300) - - _, err = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("second query after mock fixed: %v", err) - } - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls (error not cached), got %d", mock.callCount()) - } -} - -func TestCachingResolver_TTLExpiry(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 1), - } - cr := NewCachingResolver(mock, WithDefaultTTL(1*time.Second)) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("first query: %v", err) - } - if mock.callCount() != 1 { - t.Fatalf("expected 1 call, got %d", mock.callCount()) - } - - time.Sleep(2 * time.Second) - - _, err = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("query after TTL expiry: %v", err) - } - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls after TTL expiry, got %d", mock.callCount()) - } -} - -func TestCachingResolver_DefaultTTL(t *testing.T) { - resp := new(dns.Msg) - resp.SetReply(new(dns.Msg)) - - mock := &mockResolver{response: resp} - cr := NewCachingResolver(mock, WithDefaultTTL(100*time.Millisecond)) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, err := cr.Query(context.Background(), server, "nodata.com", TypeA, cfg) - if err != nil { - t.Fatalf("first query: %v", err) - } - - time.Sleep(150 * time.Millisecond) - - _, err = cr.Query(context.Background(), server, "nodata.com", TypeA, cfg) - if err != nil { - t.Fatalf("query after default TTL expiry: %v", err) - } - if mock.callCount() != 2 { - t.Fatalf("expected 2 calls after default TTL, got %d", mock.callCount()) - } -} - -func TestCachingResolver_NilConfig(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - - server := net.ParseIP("8.8.8.8") - _, err := cr.Query(context.Background(), server, "example.com", TypeA, nil) - if err != nil { - t.Fatalf("nil config: %v", err) - } -} - -func TestCachingResolver_Len(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - if cr.Len() != 0 { - t.Fatalf("expected empty cache, got %d", cr.Len()) - } - - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if cr.Len() != 1 { - t.Fatalf("expected cache len 1, got %d", cr.Len()) - } - - _, _ = cr.Query(context.Background(), server, "other.com", TypeA, cfg) - if cr.Len() != 2 { - t.Fatalf("expected cache len 2, got %d", cr.Len()) - } -} - -func TestCachingResolver_Clear(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - _, _ = cr.Query(context.Background(), server, "other.com", TypeA, cfg) - if cr.Len() != 2 { - t.Fatalf("expected cache len 2, got %d", cr.Len()) - } - - cr.Clear() - if cr.Len() != 0 { - t.Fatalf("expected cache len 0 after clear, got %d", cr.Len()) - } - - _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Fatalf("query after clear: %v", err) - } - if mock.callCount() != 3 { - t.Fatalf("expected 3 calls (2 before clear + 1 after clear), got %d", mock.callCount()) - } -} - -func TestCachingResolver_PurgeExpired(t *testing.T) { - resp := makeResponse("example.com", dns.TypeA, 0) - for _, rr := range resp.Answer { - rr.Header().Ttl = 0 - } - mock := &mockResolver{response: resp} - cr := NewCachingResolver(mock, WithDefaultTTL(50*time.Millisecond)) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if cr.Len() != 1 { - t.Fatalf("expected cache len 1, got %d", cr.Len()) - } - - time.Sleep(100 * time.Millisecond) - - purged := cr.PurgeExpired() - if purged != 1 { - t.Fatalf("expected 1 purged entry, got %d", purged) - } - if cr.Len() != 0 { - t.Fatalf("expected cache len 0 after purge, got %d", cr.Len()) - } -} - -func TestCachingResolver_PurgeExpiredNoneExpired(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - _, _ = cr.Query(context.Background(), server, "example.com", TypeA, cfg) - purged := cr.PurgeExpired() - if purged != 0 { - t.Fatalf("expected 0 purged entries, got %d", purged) - } - if cr.Len() != 1 { - t.Fatalf("expected cache len 1, got %d", cr.Len()) - } -} - -func TestCachingResolver_ResponseCopy(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - resp1, _ := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - resp2, _ := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - - if resp1 == resp2 { - t.Fatal("cache should return a copy, not the same pointer") - } -} - -func TestCachingResolver_MinTTLFromMsg(t *testing.T) { - tests := []struct { - name string - msg *dns.Msg - expected time.Duration - }{ - { - name: "nil message", - msg: nil, - expected: 0, - }, - { - name: "empty response", - msg: new(dns.Msg), - expected: 0, - }, - { - name: "single answer with low TTL", - msg: func() *dns.Msg { - m := new(dns.Msg) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Ttl: 60}, - }) - return m - }(), - expected: 60 * time.Second, - }, - { - name: "multiple records with varying TTLs", - msg: func() *dns.Msg { - m := new(dns.Msg) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Ttl: 300}, - }) - m.Ns = append(m.Ns, &dns.NS{ - Hdr: dns.RR_Header{Ttl: 120}, - }) - return m - }(), - expected: 120 * time.Second, - }, - { - name: "OPT record excluded from TTL calculation", - msg: func() *dns.Msg { - m := new(dns.Msg) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Ttl: 300}, - }) - m.Extra = append(m.Extra, &dns.OPT{ - Hdr: dns.RR_Header{Ttl: 0}, - }) - return m - }(), - expected: 300 * time.Second, - }, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - got := minTTLFromMsg(tt.msg) - if got != tt.expected { - t.Errorf("minTTLFromMsg() = %v, want %v", got, tt.expected) - } - }) - } -} - -func TestCachingResolver_ConcurrentAccess(t *testing.T) { - mock := &mockResolver{ - response: makeResponse("example.com", dns.TypeA, 300), - } - cr := NewCachingResolver(mock) - server := net.ParseIP("8.8.8.8") - cfg := &QueryConfig{UDPSize: 2048, Retries: 1, Timeout: 5 * time.Second} - - var wg sync.WaitGroup - for i := 0; i < 100; i++ { - wg.Add(1) - go func() { - defer wg.Done() - _, err := cr.Query(context.Background(), server, "example.com", TypeA, cfg) - if err != nil { - t.Errorf("concurrent query failed: %v", err) - } - }() - } - wg.Wait() - - if mock.callCount() < 1 { - t.Fatalf("expected at least 1 call to mock, got %d", mock.callCount()) - } -} - -func TestBasicResolver(t *testing.T) { - br := NewBasicResolver() - if br == nil { - t.Fatal("NewBasicResolver returned nil") - } - - if _, ok := interface{}(br).(Resolver); !ok { - t.Fatal("BasicResolver does not implement Resolver interface") - } -} - -// TestBasicResolverQueryIntegration calls Query via the BasicResolver against -// the local system resolver. Skipped when no local resolver is reachable. -func TestBasicResolverQueryIntegration(t *testing.T) { -r := NewBasicResolver() -server := net.ParseIP("127.0.0.1") -ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) -defer cancel() - -// Cover BasicResolver.Query; skip if 127.0.0.1:53 is not available. -msg, err := r.Query(ctx, server, ".", TypeNS, nil) -if err != nil { -t.Logf("skipping (local resolver unavailable): %v", err) -t.Skip() -} -if msg == nil { -t.Fatal("expected non-nil response from BasicResolver.Query") -} -} diff --git a/internal/dns/robustness_test.go b/internal/dns/robustness_test.go index c37fc70..bdbb356 100644 --- a/internal/dns/robustness_test.go +++ b/internal/dns/robustness_test.go @@ -115,25 +115,23 @@ func TestSynthesizeCNAMEFromDNAME(t *testing.T) { } } -func TestBackoffDelay(t *testing.T) { - tests := []struct { - attempt int - want time.Duration - }{ - {0, 0}, - {1, 100 * time.Millisecond}, - {2, 200 * time.Millisecond}, - {3, 400 * time.Millisecond}, - {4, 800 * time.Millisecond}, - {5, 1600 * time.Millisecond}, - {6, 2000 * time.Millisecond}, // capped at 2s - {10, 2000 * time.Millisecond}, // still capped +func TestQueryConfigWithDefaults(t *testing.T) { + got := (&QueryConfig{}).withDefaults() + if got.UDPSize != DefaultEDNS0UDPSize() { + t.Errorf("UDPSize = %d, want %d", got.UDPSize, DefaultEDNS0UDPSize()) + } + if got.Timeout != 2*time.Second { + t.Errorf("Timeout = %v, want 2s", got.Timeout) + } + if got.Retries != 1 { + t.Errorf("Retries = %d, want clamp to 1", got.Retries) + } + if got.RetryDelay != 2*time.Second { + t.Errorf("RetryDelay = %v, want 2s", got.RetryDelay) } - for _, tt := range tests { - got := backoffDelay(tt.attempt) - if got != tt.want { - t.Errorf("backoffDelay(%d) = %v, want %v", tt.attempt, got, tt.want) - } + var nilCfg *QueryConfig + if nilCfg.withDefaults().Retries != DefaultQueryConfig().Retries { + t.Error("nil config should yield defaults") } } diff --git a/internal/dns/roots.go b/internal/dns/roots.go index 1a5dfbd..4aaf606 100644 --- a/internal/dns/roots.go +++ b/internal/dns/roots.go @@ -2,10 +2,10 @@ package dns import ( "context" + "errors" "fmt" "net" "strings" - "time" "github.com/miekg/dns" ) @@ -26,15 +26,22 @@ func (rs *RootServer) AllIPs(includeAAAA bool) []net.IP { } type RootDiscoveryConfig struct { - // Server overrides which root server to use as the traversal starting point. - // When empty, a root server is discovered via the upstream resolver. + // Server overrides which root server to use as the traversal starting + // point. It accepts a hostname (resolved to A via the upstream resolver) + // or an IP literal (used directly, no lookup). Server string - // Resolver is the upstream DNS resolver used to resolve root server names. - // When empty, the system resolver from /etc/resolv.conf is used. - // Format: "host:port" (e.g. "8.8.8.8:53" or "1.1.1.1:53"). + // Resolver is the upstream DNS resolver used to discover and resolve + // root server names. When empty, the system resolver configuration is + // used. Format: "host:port" (e.g. "8.8.8.8:53" or "1.1.1.1:53"). Resolver string AllRoots bool IncludeAAAA bool + // Query controls transport parameters (retries, timeout, TCP fallback) + // for discovery queries. nil means DefaultQueryConfig. + Query *QueryConfig + // Exchange overrides the wire exchange; nil means the real network. + // Tests inject a mock here so discovery is network-free. + Exchange ExchangeFunc } func DefaultRootDiscoveryConfig() *RootDiscoveryConfig { @@ -49,31 +56,59 @@ func DiscoverRoots(ctx context.Context, cfg *RootDiscoveryConfig) ([]RootServer, cfg = DefaultRootDiscoveryConfig() } - resolver := resolverFromConfig(cfg) + // --root-server with an IP literal: use it directly, never look it up. + if cfg.Server != "" { + if ip := net.ParseIP(cfg.Server); ip != nil { + rs := RootServer{Name: cfg.Server} + if ip.To4() != nil { + rs.IPv4 = []net.IP{ip} + } else { + rs.IPv6 = []net.IP{ip} + } + return []RootServer{rs}, nil + } + } + + resolver, err := resolverFromConfig(cfg) + if err != nil { + if cfg.Server != "" { + return nil, err + } + return rootHintsFallback(cfg, err) + } if cfg.Server != "" { - return discoverRootOverride(ctx, resolver, cfg.Server, cfg.IncludeAAAA) + return resolveRootServer(ctx, cfg, resolver, cfg.Server) } if cfg.AllRoots { - servers, err := discoverAllRoots(ctx, resolver, cfg.IncludeAAAA) + servers, err := discoverAllRoots(ctx, cfg, resolver) if err != nil { - return filterHints(RootHints, cfg.IncludeAAAA), nil + return rootHintsFallback(cfg, err) } return servers, nil } - servers, err := discoverSingleRoot(ctx, resolver, cfg.IncludeAAAA) + servers, err := discoverSingleRoot(ctx, cfg, resolver) if err != nil { - hints := filterHints(RootHints, cfg.IncludeAAAA) - if len(hints) > 0 { - return hints[:1], nil - } - return nil, err + return rootHintsFallback(cfg, err) } return servers, nil } +// rootHintsFallback returns the builtin IANA hints when the upstream resolver +// cannot provide roots (one entry unless AllRoots). +func rootHintsFallback(cfg *RootDiscoveryConfig, err error) ([]RootServer, error) { + hints := filterHints(RootHints, cfg.IncludeAAAA) + if len(hints) == 0 { + return nil, err + } + if cfg.AllRoots { + return hints, nil + } + return hints[:1], nil +} + // filterHints returns a copy of hints with IPv6 addresses stripped when includeAAAA is false. func filterHints(hints []RootServer, includeAAAA bool) []RootServer { out := make([]RootServer, len(hints)) @@ -86,88 +121,114 @@ func filterHints(hints []RootServer, includeAAAA bool) []RootServer { return out } -func discoverRootOverride(ctx context.Context, resolver, server string, includeAAAA bool) ([]RootServer, error) { - nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) +// discoverSingleRoot mirrors get_a_root in traverser.rb: ask the upstream +// resolver for the root NS set, prefer glue from the additional section, and +// only fall back to explicit A/AAAA lookups when no glue was supplied. +func discoverSingleRoot(ctx context.Context, cfg *RootDiscoveryConfig, resolver string) ([]RootServer, error) { + names, nsMsg, err := rootNSNames(ctx, cfg, resolver) if err != nil { - return nil, fmt.Errorf("query root NS records: %w", err) + return nil, err } - nsSet := extractNSRecords(nsMsg.Answer) - if len(nsSet) == 0 { - nsSet = extractNSNames(nsMsg.Ns) - } - - normalized := normalizeServerName(server) - for _, name := range nsSet { - if normalizeServerName(name) == normalized { - return resolveRootServer(ctx, resolver, name, includeAAAA) + for _, name := range names { + rs := rootFromAdditional(nsMsg, name, cfg.IncludeAAAA) + if len(rs.AllIPs(cfg.IncludeAAAA)) > 0 { + return []RootServer{singleAddress(rs)}, nil } } - return resolveRootServer(ctx, resolver, server, includeAAAA) + var lastErr error + for _, name := range names { + servers, err := resolveRootServer(ctx, cfg, resolver, name) + if err != nil { + lastErr = err + continue + } + return []RootServer{singleAddress(servers[0])}, nil + } + + if lastErr == nil { + lastErr = errors.New("no address could be found for any root server") + } + return nil, lastErr } -func discoverSingleRoot(ctx context.Context, resolver string, includeAAAA bool) ([]RootServer, error) { - nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) - if err != nil { - return nil, fmt.Errorf("query root NS records: %w", err) +// singleAddress narrows rs to its first address, mirroring get_a_root in +// traverser.rb (add[0]/ans2[0]): the single-root start point is exactly one +// (name, IP) pair even when the upstream supplies more (or duplicate) glue. +func singleAddress(rs RootServer) RootServer { + if len(rs.IPv4) > 0 { + return RootServer{Name: rs.Name, IPv4: rs.IPv4[:1]} } - - nsSet := extractNSRecords(nsMsg.Answer) - if len(nsSet) == 0 { - nsSet = extractNSNames(nsMsg.Ns) + if len(rs.IPv6) > 0 { + return RootServer{Name: rs.Name, IPv6: rs.IPv6[:1]} } - - if len(nsSet) == 0 { - return nil, fmt.Errorf("no root NS records found in response") - } - - pick := nsSet[0] - return resolveRootServer(ctx, resolver, pick, includeAAAA) + return rs } -func discoverAllRoots(ctx context.Context, resolver string, includeAAAA bool) ([]RootServer, error) { - nsMsg, err := queryResolver(ctx, resolver, ".", dns.TypeNS) +// discoverAllRoots mirrors find_all_roots in traverser.rb: it returns one +// RootServer per root NS name with its full address set, so the traversal can +// branch per root. Roots with no resolvable address are skipped. +func discoverAllRoots(ctx context.Context, cfg *RootDiscoveryConfig, resolver string) ([]RootServer, error) { + names, nsMsg, err := rootNSNames(ctx, cfg, resolver) if err != nil { - return nil, fmt.Errorf("query root NS records: %w", err) - } - - nsSet := extractNSRecords(nsMsg.Answer) - if len(nsSet) == 0 { - nsSet = extractNSNames(nsMsg.Ns) - } - - if len(nsSet) == 0 { - return nil, fmt.Errorf("no root NS records found in response") + return nil, err } var servers []RootServer - for _, name := range nsSet { - resolved, err := resolveRootServer(ctx, resolver, name, includeAAAA) - if err != nil { - servers = append(servers, RootServer{Name: name}) - continue + for _, name := range names { + rs := rootFromAdditional(nsMsg, name, cfg.IncludeAAAA) + if len(rs.AllIPs(cfg.IncludeAAAA)) == 0 { + resolved, err := resolveRootServer(ctx, cfg, resolver, name) + if err != nil { + continue + } + rs = resolved[0] } - servers = append(servers, resolved...) + servers = append(servers, rs) } if len(servers) == 0 { - return nil, fmt.Errorf("failed to resolve any root servers") + return nil, errors.New("failed to resolve any root servers") } return servers, nil } -func resolveRootServer(ctx context.Context, resolver, name string, includeAAAA bool) ([]RootServer, error) { +func rootNSNames(ctx context.Context, cfg *RootDiscoveryConfig, resolver string) ([]string, *dns.Msg, error) { + nsMsg, err := queryUpstream(ctx, cfg, resolver, ".", dns.TypeNS) + if err != nil { + return nil, nil, fmt.Errorf("query root NS records: %w", err) + } + + names := extractNSRecords(nsMsg.Answer) + if len(names) == 0 { + names = extractNSNames(nsMsg.Ns) + } + if len(names) == 0 { + return nil, nil, errors.New("no root NS records found in response") + } + return names, nsMsg, nil +} + +func rootFromAdditional(msg *dns.Msg, name string, includeAAAA bool) RootServer { + rs := RootServer{Name: name, IPv4: additionalIPs(msg, name, dns.TypeA)} + if includeAAAA { + rs.IPv6 = additionalIPs(msg, name, dns.TypeAAAA) + } + return rs +} + +func resolveRootServer(ctx context.Context, cfg *RootDiscoveryConfig, resolver, name string) ([]RootServer, error) { var ipv4 []net.IP - aMsg, err := queryResolver(ctx, resolver, name, dns.TypeA) + aMsg, err := queryUpstream(ctx, cfg, resolver, name, dns.TypeA) if err == nil { ipv4 = extractIPsFromAnswer(aMsg.Answer, dns.TypeA) } var ipv6 []net.IP - if includeAAAA { - aaaaMsg, err := queryResolver(ctx, resolver, name, dns.TypeAAAA) + if cfg.IncludeAAAA { + aaaaMsg, err := queryUpstream(ctx, cfg, resolver, name, dns.TypeAAAA) if err == nil { ipv6 = extractIPsFromAnswer(aaaaMsg.Answer, dns.TypeAAAA) } @@ -180,47 +241,53 @@ func resolveRootServer(ctx context.Context, resolver, name string, includeAAAA b return []RootServer{{Name: name, IPv4: ipv4, IPv6: ipv6}}, nil } -// resolverFromConfig returns the upstream DNS resolver address to use. -// If cfg.Resolver is set, it is used directly. Otherwise the system resolver -// is read from /etc/resolv.conf. Falls back to 127.0.0.1:53 if neither is available. -func resolverFromConfig(cfg *RootDiscoveryConfig) string { +// resolverFromConfig returns the upstream DNS resolver address to use: +// cfg.Resolver when set, otherwise the system resolver configuration. There +// is deliberately no hardcoded address fallback. +func resolverFromConfig(cfg *RootDiscoveryConfig) (string, error) { if cfg != nil && cfg.Resolver != "" { - return cfg.Resolver + return cfg.Resolver, nil } return systemResolver() } -// systemResolver returns the first nameserver from the system DNS configuration. -// This is Unix-only: it reads /etc/resolv.conf, which does not exist on Windows. -// On Windows (or any system without /etc/resolv.conf) the fallback 127.0.0.1:53 applies. -func systemResolver() string { +// systemResolver returns the first nameserver from the system DNS +// configuration. This is Unix-only: it reads /etc/resolv.conf, which does not +// exist on Windows; there the caller falls back to the builtin root hints. +func systemResolver() (string, error) { cc, err := dns.ClientConfigFromFile("/etc/resolv.conf") - if err != nil || len(cc.Servers) == 0 { - return "127.0.0.1:53" + if err != nil { + return "", fmt.Errorf("read system resolver config: %w", err) } - return net.JoinHostPort(cc.Servers[0], cc.Port) + if len(cc.Servers) == 0 { + return "", errors.New("no nameservers found in system resolver config") + } + return net.JoinHostPort(cc.Servers[0], cc.Port), nil } -func queryResolver(ctx context.Context, resolverAddr, name string, qtype uint16) (*dns.Msg, error) { - c := &dns.Client{ - Net: "udp", - ReadTimeout: 5 * time.Second, - WriteTimeout: 5 * time.Second, +// queryUpstream asks the upstream resolver with recursion desired — the only +// RD=1 path in the program — using the same retry/TCP-fallback machinery as +// traversal queries. +func queryUpstream(ctx context.Context, cfg *RootDiscoveryConfig, resolverAddr, name string, qtype uint16) (*dns.Msg, error) { + var qcfg *QueryConfig + var exchange ExchangeFunc + if cfg != nil { + qcfg = cfg.Query + exchange = cfg.Exchange } - if deadline, ok := ctx.Deadline(); ok { - c.ReadTimeout = time.Until(deadline) - c.WriteTimeout = time.Until(deadline) + qcfg = qcfg.withDefaults() + if exchange == nil { + exchange = realExchange } m := new(dns.Msg) m.SetQuestion(dns.Fqdn(name), qtype) m.RecursionDesired = true - - r, _, err := c.ExchangeContext(ctx, m, resolverAddr) - if err != nil { - return nil, fmt.Errorf("resolver exchange %s %s: %w", name, QNameType(qtype), err) + if qcfg.UDPSize > MinEDNS0UDPSize() { + m.SetEdns0(uint16(qcfg.UDPSize), false) } - return r, nil + + return exchangeWithRetry(ctx, exchange, resolverAddr, m, qcfg) } func extractNSRecords(rrs []dns.RR) []string { @@ -259,6 +326,23 @@ func extractIPsFromAnswer(rrs []dns.RR, qtype uint16) []net.IP { return ips } -func normalizeServerName(name string) string { - return strings.TrimSuffix(strings.ToLower(name), ".") +// additionalIPs returns glue addresses for name from the additional section. +func additionalIPs(msg *dns.Msg, name string, qtype uint16) []net.IP { + var ips []net.IP + for _, rr := range msg.Extra { + if !strings.EqualFold(rr.Header().Name, name) { + continue + } + switch v := rr.(type) { + case *dns.A: + if qtype == dns.TypeA { + ips = append(ips, v.A) + } + case *dns.AAAA: + if qtype == dns.TypeAAAA { + ips = append(ips, v.AAAA) + } + } + } + return ips } diff --git a/internal/dns/roots_test.go b/internal/dns/roots_test.go index bc93467..b504087 100644 --- a/internal/dns/roots_test.go +++ b/internal/dns/roots_test.go @@ -2,9 +2,10 @@ package dns import ( "context" + "errors" "net" + "strings" "testing" - "time" "github.com/miekg/dns" ) @@ -35,9 +36,8 @@ func TestRootServerAllIPs(t *testing.T) { t.Run("no addresses", func(t *testing.T) { empty := RootServer{Name: "empty.root-servers.net."} - ips := empty.AllIPs(false) - if len(ips) != 0 { - t.Errorf("expected 0 IPs, got %d", len(ips)) + if len(empty.AllIPs(false)) != 0 { + t.Error("expected 0 IPs") } }) } @@ -55,31 +55,317 @@ func TestDefaultRootDiscoveryConfig(t *testing.T) { } } -func TestNormalizeServerName(t *testing.T) { - tests := []struct { - input string - want string - }{ - {"a.root-servers.net.", "a.root-servers.net"}, - {"A.ROOT-SERVERS.NET.", "a.root-servers.net"}, - {"b.root-servers.net", "b.root-servers.net"}, - {"root-servers.net.", "root-servers.net"}, +// mockUpstream builds an ExchangeFunc that answers root discovery queries. +// glue controls whether A records are put in the additional section of the +// NS response. +func mockUpstream(t *testing.T, roots map[string]string, glue bool) ExchangeFunc { + t.Helper() + return func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if !msg.RecursionDesired { + t.Error("root discovery must query the upstream resolver with RD=1") + } + q := msg.Question[0] + m := new(dns.Msg) + m.SetReply(msg) + switch { + case q.Name == "." && q.Qtype == dns.TypeNS: + for name, ip := range roots { + m.Answer = append(m.Answer, &dns.NS{ + Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: name, + }) + if glue { + m.Extra = append(m.Extra, &dns.A{ + Hdr: dns.RR_Header{Name: name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), + }) + } + } + case q.Qtype == dns.TypeA: + if ip, ok := roots[q.Name]; ok { + m.Answer = append(m.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), + }) + } + } + return m, nil + } +} + +func TestDiscoverRootsIPLiteral(t *testing.T) { + cfg := &RootDiscoveryConfig{ + Server: "192.203.230.10", + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + t.Error("IP literal root must not trigger any lookup") + return nil, errors.New("no network") + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 { + t.Fatalf("expected 1 root, got %d", len(servers)) + } + if servers[0].Name != "192.203.230.10" { + t.Errorf("Name = %q, want the IP literal", servers[0].Name) + } + if len(servers[0].IPv4) != 1 || !servers[0].IPv4[0].Equal(net.ParseIP("192.203.230.10")) { + t.Errorf("IPv4 = %v, want [192.203.230.10]", servers[0].IPv4) + } +} + +func TestDiscoverRootsIPv6Literal(t *testing.T) { + cfg := &RootDiscoveryConfig{Server: "2001:503:ba3e::2:30"} + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 || len(servers[0].IPv6) != 1 { + t.Fatalf("expected 1 root with 1 IPv6 address, got %+v", servers) + } +} + +func TestDiscoverRootsHostnameOverride(t *testing.T) { + roots := map[string]string{"e.root-servers.net.": "192.203.230.10"} + cfg := &RootDiscoveryConfig{ + Server: "e.root-servers.net", + Resolver: "192.0.2.53:53", + Exchange: mockUpstream(t, roots, false), + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 { + t.Fatalf("expected 1 root, got %d", len(servers)) + } + if servers[0].Name != "e.root-servers.net" { + t.Errorf("Name = %q", servers[0].Name) + } + if len(servers[0].IPv4) != 1 || !servers[0].IPv4[0].Equal(net.ParseIP("192.203.230.10")) { + t.Errorf("IPv4 = %v", servers[0].IPv4) + } +} + +func TestDiscoverRootsHostnameOverrideUnresolvable(t *testing.T) { + cfg := &RootDiscoveryConfig{ + Server: "nonexistent.root-servers.net", + Resolver: "192.0.2.53:53", + Exchange: mockUpstream(t, map[string]string{}, false), + } + if _, err := DiscoverRoots(context.Background(), cfg); err == nil { + t.Fatal("expected error for unresolvable root override") + } +} + +func TestDiscoverRootsSingleFromGlue(t *testing.T) { + roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} + var wireQueries int + inner := mockUpstream(t, roots, true) + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + wireQueries++ + return inner(ctx, server, msg, useTCP) + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 { + t.Fatalf("expected exactly one root, got %d", len(servers)) + } + if len(servers[0].IPv4) == 0 { + t.Error("expected glue A address") + } + if wireQueries != 1 { + t.Errorf("glue should satisfy discovery in one query, got %d", wireQueries) + } +} + +func TestDiscoverRootsSingleWithoutGlue(t *testing.T) { + roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Exchange: mockUpstream(t, roots, false), + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 || len(servers[0].IPv4) != 1 { + t.Fatalf("expected one root resolved via A lookup, got %+v", servers) + } +} + +func TestDiscoverAllRootsPerRootSet(t *testing.T) { + roots := map[string]string{ + "a.mock-roots.test.": "192.0.2.1", + "b.mock-roots.test.": "192.0.2.2", + "c.mock-roots.test.": "192.0.2.3", + } + cfg := &RootDiscoveryConfig{ + AllRoots: true, + Resolver: "192.0.2.53:53", + Exchange: mockUpstream(t, roots, true), + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 3 { + t.Fatalf("expected 3 roots, got %d", len(servers)) + } + seen := map[string]bool{} + for _, rs := range servers { + seen[rs.Name] = true + want, ok := roots[rs.Name] + if !ok { + t.Errorf("unexpected root %q", rs.Name) + continue + } + if len(rs.IPv4) != 1 || !rs.IPv4[0].Equal(net.ParseIP(want)) { + t.Errorf("root %s IPv4 = %v, want [%s]", rs.Name, rs.IPv4, want) + } + } + if len(seen) != 3 { + t.Errorf("roots not distinct: %v", seen) + } +} + +func TestDiscoverAllRootsSkipsUnresolvable(t *testing.T) { + // b has neither glue nor an A record: it must be skipped, like + // find_all_roots in traverser.rb. + roots := map[string]string{"a.mock-roots.test.": "192.0.2.1"} + exchange := func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + q := msg.Question[0] + m := new(dns.Msg) + m.SetReply(msg) + switch { + case q.Name == "." && q.Qtype == dns.TypeNS: + for _, name := range []string{"a.mock-roots.test.", "b.mock-roots.test."} { + m.Answer = append(m.Answer, &dns.NS{ + Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: name, + }) + } + case q.Qtype == dns.TypeA: + if ip, ok := roots[q.Name]; ok { + m.Answer = append(m.Answer, &dns.A{ + Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip), + }) + } + } + return m, nil + } + cfg := &RootDiscoveryConfig{ + AllRoots: true, + Resolver: "192.0.2.53:53", + Query: &QueryConfig{Retries: 1}, + Exchange: exchange, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if len(servers) != 1 || servers[0].Name != "a.mock-roots.test." { + t.Fatalf("expected only the resolvable root, got %+v", servers) + } +} + +func TestDiscoverRootsFallsBackToHints(t *testing.T) { + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Query: &QueryConfig{Retries: 1}, + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + return nil, errors.New("upstream unreachable") + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("expected hints fallback, got error: %v", err) + } + if len(servers) != 1 { + t.Fatalf("single-root mode must fall back to one hint, got %d", len(servers)) + } + if !strings.HasSuffix(servers[0].Name, ".root-servers.net.") { + t.Errorf("expected an IANA hint, got %q", servers[0].Name) } - for _, tt := range tests { - t.Run(tt.input, func(t *testing.T) { - got := normalizeServerName(tt.input) - if got != tt.want { - t.Errorf("normalizeServerName(%q) = %q, want %q", tt.input, got, tt.want) + cfg.AllRoots = true + servers, err = DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("expected hints fallback, got error: %v", err) + } + if len(servers) != len(RootHints) { + t.Errorf("all-roots fallback should return all %d hints, got %d", len(RootHints), len(servers)) + } +} + +func TestDiscoverRootsRetriesUpstream(t *testing.T) { + roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} + inner := mockUpstream(t, roots, true) + var calls int + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Query: &QueryConfig{Retries: 2, RetryDelay: 1}, + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + calls++ + if calls == 1 { + return nil, errors.New("transient failure") } - }) + return inner(ctx, server, msg, useTCP) + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots should retry: %v", err) + } + if len(servers) != 1 { + t.Fatalf("expected 1 root after retry, got %d", len(servers)) + } + if calls != 2 { + t.Errorf("expected 2 attempts, got %d", calls) + } +} + +func TestDiscoverRootsTruncationTCPFallback(t *testing.T) { + roots := map[string]string{"a.root-servers.net.": "198.41.0.4"} + inner := mockUpstream(t, roots, true) + var sawTCP bool + cfg := &RootDiscoveryConfig{ + Resolver: "192.0.2.53:53", + Query: &QueryConfig{Retries: 1, AllowTCP: true}, + Exchange: func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { + if !useTCP { + m := new(dns.Msg) + m.SetReply(msg) + m.Truncated = true + return m, nil + } + sawTCP = true + return inner(ctx, server, msg, useTCP) + }, + } + servers, err := DiscoverRoots(context.Background(), cfg) + if err != nil { + t.Fatalf("DiscoverRoots: %v", err) + } + if !sawTCP { + t.Error("expected TCP fallback on truncated upstream response") + } + if len(servers) != 1 || len(servers[0].IPv4) == 0 { + t.Fatalf("expected root from TCP response, got %+v", servers) } } func TestExtractNSRecords(t *testing.T) { t.Run("empty", func(t *testing.T) { - names := extractNSRecords(nil) - if len(names) != 0 { + if names := extractNSRecords(nil); len(names) != 0 { t.Errorf("expected 0 names, got %d", len(names)) } }) @@ -103,18 +389,24 @@ func TestExtractNSRecords(t *testing.T) { &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."}, &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."}, } - names := extractNSRecords(rrs) - if len(names) != 1 { + if names := extractNSRecords(rrs); len(names) != 1 { t.Errorf("expected 1 deduped name, got %d", len(names)) } + }) + t.Run("non-NS ignored", func(t *testing.T) { + rrs := []dns.RR{ + &dns.A{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")}, + } + if names := extractNSNames(rrs); len(names) != 0 { + t.Errorf("expected 0 names, got %d", len(names)) + } }) } func TestExtractIPsFromAnswer(t *testing.T) { t.Run("empty", func(t *testing.T) { - ips := extractIPsFromAnswer(nil, dns.TypeA) - if len(ips) != 0 { + if ips := extractIPsFromAnswer(nil, dns.TypeA); len(ips) != 0 { t.Errorf("expected 0 IPs, got %d", len(ips)) } }) @@ -124,347 +416,41 @@ func TestExtractIPsFromAnswer(t *testing.T) { &dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")}, &dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("199.9.14.201")}, } - ips := extractIPsFromAnswer(rrs, dns.TypeA) - if len(ips) != 2 { + if ips := extractIPsFromAnswer(rrs, dns.TypeA); len(ips) != 2 { t.Fatalf("expected 2 IPs, got %d", len(ips)) } }) - t.Run("AAAA records", func(t *testing.T) { - rrs := []dns.RR{ - &dns.AAAA{Hdr: dns.RR_Header{Rrtype: dns.TypeAAAA}, AAAA: net.ParseIP("2001:503:ba3e::2:30")}, - } - ips := extractIPsFromAnswer(rrs, dns.TypeAAAA) - if len(ips) != 1 { - t.Fatalf("expected 1 IP, got %d", len(ips)) - } - }) - t.Run("filter by type", func(t *testing.T) { rrs := []dns.RR{ &dns.A{Hdr: dns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")}, &dns.AAAA{Hdr: dns.RR_Header{Rrtype: dns.TypeAAAA}, AAAA: net.ParseIP("2001:503:ba3e::2:30")}, } - ips := extractIPsFromAnswer(rrs, dns.TypeA) - if len(ips) != 1 { + if ips := extractIPsFromAnswer(rrs, dns.TypeA); len(ips) != 1 { t.Fatalf("expected 1 A IP, got %d", len(ips)) } + if ips := extractIPsFromAnswer(rrs, dns.TypeAAAA); len(ips) != 1 { + t.Fatalf("expected 1 AAAA IP, got %d", len(ips)) + } }) } -func TestDiscoverRootsOverride(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - cfg := &RootDiscoveryConfig{ - Server: "a.root-servers.net", - IncludeAAAA: false, - } - - servers, err := DiscoverRoots(ctx, cfg) - if err != nil { - t.Logf("skipping (no resolver available): %v", err) - t.Skip() - } - if len(servers) == 0 { - t.Fatal("expected at least one root server") - } - if len(servers[0].IPv4) == 0 { - t.Error("expected IPv4 addresses for a.root-servers.net") - } -} - -func TestDiscoverRootsSingle(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - servers, err := DiscoverRoots(ctx, nil) - if err != nil { - t.Logf("skipping (no resolver available): %v", err) - t.Skip() - } - if len(servers) == 0 { - t.Fatal("expected at least one root server") - } - if servers[0].Name == "" { - t.Error("root server name should not be empty") - } -} - -func TestDiscoverRootsNilConfig(t *testing.T) { - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - servers, err := DiscoverRoots(ctx, nil) - if err != nil { - t.Logf("skipping (no resolver available): %v", err) - t.Skip() - } - if len(servers) == 0 { - t.Fatal("expected at least one root server with nil config") - } -} - -func TestBuildNSResponse(t *testing.T) { +func TestAdditionalIPs(t *testing.T) { msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "a.root-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "b.root-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, Ns: "c.root-servers.net."}, + msg.Extra = append(msg.Extra, + &dns.A{Hdr: dns.RR_Header{Name: "A.ROOT-SERVERS.NET.", Rrtype: dns.TypeA, Class: dns.ClassINET}, A: net.ParseIP("198.41.0.4")}, + &dns.A{Hdr: dns.RR_Header{Name: "b.root-servers.net.", Rrtype: dns.TypeA, Class: dns.ClassINET}, A: net.ParseIP("170.247.170.2")}, + &dns.AAAA{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dns.TypeAAAA, Class: dns.ClassINET}, AAAA: net.ParseIP("2001:503:ba3e::2:30")}, ) - names := extractNSRecords(msg.Answer) - if len(names) != 3 { - t.Fatalf("expected 3 NS records, got %d", len(names)) + ips := additionalIPs(msg, "a.root-servers.net.", dns.TypeA) + if len(ips) != 1 || !ips[0].Equal(net.ParseIP("198.41.0.4")) { + t.Errorf("case-insensitive glue match failed: %v", ips) } - - for _, name := range names { - if len(name) == 0 || name[len(name)-1] != '.' { - t.Errorf("expected FQDN, got %q", name) - } + if ips := additionalIPs(msg, "a.root-servers.net.", dns.TypeAAAA); len(ips) != 1 { + t.Errorf("expected 1 AAAA glue, got %v", ips) + } + if ips := additionalIPs(msg, "c.root-servers.net.", dns.TypeA); len(ips) != 0 { + t.Errorf("expected no glue for c, got %v", ips) } } - -func TestExtractNSNames(t *testing.T) { -rrs := []dns.RR{ -&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "a.root-servers.net."}, -&dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, Ns: "b.root-servers.net."}, -} -names := extractNSNames(rrs) -if len(names) != 2 { -t.Fatalf("extractNSNames: expected 2 names, got %d", len(names)) -} -} - -func TestExtractNSNamesEmpty(t *testing.T) { -names := extractNSNames(nil) -if len(names) != 0 { -t.Errorf("extractNSNames(nil): expected 0 names, got %d", len(names)) -} -} - -func TestExtractNSNamesNonNS(t *testing.T) { -rrs := []dns.RR{ -&dns.A{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("198.41.0.4")}, -} -names := extractNSNames(rrs) -if len(names) != 0 { -t.Errorf("extractNSNames with A records: expected 0 names, got %d", len(names)) -} -} - -func TestDiscoverRootsAllRoots(t *testing.T) { -ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) -defer cancel() - -cfg := &RootDiscoveryConfig{ -AllRoots: true, -IncludeAAAA: false, -} - -servers, err := DiscoverRoots(ctx, cfg) -if err != nil { -t.Logf("skipping (no resolver available): %v", err) -t.Skip() -} -if len(servers) == 0 { -t.Fatal("expected root servers with AllRoots=true") -} -t.Logf("discovered %d root servers", len(servers)) -} - -func TestDiscoverRootsAllRootsIncludeAAAA(t *testing.T) { -ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) -defer cancel() - -cfg := &RootDiscoveryConfig{ -AllRoots: true, -IncludeAAAA: true, -} - -servers, err := DiscoverRoots(ctx, cfg) -if err != nil { -t.Logf("skipping (no resolver available): %v", err) -t.Skip() -} -if len(servers) == 0 { -t.Fatal("expected root servers") -} -} - -// startMockDNSServer starts a UDP DNS server on a random port that serves -// pre-configured responses. It returns the server address and a stop function. -func startMockDNSServer(t *testing.T, handlerFn dns.HandlerFunc) string { -t.Helper() - -mux := dns.NewServeMux() -mux.HandleFunc(".", handlerFn) - -srv := &dns.Server{ -Addr: "127.0.0.1:0", -Net: "udp", -Handler: mux, -} - -started := make(chan struct{}) -srv.NotifyStartedFunc = func() { close(started) } - -go func() { -if err := srv.ListenAndServe(); err != nil && t.Failed() { -return -} -}() - -select { -case <-started: -case <-time.After(2 * time.Second): -t.Fatal("mock DNS server did not start in time") -} - -// Retrieve the actual bound address from the server's PacketConn. -addr := srv.PacketConn.LocalAddr().String() -t.Cleanup(func() { _ = srv.Shutdown() }) -return addr -} - -func TestQueryResolverSuccess(t *testing.T) { - addr := startMockDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - m := new(dns.Msg) - m.SetReply(r) - m.Answer = append(m.Answer, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "a.root-servers.net."}, - ) - _ = w.WriteMsg(m) - }) - - // queryResolver uses 5ns timeout without deadline; provide one - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - msg, err := queryResolver(ctx, addr, ".", dns.TypeNS) - if err != nil { - t.Fatalf("queryResolver: %v", err) - } - names := extractNSRecords(msg.Answer) - if len(names) == 0 { - t.Fatal("expected NS records in answer") - } -} - -func TestQueryResolverError(t *testing.T) { -ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) -defer cancel() -// Use an address nothing is listening on -_, err := queryResolver(ctx, "127.0.0.1:19999", ".", dns.TypeNS) -if err == nil { -t.Fatal("expected error for unreachable resolver") -} -} - -func TestDiscoverSingleRootWithMock(t *testing.T) { - addr := startMockDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - m := new(dns.Msg) - m.SetReply(r) - q := r.Question[0] - switch q.Qtype { - case dns.TypeNS: - m.Answer = append(m.Answer, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "mock.root-servers.test."}, - ) - case dns.TypeA: - m.Answer = append(m.Answer, - &dns.A{Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("127.0.0.1")}, - ) - } - _ = w.WriteMsg(m) - }) - - ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second) - defer cancel() - msg, err := queryResolver(ctx, addr, ".", dns.TypeNS) - if err != nil { - t.Fatalf("queryResolver: %v", err) - } - names := extractNSRecords(msg.Answer) - if len(names) == 0 { - names = extractNSNames(msg.Ns) - } - if len(names) == 0 { - t.Skip("mock NS query returned no NS records") - } - t.Logf("found %d root NS names from mock: %v", len(names), names) -} - -func TestDiscoverAllRootsWithMockServer(t *testing.T) { - addr := startMockDNSServer(t, func(w dns.ResponseWriter, r *dns.Msg) { - m := new(dns.Msg) - m.SetReply(r) - q := r.Question[0] - switch q.Qtype { - case dns.TypeNS: - m.Answer = append(m.Answer, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "a.mock-roots.test."}, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, Ns: "b.mock-roots.test."}, - ) - case dns.TypeA: - m.Answer = append(m.Answer, - &dns.A{Hdr: dns.RR_Header{Name: q.Name, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, A: net.ParseIP("127.0.0.1")}, - ) - } - _ = w.WriteMsg(m) - }) - - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - msg, err := queryResolver(ctx, addr, ".", dns.TypeNS) - if err != nil { - t.Fatalf("queryResolver: %v", err) - } - names := extractNSRecords(msg.Answer) - if len(names) < 2 { - t.Fatalf("expected 2 NS names, got %d", len(names)) - } - - // Also cover AAAA path - aaaaMsg, err := queryResolver(ctx, addr, "a.mock-roots.test.", dns.TypeAAAA) - if err != nil { - t.Logf("AAAA query error (acceptable): %v", err) - } else { - t.Logf("AAAA query returned %d answers", len(aaaaMsg.Answer)) - } -} - -func TestMinTTLFromMsgWithExtraRecords(t *testing.T) { -msg := new(dns.Msg) -msg.Answer = append(msg.Answer, &dns.A{ -Hdr: dns.RR_Header{Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) -// Extra record (non-OPT) with smaller TTL -msg.Extra = append(msg.Extra, &dns.NS{ -Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Ttl: 60}, -Ns: "a.root-servers.net.", -}) - -ttl := minTTLFromMsg(msg) -if ttl != 60*time.Second { -t.Errorf("minTTLFromMsg = %v, want 60s", ttl) -} -} - -func TestMinTTLFromMsgOPTIgnored(t *testing.T) { -msg := new(dns.Msg) -msg.Answer = append(msg.Answer, &dns.A{ -Hdr: dns.RR_Header{Ttl: 300}, -A: net.ParseIP("1.2.3.4"), -}) -// OPT record should be ignored -msg.Extra = append(msg.Extra, &dns.OPT{ -Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeOPT}, -}) - -ttl := minTTLFromMsg(msg) -if ttl != 300*time.Second { -t.Errorf("minTTLFromMsg with OPT = %v, want 300s", ttl) -} -} diff --git a/internal/dns/types.go b/internal/dns/types.go index 3a2353d..b5aad48 100644 --- a/internal/dns/types.go +++ b/internal/dns/types.go @@ -5,35 +5,12 @@ import ( ) const ( - TypeA uint16 = dns.TypeA - TypeAAAA uint16 = dns.TypeAAAA - TypeNS uint16 = dns.TypeNS - TypeCNAME uint16 = dns.TypeCNAME - TypeSOA uint16 = dns.TypeSOA - TypeMX uint16 = dns.TypeMX - TypeTXT uint16 = dns.TypeTXT - TypeSRV uint16 = dns.TypeSRV - TypePTR uint16 = dns.TypePTR - TypeANY uint16 = dns.TypeANY + TypeA uint16 = dns.TypeA + TypeAAAA uint16 = dns.TypeAAAA + TypeNS uint16 = dns.TypeNS ) -var QNameTypes = map[uint16]string{ - TypeA: "A", - TypeAAAA: "AAAA", - TypeNS: "NS", - TypeCNAME: "CNAME", - TypeSOA: "SOA", - TypeMX: "MX", - TypeTXT: "TXT", - TypeSRV: "SRV", - TypePTR: "PTR", - TypeANY: "ANY", -} - func QNameType(qtype uint16) string { - if name, ok := QNameTypes[qtype]; ok { - return name - } return dns.TypeToString[qtype] } diff --git a/internal/dns/types_test.go b/internal/dns/types_test.go index bd75ab7..1a81ba7 100644 --- a/internal/dns/types_test.go +++ b/internal/dns/types_test.go @@ -13,13 +13,6 @@ func TestQNameType(t *testing.T) { {"A record", TypeA, "A"}, {"AAAA record", TypeAAAA, "AAAA"}, {"NS record", TypeNS, "NS"}, - {"CNAME record", TypeCNAME, "CNAME"}, - {"SOA record", TypeSOA, "SOA"}, - {"MX record", TypeMX, "MX"}, - {"TXT record", TypeTXT, "TXT"}, - {"SRV record", TypeSRV, "SRV"}, - {"PTR record", TypePTR, "PTR"}, - {"ANY record", TypeANY, "ANY"}, {"unknown type", uint16(9999), ""}, } @@ -33,36 +26,6 @@ func TestQNameType(t *testing.T) { } } -func TestConstantsMatchMiekg(t *testing.T) { - tests := []struct { - name string - local uint16 - }{ - {"TypeA", TypeA}, - {"TypeAAAA", TypeAAAA}, - {"TypeNS", TypeNS}, - {"TypeCNAME", TypeCNAME}, - {"TypeSOA", TypeSOA}, - {"TypeMX", TypeMX}, - {"TypeTXT", TypeTXT}, - {"TypeSRV", TypeSRV}, - {"TypePTR", TypePTR}, - {"TypeANY", TypeANY}, - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - mapped, ok := QNameTypes[tt.local] - if !ok { - t.Errorf("QNameTypes missing entry for %s (%d)", tt.name, tt.local) - } - if mapped != QNameType(tt.local) { - t.Errorf("QNameType(%d) = %q, QNameTypes[%d] = %q", tt.local, QNameType(tt.local), tt.local, mapped) - } - }) - } -} - func TestDefaultEDNS0UDPSize(t *testing.T) { if got := DefaultEDNS0UDPSize(); got != 2048 { t.Errorf("DefaultEDNS0UDPSize() = %d, want 2048", got) diff --git a/internal/integration/integration_test.go b/internal/integration/integration_test.go index 949a3f4..f3315e7 100644 --- a/internal/integration/integration_test.go +++ b/internal/integration/integration_test.go @@ -1,9 +1,11 @@ // Package integration provides end-to-end tests for ExploreDNS using a mock -// DNS server that allows deterministic, network-independent testing. +// DNS exchange that allows deterministic, network-independent testing of the +// full engine through its exported API. package integration import ( "context" + "math" "net" "testing" "time" @@ -13,523 +15,291 @@ import ( "github.com/miekg/dns" ) -// mockZone represents a simple in-memory DNS zone for testing. -type mockZone struct { - // map[name][qtype] → []RR - records map[string]map[uint16][]dns.RR +// mockNet maps (server IP, qname, qtype) to a canned response, mirroring how +// distinct nameservers answer differently for the same question. +type mockNet struct { + responses map[string]*dns.Msg } -func newMockZone() *mockZone { - return &mockZone{records: make(map[string]map[uint16][]dns.RR)} +func newMockNet() *mockNet { + return &mockNet{responses: make(map[string]*dns.Msg)} } -func (z *mockZone) addA(name, ip string) { - fqdn := dns.Fqdn(name) - if z.records[fqdn] == nil { - z.records[fqdn] = make(map[uint16][]dns.RR) +func key(server, qname string, qtype uint16) string { + return server + "|" + dns.Fqdn(qname) + "|" + dns.TypeToString[qtype] +} + +func (m *mockNet) on(server, qname string, qtype uint16, msg *dns.Msg) { + m.responses[key(server, qname, qtype)] = msg +} + +func (m *mockNet) exchange(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h } - z.records[fqdn][dns.TypeA] = append(z.records[fqdn][dns.TypeA], &dns.A{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP(ip), - }) + q := msg.Question[0] + resp, ok := m.responses[key(host, q.Name, q.Qtype)] + if !ok { + // Unknown question: NXDOMAIN, like an authoritative miss. + out := new(dns.Msg) + out.SetReply(msg) + out.Rcode = dns.RcodeNameError + return out, nil + } + out := resp.Copy() + out.SetReply(msg) + out.Answer, out.Ns, out.Extra = resp.Answer, resp.Ns, resp.Extra + out.Rcode = resp.Rcode + return out, nil } -func (z *mockZone) addNS(zone, ns string) { - fqdn := dns.Fqdn(zone) - if z.records[fqdn] == nil { - z.records[fqdn] = make(map[uint16][]dns.RR) +func aRR(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip).To4(), } - z.records[fqdn][dns.TypeNS] = append(z.records[fqdn][dns.TypeNS], &dns.NS{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, - Ns: dns.Fqdn(ns), - }) } -func (z *mockZone) addCNAME(name, target string) { - fqdn := dns.Fqdn(name) - if z.records[fqdn] == nil { - z.records[fqdn] = make(map[uint16][]dns.RR) +func nsRR(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: dns.Fqdn(target), } - z.records[fqdn][dns.TypeCNAME] = append(z.records[fqdn][dns.TypeCNAME], &dns.CNAME{ - Hdr: dns.RR_Header{Name: fqdn, Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 300}, +} + +func cnameRR(name, target string) dns.RR { + return &dns.CNAME{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 300}, Target: dns.Fqdn(target), - }) -} - -// makeExchange creates a mock ExchangeFunc that serves responses from the zone. -// It simulates referral behavior: if a name matches a zone NS record, it returns -// a referral with glue. If it matches an A record, it returns the answer. -func (z *mockZone) makeExchange() dnsinternal.ExchangeFunc { - return func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - if len(msg.Question) == 0 { - return nil, nil - } - q := msg.Question[0] - - resp := new(dns.Msg) - resp.SetReply(msg) - resp.Authoritative = true - - // Direct answer - if rrs, ok := z.records[q.Name]; ok { - if answers, ok := rrs[q.Qtype]; ok { - resp.Answer = append(resp.Answer, answers...) - return resp, nil - } - // CNAME chain — return CNAME + answer for target if qtype != CNAME - if cnameRRs, ok := rrs[dns.TypeCNAME]; ok && q.Qtype != dns.TypeCNAME { - resp.Answer = append(resp.Answer, cnameRRs...) - return resp, nil - } - } - - // Check for zone delegation: look for NS records covering any suffix of qname - labels := dns.SplitDomainName(q.Name) - for i := 0; i < len(labels); i++ { - zone := dns.Fqdn(joinLabels(labels[i:])) - if nsRRs, ok := z.records[zone][dns.TypeNS]; ok && zone != q.Name { - // Return referral - resp.Authoritative = false - resp.Ns = append(resp.Ns, nsRRs...) - for _, ns := range nsRRs { - nsName := ns.(*dns.NS).Ns - if aRRs, ok := z.records[nsName][dns.TypeA]; ok { - resp.Extra = append(resp.Extra, aRRs...) - } - } - return resp, nil - } - } - - // NXDOMAIN - resp.Authoritative = true - resp.Rcode = dns.RcodeNameError - return resp, nil } } -func joinLabels(labels []string) string { - result := "" - for i, l := range labels { - if i > 0 { - result += "." - } - result += l - } - return result +func answerMsg(rrs ...dns.RR) *dns.Msg { + m := new(dns.Msg) + m.Answer = rrs + return m } -// setupTestZone creates a mock zone with a typical referral hierarchy: -// -// root → com (referral) → example.com (referral) → www.example.com (A) -func setupTestZone() *mockZone { - z := newMockZone() - - // Root server glue - z.addA("a.root-servers.test", "198.41.0.4") - - // com TLD referral from root - z.addNS("com", "a.gtld-servers.test") - z.addA("a.gtld-servers.test", "192.5.6.30") - - // example.com NS referral from com TLD - z.addNS("example.com", "ns1.example.com") - z.addA("ns1.example.com", "1.2.3.4") - - // Actual A records - z.addA("example.com", "93.184.216.34") - z.addA("www.example.com", "93.184.216.34") - - return z +func referralMsg(nsRRs []dns.RR, glue ...dns.RR) *dns.Msg { + m := new(dns.Msg) + m.Ns = nsRRs + m.Extra = glue + return m } -// TestIntegrationSimpleAQuery verifies end-to-end traversal with mock DNS -// that returns A record answers without network dependency. -func TestIntegrationSimpleAQuery(t *testing.T) { - z := setupTestZone() - - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, +func newTraverser(maxDepth int) *traverse.Traverser { + return traverse.NewTraverser(&traverse.TraverserConfig{ + MaxDepth: maxDepth, QueryType: dnsinternal.TypeA, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + QueryConfig: &dnsinternal.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, }) - tr.SetExchange(z.makeExchange()) +} +func run(t *testing.T, tr *traverse.Traverser, m *mockNet, qname string) *traverse.Referral { + t.Helper() + tr.SetExchange(m.exchange) ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() - - results, err := tr.Traverse(ctx, "example.com") + root, err := tr.Run(ctx, qname) if err != nil { - t.Fatalf("Traverse: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results from traversal") + t.Fatalf("Run(%q): %v", qname, err) } + assertProbabilityInvariant(t, root) + return root +} - var foundAnswer bool - for _, r := range results { - if r.Response != nil && r.Response.Type == traverse.RespAnswer { - foundAnswer = true - if r.Response.Decoded != nil && len(r.Response.Decoded.Answers) > 0 { - for _, rr := range r.Response.Decoded.Answers { - if a, ok := rr.(*dns.A); ok { - t.Logf("Found A record: %v", a.A) - } - } - } - } +// assertProbabilityInvariant checks the engine ground rule: aggregated leaf +// probabilities at the root sum to 1.0. +func assertProbabilityInvariant(t *testing.T, root *traverse.Referral) { + t.Helper() + sum := 0.0 + for _, leaf := range root.StatsList() { + sum += leaf.Prob } - if !foundAnswer { - t.Errorf("expected to find an answer response; got types: %v", responseTypes(results)) + if math.Abs(sum-1.0) > 1e-9 { + t.Errorf("leaf probabilities sum to %v, want 1.0", sum) } } -// TestIntegrationReferralChain verifies multi-hop referral traversal: -// root → com → example.com, with glue records at each step. +func statuses(root *traverse.Referral) map[traverse.Status]float64 { + out := make(map[traverse.Status]float64) + for _, leaf := range root.StatsList() { + out[leaf.Response.Status] += leaf.Prob + } + return out +} + +// TestIntegrationReferralChain verifies the classic delegation walk: +// root → com → example.com with per-server responses. func TestIntegrationReferralChain(t *testing.T) { - z := newMockZone() + m := newMockNet() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.test")}, + aRR("a.gtld-servers.test", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.2.3.4"), + )) + m.on("1.2.3.4", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "93.184.216.34"))) - // Root delegates to com - z.addNS("com", "a.gtld-servers.test") - z.addA("a.gtld-servers.test", "192.5.6.30") - - // TLD delegates to example.com - z.addNS("example.com", "ns1.example.com") - z.addA("ns1.example.com", "1.2.3.4") - - // Authoritative answer - z.addA("example.com", "93.184.216.34") - - exchange := z.makeExchange() - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(exchange) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) - } - - // Count referrals and answers - var referrals, answers int - for _, r := range results { - if r.Response == nil { - continue - } - switch r.Response.Type { - case traverse.RespReferral: - referrals++ - case traverse.RespAnswer: - answers++ - } - } - t.Logf("referrals=%d answers=%d total=%d", referrals, answers, len(results)) - if answers == 0 { - t.Errorf("expected at least one answer; types: %v", responseTypes(results)) + root := run(t, newTraverser(10), m, "www.example.com") + got := statuses(root) + if math.Abs(got[traverse.StatusAnswered]-1.0) > 1e-9 { + t.Errorf("statuses = %v, want 100%% answered", got) } } -// TestIntegrationCNAMEResolution verifies that CNAME chains are followed correctly. -func TestIntegrationCNAMEResolution(t *testing.T) { - z := newMockZone() +// TestIntegrationCNAMERestart verifies that an out-of-zone CNAME target +// restarts the traversal from the branch cache. +func TestIntegrationCNAMERestart(t *testing.T) { + m := newMockNet() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.2.3.4"), + )) + m.on("1.2.3.4", "www.example.com", dns.TypeA, answerMsg(cnameRR("www.example.com", "cdn.example.net"))) + m.on("198.41.0.4", "cdn.example.net", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.net", "ns1.example.net")}, + aRR("ns1.example.net", "5.6.7.8"), + )) + m.on("5.6.7.8", "cdn.example.net", dns.TypeA, answerMsg(aRR("cdn.example.net", "93.184.216.35"))) - // www.example.com → CNAME → example.com → A record - z.addCNAME("www.example.com", "example.com") - z.addA("example.com", "93.184.216.34") - - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(z.makeExchange()) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("Traverse CNAME: %v", err) + root := run(t, newTraverser(10), m, "www.example.com") + got := statuses(root) + if math.Abs(got[traverse.StatusAnswered]-1.0) > 1e-9 { + t.Errorf("statuses = %v, want 100%% answered via restart", got) } - if len(results) == 0 { - t.Fatal("expected results") - } - - var foundCNAME, foundAnswer bool - for _, r := range results { - if r.Response == nil { - continue - } - if r.Response.Type == traverse.RespCNAMEFollow { - foundCNAME = true - } - if r.Response.Type == traverse.RespAnswer { - foundAnswer = true + for _, leaf := range root.StatsList() { + if leaf.Response.Status == traverse.StatusAnswered && leaf.Response.Qname != "cdn.example.net" { + t.Errorf("answered qname = %q, want the CNAME target", leaf.Response.Qname) } } - t.Logf("CNAME traversal: foundCNAME=%v foundAnswer=%v types=%v", foundCNAME, foundAnswer, responseTypes(results)) } -// TestIntegrationNXDOMAIN verifies that NXDOMAIN responses are correctly classified. +// TestIntegrationNXDOMAIN verifies rcode errors surface as error leaves with +// the reference wording. func TestIntegrationNXDOMAIN(t *testing.T) { - z := newMockZone() - // Zone has no records for nonexistent.example.com - - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(z.makeExchange()) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "nonexistent.example.test") - if err != nil { - t.Fatalf("Traverse NXDOMAIN: %v", err) + m := newMockNet() + // mock returns NXDOMAIN for anything unmocked + root := run(t, newTraverser(10), m, "nonexistent.example.test") + got := statuses(root) + if math.Abs(got[traverse.StatusError]-1.0) > 1e-9 { + t.Errorf("statuses = %v, want 100%% error", got) } - if len(results) == 0 { - t.Fatal("expected at least one result for NXDOMAIN") - } - - var foundNXDOMAIN bool - for _, r := range results { - if r.Response != nil && r.Response.Type == traverse.RespNXDOMAIN { - foundNXDOMAIN = true - break + for _, leaf := range root.StatsList() { + if leaf.Response.DQ.ErrorMessage != "No such domain (NXDOMAIN)" { + t.Errorf("error message = %q", leaf.Response.DQ.ErrorMessage) } } - if !foundNXDOMAIN { - t.Errorf("expected NXDOMAIN result; got: %v", responseTypes(results)) - } } -// TestIntegrationSERVFAIL verifies that SERVFAIL responses are correctly handled. +// TestIntegrationSERVFAIL verifies SERVFAIL classification. func TestIntegrationSERVFAIL(t *testing.T) { - sfMsg := new(dns.Msg) - sfMsg.Rcode = dns.RcodeServerFailure + m := newMockNet() + sf := new(dns.Msg) + sf.Rcode = dns.RcodeServerFailure + m.on("198.41.0.4", "example.com", dns.TypeA, sf) - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return sfMsg.Copy(), nil - }) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse SERVFAIL: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least one result") - } - if results[0].Response.Type != traverse.RespSERVFAIL { - t.Errorf("expected SERVFAIL, got %v", results[0].Response.Type) + root := run(t, newTraverser(10), m, "example.com") + for _, leaf := range root.StatsList() { + if leaf.Response.Status != traverse.StatusError { + t.Errorf("status = %q, want error", leaf.Response.Status) + } + if leaf.Response.DQ.ErrorMessage != "Server failure (SERVFAIL)" { + t.Errorf("error message = %q", leaf.Response.DQ.ErrorMessage) + } } } -// TestIntegrationCNAMELoop verifies that CNAME loops are detected and reported. +// TestIntegrationCNAMELoop verifies cross-response CNAME loops terminate as +// cname_loop leaves. func TestIntegrationCNAMELoop(t *testing.T) { - callCount := 0 - // www.a.test → CNAME → www.b.test → CNAME → www.a.test (loop) - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - callCount++ - if len(msg.Question) == 0 { - return nil, nil - } - q := msg.Question[0] + m := newMockNet() + m.on("198.41.0.4", "www.a.test", dns.TypeA, answerMsg(cnameRR("www.a.test", "www.b.test"))) + m.on("198.41.0.4", "www.b.test", dns.TypeA, answerMsg(cnameRR("www.b.test", "www.a.test"))) - resp := new(dns.Msg) - resp.SetReply(msg) - resp.Authoritative = true - - switch q.Name { - case "www.a.test.": - resp.Answer = append(resp.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.a.test.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "www.b.test.", - }) - case "www.b.test.": - resp.Answer = append(resp.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.b.test.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "www.a.test.", - }) - default: - resp.Rcode = dns.RcodeNameError - } - return resp, nil - }) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "www.a.test") - if err != nil { - t.Fatalf("Traverse CNAME loop: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results from CNAME loop traversal") - } - - var foundLoop bool - for _, r := range results { - if r.Response != nil && r.Response.Type == traverse.RespCNAMELoop { - foundLoop = true - break - } - } - if !foundLoop { - t.Logf("types found: %v", responseTypes(results)) - // CNAME loop detection may vary based on implementation; warn rather than fail - t.Logf("CNAME loop not detected as RespCNAMELoop (may be handled differently)") + root := run(t, newTraverser(10), m, "www.a.test") + got := statuses(root) + if math.Abs(got[traverse.StatusCNAMELoop]-1.0) > 1e-9 { + t.Errorf("statuses = %v, want 100%% cname_loop", got) } } -// TestIntegrationMaxDepthExceeded verifies that infinite referral chains are -// cut off at the configured max depth. +// TestIntegrationMaxDepthExceeded verifies that an endless referral chain is +// cut off with a "Maxdepth N exceeded" exception leaf. func TestIntegrationMaxDepthExceeded(t *testing.T) { - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 3, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - // Always return a referral to ns.example.com - resp := new(dns.Msg) - resp.SetReply(msg) - resp.Ns = append(resp.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS, Class: dns.ClassINET}, - Ns: "ns.example.com.", - }) - resp.Extra = append(resp.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }) - return resp, nil - }) + m := newMockNet() + // Each hop delegates one label deeper: the node at depth 3 (refid 1.1.1) + // is never queried because MaxDepth 3 injects the exception first. + m.on("198.41.0.4", "www.d2.d1", dns.TypeA, referralMsg( + []dns.RR{nsRR("d1", "ns.d1")}, + aRR("ns.d1", "10.0.0.1"), + )) + m.on("10.0.0.1", "www.d2.d1", dns.TypeA, referralMsg( + []dns.RR{nsRR("d2.d1", "ns.d2.d1")}, + aRR("ns.d2.d1", "10.0.0.2"), + )) - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() + root := run(t, newTraverser(3), m, "www.d2.d1") - results, err := tr.Traverse(ctx, "deep.example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) + foundMaxdepth := false + for _, leaf := range root.StatsList() { + if leaf.Response.Status == traverse.StatusException && + leaf.Response.DQ.ExceptionMessage == "Maxdepth 3 exceeded" { + foundMaxdepth = true + } } - t.Logf("max depth test: %d results, types: %v", len(results), responseTypes(results)) - if len(results) == 0 { - t.Fatal("expected results even with max depth exceeded") + if !foundMaxdepth { + t.Errorf("expected a Maxdepth 3 exceeded exception leaf, got %v", statuses(root)) } } -// TestIntegrationHooksReceiveEvents verifies that traversal hooks receive -// the expected start and complete events. +// TestIntegrationHooksReceiveEvents verifies start/answer events pair up. func TestIntegrationHooksReceiveEvents(t *testing.T) { - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) + m := newMockNet() + m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "93.184.216.34"))) - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - var startEvents, completeEvents int + tr := newTraverser(10) + var startEvents, answerEvents int tr.SetHooks(&traverse.TraverserHooks{ OnEvent: func(e traverse.TraversalEvent) { switch e.Stage { - case traverse.EventStart: + case traverse.StageStart: startEvents++ - case traverse.EventComplete: - completeEvents++ + case traverse.StageAnswer: + answerEvents++ } }, }) - - ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) - defer cancel() - - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) - } - _ = results + run(t, tr, m, "example.com") if startEvents == 0 { t.Error("expected at least one start event") } - if completeEvents == 0 { - t.Error("expected at least one complete event") - } - if startEvents != completeEvents { - t.Errorf("start events (%d) != complete events (%d)", startEvents, completeEvents) + if startEvents != answerEvents { + t.Errorf("start events (%d) != answer events (%d)", startEvents, answerEvents) } } // TestIntegrationContextCancellation verifies that the traversal respects // context cancellation and returns an appropriate error. func TestIntegrationContextCancellation(t *testing.T) { - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 10, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - // Always return referral to keep loop going - resp := new(dns.Msg) - resp.SetReply(msg) - resp.Ns = append(resp.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, - Ns: "ns.example.com.", - }) - resp.Extra = append(resp.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA}, - A: net.ParseIP("1.2.3.4"), - }) - return resp, nil - }) + m := newMockNet() + m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "93.184.216.34"))) + tr := newTraverser(10) + tr.SetExchange(m.exchange) ctx, cancel := context.WithCancel(context.Background()) - cancel() // Cancel before traversal starts + cancel() // cancel before traversal starts - _, err := tr.Traverse(ctx, "example.com") - if err == nil { + if _, err := tr.Run(ctx, "example.com"); err == nil { t.Fatal("expected error when context is cancelled") } } - -// responseTypes returns a summary of response types for debugging. -func responseTypes(results []traverse.TraversalResult) []string { - var types []string - for _, r := range results { - if r.Response != nil { - types = append(types, r.Response.Type.String()) - } else { - types = append(types, "nil") - } - } - return types -} diff --git a/internal/output/coverage_test.go b/internal/output/coverage_test.go deleted file mode 100644 index 3f7f0d0..0000000 --- a/internal/output/coverage_test.go +++ /dev/null @@ -1,877 +0,0 @@ -package output - -import ( - "bytes" - "context" - "net" - "strings" - "testing" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" -) - -// ---- stats.go coverage ---- - -func TestRRDataString(t *testing.T) { - cases := []struct { - rr miekgdns.RR - want string - }{ - { - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - "1.2.3.4", - }, - { - &miekgdns.AAAA{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeAAAA}, AAAA: net.ParseIP("::1")}, - "::1", - }, - { - &miekgdns.CNAME{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeCNAME}, Target: "example.com."}, - "example.com.", - }, - { - &miekgdns.NS{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeNS}, Ns: "ns1.example.com."}, - "ns1.example.com.", - }, - { - &miekgdns.MX{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeMX}, Preference: 10, Mx: "mail.example.com."}, - "10 mail.example.com.", - }, - { - &miekgdns.TXT{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeTXT}, Txt: []string{"v=spf1", "include:example.com"}}, - "v=spf1 include:example.com", - }, - } - for _, tc := range cases { - got := rrDataString(tc.rr) - if got != tc.want { - t.Errorf("rrDataString(%T) = %q, want %q", tc.rr, got, tc.want) - } - } -} - -func TestSummaryTypeLabel(t *testing.T) { - cases := []struct { - input string - want string - }{ - {"nodata", "found no such record"}, - {"nxdomain", "name does not exist"}, - {"servfail", "resulted in SERVFAIL"}, - {"refused", "query refused by server"}, - {"notimp", "query type not implemented by server"}, - {"cname_loop", "resulted in a CNAME loop"}, - {"error", "resulted in an error"}, - {"referral", "resulted in a referral"}, - {"unknown_type", "unknown_type"}, - } - for _, tc := range cases { - got := summaryTypeLabel(tc.input) - if got != tc.want { - t.Errorf("summaryTypeLabel(%q) = %q, want %q", tc.input, got, tc.want) - } - } -} - -func TestCollectServers(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 0.5, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - } - results := []traverse.TraversalResult{ - {Referral: ref, Response: resp}, - } - servers := collectServers(results) - if len(servers) == 0 { - t.Fatal("expected at least one server") - } -} - -func TestCollectServersNilServer(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: nil, - Type: traverse.RespAnswer, - } - results := []traverse.TraversalResult{{Referral: ref, Response: resp}} - servers := collectServers(results) - if len(servers) != 0 { - t.Errorf("expected 0 servers with nil server, got %d", len(servers)) - } -} - -func TestCollectServersDedup(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 0.5, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - } - results := []traverse.TraversalResult{ - {Referral: ref, Response: resp}, - {Referral: ref, Response: resp}, - } - servers := collectServers(results) - for _, ips := range servers { - for _, ip := range ips { - count := 0 - for _, i := range ips { - if i == ip { - count++ - } - } - if count > 1 { - t.Errorf("duplicate IP %s in server list", ip) - } - } - } -} - -func TestServerName(t *testing.T) { - t.Run("uses bailiwick", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - result := traverse.TraversalResult{Referral: ref, Response: nil} - name := serverName(result) - if name != "com" { - t.Errorf("serverName = %q, want 'com'", name) - } - }) - - t.Run("uses NSName when bailiwick is root", func(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, ".", 0, 1.0, nil) - ref.NSName = "ns1.example.com." - result := traverse.TraversalResult{Referral: ref, Response: nil} - name := serverName(result) - if name != "ns1.example.com." { - t.Errorf("serverName = %q, want 'ns1.example.com.'", name) - } - }) - - t.Run("uses server IP from response", func(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - name := serverName(result) - if name != "1.2.3.4" { - t.Errorf("serverName = %q, want '1.2.3.4'", name) - } - }) - - t.Run("unknown fallback", func(t *testing.T) { - result := traverse.TraversalResult{Referral: nil, Response: nil} - name := serverName(result) - if name != "unknown" { - t.Errorf("serverName = %q, want 'unknown'", name) - } - }) -} - -func TestComputeSummaryNonAnswerTypes(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - for _, respType := range []traverse.ResponseType{ - traverse.RespNXDOMAIN, traverse.RespSERVFAIL, traverse.RespNODATA, - } { - resp := &traverse.Response{Referral: ref, Type: respType} - stats := ComputeSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if stats == nil { - t.Errorf("ComputeSummary returned nil for %v", respType) - continue - } - if len(stats.ByType) == 0 { - t.Errorf("expected ByType entry for %v", respType) - } - } -} - -func TestComputeSummaryNilResponse(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - stats := ComputeSummary([]traverse.TraversalResult{{Referral: ref, Response: nil}}) - if stats != nil { - t.Error("expected nil stats for nil response") - } -} - -func TestComputeSummaryNilReferral(t *testing.T) { - resp := &traverse.Response{Type: traverse.RespAnswer} - stats := ComputeSummary([]traverse.TraversalResult{{Referral: nil, Response: resp}}) - if stats != nil { - t.Error("expected nil stats for nil referral") - } -} - -func TestComputeSummaryAnswerKeyEmpty(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - // Answer with only CNAME (no data key) should go into ByType - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.CNAME{ - Hdr: miekgdns.RR_Header{Name: "www.example.com.", Rrtype: miekgdns.TypeCNAME}, - Target: "example.com.", - }, - }, - }, - } - stats := ComputeSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if stats == nil { - t.Fatal("expected non-nil stats") - } -} - -// ---- text.go coverage ---- - -func TestTextFormatterWriteResolve(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteResolve(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - if err != nil { - t.Fatalf("WriteResolve: %v", err) - } - if buf.Len() == 0 { - t.Error("expected output from WriteResolve") - } -} - -func TestTextFormatterWriteResolveNonStart(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - f := newTextFormatter(cfg, &buf) - - err := f.WriteResolve(traverse.TraversalEvent{ - Stage: traverse.EventComplete, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - if err != nil { - t.Fatalf("WriteResolve: %v", err) - } - if buf.Len() != 0 { - t.Error("expected no output for non-start resolve event") - } -} - -func TestTextFormatterWriteServers(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: miekgdns.TypeA}, - A: net.ParseIP("1.2.3.4"), - }, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.ShowServers = true - cfg.ShowResults = false - cfg.ShowSummaryResults = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if err != nil { - t.Fatalf("WriteSummary: %v", err) - } - if !strings.Contains(buf.String(), "The following servers were encountered:") { - t.Errorf("expected server list header, got %q", buf.String()) - } -} - -func TestTextFormatterWriteResults(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: miekgdns.TypeA}, - A: net.ParseIP("93.184.216.34"), - }, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.ShowServers = false - cfg.ShowResults = true - cfg.ShowSummaryResults = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if err != nil { - t.Fatalf("WriteSummary: %v", err) - } - if !strings.Contains(buf.String(), "Results:") { - t.Errorf("expected 'Results:' header, got %q", buf.String()) - } -} - -func TestTextFormatterFormatResultLineAllTypes(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - - cases := []struct { - respType traverse.ResponseType - contains string - }{ - {traverse.RespNODATA, "no such record"}, - {traverse.RespNXDOMAIN, "does not exist"}, - {traverse.RespSERVFAIL, "SERVFAIL"}, - {traverse.RespREFUSED, "refused"}, - {traverse.RespNOTIMPL, "not implemented"}, - {traverse.RespCNAMELoop, "CNAME loop"}, - {traverse.RespError, "error"}, - } - - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &bytes.Buffer{}) - - for _, tc := range cases { - resp := &traverse.Response{ - Referral: ref, - Type: tc.respType, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - line := f.formatResultLine(result) - if !strings.Contains(strings.ToLower(line), strings.ToLower(tc.contains)) { - t.Errorf("formatResultLine(%v) = %q, want substring %q", tc.respType, line, tc.contains) - } - } -} - -func TestTextFormatterFormatResultLineErrorWithMessage(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespError, - ErrorMessage: "custom error message", - } - - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &bytes.Buffer{}) - - line := f.formatResultLine(traverse.TraversalResult{Referral: ref, Response: resp}) - if !strings.Contains(line, "custom error message") { - t.Errorf("expected custom error message, got %q", line) - } -} - -func TestTextFormatterFormatResultLineCNAMELoopWithMessage(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespCNAMELoop, - ErrorMessage: "CNAME loop detected: example.com", - } - - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &bytes.Buffer{}) - - line := f.formatResultLine(traverse.TraversalResult{Referral: ref, Response: resp}) - if !strings.Contains(line, "CNAME loop detected") { - t.Errorf("expected CNAME loop message, got %q", line) - } -} - -func TestTextFormatterFormatResultLineAnswerMultiple(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.1.1.1")}, - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("2.2.2.2")}, - }, - }, - } - - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &bytes.Buffer{}) - - line := f.formatResultLine(traverse.TraversalResult{Referral: ref, Response: resp}) - if !strings.Contains(line, "/") { - t.Errorf("expected '/' separator for multiple answers, got %q", line) - } -} - -func TestTextFormatterColorize(t *testing.T) { - cfg := DefaultConfig() - cfg.Color = true - f := newTextFormatter(cfg, &bytes.Buffer{}) - - colored := f.colorize("hello", colorGreen) - if !strings.Contains(colored, "\033[") { - t.Error("expected ANSI color code in colored output") - } - - cfg.Color = false - f2 := newTextFormatter(cfg, &bytes.Buffer{}) - plain := f2.colorize("hello", colorGreen) - if plain != "hello" { - t.Errorf("expected plain text without color, got %q", plain) - } -} - -func TestTextFormatterColorizeEmpty(t *testing.T) { - cfg := DefaultConfig() - cfg.Color = true - f := newTextFormatter(cfg, &bytes.Buffer{}) - out := f.colorize("hello", "") - if out != "hello" { - t.Errorf("empty color should return plain text, got %q", out) - } -} - -func TestTextFormatterVerboseProgress(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.Verbose = true - f := newTextFormatter(cfg, &buf) - - err := f.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - if err != nil { - t.Fatalf("WriteProgress: %v", err) - } - if buf.Len() == 0 { - t.Error("expected output with verbose mode") - } - out := buf.String() - if !strings.Contains(out, "com") { - t.Errorf("expected bailiwick in verbose output, got %q", out) - } -} - -func TestTextFormatterProgressResolving(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - // no addresses = resolving - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - if err != nil { - t.Fatalf("WriteProgress: %v", err) - } - if !strings.Contains(buf.String(), "resolving") { - t.Errorf("expected 'resolving' in output, got %q", buf.String()) - } -} - -func TestTextFormatterWriteServersWithVersions(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.ShowServers = true - cfg.ShowResults = false - cfg.ShowSummaryResults = false - cfg.ShowVersions = true - cfg.Fingerprints = map[string]string{"1.2.3.4": "BIND 9.16"} - f := newTextFormatter(cfg, &buf) - - err := f.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if err != nil { - t.Fatalf("WriteSummary: %v", err) - } - if !strings.Contains(buf.String(), "BIND 9.16") { - t.Errorf("expected version string in server output, got %q", buf.String()) - } -} - -func TestTextFormatterWriteResult(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - f := newTextFormatter(cfg, &buf) - - err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}) - if err != nil { - t.Fatalf("WriteResult: %v", err) - } - if buf.Len() == 0 { - t.Error("expected output from WriteResult") - } -} - -func TestTextFormatterWriteResultNilRefs(t *testing.T) { - var buf bytes.Buffer - cfg := DefaultConfig() - f := newTextFormatter(cfg, &buf) - - err := f.WriteResult(traverse.TraversalResult{Referral: nil, Response: nil}) - if err != nil { - t.Fatalf("WriteResult: %v", err) - } - if buf.Len() != 0 { - t.Error("expected no output for nil referral/response") - } -} - -func TestReferralServerLabelVariants(t *testing.T) { - t.Run("with addresses", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")} - label := referralServerLabel(ref, nil) - if !strings.Contains(label, "1.2.3.4") { - t.Errorf("expected IP in label, got %q", label) - } - }) - - t.Run("with NSName", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - ref.NSName = "ns1.example.com." - label := referralServerLabel(ref, nil) - if label != "ns1.example.com." { - t.Errorf("expected NSName, got %q", label) - } - }) - - t.Run("with bailiwick", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - label := referralServerLabel(ref, nil) - if label != "com" { - t.Errorf("expected trimmed bailiwick, got %q", label) - } - }) - - t.Run("unknown", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - label := referralServerLabel(ref, nil) - if label != "unknown" { - t.Errorf("expected 'unknown', got %q", label) - } - }) - - t.Run("with response server", func(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{Server: net.ParseIP("5.6.7.8")} - label := referralServerLabel(ref, resp) - if label != "5.6.7.8" { - t.Errorf("expected server IP, got %q", label) - } - }) -} - -// ---- json.go coverage ---- - -func TestJSONFormatterWriteResolve(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowResolves = true - f := newJSONFormatter(cfg, &buf) - - err := f.WriteResolve(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - if err != nil { - t.Fatalf("WriteResolve: %v", err) - } - if len(f.payload.Resolves) != 1 { - t.Errorf("expected 1 resolve entry, got %d", len(f.payload.Resolves)) - } -} - -func TestJSONFormatterWriteResolveShowResolvesFalse(t *testing.T) { - ref := traverse.NewReferral("ns1.example.com.", dns.TypeA, "example.com.", 1, 0.5, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowResolves = false - f := newJSONFormatter(cfg, &buf) - - err := f.WriteResolve(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - if err != nil { - t.Fatalf("WriteResolve: %v", err) - } - if len(f.payload.Resolves) != 0 { - t.Errorf("expected 0 resolve entries when ShowResolves=false, got %d", len(f.payload.Resolves)) - } -} - -func TestJSONFormatterWriteResult(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowAllStats = true - f := newJSONFormatter(cfg, &buf) - - err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}) - if err != nil { - t.Fatalf("WriteResult: %v", err) - } - if len(f.payload.Results) != 1 { - t.Errorf("expected 1 result entry, got %d", len(f.payload.Results)) - } -} - -func TestJSONFormatterWriteResultShowAllStatsFalse(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{Referral: ref, Type: traverse.RespAnswer} - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowAllStats = false - f := newJSONFormatter(cfg, &buf) - - err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}) - if err != nil { - t.Fatalf("WriteResult: %v", err) - } - if len(f.payload.Results) != 0 { - t.Errorf("expected 0 result entries when ShowAllStats=false, got %d", len(f.payload.Results)) - } -} - -func TestJSONFormatterWriteSummaryWithServersAndVersions(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeA}, A: net.ParseIP("1.2.3.4")}, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowServers = true - cfg.ShowVersions = true - cfg.Fingerprints = map[string]string{"1.2.3.4": "BIND 9.16"} - f := newJSONFormatter(cfg, &buf) - - err := f.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if err != nil { - t.Fatalf("WriteSummary: %v", err) - } - found := false - for _, srv := range f.payload.Servers { - if srv.Version == "BIND 9.16" { - found = true - } - } - if !found { - t.Error("expected version in server list") - } -} - -func TestJSONFormatterStageName(t *testing.T) { - if stageName(traverse.EventStart) != "start" { - t.Errorf("expected 'start', got %q", stageName(traverse.EventStart)) - } - if stageName(traverse.EventComplete) != "complete" { - t.Errorf("expected 'complete', got %q", stageName(traverse.EventComplete)) - } - if stageName(traverse.EventStage(99)) != "unknown" { - t.Errorf("expected 'unknown' for unknown stage") - } -} - -func TestJSONFormatterEventToJSONNilReferral(t *testing.T) { - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - f := newJSONFormatter(cfg, &buf) - - item := f.eventToJSON(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: nil}, - }) - if item.Name != "" { - t.Errorf("expected empty name for nil referral, got %q", item.Name) - } -} - -func TestJSONFormatterEventToJSONWithResponse(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - f := newJSONFormatter(cfg, &buf) - - item := f.eventToJSON(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref, Response: resp}, - }) - if item.Server != "1.2.3.4" { - t.Errorf("expected server IP, got %q", item.Server) - } -} - -func TestJSONFormatterWriteProgressShowProgressFalse(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.ShowProgress = false - f := newJSONFormatter(cfg, &buf) - - err := f.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - if err != nil { - t.Fatalf("WriteProgress: %v", err) - } - if len(f.payload.Progress) != 0 { - t.Errorf("expected 0 progress entries when ShowProgress=false, got %d", len(f.payload.Progress)) - } -} - -// ---- runner.go coverage ---- - -func TestRunTraversalNilTraverser(t *testing.T) { - _, err := RunTraversal(context.Background(), nil, nil, nil, "example.com") - if err == nil { - t.Fatal("expected error for nil traverser") - } -} - -func TestNewFormatterNilConfig(t *testing.T) { - f := NewFormatter(nil, &bytes.Buffer{}) - if f == nil { - t.Fatal("NewFormatter(nil) should not return nil") - } -} - -func TestAttachHooksNilCfg(t *testing.T) { - h := AttachHooks(nil, nil) - if h != nil { - t.Fatal("AttachHooks(nil, nil) should return nil") - } -} - -func TestAttachHooksDebugMode(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Debug = 1 - cfg.ShowResolves = true - - // formatter that returns an error on WriteResolve - formatter := &errorFormatter{} - hooks := AttachHooks(cfg, formatter) - if hooks == nil { - t.Fatal("expected non-nil hooks") - } - - // Call OnEvent with IsResolve=true - should call WriteResolve and log error to stderr (debug>0) - hooks.OnEvent(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - IsResolve: true, - }) - _ = buf.String() // no assertion - just ensure it doesn't panic -} - -// errorFormatter is a mock formatter for testing error paths. -type errorFormatter struct{} - -func (f *errorFormatter) WriteProgress(_ traverse.TraversalEvent) error { return nil } -func (f *errorFormatter) WriteResolve(_ traverse.TraversalEvent) error { return nil } -func (f *errorFormatter) WriteResult(_ traverse.TraversalResult) error { return nil } -func (f *errorFormatter) WriteSummary(_ []traverse.TraversalResult) error { - return nil -} -func (f *errorFormatter) Flush() error { return nil } diff --git a/internal/output/formatter.go b/internal/output/formatter.go index 1fc1cd5..c7df01c 100644 --- a/internal/output/formatter.go +++ b/internal/output/formatter.go @@ -1,10 +1,11 @@ -// Package output renders ExploreDNS traversal results for human consumption or -// machine processing. +// Package output renders ExploreDNS traversal results for human consumption +// or machine processing. // // Two formats are supported: // -// - FormatText — a coloured hierarchical tree (default) -// - FormatJSON — a JSON array of traversal results +// - FormatText — dnstraverse-style header/progress/results/summary text +// (default) +// - FormatJSON — a single JSON document with the aggregated results // // Create a Formatter via NewFormatter and call RunTraversal to drive the // traversal engine and stream output incrementally. @@ -26,9 +27,19 @@ const ( ) type Config struct { - Format Format - Domain string - QueryType string + Format Format + Domain string + QueryType string + + // Engine settings echoed in the header block (bin/dnstraverse). + Fast bool + AllRootServers bool + UDPSize int + Retries int + MaxDepth int + AllowTCP bool + AlwaysTCP bool + ShowProgress bool ShowResolves bool ShowServers bool @@ -49,22 +60,50 @@ type Config struct { func DefaultConfig() *Config { return &Config{ Format: FormatText, + Fast: true, + UDPSize: 2048, + Retries: 2, + MaxDepth: 20, + AllowTCP: true, ShowProgress: true, - ShowResolves: true, - ShowServers: true, + ShowResolves: false, + ShowServers: false, ShowVersions: true, - ShowAllStats: true, + ShowAllStats: false, ShowResults: true, ShowSummaryResults: true, - Color: os.Getenv("NO_COLOR") == "", + Color: ColorEnabled(os.Stdout), } } +// ColorEnabled reports whether colour output should be used for w: only when +// NO_COLOR is unset and w is a terminal. +func ColorEnabled(w io.Writer) bool { + if os.Getenv("NO_COLOR") != "" { + return false + } + f, ok := w.(*os.File) + if !ok { + return false + } + info, err := f.Stat() + if err != nil { + return false + } + return info.Mode()&os.ModeCharDevice != 0 +} + +// Formatter renders traversal progress and the aggregated results. The header +// is written once before the run, progress arrives via the traverser hooks, +// and the aggregated leaves and the servers seen arrive once the run +// completed. type Formatter interface { + // WriteHeader renders the pre-run header block from the discovered roots + // (suppressed entirely by --quiet in text mode). + WriteHeader(roots []traverse.StartServer) error WriteProgress(event traverse.TraversalEvent) error WriteResolve(event traverse.TraversalEvent) error - WriteResult(result traverse.TraversalResult) error - WriteSummary(results []traverse.TraversalResult) error + WriteSummary(root *traverse.Referral, servers map[string][]string) error Flush() error } @@ -85,6 +124,11 @@ func AttachHooks(cfg *Config, formatter Formatter) *traverse.TraverserHooks { if cfg == nil || formatter == nil { return nil } + if !cfg.ShowProgress { + // Ruby registers no progress callbacks at all without show-progress; + // resolve display additionally requires show-resolves. + return nil + } logErr := func(context string, err error) { if err != nil && cfg.Debug > 0 { fmt.Fprintf(os.Stderr, "Debug: formatter %s: %v\n", context, err) @@ -92,15 +136,13 @@ func AttachHooks(cfg *Config, formatter Formatter) *traverse.TraverserHooks { } return &traverse.TraverserHooks{ OnEvent: func(event traverse.TraversalEvent) { - switch { - case event.IsResolve && cfg.ShowResolves: - logErr("WriteResolve", formatter.WriteResolve(event)) - case !event.IsResolve && cfg.ShowProgress: - logErr("WriteProgress", formatter.WriteProgress(event)) - } - if event.Stage == traverse.EventComplete && cfg.ShowAllStats { - logErr("WriteResult", formatter.WriteResult(event.Result)) + if event.IsResolve { + if cfg.ShowResolves { + logErr("WriteResolve", formatter.WriteResolve(event)) + } + return } + logErr("WriteProgress", formatter.WriteProgress(event)) }, } } diff --git a/internal/output/formatter_test.go b/internal/output/formatter_test.go index 34fcdc7..c4393d6 100644 --- a/internal/output/formatter_test.go +++ b/internal/output/formatter_test.go @@ -7,369 +7,420 @@ import ( "net" "strings" "testing" + "time" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" + idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" + "github.com/miekg/dns" ) +// mockDelegation wires root → com → example.com (2 NS, one glueless answer +// path) through the single injected exchange. +func mockDelegation() idns.ExchangeFunc { + responses := map[string]*dns.Msg{} + set := func(server, qname string, msg *dns.Msg) { + responses[server+"/"+dns.Fqdn(qname)] = msg + } + a := func(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip).To4(), + } + } + ns := func(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: dns.Fqdn(target), + } + } + + rootMsg := new(dns.Msg) + rootMsg.Ns = []dns.RR{ns("com", "a.gtld-servers.net")} + rootMsg.Extra = []dns.RR{a("a.gtld-servers.net", "192.5.6.30")} + set("198.41.0.4", "www.example.com", rootMsg) + + comMsg := new(dns.Msg) + comMsg.Ns = []dns.RR{ns("example.com", "ns1.example.com"), ns("example.com", "ns2.example.com")} + comMsg.Extra = []dns.RR{a("ns1.example.com", "1.1.1.1"), a("ns2.example.com", "2.2.2.2")} + set("192.5.6.30", "www.example.com", comMsg) + + answer := new(dns.Msg) + answer.Answer = []dns.RR{a("www.example.com", "9.9.9.9")} + set("1.1.1.1", "www.example.com", answer) + set("2.2.2.2", "www.example.com", answer) + + return func(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h + } + resp, ok := responses[host+"/"+msg.Question[0].Name] + if !ok { + return nil, &net.DNSError{Err: "no mock", Name: msg.Question[0].Name} + } + out := resp.Copy() + out.SetReply(msg) + out.Answer, out.Ns, out.Extra = resp.Answer, resp.Ns, resp.Extra + return out, nil + } +} + +func newMockTraverser() *traverse.Traverser { + tr := traverse.NewTraverser(&traverse.TraverserConfig{ + MaxDepth: traverse.DefaultMaxDepth, + QueryType: dns.TypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + QueryConfig: &idns.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, + }) + tr.SetExchange(mockDelegation()) + return tr +} + func TestDefaultConfig(t *testing.T) { cfg := DefaultConfig() - if !cfg.ShowProgress { - t.Fatal("expected ShowProgress default true") - } - if cfg.Format != FormatText { - t.Fatalf("Format = %v, want text", cfg.Format) + if cfg.Format != FormatText || !cfg.ShowProgress || !cfg.ShowResults { + t.Errorf("unexpected defaults: %+v", cfg) } } func TestNewFormatterSelectsImplementation(t *testing.T) { - text := NewFormatter(DefaultConfig(), &bytes.Buffer{}) - if _, ok := text.(*textFormatter); !ok { - t.Fatalf("expected text formatter, got %T", text) + if _, ok := NewFormatter(&Config{Format: FormatJSON}, &bytes.Buffer{}).(*jsonFormatter); !ok { + t.Error("FormatJSON should select the JSON formatter") } - - jsonCfg := DefaultConfig() - jsonCfg.Format = FormatJSON - jsonFmt := NewFormatter(jsonCfg, &bytes.Buffer{}) - if _, ok := jsonFmt.(*jsonFormatter); !ok { - t.Fatalf("expected json formatter, got %T", jsonFmt) + if _, ok := NewFormatter(&Config{Format: FormatText}, &bytes.Buffer{}).(*textFormatter); !ok { + t.Error("FormatText should select the text formatter") + } + if NewFormatter(nil, &bytes.Buffer{}) == nil { + t.Error("nil config should still produce a formatter") } } -func TestComputeSummaryAggregatesAnswers(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("93.184.216.34"), - }, - }, - }, - } - - stats := ComputeSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}) - if len(stats.Answers) != 1 { - t.Fatalf("answers = %d, want 1", len(stats.Answers)) - } - if stats.Answers[0].Prob != 1.0 { - t.Fatalf("prob = %v, want 1.0", stats.Answers[0].Prob) - } -} - -func TestTextFormatterSummaryOutput(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("93.184.216.34"), - }, - }, - }, - } - +func TestRunTraversalTextOutput(t *testing.T) { var buf bytes.Buffer cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowServers = true + cfg.ShowVersions = false // no fingerprint network calls in tests cfg.Color = false - cfg.ShowServers = false - cfg.ShowResults = false formatter := NewFormatter(cfg, &buf) - if err := formatter.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}); err != nil { - t.Fatalf("WriteSummary: %v", err) - } - - out := buf.String() - if !strings.Contains(out, "Summary:") { - t.Fatalf("expected summary header, got %q", out) - } - if !strings.Contains(out, "100%") { - t.Fatalf("expected probability in summary, got %q", out) - } - if !strings.Contains(out, "93.184.216.34") { - t.Fatalf("expected answer IP in summary, got %q", out) - } -} - -func TestJSONFormatterProducesValidOutput(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("198.41.0.4"), - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("93.184.216.34"), - }, - }, - }, - } - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Format = FormatJSON - cfg.Domain = "example.com" - cfg.QueryType = "A" - cfg.ShowServers = false - cfg.ShowResults = true - cfg.ShowSummaryResults = true - formatter := NewFormatter(cfg, &buf) - - if err := formatter.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }); err != nil { - t.Fatalf("WriteProgress: %v", err) - } - if err := formatter.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}); err != nil { - t.Fatalf("WriteSummary: %v", err) - } - if err := formatter.Flush(); err != nil { - t.Fatalf("Flush: %v", err) - } - - var payload map[string]any - if err := json.Unmarshal(buf.Bytes(), &payload); err != nil { - t.Fatalf("invalid json: %v\n%s", err, buf.String()) - } - if payload["domain"] != "example.com" { - t.Fatalf("domain = %v", payload["domain"]) - } - if _, ok := payload["summary"]; !ok { - t.Fatalf("expected summary in json output") - } -} - -func TestRunTraversalUsesHooks(t *testing.T) { - answerResp := func() *miekgdns.Msg { - m := new(miekgdns.Msg) - m.SetReply(new(miekgdns.Msg)) - m.Answer = append(m.Answer, &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - return m - }() - - tr := traverse.NewTraverser(&traverse.TraverserConfig{ - MaxDepth: 5, - QueryType: dns.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *miekgdns.Msg, useTCP bool) (*miekgdns.Msg, error) { - return answerResp.Copy(), nil - }) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - cfg.ShowServers = false - cfg.ShowResults = false - cfg.ShowSummaryResults = true - formatter := NewFormatter(cfg, &buf) - - _, err := RunTraversal(context.Background(), tr, cfg, formatter, "example.com") + root, err := RunTraversal(context.Background(), newMockTraverser(), cfg, formatter, "www.example.com") if err != nil { t.Fatalf("RunTraversal: %v", err) } - if !strings.Contains(buf.String(), "Summary:") { - t.Fatalf("expected formatted summary output, got %q", buf.String()) + if root == nil || len(root.Stats) == 0 { + t.Fatal("expected aggregated stats on the root referral") + } + + out := buf.String() + for _, want := range []string{ + "# Using fast mode", + "Using 198.41.0.4 (198.41.0.4) as initial root", + "Running query www.example.com type a", + "1 198.41.0.4 (198.41.0.4)", + "1.1 a.gtld-servers.net (192.5.6.30)", + "1.1.1 ns1.example.com (1.1.1.1)", + "Results:", + " 50.0%: Answer from ns1.example.com (1.1.1.1)", + " 50.0%: Answer from ns2.example.com (2.2.2.2)", + "Summary Results:", + " 100% answered with www.example.com. 300 IN A 9.9.9.9", + "The following servers were encountered:", + } { + if !strings.Contains(out, want) { + t.Errorf("output missing %q\n---\n%s", want, out) + } } } -func TestJSONFormatterWriteResolveAndResult(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("198.41.0.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{ -Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, -A: net.ParseIP("1.2.3.4"), -}, -}, -}, +func TestVerboseProgressFormat(t *testing.T) { + var buf bytes.Buffer + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowVersions = false + cfg.Verbose = true + cfg.Color = false + formatter := NewFormatter(cfg, &buf) + + if _, err := RunTraversal(context.Background(), newMockTraverser(), cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + out := buf.String() + // Verbose rows are " [qname] () "; the + // root bailiwick renders as "<>". + for _, want := range []string{ + "1 [www.example.com] 198.41.0.4 (198.41.0.4) <>", + "1.1 [www.example.com] a.gtld-servers.net (192.5.6.30) ", + "1.1.1 [www.example.com] ns1.example.com (1.1.1.1) ", + } { + if !strings.Contains(out, want) { + t.Errorf("output missing %q\n---\n%s", want, out) + } + } } -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Format = FormatJSON -cfg.Domain = "example.com" -cfg.QueryType = "A" -cfg.ShowResolves = true -cfg.ShowAllStats = true -cfg.ShowProgress = true -f := NewFormatter(cfg, &buf).(*jsonFormatter) +func TestRunTraversalQuietSuppressesHeader(t *testing.T) { + var buf bytes.Buffer + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowVersions = false + cfg.Quiet = true + cfg.Color = false + formatter := NewFormatter(cfg, &buf) -// WriteResolve -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteResolve: %v", err) + if _, err := RunTraversal(context.Background(), newMockTraverser(), cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + out := buf.String() + for _, banned := range []string{"# Using fast mode", "as initial root", "Running query"} { + if strings.Contains(out, banned) { + t.Errorf("quiet output must not contain %q\n---\n%s", banned, out) + } + } + if !strings.Contains(out, "Results:") { + t.Errorf("quiet must still print results\n---\n%s", out) + } } -// WriteResult -if err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}); err != nil { -t.Fatalf("WriteResult: %v", err) +func TestRunTraversalJSONOutput(t *testing.T) { + var buf bytes.Buffer + cfg := DefaultConfig() + cfg.Format = FormatJSON + cfg.Domain = "www.example.com" + cfg.QueryType = "A" + cfg.ShowVersions = false + formatter := NewFormatter(cfg, &buf) + + if _, err := RunTraversal(context.Background(), newMockTraverser(), cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + + var doc map[string]any + if err := json.Unmarshal(buf.Bytes(), &doc); err != nil { + t.Fatalf("invalid JSON: %v\n%s", err, buf.String()) + } + if doc["domain"] != "www.example.com" { + t.Errorf("domain = %v", doc["domain"]) + } + if doc["qtype"] != "A" { + t.Errorf("qtype = %v", doc["qtype"]) + } + root, ok := doc["root"].(map[string]any) + if !ok || root["ip"] != "198.41.0.4" { + t.Errorf("root = %v", doc["root"]) + } + for _, banned := range []string{"progress", "resolves"} { + if _, present := doc[banned]; present { + t.Errorf("JSON document must not contain %q", banned) + } + } + results, ok := doc["results"].([]any) + if !ok || len(results) != 2 { + t.Fatalf("results = %v", doc["results"]) + } + summary, ok := doc["summary"].(map[string]any) + if !ok { + t.Fatalf("summary missing: %v", doc) + } + byStatus := summary["by_status"].(map[string]any) + if prob := byStatus["answered"].(float64); prob < 0.999 || prob > 1.001 { + t.Errorf("answered summary prob = %v", prob) + } } -// WriteProgress with EventComplete to cover stageName "complete" -if err := f.WriteProgress(traverse.TraversalEvent{ -Stage: traverse.EventComplete, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteProgress EventComplete: %v", err) +// mockGluelessDelegation wires root → com → example.com where the single NS +// (ns1.example.net) comes without glue, forcing a resolve subtree that walks +// root → net → answer. +func mockGluelessDelegation() idns.ExchangeFunc { + responses := map[string]*dns.Msg{} + set := func(server, qname string, msg *dns.Msg) { + responses[server+"/"+dns.Fqdn(qname)] = msg + } + a := func(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip).To4(), + } + } + ns := func(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: dns.Fqdn(target), + } + } + + comRef := new(dns.Msg) + comRef.Ns = []dns.RR{ns("com", "a.gtld-servers.net")} + comRef.Extra = []dns.RR{a("a.gtld-servers.net", "192.5.6.30")} + set("198.41.0.4", "www.example.com", comRef) + + glueless := new(dns.Msg) + glueless.Ns = []dns.RR{ns("example.com", "ns1.example.net")} + set("192.5.6.30", "www.example.com", glueless) + + netRef := new(dns.Msg) + netRef.Ns = []dns.RR{ns("net", "b.gtld-servers.net")} + netRef.Extra = []dns.RR{a("b.gtld-servers.net", "192.33.14.31")} + set("198.41.0.4", "ns1.example.net", netRef) + + nsAnswer := new(dns.Msg) + nsAnswer.Answer = []dns.RR{a("ns1.example.net", "3.3.3.3")} + set("192.33.14.31", "ns1.example.net", nsAnswer) + + answer := new(dns.Msg) + answer.Answer = []dns.RR{a("www.example.com", "9.9.9.9")} + set("3.3.3.3", "www.example.com", answer) + + return func(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h + } + resp, ok := responses[host+"/"+msg.Question[0].Name] + if !ok { + return nil, &net.DNSError{Err: "no mock", Name: msg.Question[0].Name} + } + out := resp.Copy() + out.SetReply(msg) + out.Answer, out.Ns, out.Extra = resp.Answer, resp.Ns, resp.Extra + return out, nil + } } -if err := f.Flush(); err != nil { -t.Fatalf("Flush: %v", err) -} +func runGlueless(t *testing.T, cfg *Config) string { + t.Helper() + tr := traverse.NewTraverser(&traverse.TraverserConfig{ + MaxDepth: traverse.DefaultMaxDepth, + QueryType: dns.TypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + QueryConfig: &idns.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, + }) + tr.SetExchange(mockGluelessDelegation()) + + var buf bytes.Buffer + formatter := NewFormatter(cfg, &buf) + if _, err := RunTraversal(context.Background(), tr, cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + return buf.String() } -func TestJSONFormatterWriteResolveFlagOff(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) +func TestResolveProgressHiddenByDefault(t *testing.T) { + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowVersions = false + cfg.Color = false + out := runGlueless(t, cfg) -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Format = FormatJSON -cfg.ShowResolves = false -cfg.ShowAllStats = false -f := NewFormatter(cfg, &buf).(*jsonFormatter) - -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref}, -}); err != nil { -t.Fatalf("WriteResolve: %v", err) -} -if err := f.WriteResult(traverse.TraversalResult{Referral: ref}); err != nil { -t.Fatalf("WriteResult: %v", err) -} + for _, want := range []string{ + "1.1.1 ns1.example.net -- resolving", + "1.1.1 ns1.example.net (3.3.3.3)", + "100.0%: Answer from ns1.example.net (3.3.3.3)", + } { + if !strings.Contains(out, want) { + t.Errorf("output missing %q\n---\n%s", want, out) + } + } + // Resolve subtree nodes (".0." refids) render only under --show-resolves, + // and resolve outcomes never appear as separate Results entries. + for _, line := range strings.Split(out, "\n") { + if strings.HasPrefix(line, "1.1.1.0") { + t.Errorf("resolve subtree must be hidden by default: %q", line) + } + } + if strings.Contains(out, "ns1.example.net./IN/A") { + t.Errorf("resolve leaves must not pollute Results\n---\n%s", out) + } } -func TestJSONFormatterWriteSummaryWithServers(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{ -Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, -A: net.ParseIP("1.2.3.4"), -}, -}, -}, -} -results := []traverse.TraversalResult{{Referral: ref, Response: resp}} +func TestResolveProgressShownWithShowResolves(t *testing.T) { + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowVersions = false + cfg.ShowResolves = true + cfg.Color = false + out := runGlueless(t, cfg) -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Format = FormatJSON -cfg.Domain = "example.com" -cfg.QueryType = "A" -cfg.ShowServers = true -cfg.ShowVersions = false -cfg.ShowResults = true -cfg.ShowSummaryResults = true -f := NewFormatter(cfg, &buf).(*jsonFormatter) - -if err := f.WriteSummary(results); err != nil { -t.Fatalf("WriteSummary: %v", err) -} -if err := f.Flush(); err != nil { -t.Fatalf("Flush: %v", err) + for _, want := range []string{ + "1.1.1 ns1.example.net -- resolving", + "1.1.1.0.1 198.41.0.4 (198.41.0.4)", + "1.1.1.0.1.1 b.gtld-servers.net (192.33.14.31)", + "1.1.1 ns1.example.net (3.3.3.3)", + } { + if !strings.Contains(out, want) { + t.Errorf("output missing %q\n---\n%s", want, out) + } + } } -var payload map[string]any -if err := json.Unmarshal(buf.Bytes(), &payload); err != nil { -t.Fatalf("invalid JSON: %v\n%s", err, buf.String()) -} -if _, ok := payload["servers"]; !ok { -t.Error("expected 'servers' field in JSON output") -} +func TestRunTraversalRequiresTraverser(t *testing.T) { + if _, err := RunTraversal(context.Background(), nil, DefaultConfig(), nil, "example.com"); err == nil { + t.Fatal("expected error for nil traverser") + } } -func TestNewFormatterNilWriter(t *testing.T) { -// Should not panic with nil writer -cfg := DefaultConfig() -f := NewFormatter(cfg, nil) -if f == nil { -t.Error("NewFormatter should not return nil") -} +func TestFormatProbability(t *testing.T) { + tests := []struct { + prob float64 + want string + }{ + {1.0, " 100%"}, + {0.5, " 50%"}, + {0.933, "93.3%"}, + {0.067, " 6.7%"}, + {1.0 / 3, "33.3%"}, + } + for _, tt := range tests { + if got := formatProbability(tt.prob); got != tt.want { + t.Errorf("formatProbability(%v) = %q, want %q", tt.prob, got, tt.want) + } + } } -func TestAttachHooksShowResolves(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, +func TestSummaryStatusLabels(t *testing.T) { + tests := map[traverse.Status]string{ + traverse.StatusNoData: "found no such record", + traverse.StatusReferralLame: "resulted in a lame referral", + traverse.StatusException: "resulted in an exception", + traverse.StatusError: "resulted in an error", + traverse.StatusNoGlue: "found no glue", + traverse.StatusLoop: "resulted in a loop", + traverse.StatusCNAMELoop: "resulted in a CNAME loop", + traverse.Status("odd"): "odd", + } + for status, want := range tests { + if got := summaryStatusLabel(status); got != want { + t.Errorf("summaryStatusLabel(%q) = %q, want %q", status, got, want) + } + } } -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.ShowProgress = false -cfg.ShowResolves = true -cfg.ShowAllStats = true -cfg.Color = false -formatter := NewFormatter(cfg, &buf) -hooks := AttachHooks(cfg, formatter) - -// Trigger a resolve event -hooks.OnEvent(traverse.TraversalEvent{ -Stage: traverse.EventStart, -IsResolve: true, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}) - -if buf.Len() == 0 { -t.Error("expected resolve output when ShowResolves is true") -} +func TestCollectUniqueServerIPs(t *testing.T) { + servers := map[string][]string{ + "ns1.example.com": {"1.1.1.1", "2.2.2.2"}, + "ns2.example.com": {"1.1.1.1"}, + } + ips := collectUniqueServerIPs(servers) + if len(ips) != 2 { + t.Errorf("unique IPs = %v", ips) + } } -func TestAttachHooksShowAllStats(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("1.2.3.4")}, -}, -}, -} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.ShowProgress = false -cfg.ShowResolves = false -cfg.ShowAllStats = true -cfg.Color = false -formatter := NewFormatter(cfg, &buf) -hooks := AttachHooks(cfg, formatter) - -hooks.OnEvent(traverse.TraversalEvent{ -Stage: traverse.EventComplete, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}) +func TestReverseString(t *testing.T) { + if got := reverseString("abc"); got != "cba" { + t.Errorf("reverseString = %q", got) + } } diff --git a/internal/output/golden_test.go b/internal/output/golden_test.go new file mode 100644 index 0000000..8675fc4 --- /dev/null +++ b/internal/output/golden_test.go @@ -0,0 +1,182 @@ +package output + +import ( + "bytes" + "context" + "net" + "strings" + "testing" + "time" + + idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" + "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" + "github.com/miekg/dns" +) + +// TestTextOutputMatchesReferenceCapture rebuilds the topology of +// docs/captures/dnstraverse-ruby-www.example.com-A.txt (root → com → the two +// cloudflare NS, three IPs each) through the mock exchange and asserts the +// complete text output byte-for-byte against the reference format. It differs +// from the capture only in volatile values: the root chosen, the number of +// gTLD servers, and the server fingerprints (versions are disabled so no +// network is touched). +func TestTextOutputMatchesReferenceCapture(t *testing.T) { + responses := map[string]*dns.Msg{} + set := func(server, qname string, msg *dns.Msg) { + responses[server+"/"+dns.Fqdn(qname)] = msg + } + a := func(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, + A: net.ParseIP(ip).To4(), + } + } + ns := func(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 300}, + Ns: dns.Fqdn(target), + } + } + + // Upstream resolver: ". NS" returns one root with glue (root discovery). + rootNS := new(dns.Msg) + rootNS.Answer = []dns.RR{ns(".", "m.root-servers.net")} + rootNS.Extra = []dns.RR{a("m.root-servers.net", "202.12.27.33")} + set("10.0.0.53", ".", rootNS) + + // Root referral to com (three gTLD servers, all glued). + comRef := new(dns.Msg) + comRef.Ns = []dns.RR{ + ns("com", "a.gtld-servers.net"), + ns("com", "b.gtld-servers.net"), + ns("com", "c.gtld-servers.net"), + } + comRef.Extra = []dns.RR{ + a("a.gtld-servers.net", "192.5.6.30"), + a("b.gtld-servers.net", "192.33.14.30"), + a("c.gtld-servers.net", "192.26.92.30"), + } + set("202.12.27.33", "www.example.com", comRef) + + // gTLD referral to example.com: two NS, three glue addresses each. + heraIPs := []string{"108.162.192.162", "172.64.32.162", "173.245.58.162"} + elliottIPs := []string{"108.162.195.228", "162.159.44.228", "172.64.35.228"} + exampleRef := new(dns.Msg) + exampleRef.Ns = []dns.RR{ + ns("example.com", "hera.ns.cloudflare.com"), + ns("example.com", "elliott.ns.cloudflare.com"), + } + for _, ip := range heraIPs { + exampleRef.Extra = append(exampleRef.Extra, a("hera.ns.cloudflare.com", ip)) + } + for _, ip := range elliottIPs { + exampleRef.Extra = append(exampleRef.Extra, a("elliott.ns.cloudflare.com", ip)) + } + for _, ip := range []string{"192.5.6.30", "192.33.14.30", "192.26.92.30"} { + set(ip, "www.example.com", exampleRef) + } + + answer := new(dns.Msg) + answer.Answer = []dns.RR{ + a("www.example.com", "104.20.23.154"), + a("www.example.com", "172.66.147.243"), + } + for _, ip := range append(append([]string{}, heraIPs...), elliottIPs...) { + set(ip, "www.example.com", answer) + } + + exchange := func(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h + } + resp, ok := responses[host+"/"+msg.Question[0].Name] + if !ok { + return nil, &net.DNSError{Err: "no mock", Name: msg.Question[0].Name} + } + out := resp.Copy() + out.SetReply(msg) + out.Answer, out.Ns, out.Extra = resp.Answer, resp.Ns, resp.Extra + return out, nil + } + + tr := traverse.NewTraverser(&traverse.TraverserConfig{ + MaxDepth: traverse.DefaultMaxDepth, + QueryType: dns.TypeA, + Fast: true, + RootConfig: &idns.RootDiscoveryConfig{Resolver: "10.0.0.53:53"}, + QueryConfig: &idns.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, + }) + tr.SetExchange(exchange) + + cfg := DefaultConfig() + cfg.Domain = "www.example.com" + cfg.QueryType = "a" + cfg.ShowServers = true + cfg.ShowVersions = false + cfg.Color = false + + var buf bytes.Buffer + formatter := NewFormatter(cfg, &buf) + if _, err := RunTraversal(context.Background(), tr, cfg, formatter, "www.example.com"); err != nil { + t.Fatalf("RunTraversal: %v", err) + } + + answerBlock := func(server, ip string) string { + return " 16.7%: Answer from " + server + " (" + ip + ")\n" + + " www.example.com.\t300\tIN\tA\t104.20.23.154\n" + + " www.example.com.\t300\tIN\tA\t172.66.147.243\n" + } + + want := strings.Join([]string{ + "# Using fast mode", + "# Limiting traverse to one root", + "# UDP size 2048 (EDNS0 is on)", + "# Retries 2, max depth 20", + "# Allow TCP is true, always TCP is false", + "Using m.root-servers.net (202.12.27.33) as initial root", + "Running query www.example.com type a", + "1 m.root-servers.net (202.12.27.33)", + "1.1 a.gtld-servers.net (192.5.6.30)", + "1.1.1 hera.ns.cloudflare.com (108.162.192.162,172.64.32.162,173.245.58.162)", + "1.1.2 elliott.ns.cloudflare.com (108.162.195.228,162.159.44.228,172.64.35.228)", + "1.2 b.gtld-servers.net (192.33.14.30)", + "1.2.1 hera.ns.cloudflare.com (108.162.192.162,172.64.32.162,173.245.58.162) -- completed earlier (1.1.1)", + "1.2.2 elliott.ns.cloudflare.com (108.162.195.228,162.159.44.228,172.64.35.228) -- completed earlier (1.1.2)", + "1.3 c.gtld-servers.net (192.26.92.30)", + "1.3.1 hera.ns.cloudflare.com (108.162.192.162,172.64.32.162,173.245.58.162) -- completed earlier (1.1.1)", + "1.3.2 elliott.ns.cloudflare.com (108.162.195.228,162.159.44.228,172.64.35.228) -- completed earlier (1.1.2)", + "", + "The following servers were encountered:", + " hera.ns.cloudflare.com: 108.162.192.162", + " hera.ns.cloudflare.com: 172.64.32.162", + " hera.ns.cloudflare.com: 173.245.58.162", + "elliott.ns.cloudflare.com: 108.162.195.228", + "elliott.ns.cloudflare.com: 162.159.44.228", + "elliott.ns.cloudflare.com: 172.64.35.228", + " a.gtld-servers.net: 192.5.6.30", + " b.gtld-servers.net: 192.33.14.30", + " c.gtld-servers.net: 192.26.92.30", + " m.root-servers.net: 202.12.27.33", + "", + "Results:", + answerBlock("hera.ns.cloudflare.com", "108.162.192.162"), + answerBlock("elliott.ns.cloudflare.com", "108.162.195.228"), + answerBlock("elliott.ns.cloudflare.com", "162.159.44.228"), + answerBlock("hera.ns.cloudflare.com", "172.64.32.162"), + answerBlock("elliott.ns.cloudflare.com", "172.64.35.228"), + answerBlock("hera.ns.cloudflare.com", "173.245.58.162") + + "\n" + + "Summary Results:\n" + + " 100% answered with www.example.com. 300 IN A 104.20.23.154\n" + + " www.example.com. 300 IN A 172.66.147.243\n", + }, "\n") + + if got := buf.String(); got != want { + t.Errorf("output does not match the reference capture format\n--- got ---\n%s\n--- want ---\n%s", got, want) + } +} diff --git a/internal/output/json.go b/internal/output/json.go index d1a85d1..0306f52 100644 --- a/internal/output/json.go +++ b/internal/output/json.go @@ -3,11 +3,15 @@ package output import ( "encoding/json" "io" + "sort" + "strings" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" ) +// jsonFormatter accumulates the run into a single document emitted once by +// Flush: {domain, qtype, root, results, summary, servers}. Aggregated leaves +// appear exactly once (in results); progress events are not recorded. type jsonFormatter struct { cfg *Config w io.Writer @@ -15,32 +19,31 @@ type jsonFormatter struct { } type jsonDocument struct { - Domain string `json:"domain"` - QueryType string `json:"query_type"` - Progress []jsonProgressEvent `json:"progress,omitempty"` - Resolves []jsonProgressEvent `json:"resolves,omitempty"` - Results []jsonResult `json:"results,omitempty"` - Servers []jsonServer `json:"servers,omitempty"` - Summary jsonSummary `json:"summary,omitempty"` + Domain string `json:"domain"` + QueryType string `json:"qtype"` + Root *jsonRoot `json:"root,omitempty"` + Results []jsonResult `json:"results,omitempty"` + Summary *jsonSummary `json:"summary,omitempty"` + Servers []jsonServer `json:"servers,omitempty"` } -type jsonProgressEvent struct { - Stage string `json:"stage"` - Depth int `json:"depth"` - Name string `json:"name"` - QType string `json:"qtype"` - Server string `json:"server,omitempty"` - Bailiwick string `json:"bailiwick,omitempty"` - Resolving bool `json:"resolving,omitempty"` +// jsonRoot is the initial root the traversal started from. +type jsonRoot struct { + Name string `json:"name"` + IP string `json:"ip,omitempty"` } +// jsonResult is one aggregated leaf outcome, emitted exactly once. type jsonResult struct { - Depth int `json:"depth"` - Probability float64 `json:"probability"` - ResponseType string `json:"response_type"` - Server string `json:"server,omitempty"` - Answers []string `json:"answers,omitempty"` - CNAMEChain []string `json:"cname_chain,omitempty"` + RefID string `json:"refid,omitempty"` + Probability float64 `json:"probability"` + Status string `json:"status"` + Server string `json:"server,omitempty"` + IP string `json:"ip,omitempty"` + Qname string `json:"qname,omitempty"` + Qtype string `json:"qtype,omitempty"` + Answers []string `json:"answers,omitempty"` + Message string `json:"message,omitempty"` } type jsonServer struct { @@ -50,12 +53,11 @@ type jsonServer struct { } type jsonSummary struct { - ByType map[string]float64 `json:"by_type,omitempty"` - Answers []jsonAnswerStat `json:"answers,omitempty"` + ByStatus map[string]float64 `json:"by_status,omitempty"` + Answers []jsonAnswerStat `json:"answers,omitempty"` } type jsonAnswerStat struct { - RData string `json:"rdata"` Probability float64 `json:"probability"` Records []string `json:"records,omitempty"` } @@ -71,43 +73,48 @@ func newJSONFormatter(cfg *Config, w io.Writer) *jsonFormatter { } } -func (f *jsonFormatter) WriteProgress(event traverse.TraversalEvent) error { - if !f.cfg.ShowProgress { +func (f *jsonFormatter) WriteHeader(roots []traverse.StartServer) error { + if len(roots) == 0 { return nil } - f.payload.Progress = append(f.payload.Progress, f.eventToJSON(event)) + root := &jsonRoot{Name: roots[0].Name} + if len(roots[0].IPs) > 0 { + root.IP = roots[0].IPs[0] + } + f.payload.Root = root return nil } -func (f *jsonFormatter) WriteResolve(event traverse.TraversalEvent) error { - if !f.cfg.ShowResolves { - return nil - } - f.payload.Resolves = append(f.payload.Resolves, f.eventToJSON(event)) +// WriteProgress is a no-op: the JSON document contains only the aggregated +// outcome, never per-event duplicates. +func (f *jsonFormatter) WriteProgress(traverse.TraversalEvent) error { return nil } -func (f *jsonFormatter) WriteResult(result traverse.TraversalResult) error { - if !f.cfg.ShowAllStats { - return nil - } - f.payload.Results = append(f.payload.Results, f.resultToJSON(result)) +func (f *jsonFormatter) WriteResolve(traverse.TraversalEvent) error { return nil } -func (f *jsonFormatter) WriteSummary(results []traverse.TraversalResult) error { - if f.cfg.ShowResults { - for _, result := range terminalResults(results) { - f.payload.Results = append(f.payload.Results, f.resultToJSON(result)) +func (f *jsonFormatter) WriteSummary(root *traverse.Referral, servers map[string][]string) error { + if root != nil && f.cfg.ShowResults { + // StatsList is sorted by stats key: deterministic ordering. + for _, leaf := range root.StatsList() { + f.payload.Results = append(f.payload.Results, leafToJSON(leaf)) } } if f.cfg.ShowServers { - servers := collectServers(results) - for name, ips := range servers { - srv := jsonServer{Name: name, IPs: ips} + names := make([]string, 0, len(servers)) + for name := range servers { + names = append(names, name) + } + sort.Slice(names, func(i, j int) bool { + return reverseString(strings.ToLower(names[i])) < reverseString(strings.ToLower(names[j])) + }) + for _, name := range names { + srv := jsonServer{Name: name, IPs: servers[name]} if f.cfg.ShowVersions && f.cfg.Fingerprints != nil { - for _, ip := range ips { + for _, ip := range srv.IPs { if v := f.cfg.Fingerprints[ip]; v != "" { srv.Version = v break @@ -119,18 +126,19 @@ func (f *jsonFormatter) WriteSummary(results []traverse.TraversalResult) error { } if f.cfg.ShowSummaryResults { - stats := ComputeSummary(results) - if stats != nil { - f.payload.Summary = jsonSummary{ - ByType: stats.ByType, + if stats := root.SummaryStats(); stats != nil { + summary := &jsonSummary{ByStatus: make(map[string]float64)} + for status, prob := range stats.ByStatus { + summary.ByStatus[string(status)] = prob } for _, answer := range stats.Answers { - f.payload.Summary.Answers = append(f.payload.Summary.Answers, jsonAnswerStat{ - RData: answer.RData, - Probability: answer.Prob, - Records: answer.RRs, - }) + stat := jsonAnswerStat{Probability: answer.Prob} + for _, rr := range answer.RRs { + stat.Records = append(stat.Records, collapseWhitespace(rr.String())) + } + summary.Answers = append(summary.Answers, stat) } + f.payload.Summary = summary } } @@ -143,56 +151,29 @@ func (f *jsonFormatter) Flush() error { return enc.Encode(f.payload) } -func (f *jsonFormatter) eventToJSON(event traverse.TraversalEvent) jsonProgressEvent { - ref := event.Result.Referral - if ref == nil { - return jsonProgressEvent{} +func leafToJSON(leaf *traverse.StatsEntry) jsonResult { + resp := leaf.Response + item := jsonResult{ + Probability: leaf.Prob, + Status: string(resp.Status), + IP: resp.IP, + Qname: resp.Qname, + Qtype: traverse.TypeToString(resp.Qtype), } - - item := jsonProgressEvent{ - Stage: stageName(event.Stage), - Depth: ref.Depth, - Name: trimDomain(ref.Name), - QType: dns.QNameType(ref.Qtype), - Bailiwick: trimDomain(ref.Bailiwick), - Resolving: !ref.HasAddresses(), + if leaf.Referral != nil { + item.RefID = leaf.Referral.RefID + item.Server = leaf.Referral.Server } - if event.Result.Response != nil && event.Result.Response.Server != nil { - item.Server = event.Result.Response.Server.String() - } else { - item.Server = referralServerLabel(ref, event.Result.Response) - } - return item -} - -func (f *jsonFormatter) resultToJSON(result traverse.TraversalResult) jsonResult { - item := jsonResult{} - if result.Referral != nil { - item.Depth = result.Referral.Depth - item.Probability = result.Referral.Prob - } - if result.Response != nil { - item.ResponseType = result.Response.Type.String() - if result.Response.Server != nil { - item.Server = result.Response.Server.String() + if resp.DQ != nil { + for _, rr := range resp.DQ.Answers { + item.Answers = append(item.Answers, collapseWhitespace(rr.String())) } - if result.Response.Decoded != nil { - for _, rr := range result.Response.Decoded.Answers { - item.Answers = append(item.Answers, dns.FormatRecord(rr)) - } - item.CNAMEChain = append(item.CNAMEChain, result.Response.Decoded.CNAMEChain...) + switch resp.Status { + case traverse.StatusError: + item.Message = resp.DQ.ErrorMessage + case traverse.StatusException: + item.Message = resp.DQ.ExceptionMessage } } return item } - -func stageName(stage traverse.EventStage) string { - switch stage { - case traverse.EventStart: - return "start" - case traverse.EventComplete: - return "complete" - default: - return "unknown" - } -} diff --git a/internal/output/runner.go b/internal/output/runner.go index f69b119..eba67b4 100644 --- a/internal/output/runner.go +++ b/internal/output/runner.go @@ -9,7 +9,10 @@ import ( "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" ) -func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Config, formatter Formatter, domain string) ([]traverse.TraversalResult, error) { +// RunTraversal drives one traversal and streams its output through the +// formatter. It returns the synthetic root referral whose Stats aggregate +// every leaf outcome. +func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Config, formatter Formatter, domain string) (*traverse.Referral, error) { if traverser == nil { return nil, fmt.Errorf("traverser is required") } @@ -22,40 +25,53 @@ func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Confi traverser.SetHooks(AttachHooks(cfg, formatter)) - results, err := traverser.Traverse(ctx, domain) + // Discover the roots up front so the header can report the initial root; + // Run reuses the memoised discovery. + roots, err := traverser.Roots(ctx) if err != nil { - return results, err + return nil, err } + if err := formatter.WriteHeader(roots); err != nil { + return nil, err + } + + root, err := traverser.Run(ctx, domain) + if err != nil { + return root, err + } + + servers := traverser.ServersEncountered() // Fingerprint servers when both ShowVersions and ShowServers are enabled. // Gating on ShowServers avoids unnecessary network calls when versions // would not be displayed anyway. if cfg.ShowVersions && cfg.ShowServers { - cfg.Fingerprints = fingerprint.New().FingerprintAll(ctx, collectUniqueServerIPs(results)) + cfg.Fingerprints = fingerprint.New().FingerprintAll(ctx, collectUniqueServerIPs(servers)) } - if err := formatter.WriteSummary(results); err != nil { - return results, err + if err := formatter.WriteSummary(root, servers); err != nil { + return root, err } if err := formatter.Flush(); err != nil { - return results, err + return root, err } - return results, nil + return root, nil } -// collectUniqueServerIPs returns the set of unique server IPs seen in results. -func collectUniqueServerIPs(results []traverse.TraversalResult) []net.IP { +// collectUniqueServerIPs returns the set of unique server IPs encountered. +func collectUniqueServerIPs(servers map[string][]string) []net.IP { seen := make(map[string]bool) var ips []net.IP - for _, r := range results { - if r.Response == nil || r.Response.Server == nil { - continue - } - key := r.Response.Server.String() - if !seen[key] { - seen[key] = true - ips = append(ips, r.Response.Server) + for _, addrs := range servers { + for _, addr := range addrs { + if seen[addr] { + continue + } + seen[addr] = true + if ip := net.ParseIP(addr); ip != nil { + ips = append(ips, ip) + } } } return ips diff --git a/internal/output/stats.go b/internal/output/stats.go index 45b7e8b..f20a4fb 100644 --- a/internal/output/stats.go +++ b/internal/output/stats.go @@ -2,248 +2,44 @@ package output import ( "fmt" - "sort" "strings" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" ) -type summaryEntry struct { - Type string - Prob float64 -} - -type answerEntry struct { - RData string - Prob float64 - RRs []string -} - -type SummaryStats struct { - ByType map[string]float64 - Answers []answerEntry -} - -func ComputeSummary(results []traverse.TraversalResult) *SummaryStats { - stats := &SummaryStats{ - ByType: make(map[string]float64), - } - - for _, result := range results { - if result.Response == nil || !result.Response.IsTerminal() { - continue - } - if result.Referral == nil { - continue - } - - prob := result.Referral.Prob - respType := result.Response.Type.String() - - switch result.Response.Type { - case traverse.RespAnswer: - key, rrs := answerKey(result.Response) - if key == "" { - stats.ByType[respType] += prob - continue - } - found := false - for i := range stats.Answers { - if stats.Answers[i].RData == key { - stats.Answers[i].Prob += prob - found = true - break - } - } - if !found { - stats.Answers = append(stats.Answers, answerEntry{ - RData: key, - Prob: prob, - RRs: rrs, - }) - } - default: - stats.ByType[respType] += prob - } - } - - sort.Slice(stats.Answers, func(i, j int) bool { - return stats.Answers[i].RData < stats.Answers[j].RData - }) - - if len(stats.Answers) == 0 && len(stats.ByType) == 0 { - return nil - } - return stats -} - -func answerKey(resp *traverse.Response) (string, []string) { - if resp == nil || resp.Decoded == nil { - return "", nil - } - - var rdatas []string - var formatted []string - for _, rr := range resp.Decoded.Answers { - if _, ok := rr.(*miekgdns.CNAME); ok { - continue - } - rdata := rrDataString(rr) - if rdata == "" { - continue - } - rdatas = append(rdatas, rdata) - formatted = append(formatted, dns.FormatRecord(rr)) - } - - if len(rdatas) == 0 { - return "", nil - } - sort.Strings(rdatas) - return strings.Join(rdatas, " / "), formatted -} - -func rrDataString(rr miekgdns.RR) string { - switch v := rr.(type) { - case *miekgdns.A: - return v.A.String() - case *miekgdns.AAAA: - return v.AAAA.String() - case *miekgdns.CNAME: - return v.Target - case *miekgdns.NS: - return v.Ns - case *miekgdns.MX: - return fmt.Sprintf("%d %s", v.Preference, v.Mx) - case *miekgdns.TXT: - return strings.Join(v.Txt, " ") - default: - return rr.String() - } -} - +// formatProbability renders txt_prob (summary_stats.rb): %5.1f%% with a +// trailing ".0" trimmed, right-justified to width 5. func formatProbability(prob float64) string { text := fmt.Sprintf("%.1f%%", prob*100) text = strings.Replace(text, ".0%", "%", 1) return fmt.Sprintf("%5s", text) } -func summaryTypeLabel(respType string) string { - switch respType { - case "nodata": +// summaryStatusLabel returns the Summary Results wording per status +// (summary_stats.rb text; answered is handled separately). +func summaryStatusLabel(status traverse.Status) string { + switch status { + case traverse.StatusNoData: return "found no such record" - case "nxdomain": - return "name does not exist" - case "servfail": - return "resulted in SERVFAIL" - case "refused": - return "query refused by server" - case "notimp": - return "query type not implemented by server" - case "cname_loop": - return "resulted in a CNAME loop" - case "ns_error": - return "nameserver lookup failed" - case "error": + case traverse.StatusReferralLame: + return "resulted in a lame referral" + case traverse.StatusException: + return "resulted in an exception" + case traverse.StatusError: return "resulted in an error" - case "referral": - return "resulted in a referral" + case traverse.StatusNoGlue: + return "found no glue" + case traverse.StatusLoop: + return "resulted in a loop" + case traverse.StatusCNAMELoop: + return "resulted in a CNAME loop" default: - return respType + return string(status) } } -func trimDomain(name string) string { - return strings.TrimSuffix(name, ".") -} - -func collectServers(results []traverse.TraversalResult) map[string][]string { - servers := make(map[string][]string) - for _, result := range results { - if result.Response == nil || result.Response.Server == nil { - continue - } - name := serverName(result) - ip := result.Response.Server.String() - if containsString(servers[name], ip) { - continue - } - servers[name] = append(servers[name], ip) - } - return servers -} - -func serverName(result traverse.TraversalResult) string { - if result.Referral != nil && result.Referral.Bailiwick != "" && result.Referral.Bailiwick != "." { - return trimDomain(result.Referral.Bailiwick) - } - if result.Referral != nil && result.Referral.NSName != "" { - return result.Referral.NSName - } - if result.Response != nil && result.Response.Server != nil { - return result.Response.Server.String() - } - return "unknown" -} - -func containsString(items []string, target string) bool { - for _, item := range items { - if item == target { - return true - } - } - return false -} - -// DeduplicateResults collapses terminal results that represent the same -// outcome from the same server into a single entry with summed probability. -// This prevents the same nameserver failure (or answer) from appearing once -// per delegation path when several parent servers all refer to the same child. -func DeduplicateResults(results []traverse.TraversalResult) []traverse.TraversalResult { - type entry struct { - result traverse.TraversalResult - prob float64 - } - keys := make(map[string]*entry) - var order []string - - for _, r := range results { - if r.Response == nil || r.Referral == nil { - continue - } - key := resultDeduplicationKey(r) - if e, ok := keys[key]; ok { - e.prob += r.Referral.Prob - } else { - keys[key] = &entry{result: r, prob: r.Referral.Prob} - order = append(order, key) - } - } - - deduped := make([]traverse.TraversalResult, 0, len(order)) - for _, key := range order { - e := keys[key] - refCopy := *e.result.Referral - refCopy.Prob = e.prob - deduped = append(deduped, traverse.TraversalResult{ - Referral: &refCopy, - Response: e.result.Response, - }) - } - return deduped -} - -func resultDeduplicationKey(r traverse.TraversalResult) string { - bailiwick := strings.TrimSuffix(r.Referral.Bailiwick, ".") - switch r.Response.Type { - case traverse.RespAnswer: - key, _ := answerKey(r.Response) - return "answer:" + bailiwick + ":" + key - case traverse.RespNSResolutionFailed: - return "ns_error:" + r.Response.ErrorMessage - default: - return r.Response.Type.String() + ":" + bailiwick + ":" + r.Response.ErrorMessage - } +// collapseWhitespace renders an RR on one line with runs of whitespace +// collapsed to single spaces (summary_stats.rb text). +func collapseWhitespace(s string) string { + return strings.Join(strings.Fields(s), " ") } diff --git a/internal/output/stats_test.go b/internal/output/stats_test.go deleted file mode 100644 index 01c21ae..0000000 --- a/internal/output/stats_test.go +++ /dev/null @@ -1,327 +0,0 @@ -package output - -import ( - "net" - "testing" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" -) - -func makeAnswerResult(name string, ip string, prob float64) traverse.TraversalResult { - ref := traverse.NewReferral(name, dns.TypeA, ".", 0, prob, nil) - server := net.ParseIP("198.41.0.4") - resp := &traverse.Response{ - Referral: ref, - Server: server, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: name + ".", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP(ip), - }, - }, - }, - } - return traverse.TraversalResult{Referral: ref, Response: resp} -} - -func TestRRDataStringAllTypes(t *testing.T) { - tests := []struct { - rr miekgdns.RR - want string - }{ - { - &miekgdns.A{Hdr: miekgdns.RR_Header{Rrtype: dns.TypeA}, A: net.ParseIP("1.2.3.4")}, - "1.2.3.4", - }, - { - &miekgdns.AAAA{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeAAAA}, AAAA: net.ParseIP("::1")}, - "::1", - }, - { - &miekgdns.CNAME{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeCNAME}, Target: "example.com."}, - "example.com.", - }, - { - &miekgdns.NS{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeNS}, Ns: "ns1.example.com."}, - "ns1.example.com.", - }, - { - &miekgdns.MX{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeMX}, Preference: 10, Mx: "mail.example.com."}, - "10 mail.example.com.", - }, - { - &miekgdns.TXT{Hdr: miekgdns.RR_Header{Rrtype: miekgdns.TypeTXT}, Txt: []string{"v=spf1", "include:example.com"}}, - "v=spf1 include:example.com", - }, - } - - for _, tc := range tests { - got := rrDataString(tc.rr) - if got != tc.want { - t.Errorf("rrDataString(%T) = %q, want %q", tc.rr, got, tc.want) - } - } -} - -func TestRRDataStringDefault(t *testing.T) { - // SOA record hits the default case - rr := &miekgdns.SOA{ - Hdr: miekgdns.RR_Header{Name: ".", Rrtype: miekgdns.TypeSOA, Class: miekgdns.ClassINET}, - Ns: "a.root-servers.net.", - Mbox: "nstld.verisign-grs.com.", - } - got := rrDataString(rr) - if got == "" { - t.Error("rrDataString(SOA) should return non-empty string via default case") - } -} - -func TestSummaryTypeLabelAllTypes(t *testing.T) { - cases := map[string]string{ - "nodata": "found no such record", - "nxdomain": "name does not exist", - "servfail": "resulted in SERVFAIL", - "refused": "query refused by server", - "notimp": "query type not implemented by server", - "cname_loop": "resulted in a CNAME loop", - "ns_error": "nameserver lookup failed", - "error": "resulted in an error", - "referral": "resulted in a referral", - "unknown_type": "unknown_type", - } - for input, want := range cases { - got := summaryTypeLabel(input) - if got != want { - t.Errorf("summaryTypeLabel(%q) = %q, want %q", input, got, want) - } - } -} - -func TestCollectServersEmpty(t *testing.T) { - servers := collectServers(nil) - if len(servers) != 0 { - t.Errorf("collectServers(nil) = %v, want empty", servers) - } -} - -func TestCollectServersDeduplication(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - server := net.ParseIP("1.2.3.4") - resp := &traverse.Response{ - Referral: ref, - Server: server, - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - - servers := collectServers([]traverse.TraversalResult{result, result}) - name := "com" - ips := servers[name] - if len(ips) != 1 { - t.Errorf("expected deduplication: got %d IPs, want 1", len(ips)) - } -} - -func TestCollectServersWithBailiwick(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) - server := net.ParseIP("1.2.3.4") - resp := &traverse.Response{ - Referral: ref, - Server: server, - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - - servers := collectServers([]traverse.TraversalResult{result}) - if len(servers) == 0 { - t.Fatal("expected at least one server entry") - } - if _, ok := servers["com"]; !ok { - t.Errorf("expected server name 'com', got keys: %v", servers) - } -} - -func TestServerNameFallbacks(t *testing.T) { - // No bailiwick, no NSName, with server IP - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - name := serverName(result) - if name != "1.2.3.4" { - t.Errorf("serverName with root bailiwick = %q, want IP", name) - } -} - -func TestServerNameWithNSName(t *testing.T) { - ref := &traverse.Referral{ - Name: "example.com.", - Qtype: dns.TypeA, - Bailiwick: ".", - NSName: "ns1.example.com.", - } - resp := &traverse.Response{ - Referral: ref, - Server: net.ParseIP("5.5.5.5"), - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: ref, Response: resp} - // Bailiwick is "." so falls through to NSName - name := serverName(result) - if name == "" { - t.Error("serverName should return non-empty string") - } -} - -func TestServerNameNilReferral(t *testing.T) { - resp := &traverse.Response{ - Server: net.ParseIP("1.2.3.4"), - Type: traverse.RespAnswer, - } - result := traverse.TraversalResult{Referral: nil, Response: resp} - name := serverName(result) - if name == "" { - t.Error("serverName with nil referral should return non-empty string") - } -} - -func TestContainsString(t *testing.T) { - items := []string{"a", "b", "c"} - if !containsString(items, "b") { - t.Error("containsString should find 'b' in slice") - } - if containsString(items, "d") { - t.Error("containsString should not find 'd' in slice") - } - if containsString(nil, "a") { - t.Error("containsString on nil slice should return false") - } -} - -func TestComputeSummaryMixedResults(t *testing.T) { - refAnswer := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 0.6, nil) - respAnswer := &traverse.Response{ - Referral: refAnswer, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }, - }, - }, - } - - refNXD := traverse.NewReferral("notexist.com.", dns.TypeA, ".", 0, 0.4, nil) - respNXD := &traverse.Response{ - Referral: refNXD, - Type: traverse.RespNXDOMAIN, - } - - results := []traverse.TraversalResult{ - {Referral: refAnswer, Response: respAnswer}, - {Referral: refNXD, Response: respNXD}, - } - - stats := ComputeSummary(results) - if stats == nil { - t.Fatal("ComputeSummary returned nil for non-empty results") - } - if len(stats.Answers) != 1 { - t.Errorf("expected 1 answer entry, got %d", len(stats.Answers)) - } - if _, ok := stats.ByType["nxdomain"]; !ok { - t.Error("expected nxdomain in ByType") - } -} - -func TestComputeSummaryAnswerWithCNAMEOnly(t *testing.T) { - // Answer with only CNAME record - no final answer, should be in ByType - ref := traverse.NewReferral("www.example.com.", dns.TypeA, ".", 0, 1.0, nil) - resp := &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.CNAME{ - Hdr: miekgdns.RR_Header{Name: "www.example.com.", Rrtype: miekgdns.TypeCNAME, Class: miekgdns.ClassINET}, - Target: "example.com.", - }, - }, - }, - } - results := []traverse.TraversalResult{{Referral: ref, Response: resp}} - stats := ComputeSummary(results) - if stats == nil { - t.Fatal("ComputeSummary returned nil") - } -} - -func TestComputeSummaryAccumulates(t *testing.T) { - // Two answers with the same IP should accumulate probability - ref1 := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 0.5, nil) - ref2 := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 0.5, nil) - - makeResp := func(ref *traverse.Referral) *traverse.Response { - return &traverse.Response{ - Referral: ref, - Type: traverse.RespAnswer, - Decoded: &dns.DecodedResponse{ - Answers: []miekgdns.RR{ - &miekgdns.A{ - Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }, - }, - }, - } - } - - results := []traverse.TraversalResult{ - {Referral: ref1, Response: makeResp(ref1)}, - {Referral: ref2, Response: makeResp(ref2)}, - } - stats := ComputeSummary(results) - if stats == nil { - t.Fatal("ComputeSummary returned nil") - } - if len(stats.Answers) != 1 { - t.Fatalf("expected 1 answer after accumulation, got %d", len(stats.Answers)) - } - if stats.Answers[0].Prob < 0.99 { - t.Errorf("accumulated prob = %.2f, want ~1.0", stats.Answers[0].Prob) - } -} - -func TestCollectUniqueServerIPs(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - ip1 := net.ParseIP("1.2.3.4") - ip2 := net.ParseIP("5.6.7.8") - - results := []traverse.TraversalResult{ - {Referral: ref, Response: &traverse.Response{Server: ip1, Type: traverse.RespAnswer}}, - {Referral: ref, Response: &traverse.Response{Server: ip1, Type: traverse.RespAnswer}}, // dup - {Referral: ref, Response: &traverse.Response{Server: ip2, Type: traverse.RespAnswer}}, - {Referral: ref, Response: nil}, // nil response - } - - ips := collectUniqueServerIPs(results) - if len(ips) != 2 { - t.Errorf("expected 2 unique IPs, got %d", len(ips)) - } -} - -func TestCollectUniqueServerIPsEmpty(t *testing.T) { - ips := collectUniqueServerIPs(nil) - if len(ips) != 0 { - t.Errorf("expected 0 IPs for nil results, got %d", len(ips)) - } -} diff --git a/internal/output/text.go b/internal/output/text.go index c7482d6..dd7b18b 100644 --- a/internal/output/text.go +++ b/internal/output/text.go @@ -6,7 +6,6 @@ import ( "sort" "strings" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" ) @@ -19,52 +18,143 @@ func newTextFormatter(cfg *Config, w io.Writer) *textFormatter { return &textFormatter{cfg: cfg, w: w} } +// WriteHeader renders the pre-run header block (bin/dnstraverse): the "#" +// settings lines, the initial root, and the "Running query" line. --quiet +// suppresses the whole block. The EDNS0 state reflects the UDP size (the Ruby +// source always printed "on" — a documented deviation we fix). +func (f *textFormatter) WriteHeader(roots []traverse.StartServer) error { + if f.cfg.Quiet { + return nil + } + var b strings.Builder + if f.cfg.Fast { + b.WriteString("# Using fast mode\n") + } + if !f.cfg.AllRootServers { + b.WriteString("# Limiting traverse to one root\n") + } + edns := "on" + if f.cfg.UDPSize <= 512 { + edns = "off" + } + fmt.Fprintf(&b, "# UDP size %d (EDNS0 is %s)\n", f.cfg.UDPSize, edns) + fmt.Fprintf(&b, "# Retries %d, max depth %d\n", f.cfg.Retries, f.cfg.MaxDepth) + fmt.Fprintf(&b, "# Allow TCP is %t, always TCP is %t\n", f.cfg.AllowTCP, f.cfg.AlwaysTCP) + if len(roots) > 0 { + ip := "" + if len(roots[0].IPs) > 0 { + ip = roots[0].IPs[0] + } + fmt.Fprintf(&b, "Using %s (%s) as initial root\n", roots[0].Name, ip) + if f.cfg.AllRootServers { + b.WriteString("All roots:\n") + for _, root := range roots { + fmt.Fprintf(&b, " %s %s\n", root.Name, strings.Join(root.IPs, ", ")) + } + } + } + fmt.Fprintf(&b, "Running query %s type %s\n", f.cfg.Domain, f.cfg.QueryType) + _, err := io.WriteString(f.w, b.String()) + return err +} + func (f *textFormatter) WriteProgress(event traverse.TraversalEvent) error { - if event.Stage != traverse.EventStart { + ref := event.Referral + if ref == nil || ref.IsRootRoot() { return nil } - line := f.formatReferralLine(event.Result, false) - if !event.Result.Referral.HasAddresses() { - line += " -- resolving" + switch event.Stage { + case traverse.StageStart: + line := f.referralTxt(ref) + if !ref.Resolved() { + line += " -- resolving" + } + return f.writeLine(line) + case traverse.StageAnswerFast: + return f.writeLine(fmt.Sprintf("%s -- completed earlier (%s)", + f.referralTxt(ref), event.CompletedEarlier)) + case traverse.StageNewReferralSet: + // One line per extra childset: the parent refid and the IP that + // produced the children (progress_main :new_referral_set). + refid := event.RefID + if i := strings.LastIndex(refid, "."); i >= 0 { + refid = refid[:i] + } + return f.writeLine(fmt.Sprintf("%s %s", refid, ref.ParentIP)) + case traverse.StageAnswer: + if f.cfg.Verbose { + for _, warning := range ref.Warnings { + if err := f.writeLine(fmt.Sprintf("%s WARNING: %s", event.RefID, warning)); err != nil { + return err + } + } + } + if f.cfg.ShowAllStats { + return f.writeStatsBlocks(ref, fmt.Sprintf("%s Results:", event.RefID), false) + } } - return f.writeLine(line) + return nil } +// WriteResolve renders progress for glue-resolution subtree nodes +// (progress_resolves): like main progress, but the fast-mode marker carries +// no refid. func (f *textFormatter) WriteResolve(event traverse.TraversalEvent) error { - if event.Stage != traverse.EventStart { + ref := event.Referral + if ref == nil || ref.IsRootRoot() { return nil } - return f.writeLine(f.formatReferralLine(event.Result, true)) -} - -func (f *textFormatter) WriteResult(result traverse.TraversalResult) error { - if result.Response == nil || result.Referral == nil { - return nil + switch event.Stage { + case traverse.StageStart: + return f.writeLine(f.referralTxt(ref)) + case traverse.StageAnswerFast: + return f.writeLine(f.referralTxt(ref) + " -- completed earlier") } - prefix := strings.Repeat(" ", result.Referral.Depth+1) - line := prefix + f.formatResultLine(result) - return f.writeLine(line) + return nil } -func (f *textFormatter) WriteSummary(results []traverse.TraversalResult) error { +// referralTxt renders one progress row: " ()"; verbose +// adds "[qname]" and "". Unresolved servers have no parens +// (referral_txt_normal / referral_txt_verbose in bin/dnstraverse). +func (f *textFormatter) referralTxt(ref *traverse.Referral) string { + var b strings.Builder + b.WriteString(ref.RefID) + if f.cfg.Verbose { + fmt.Fprintf(&b, " [%s]", ref.Qname) + } + b.WriteString(" " + ref.Server) + if ref.Resolved() { + fmt.Fprintf(&b, " (%s)", ref.TxtIPs()) + } + if f.cfg.Verbose { + fmt.Fprintf(&b, " <%s>", ref.Bailiwick) + } + return b.String() +} + +func (f *textFormatter) WriteSummary(root *traverse.Referral, servers map[string][]string) error { + // Blank line separating progress from the sections (bin/dnstraverse: + // "puts if options[:progress]"). + if f.cfg.ShowProgress { + if _, err := fmt.Fprintln(f.w); err != nil { + return err + } + } if f.cfg.ShowServers { - if err := f.writeServers(results); err != nil { + if err := f.writeServers(servers); err != nil { return err } } - if f.cfg.ShowResults { - if err := f.writeResults(results); err != nil { + if err := f.writeResults(root); err != nil { return err } } - if f.cfg.ShowSummaryResults { - if err := f.writeSummaryResults(results); err != nil { + if err := f.writeSummaryResults(root); err != nil { return err } } - return nil } @@ -72,8 +162,10 @@ func (f *textFormatter) Flush() error { return nil } -func (f *textFormatter) writeServers(results []traverse.TraversalResult) error { - servers := collectServers(results) +// writeServers renders "The following servers were encountered:" sorted by +// lowercased reversed name (bin/dnstraverse); the name column is at least 16 +// characters wide. +func (f *textFormatter) writeServers(servers map[string][]string) error { if len(servers) == 0 { return nil } @@ -83,29 +175,26 @@ func (f *textFormatter) writeServers(results []traverse.TraversalResult) error { } names := make([]string, 0, len(servers)) + width := 16 for name := range servers { names = append(names, name) - } - sort.Slice(names, func(i, j int) bool { - return strings.ToLower(names[i]) > strings.ToLower(names[j]) - }) - - width := 16 - for _, name := range names { if len(name) > width { width = len(name) } } + sort.Slice(names, func(i, j int) bool { + return reverseString(strings.ToLower(names[i])) < reverseString(strings.ToLower(names[j])) + }) for _, name := range names { for _, ip := range servers[name] { line := fmt.Sprintf("%*s: %-15s", width, name, ip) if f.cfg.ShowVersions { if version, ok := f.cfg.Fingerprints[ip]; ok && version != "" { - line += " " + version + line += " " + version } } - if _, err := fmt.Fprintln(f.w, line); err != nil { + if _, err := fmt.Fprintln(f.w, strings.TrimRight(line, " ")); err != nil { return err } } @@ -114,129 +203,167 @@ func (f *textFormatter) writeServers(results []traverse.TraversalResult) error { return err } -func (f *textFormatter) writeResults(results []traverse.TraversalResult) error { +func reverseString(s string) string { + runes := []rune(s) + for i, j := 0, len(runes)-1; i < j; i, j = i+1, j-1 { + runes[i], runes[j] = runes[j], runes[i] + } + return string(runes) +} + +func (f *textFormatter) writeResults(root *traverse.Referral) error { + if root == nil { + return nil + } if _, err := fmt.Fprintln(f.w, "Results:"); err != nil { return err } - - terminal := terminalResults(results) - deduped := DeduplicateResults(terminal) - for _, result := range deduped { - prefix := strings.Repeat(" ", result.Referral.Depth+1) - line := prefix + f.formatResultLine(result) - if _, err := fmt.Fprintln(f.w, line); err != nil { - return err - } + if err := f.writeStatsBlocks(root, "", true); err != nil { + return err } _, err := fmt.Fprintln(f.w) return err } -func (f *textFormatter) writeSummaryResults(results []traverse.TraversalResult) error { - stats := ComputeSummary(results) +// writeStatsBlocks renders every aggregated leaf of ref, sorted by stats key +// (referral.rb stats_display). With spacing, blocks are separated by blank +// lines. +func (f *textFormatter) writeStatsBlocks(ref *traverse.Referral, prefix string, spacing bool) error { + first := true + for _, leaf := range ref.StatsList() { + if spacing && !first { + if _, err := fmt.Fprintln(f.w); err != nil { + return err + } + } + first = false + for _, line := range f.formatLeaf(leaf, prefix) { + if err := f.writeLine(line); err != nil { + return err + } + } + } + return nil +} + +func (f *textFormatter) writeSummaryResults(root *traverse.Referral) error { + stats := root.SummaryStats() if stats == nil { return nil } - if _, err := fmt.Fprintln(f.w, "Summary:"); err != nil { + if _, err := fmt.Fprintln(f.w, "Summary Results:"); err != nil { return err } prefix := " " for _, answer := range stats.Answers { - line := fmt.Sprintf("%s%s answered with %s", prefix, formatProbability(answer.Prob), answer.RData) + initial := fmt.Sprintf("%s%s answered with ", prefix, formatProbability(answer.Prob)) + var rrs []string + for _, rr := range answer.RRs { + rrs = append(rrs, collapseWhitespace(rr.String())) + } + line := initial + strings.Join(rrs, "\n"+strings.Repeat(" ", len(initial))) if _, err := fmt.Fprintln(f.w, f.colorize(line, colorGreen)); err != nil { return err } } - types := make([]string, 0, len(stats.ByType)) - for respType := range stats.ByType { - types = append(types, respType) + statuses := make([]traverse.Status, 0, len(stats.ByStatus)) + for status := range stats.ByStatus { + if status != traverse.StatusAnswered { + statuses = append(statuses, status) + } } - sort.Strings(types) + sort.Slice(statuses, func(i, j int) bool { return statuses[i] < statuses[j] }) - for _, respType := range types { - line := fmt.Sprintf("%s%s %s", prefix, formatProbability(stats.ByType[respType]), summaryTypeLabel(respType)) + for _, status := range statuses { + line := fmt.Sprintf("%s%s %s", prefix, formatProbability(stats.ByStatus[status]), summaryStatusLabel(status)) if _, err := fmt.Fprintln(f.w, line); err != nil { return err } } - _, err := fmt.Fprintln(f.w) - return err + return nil } -func (f *textFormatter) formatReferralLine(result traverse.TraversalResult, isResolve bool) string { - ref := result.Referral - if ref == nil { - return "" +// formatLeaf renders one aggregated leaf per referral.rb stats_display: +// "%5.1f%%: " plus indented RRs for answers and the +// "While querying" line when the failing query differs from the original. +func (f *textFormatter) formatLeaf(leaf *traverse.StatsEntry, prefix string) []string { + resp := leaf.Response + ref := leaf.Referral + if resp == nil || ref == nil { + return nil } + indent := prefix + strings.Repeat(" ", 12) + where := fmt.Sprintf("%s (%s)", ref.Server, resp.IP) + head := fmt.Sprintf("%s%5.1f%%: ", prefix, leaf.Prob*100) - indent := strings.Repeat(" ", ref.Depth) - refID := referralID(ref) - server := referralServerLabel(ref, result.Response) - qtype := dns.QNameType(ref.Qtype) - qname := trimDomain(ref.Name) - - if f.cfg.Verbose { - bailiwick := trimDomain(ref.Bailiwick) - if isResolve { - return fmt.Sprintf("%s%s [%s] %s <%s>", indent, refID, qname, server, bailiwick) + var lines []string + switch resp.Status { + case traverse.StatusException: + msg := "" + if resp.DQ != nil { + msg = resp.DQ.ExceptionMessage } - return fmt.Sprintf("%s%s [%s] %s <%s> (%s)", indent, refID, qname, server, bailiwick, qtype) - } - - if isResolve { - return fmt.Sprintf("%s%s %s", indent, refID, server) - } - return fmt.Sprintf("%s%s %s (%s)", indent, refID, server, qtype) -} - -func (f *textFormatter) formatResultLine(result traverse.TraversalResult) string { - prob := formatProbability(result.Referral.Prob) - switch result.Response.Type { - case traverse.RespAnswer: - key, _ := answerKey(result.Response) - if key == "" { - return fmt.Sprintf("%s resulted in answer", prob) + lines = append(lines, head+f.colorize(fmt.Sprintf("%s at %s", msg, where), colorRed)) + case traverse.StatusNoGlue: + parent := "" + if ref.Parent != nil { + parent = ref.Parent.Server } - nsLabel := "" - if result.Referral.Bailiwick != "" && result.Referral.Bailiwick != "." { - nsLabel = trimDomain(result.Referral.Bailiwick) + " " + lines = append(lines, head+f.colorize(fmt.Sprintf("No glue at %s (%s) for %s", parent, resp.IP, ref.Server), colorYellow)) + case traverse.StatusReferralLame: + parent := "" + if ref.Parent != nil { + parent = ref.Parent.Server } - return f.colorize(fmt.Sprintf("%s %sanswered with %s", prob, nsLabel, key), colorGreen) - case traverse.RespNODATA: - return fmt.Sprintf("%s found no such record", prob) - case traverse.RespNXDOMAIN: - return f.colorize(fmt.Sprintf("%s name does not exist", prob), colorYellow) - case traverse.RespSERVFAIL: - return f.colorize(fmt.Sprintf("%s resulted in SERVFAIL", prob), colorRed) - case traverse.RespREFUSED: - return f.colorize(fmt.Sprintf("%s query refused by server", prob), colorRed) - case traverse.RespNOTIMPL: - return f.colorize(fmt.Sprintf("%s query type not implemented by server", prob), colorRed) - case traverse.RespCNAMELoop: - msg := "CNAME loop detected" - if result.Response.ErrorMessage != "" { - msg = result.Response.ErrorMessage + lines = append(lines, head+f.colorize(fmt.Sprintf("Lame referral from %s (%s) to %s", parent, ref.ParentIP, where), colorYellow)) + case traverse.StatusLoop: + lines = append(lines, head+f.colorize(fmt.Sprintf("Loop encountered at %s", resp.Server), colorRed)) + case traverse.StatusCNAMELoop: + lines = append(lines, head+f.colorize(fmt.Sprintf("CNAME loop encountered at %s", resp.Server), colorRed)) + case traverse.StatusError: + msg := "" + if resp.DQ != nil { + msg = resp.DQ.ErrorMessage } - return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorRed) - case traverse.RespNSResolutionFailed: - msg := "nameserver lookup failed" - if result.Response.ErrorMessage != "" { - msg = result.Response.ErrorMessage + lines = append(lines, head+f.colorize(fmt.Sprintf("%s at %s", msg, where), colorRed)) + case traverse.StatusNoData: + lines = append(lines, head+fmt.Sprintf("NODATA (for this type) at %s", where)) + case traverse.StatusAnswered: + lines = append(lines, head+f.colorize(fmt.Sprintf("Answer from %s", where), colorGreen)) + if resp.DQ != nil { + for _, rr := range resp.DQ.Answers { + lines = append(lines, indent+rr.String()) + } } - return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorYellow) - case traverse.RespError: - msg := "resulted in an error" - if result.Response.ErrorMessage != "" { - msg = result.Response.ErrorMessage - } - return f.colorize(fmt.Sprintf("%s %s", prob, msg), colorRed) default: - return fmt.Sprintf("%s %s", prob, result.Response.Type) + // The Ruby fallback prints "Stopped at ())" with a stray + // paren — a documented deviation we fix. + lines = append(lines, head+fmt.Sprintf("Stopped at %s", where)) + lines = append(lines, indent+leaf.Key) } + + if resp.Status != traverse.StatusAnswered { + origQname, origQclass, origQtype := originalQuery(ref) + if resp.Qname != origQname || resp.Qclass != origQclass || resp.Qtype != origQtype { + lines = append(lines, indent+fmt.Sprintf("While querying %s/%s/%s", + resp.Qname, traverse.ClassToString(resp.Qclass), traverse.TypeToString(resp.Qtype))) + } + } + return lines +} + +// originalQuery walks to the rootroot node to find the query the whole +// traversal was started for. +func originalQuery(ref *traverse.Referral) (string, uint16, uint16) { + top := ref + for top.Parent != nil { + top = top.Parent + } + return top.Qname, top.Qclass, top.Qtype } func (f *textFormatter) writeLine(line string) error { @@ -254,43 +381,6 @@ func (f *textFormatter) colorize(text, color string) string { return color + text + colorReset } -func referralID(ref *traverse.Referral) string { - if ref == nil { - return "" - } - return fmt.Sprintf("%d", ref.Depth+1) -} - -func referralServerLabel(ref *traverse.Referral, resp *traverse.Response) string { - if resp != nil && resp.Server != nil { - return resp.Server.String() - } - if ref.HasAddresses() { - ips := make([]string, 0, len(ref.Addresses)) - for _, addr := range ref.Addresses { - ips = append(ips, addr.String()) - } - return strings.Join(ips, ", ") - } - if ref.NSName != "" { - return ref.NSName - } - if ref.Bailiwick != "" && ref.Bailiwick != "." { - return trimDomain(ref.Bailiwick) - } - return "unknown" -} - -func terminalResults(results []traverse.TraversalResult) []traverse.TraversalResult { - var terminal []traverse.TraversalResult - for _, result := range results { - if result.Response != nil && result.Response.IsTerminal() { - terminal = append(terminal, result) - } - } - return terminal -} - const ( colorReset = "\033[0m" colorGreen = "\033[32m" diff --git a/internal/output/text_test.go b/internal/output/text_test.go deleted file mode 100644 index c11342f..0000000 --- a/internal/output/text_test.go +++ /dev/null @@ -1,379 +0,0 @@ -package output - -import ( - "bytes" - "net" - "strings" - "testing" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" - miekgdns "github.com/miekg/dns" -) - -func TestTextFormatterProgressIndentation(t *testing.T) { - root := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - child := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 0.5, root) - - var buf bytes.Buffer - cfg := DefaultConfig() - cfg.Color = false - formatter := newTextFormatter(cfg, &buf) - - if err := formatter.WriteProgress(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: child}, - }); err != nil { - t.Fatalf("WriteProgress: %v", err) - } - - out := buf.String() - if !strings.HasPrefix(out, " 2 ") { - t.Fatalf("expected depth-based indentation, got %q", out) - } -} - -func TestAttachHooksRespectsShowFlags(t *testing.T) { - ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) - var progressCount int - - cfg := DefaultConfig() - cfg.ShowProgress = false - cfg.ShowResolves = false - cfg.ShowAllStats = false - - var buf bytes.Buffer - formatter := NewFormatter(cfg, &buf) - hooks := AttachHooks(cfg, formatter) - hooks.OnEvent(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - - if buf.Len() != 0 { - t.Fatalf("expected no output when ShowProgress is false") - } - - cfg.ShowProgress = true - hooks = AttachHooks(cfg, formatter) - hooks.OnEvent(traverse.TraversalEvent{ - Stage: traverse.EventStart, - Result: traverse.TraversalResult{Referral: ref}, - }) - progressCount = strings.Count(buf.String(), "\n") - if progressCount == 0 { - t.Fatal("expected progress output when ShowProgress is true") - } -} - -func TestTextWriteResolve(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("198.41.0.4") -resp := &traverse.Response{Server: server, Type: traverse.RespAnswer} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -f := newTextFormatter(cfg, &buf) - -// EventStart - should write line -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteResolve EventStart: %v", err) -} -if buf.Len() == 0 { -t.Error("expected output for WriteResolve EventStart") -} - -buf.Reset() -// EventComplete - should write nothing -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventComplete, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteResolve EventComplete: %v", err) -} -if buf.Len() != 0 { -t.Error("expected no output for WriteResolve EventComplete") -} -} - -func TestTextWriteResult(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -server := net.ParseIP("198.41.0.4") - -tests := []struct { -name string -respType traverse.ResponseType -msg *dns.DecodedResponse -errorMsg string -}{ -{"answer", traverse.RespAnswer, &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("1.2.3.4")}, -}, -}, ""}, -{"nodata", traverse.RespNODATA, nil, ""}, -{"nxdomain", traverse.RespNXDOMAIN, nil, ""}, -{"servfail", traverse.RespSERVFAIL, nil, ""}, -{"refused", traverse.RespREFUSED, nil, ""}, -{"notimp", traverse.RespNOTIMPL, nil, ""}, -{"cname_loop", traverse.RespCNAMELoop, nil, "loop detected"}, -{"ns_error", traverse.RespNSResolutionFailed, nil, "nameserver ns1.example.com could not be resolved"}, -{"error", traverse.RespError, nil, "something went wrong"}, -} - -for _, tc := range tests { -t.Run(tc.name, func(t *testing.T) { -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -f := newTextFormatter(cfg, &buf) - -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: tc.respType, -Decoded: tc.msg, -ErrorMessage: tc.errorMsg, -} -result := traverse.TraversalResult{Referral: ref, Response: resp} -if err := f.WriteResult(result); err != nil { -t.Fatalf("WriteResult %q: %v", tc.name, err) -} -}) -} -} - -func TestTextWriteResultNilResponse(t *testing.T) { -var buf bytes.Buffer -cfg := DefaultConfig() -f := newTextFormatter(cfg, &buf) -if err := f.WriteResult(traverse.TraversalResult{Referral: nil, Response: nil}); err != nil { -t.Fatalf("WriteResult nil: %v", err) -} -if buf.Len() != 0 { -t.Error("expected no output for nil result") -} -} - -func TestTextWriteResultAnswerMultipleRRs(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 0.5, nil) -resp := &traverse.Response{ -Referral: ref, -Server: net.ParseIP("1.2.3.4"), -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("1.2.3.4")}, -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("5.6.7.8")}, -}, -}, -} -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -f := newTextFormatter(cfg, &buf) -if err := f.WriteResult(traverse.TraversalResult{Referral: ref, Response: resp}); err != nil { -t.Fatalf("WriteResult: %v", err) -} -if !strings.Contains(buf.String(), "/") { -t.Errorf("expected '/' separator for multiple answers, got: %q", buf.String()) -} -} - -func TestTextWriteSummaryWithServersAndResults(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespAnswer, -Decoded: &dns.DecodedResponse{ -Answers: []miekgdns.RR{ -&miekgdns.A{Hdr: miekgdns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: miekgdns.ClassINET}, A: net.ParseIP("1.2.3.4")}, -}, -}, -} -results := []traverse.TraversalResult{{Referral: ref, Response: resp}} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -cfg.ShowServers = true -cfg.ShowResults = true -cfg.ShowSummaryResults = true -f := newTextFormatter(cfg, &buf) -if err := f.WriteSummary(results); err != nil { -t.Fatalf("WriteSummary: %v", err) -} -out := buf.String() -if !strings.Contains(out, "Summary:") { -t.Errorf("expected Summary: in output, got: %q", out) -} -if !strings.Contains(out, "Results:") { -t.Errorf("expected Results: in output, got: %q", out) -} -} - -func TestTextWriteSummaryNoResults(t *testing.T) { -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.ShowServers = false -cfg.ShowResults = false -cfg.ShowSummaryResults = false -f := newTextFormatter(cfg, &buf) -if err := f.WriteSummary(nil); err != nil { -t.Fatalf("WriteSummary nil: %v", err) -} -} - -func TestTextWriteSummaryNXDOMAIN(t *testing.T) { -ref := traverse.NewReferral("gone.example.com.", dns.TypeA, "com.", 1, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{ -Referral: ref, -Server: server, -Type: traverse.RespNXDOMAIN, -} -results := []traverse.TraversalResult{{Referral: ref, Response: resp}} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -cfg.ShowServers = true -cfg.ShowResults = true -cfg.ShowSummaryResults = true -f := newTextFormatter(cfg, &buf) -if err := f.WriteSummary(results); err != nil { -t.Fatalf("WriteSummary: %v", err) -} -} - -func TestFormatReferralLineVerbose(t *testing.T) { -root := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -child := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, root) - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -cfg.Verbose = true -f := newTextFormatter(cfg, &buf) - -event := traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: child}, -} -if err := f.WriteProgress(event); err != nil { -t.Fatalf("WriteProgress verbose: %v", err) -} -out := buf.String() -if !strings.Contains(out, "com") { -t.Errorf("expected bailiwick in verbose output, got: %q", out) -} -} - -func TestFormatReferralLineVerboseResolve(t *testing.T) { -ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 1, 1.0, nil) -server := net.ParseIP("1.2.3.4") -resp := &traverse.Response{Server: server, Type: traverse.RespAnswer} - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -cfg.Verbose = true -f := newTextFormatter(cfg, &buf) - -if err := f.WriteResolve(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref, Response: resp}, -}); err != nil { -t.Fatalf("WriteResolve verbose: %v", err) -} -if buf.Len() == 0 { -t.Error("expected output for verbose WriteResolve") -} -} - -func TestTextWriteProgressNoAddresses(t *testing.T) { -// Test the "resolving" suffix when referral has no addresses -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -// No addresses set, so HasAddresses() returns false - -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = false -f := newTextFormatter(cfg, &buf) -if err := f.WriteProgress(traverse.TraversalEvent{ -Stage: traverse.EventStart, -Result: traverse.TraversalResult{Referral: ref}, -}); err != nil { -t.Fatalf("WriteProgress: %v", err) -} -if !strings.Contains(buf.String(), "resolving") { -t.Errorf("expected 'resolving' suffix when no addresses, got: %q", buf.String()) -} -} - -func TestColorize(t *testing.T) { -var buf bytes.Buffer -cfg := DefaultConfig() -cfg.Color = true -f := newTextFormatter(cfg, &buf) - -colored := f.colorize("hello", colorGreen) -if colored == "hello" { -t.Error("expected colorized output with Color=true") -} - -cfg.Color = false -f2 := newTextFormatter(cfg, &buf) -plain := f2.colorize("hello", colorGreen) -if plain != "hello" { -t.Errorf("expected plain text with Color=false, got %q", plain) -} - -// Empty color -empty := f.colorize("hello", "") -if empty != "hello" { -t.Errorf("expected plain text for empty color, got %q", empty) -} -} - -func TestReferralServerLabelFallbacks(t *testing.T) { -// With server IP in response -ref := traverse.NewReferral("example.com.", dns.TypeA, ".", 0, 1.0, nil) -resp := &traverse.Response{Server: net.ParseIP("1.2.3.4")} -label := referralServerLabel(ref, resp) -if label != "1.2.3.4" { -t.Errorf("expected '1.2.3.4', got %q", label) -} - -// With addresses in referral, no response server -ref2 := traverse.NewReferral("example.com.", dns.TypeA, "ns1.example.com.", 0, 1.0, nil) -ref2.Addresses = []net.IP{net.ParseIP("5.6.7.8")} -label2 := referralServerLabel(ref2, nil) -if label2 != "5.6.7.8" { -t.Errorf("expected '5.6.7.8', got %q", label2) -} - -// With NSName -ref3 := &traverse.Referral{ -Name: "example.com.", -NSName: "ns1.example.com.", -Bailiwick: ".", -} -label3 := referralServerLabel(ref3, nil) -if label3 != "ns1.example.com." { -t.Errorf("expected NSName, got %q", label3) -} - -// With non-root bailiwick, no addresses, no NSName -ref4 := traverse.NewReferral("example.com.", dns.TypeA, "com.", 0, 1.0, nil) -label4 := referralServerLabel(ref4, nil) -if label4 != "com" { -t.Errorf("expected 'com', got %q", label4) -} -} diff --git a/internal/traverse/cache.go b/internal/traverse/cache.go index 9894324..b7ca65e 100644 --- a/internal/traverse/cache.go +++ b/internal/traverse/cache.go @@ -1,6 +1,7 @@ package traverse import ( + "fmt" "net" "strings" "sync" @@ -8,117 +9,157 @@ import ( miekgdns "github.com/miekg/dns" ) +// StartServer is one entry returned by GetStartServers: a nameserver hostname +// plus its cached IPv4 addresses. IPs == nil means no addresses are cached and +// the caller must resolve the name itself (glueless). +type StartServer struct { + Name string + IPs []string +} + +// InfoCache is the hierarchical per-branch record cache (info_cache.rb). Each +// response wraps its parent's cache in a child so sibling branches never see +// each other's records; lookups recurse towards the root cache. type InfoCache struct { parent *InfoCache mu sync.RWMutex - ns map[string][]string - glue map[string][]net.IP + data map[string][]miekgdns.RR } func NewInfoCache(parent *InfoCache) *InfoCache { return &InfoCache{ parent: parent, - ns: make(map[string][]string), - glue: make(map[string][]net.IP), + data: make(map[string][]miekgdns.RR), } } -func (c *InfoCache) StoreNS(zone string, nameservers []string) { - if len(nameservers) == 0 { - return - } - zone = normalize(zone) - c.mu.Lock() - seen := make(map[string]bool) - for _, ns := range nameservers { - ns = normalize(ns) - if !seen[ns] { - seen[ns] = true - c.ns[zone] = append(c.ns[zone], ns) - } - } - c.mu.Unlock() -} - -func (c *InfoCache) LookupNS(zone string) []string { - zone = normalize(zone) - if names := c.localNS(zone); len(names) > 0 { - return names - } - if c.parent != nil { - return c.parent.LookupNS(zone) - } - return nil -} - -func (c *InfoCache) localNS(zone string) []string { - c.mu.RLock() - defer c.mu.RUnlock() - names, ok := c.ns[zone] - if !ok { - return nil - } - result := make([]string, len(names)) - copy(result, names) - return result -} - -func (c *InfoCache) StoreGlue(name string, addrs []net.IP) { - if len(addrs) == 0 { - return - } - name = normalize(name) - c.mu.Lock() - seen := make(map[string]bool) - for _, addr := range addrs { - key := addr.String() - if !seen[key] { - seen[key] = true - c.glue[name] = append(c.glue[name], addr) - } - } - c.mu.Unlock() -} - -func (c *InfoCache) LookupGlue(name string) []net.IP { - name = normalize(name) - if addrs := c.localGlue(name); len(addrs) > 0 { - return addrs - } - if c.parent != nil { - return c.parent.LookupGlue(name) - } - return nil -} - -func (c *InfoCache) localGlue(name string) []net.IP { - c.mu.RLock() - defer c.mu.RUnlock() - addrs, ok := c.glue[name] - if !ok { - return nil - } - result := make([]net.IP, len(addrs)) - copy(result, addrs) - return result -} - func (c *InfoCache) Child() *InfoCache { return NewInfoCache(c) } -func (c *InfoCache) NSCount() int { - c.mu.RLock() - defer c.mu.RUnlock() - return len(c.ns) +// canonicalName lowercases a DNS name and strips the trailing dot; the root +// (and empty string) canonicalises to "", matching the Ruby engine's +// representation of "no bailiwick". +func canonicalName(name string) string { + return strings.TrimSuffix(strings.ToLower(name), ".") } -func (c *InfoCache) GlueCount() int { - c.mu.RLock() - defer c.mu.RUnlock() - return len(c.glue) +func cacheKey(name string, qclass, qtype uint16) string { + return fmt.Sprintf("%s:%d:%d", canonicalName(name), qclass, qtype) } -func normalize(name string) string { - return strings.ToLower(miekgdns.Fqdn(name)) +func rrCacheKey(rr miekgdns.RR) string { + h := rr.Header() + return cacheKey(h.Name, h.Class, h.Rrtype) +} + +// Add stores resource records, REPLACING any existing entries that share a +// name:class:type key (info_cache.rb add: clear pass, then append pass, so +// several records under one key in a single call are all kept). +func (c *InfoCache) Add(rrs []miekgdns.RR) { + c.mu.Lock() + defer c.mu.Unlock() + for _, rr := range rrs { + c.data[rrCacheKey(rr)] = nil + } + for _, rr := range rrs { + key := rrCacheKey(rr) + c.data[key] = append(c.data[key], rr) + } +} + +// AddHints seeds NS records for domain ("" = root hints) plus A/AAAA records +// for each server that has known addresses (info_cache.rb add_hints). +func (c *InfoCache) AddHints(domain string, servers []StartServer) { + var rrs []miekgdns.RR + owner := miekgdns.Fqdn(canonicalName(domain)) + for _, srv := range servers { + name := miekgdns.Fqdn(canonicalName(srv.Name)) + rrs = append(rrs, &miekgdns.NS{ + Hdr: miekgdns.RR_Header{Name: owner, Rrtype: miekgdns.TypeNS, Class: miekgdns.ClassINET}, + Ns: name, + }) + for _, ip := range srv.IPs { + addr := net.ParseIP(ip) + if addr == nil { + continue + } + if v4 := addr.To4(); v4 != nil { + rrs = append(rrs, &miekgdns.A{ + Hdr: miekgdns.RR_Header{Name: name, Rrtype: miekgdns.TypeA, Class: miekgdns.ClassINET}, + A: v4, + }) + } else { + rrs = append(rrs, &miekgdns.AAAA{ + Hdr: miekgdns.RR_Header{Name: name, Rrtype: miekgdns.TypeAAAA, Class: miekgdns.ClassINET}, + AAAA: addr, + }) + } + } + } + c.Add(rrs) +} + +// Get returns the cached RRset for name/class/type, consulting parent caches +// on a local miss. Returns nil when nothing is cached anywhere in the chain. +func (c *InfoCache) Get(name string, qclass, qtype uint16) []miekgdns.RR { + key := cacheKey(name, qclass, qtype) + c.mu.RLock() + rrs, ok := c.data[key] + c.mu.RUnlock() + if ok { + out := make([]miekgdns.RR, len(rrs)) + copy(out, rrs) + return out + } + if c.parent != nil { + return c.parent.Get(name, qclass, qtype) + } + return nil +} + +// getNS finds the nearest cached NS RRset at or above domain, walking labels +// upward to the root (info_cache.rb get_ns?). +func (c *InfoCache) getNS(domain string) ([]miekgdns.RR, error) { + domain = canonicalName(domain) + for { + if rrs := c.Get(domain, miekgdns.ClassINET, miekgdns.TypeNS); len(rrs) > 0 { + return rrs, nil + } + if domain == "" { + return nil, fmt.Errorf("no nameservers available for %q -- no root hints set??", domain) + } + if i := strings.Index(domain, "."); i >= 0 { + domain = domain[i+1:] + } else { + domain = "" + } + } +} + +// GetStartServers returns the servers to start querying for domain: the +// nearest cached NS RRset walking labels upward, each nameserver paired with +// its cached A addresses (nil when unknown). newbailiwick is the owner name of +// that NS RRset ("" for root). +func (c *InfoCache) GetStartServers(domain string) (starters []StartServer, newbailiwick string, err error) { + ns, err := c.getNS(domain) + if err != nil { + return nil, "", err + } + for _, rr := range ns { + nsrr, ok := rr.(*miekgdns.NS) + if !ok { + continue + } + name := canonicalName(nsrr.Ns) + var ips []string + for _, iprr := range c.Get(name, miekgdns.ClassINET, miekgdns.TypeA) { + if a, ok := iprr.(*miekgdns.A); ok { + ips = append(ips, a.A.String()) + } + } + starters = append(starters, StartServer{Name: name, IPs: ips}) + } + newbailiwick = canonicalName(ns[0].Header().Name) + return starters, newbailiwick, nil } diff --git a/internal/traverse/cache_test.go b/internal/traverse/cache_test.go index cc024fa..0b938ef 100644 --- a/internal/traverse/cache_test.go +++ b/internal/traverse/cache_test.go @@ -3,201 +3,258 @@ package traverse import ( "fmt" "net" - "strings" "sync" "testing" "github.com/miekg/dns" ) -func TestNewInfoCache(t *testing.T) { - c := NewInfoCache(nil) - if c.parent != nil { - t.Error("root cache should have nil parent") - } - if c.NSCount() != 0 { - t.Errorf("NSCount = %d, want 0", c.NSCount()) - } - if c.GlueCount() != 0 { - t.Errorf("GlueCount = %d, want 0", c.GlueCount()) +func nsRR(zone, target string) dns.RR { + return &dns.NS{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeNS, Class: dns.ClassINET}, + Ns: dns.Fqdn(target), } } -func TestInfoCacheStoreAndLookupNS(t *testing.T) { - c := NewInfoCache(nil) - - c.StoreNS("com.", []string{"a.gtld-servers.net.", "b.gtld-servers.net."}) - if c.NSCount() != 1 { - t.Errorf("NSCount = %d, want 1", c.NSCount()) - } - - names := c.LookupNS("com.") - if len(names) != 2 { - t.Fatalf("expected 2 nameservers, got %d", len(names)) - } - if names[0] != "a.gtld-servers.net." { - t.Errorf("nameserver[0] = %q, want %q", names[0], "a.gtld-servers.net.") +func aRR(name, ip string) dns.RR { + return &dns.A{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeA, Class: dns.ClassINET}, + A: net.ParseIP(ip).To4(), } } -func TestInfoCacheNSDedup(t *testing.T) { - c := NewInfoCache(nil) - c.StoreNS("com.", []string{"a.gtld-servers.net.", "a.gtld-servers.net."}) - names := c.LookupNS("com.") - if len(names) != 1 { - t.Errorf("expected 1 deduped NS, got %d", len(names)) +func aaaaRR(name, ip string) dns.RR { + return &dns.AAAA{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeAAAA, Class: dns.ClassINET}, + AAAA: net.ParseIP(ip), } } -func TestInfoCacheNSCaseInsensitive(t *testing.T) { - c := NewInfoCache(nil) - c.StoreNS("COM.", []string{"A.GTLD-SERVERS.NET."}) - names := c.LookupNS("com.") - if len(names) != 1 { - t.Fatalf("expected 1 NS, got %d", len(names)) +func TestCanonicalName(t *testing.T) { + tests := []struct{ in, want string }{ + {"example.com", "example.com"}, + {"Example.COM.", "example.com"}, + {".", ""}, + {"", ""}, + {"WWW.Example.Com", "www.example.com"}, } - if names[0] != "a.gtld-servers.net." { - t.Errorf("NS = %q, want %q", names[0], "a.gtld-servers.net.") + for _, tt := range tests { + if got := canonicalName(tt.in); got != tt.want { + t.Errorf("canonicalName(%q) = %q, want %q", tt.in, got, tt.want) + } } } -func TestInfoCacheNSLookupMiss(t *testing.T) { +func TestInfoCacheAddAndGet(t *testing.T) { c := NewInfoCache(nil) - names := c.LookupNS("org.") - if names != nil { - t.Errorf("expected nil for miss, got %v", names) + c.Add([]dns.RR{nsRR("com", "a.gtld-servers.net"), nsRR("com", "b.gtld-servers.net")}) + + rrs := c.Get("com", dns.ClassINET, dns.TypeNS) + if len(rrs) != 2 { + t.Fatalf("expected 2 NS records, got %d", len(rrs)) } } -func TestInfoCacheNSStoreEmpty(t *testing.T) { +func TestInfoCacheAddReplacesSameKey(t *testing.T) { c := NewInfoCache(nil) - c.StoreNS("com.", nil) - if c.NSCount() != 0 { - t.Errorf("expected 0 after empty store, got %d", c.NSCount()) + c.Add([]dns.RR{nsRR("com", "old1.example.net"), nsRR("com", "old2.example.net")}) + c.Add([]dns.RR{nsRR("com", "new.example.net")}) + + rrs := c.Get("com", dns.ClassINET, dns.TypeNS) + if len(rrs) != 1 { + t.Fatalf("add should replace same name:class:type entry, got %d records", len(rrs)) + } + if rrs[0].(*dns.NS).Ns != "new.example.net." { + t.Errorf("NS = %q, want new.example.net.", rrs[0].(*dns.NS).Ns) } } -func TestInfoCacheChainedNS(t *testing.T) { +func TestInfoCacheAddKeepsDistinctKeys(t *testing.T) { + c := NewInfoCache(nil) + c.Add([]dns.RR{nsRR("com", "a.gtld-servers.net"), aRR("a.gtld-servers.net", "192.5.6.30")}) + c.Add([]dns.RR{nsRR("org", "a0.org-servers.net")}) + + if got := c.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 { + t.Errorf("com NS lost after unrelated add: %v", got) + } + if got := c.Get("a.gtld-servers.net", dns.ClassINET, dns.TypeA); len(got) != 1 { + t.Errorf("glue lost after unrelated add: %v", got) + } +} + +func TestInfoCacheGetCaseInsensitive(t *testing.T) { + c := NewInfoCache(nil) + c.Add([]dns.RR{nsRR("COM.", "A.GTLD-SERVERS.NET.")}) + if got := c.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 { + t.Fatalf("expected case-insensitive hit, got %v", got) + } +} + +func TestInfoCacheGetMiss(t *testing.T) { + c := NewInfoCache(nil) + if got := c.Get("org", dns.ClassINET, dns.TypeNS); got != nil { + t.Errorf("expected nil for miss, got %v", got) + } +} + +func TestInfoCacheGetRecursesToParent(t *testing.T) { parent := NewInfoCache(nil) - parent.StoreNS("com.", []string{"a.gtld-servers.net."}) - + parent.Add([]dns.RR{nsRR("com", "a.gtld-servers.net")}) child := parent.Child() if child.parent != parent { - t.Error("child parent should be the parent cache") + t.Fatal("Child() should link to parent") } - - names := child.LookupNS("com.") - if len(names) != 1 { - t.Fatalf("expected 1 NS from parent, got %d", len(names)) - } - if child.NSCount() != 0 { - t.Errorf("child NSCount = %d, want 0", child.NSCount()) + if got := child.Get("com", dns.ClassINET, dns.TypeNS); len(got) != 1 { + t.Fatalf("expected parent hit through child, got %v", got) } } -func TestInfoCacheChildOverridesParent(t *testing.T) { +func TestInfoCacheChildShadowsParent(t *testing.T) { parent := NewInfoCache(nil) - parent.StoreNS("com.", []string{"a.gtld-servers.net."}) - + parent.Add([]dns.RR{nsRR("com", "parent.example.net")}) child := parent.Child() - child.StoreNS("com.", []string{"b.gtld-servers.net."}) + child.Add([]dns.RR{nsRR("com", "child.example.net")}) - names := child.LookupNS("com.") - if len(names) != 1 { - t.Fatalf("expected 1 NS, got %d", len(names)) + got := child.Get("com", dns.ClassINET, dns.TypeNS) + if len(got) != 1 || got[0].(*dns.NS).Ns != "child.example.net." { + t.Errorf("child entry should shadow parent, got %v", got) } - if names[0] != "b.gtld-servers.net." { - t.Errorf("expected child's NS to override, got %q", names[0]) + // The parent must be untouched. + got = parent.Get("com", dns.ClassINET, dns.TypeNS) + if len(got) != 1 || got[0].(*dns.NS).Ns != "parent.example.net." { + t.Errorf("parent entry modified, got %v", got) } } -func TestInfoCacheStoreAndLookupGlue(t *testing.T) { +func TestGetStartServersWalksLabels(t *testing.T) { c := NewInfoCache(nil) - addrs := []net.IP{net.ParseIP("1.2.3.4"), net.ParseIP("5.6.7.8")} - c.StoreGlue("ns1.example.com.", addrs) + c.Add([]dns.RR{ + nsRR("com", "a.gtld-servers.net"), + nsRR("com", "b.gtld-servers.net"), + aRR("a.gtld-servers.net", "192.5.6.30"), + }) - result := c.LookupGlue("ns1.example.com.") - if len(result) != 2 { - t.Fatalf("expected 2 glue addresses, got %d", len(result)) + starters, bw, err := c.GetStartServers("www.deep.example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "com" { + t.Errorf("newbailiwick = %q, want com", bw) + } + if len(starters) != 2 { + t.Fatalf("expected 2 starters, got %d", len(starters)) + } + if starters[0].Name != "a.gtld-servers.net" { + t.Errorf("starter[0] = %q", starters[0].Name) + } + if len(starters[0].IPs) != 1 || starters[0].IPs[0] != "192.5.6.30" { + t.Errorf("starter[0] IPs = %v, want [192.5.6.30]", starters[0].IPs) + } + if starters[1].IPs != nil { + t.Errorf("glueless starter should have nil IPs, got %v", starters[1].IPs) } } -func TestInfoCacheGlueDedup(t *testing.T) { +func TestGetStartServersPrefersDeepestZone(t *testing.T) { c := NewInfoCache(nil) - ip := net.ParseIP("1.2.3.4") - c.StoreGlue("ns1.example.com.", []net.IP{ip, ip}) - result := c.LookupGlue("ns1.example.com.") - if len(result) != 1 { - t.Errorf("expected 1 deduped glue, got %d", len(result)) + c.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}}) + c.Add([]dns.RR{nsRR("example.com", "ns1.example.com"), aRR("ns1.example.com", "1.2.3.4")}) + + starters, bw, err := c.GetStartServers("www.example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "example.com" { + t.Errorf("newbailiwick = %q, want example.com", bw) + } + if len(starters) != 1 || starters[0].Name != "ns1.example.com" { + t.Errorf("starters = %v", starters) } } -func TestInfoCacheChainedGlue(t *testing.T) { +func TestGetStartServersRootHints(t *testing.T) { + c := NewInfoCache(nil) + c.AddHints("", []StartServer{ + {Name: "a.root-servers.net", IPs: []string{"198.41.0.4", "2001:503:ba3e::2:30"}}, + {Name: "b.root-servers.net", IPs: []string{"170.247.170.2"}}, + }) + + starters, bw, err := c.GetStartServers("anything.example.org") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "" { + t.Errorf("root bailiwick should be \"\", got %q", bw) + } + if len(starters) != 2 { + t.Fatalf("expected 2 root starters, got %d", len(starters)) + } + // Only the IPv4 address surfaces (IPv4-only transport); the AAAA is + // cached but not returned as a start address. + if len(starters[0].IPs) != 1 || starters[0].IPs[0] != "198.41.0.4" { + t.Errorf("starter[0].IPs = %v, want [198.41.0.4]", starters[0].IPs) + } + if got := c.Get("a.root-servers.net", dns.ClassINET, dns.TypeAAAA); len(got) != 1 { + t.Errorf("AAAA hint should be cached, got %v", got) + } +} + +func TestGetStartServersNoRootHints(t *testing.T) { + c := NewInfoCache(nil) + if _, _, err := c.GetStartServers("example.com"); err == nil { + t.Fatal("expected error with no NS cached anywhere") + } +} + +func TestGetStartServersExactDomainMatch(t *testing.T) { + c := NewInfoCache(nil) + c.Add([]dns.RR{nsRR("example.com", "ns1.example.net")}) + _, bw, err := c.GetStartServers("example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "example.com" { + t.Errorf("newbailiwick = %q, want example.com", bw) + } +} + +func TestGetStartServersUsesBranchCache(t *testing.T) { parent := NewInfoCache(nil) - parent.StoreGlue("ns1.example.com.", []net.IP{net.ParseIP("1.2.3.4")}) - + parent.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}}) child := parent.Child() - result := child.LookupGlue("ns1.example.com.") - if len(result) != 1 { - t.Fatalf("expected 1 glue from parent, got %d", len(result)) - } - if child.GlueCount() != 0 { - t.Errorf("child GlueCount = %d, want 0", child.GlueCount()) - } -} + child.Add([]dns.RR{nsRR("com", "a.gtld-servers.net")}) -func TestInfoCacheGlueLookupMiss(t *testing.T) { - c := NewInfoCache(nil) - result := c.LookupGlue("nonexistent.example.com.") - if result != nil { - t.Errorf("expected nil for miss, got %v", result) + _, bw, err := child.GetStartServers("www.example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "com" { + t.Errorf("newbailiwick = %q, want com (child cache hit)", bw) + } + + // A sibling branch must not see the child's records. + sibling := parent.Child() + _, bw, err = sibling.GetStartServers("www.example.com") + if err != nil { + t.Fatalf("GetStartServers: %v", err) + } + if bw != "" { + t.Errorf("sibling newbailiwick = %q, want \"\" (root only)", bw) } } func TestInfoCacheConcurrentAccess(t *testing.T) { c := NewInfoCache(nil) var wg sync.WaitGroup - for i := 0; i < 100; i++ { wg.Add(1) go func(i int) { defer wg.Done() - name := strings.ToLower(dns.Fqdn(fmt.Sprintf("ns%d.example.com.", i))) - c.StoreNS("example.com.", []string{name}) - c.StoreGlue(name, []net.IP{net.ParseIP(fmt.Sprintf("1.2.3.%d", i%256))}) - _ = c.LookupNS("example.com.") - _ = c.LookupGlue(name) + name := fmt.Sprintf("ns%d.example.com", i) + c.Add([]dns.RR{nsRR("example.com", name), aRR(name, fmt.Sprintf("1.2.3.%d", i%256))}) + _, _, _ = c.GetStartServers("www.example.com") + _ = c.Get(name, dns.ClassINET, dns.TypeA) }(i) } wg.Wait() } - -func TestInfoCacheNilParent(t *testing.T) { - c := NewInfoCache(nil) - if c.LookupNS("com.") != nil { - t.Error("root cache should return nil for miss") - } - if c.LookupGlue("ns.example.com.") != nil { - t.Error("root cache should return nil for glue miss") - } -} - -func TestNormalize(t *testing.T) { - tests := []struct { - input string - want string - }{ - {"example.com", "example.com."}, - {"Example.COM.", "example.com."}, - {"EXAMPLE.COM", "example.com."}, - } - - for _, tt := range tests { - t.Run(tt.input, func(t *testing.T) { - got := normalize(tt.input) - if got != tt.want { - t.Errorf("normalize(%q) = %q, want %q", tt.input, got, tt.want) - } - }) - } -} diff --git a/internal/traverse/coverage_test.go b/internal/traverse/coverage_test.go deleted file mode 100644 index 63d7817..0000000 --- a/internal/traverse/coverage_test.go +++ /dev/null @@ -1,727 +0,0 @@ -package traverse - -import ( - "context" - "net" - "testing" - "time" - - dnsinternal "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - "github.com/miekg/dns" -) - -func TestSetHooks(t *testing.T) { - tr := NewTraverser(nil) - hooks := &TraverserHooks{ - OnEvent: func(event TraversalEvent) {}, - } - tr.SetHooks(hooks) - if tr.config.Hooks != hooks { - t.Error("SetHooks should set config.Hooks") - } - - // SetHooks on nil config traverser (initializes config) - tr2 := &Traverser{} - tr2.SetHooks(hooks) - if tr2.config == nil || tr2.config.Hooks != hooks { - t.Error("SetHooks should initialize config when nil") - } -} - -func TestNewAQuery(t *testing.T) { - msg := newAQuery("example.com.") - if msg == nil { - t.Fatal("newAQuery returned nil") - } - if !msg.RecursionDesired { - t.Error("expected RD=true in newAQuery") - } - if len(msg.Question) == 0 { - t.Fatal("expected question in newAQuery") - } - if msg.Question[0].Qtype != dns.TypeA { - t.Errorf("expected TypeA, got %d", msg.Question[0].Qtype) - } -} - -func TestResolveGlueViaSystemCacheHit(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - cache := NewInfoCache(nil) - expected := []net.IP{net.ParseIP("1.2.3.4")} - cache.StoreGlue("ns1.example.com.", expected) - - ctx := context.Background() - addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", cache) - if len(addrs) == 0 { - t.Error("expected addresses from cache hit") - } -} - -func TestResolveGlueViaSystemExpiredContext(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - - // Expired context → remaining <= 0 → returns nil immediately - ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(-time.Second)) - defer cancel() - - addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", nil) - if len(addrs) != 0 { - t.Errorf("expected nil from expired context, got %v", addrs) - } -} - -func TestResolveGlueViaSystemTimeout(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - - // Very short timeout will fail the DNS query to 127.0.0.1:53 - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Millisecond) - defer cancel() - time.Sleep(15 * time.Millisecond) // ensure it's expired - - addrs := tr.resolveGlueViaSystem(ctx, "ns1.example.com.", nil) - // May return nil (timeout) or addresses (if local resolver responds instantly) - t.Logf("resolveGlueViaSystem returned %d addresses", len(addrs)) -} - -func TestEnsureRDFalseWithExchange(t *testing.T) { - rdFalseMsg := new(dns.Msg) - rdFalseMsg.SetReply(new(dns.Msg)) - rdFalseMsg.RecursionDesired = false - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return rdFalseMsg.Copy(), nil - }) - - rdTrueMsg := new(dns.Msg) - rdTrueMsg.SetReply(new(dns.Msg)) - rdTrueMsg.RecursionDesired = true - - result := tr.ensureRDFalse(rdTrueMsg, net.ParseIP("198.41.0.4"), "example.com.", dnsinternal.TypeA, nil) - if result == nil { - t.Fatal("ensureRDFalse with exchange should return non-nil") - } -} - -func TestEnsureRDFalseNilMsg(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - result := tr.ensureRDFalse(nil, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil) - if result != nil { - t.Error("ensureRDFalse(nil) should return nil") - } -} - -func TestEnsureRDFalseRDAlreadyFalse(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - - msg := new(dns.Msg) - msg.RecursionDesired = false - result := tr.ensureRDFalse(msg, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil) - if result != msg { - t.Error("ensureRDFalse should return same msg when RD=false") - } -} - -func TestEnsureRDFalseNoExchange(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - // No exchange set - - msg := new(dns.Msg) - msg.RecursionDesired = true - result := tr.ensureRDFalse(msg, net.ParseIP("1.2.3.4"), "example.com.", dnsinternal.TypeA, nil) - if result == nil { - t.Fatal("ensureRDFalse without exchange should return msg with RD cleared") - } - if result.RecursionDesired { - t.Error("expected RD=false after ensureRDFalse without exchange") - } -} - -func TestResolveNSFromCache(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - cache := NewInfoCache(nil) - expected := []net.IP{net.ParseIP("1.2.3.4")} - cache.StoreGlue("ns1.example.com.", expected) - - ctx := context.Background() - addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", cache, nil, 0) - if err != nil { - t.Fatalf("ResolveNS cache hit: %v", err) - } - if len(addrs) == 0 { - t.Error("expected addresses from cache") - } -} - -func TestResolveNSCircularReferral(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - visited := map[string]bool{"ns1.example.com.": true} - ctx := context.Background() - _, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, visited, 0) - if err == nil { - t.Fatal("expected circular referral error") - } - var circErr *CircularReferralError - if _, ok := err.(*CircularReferralError); !ok { - t.Errorf("expected CircularReferralError, got %T: %v", err, err) - } - _ = circErr -} - -func TestResolveNSMaxDepth(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - ctx := context.Background() - _, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, DefaultMaxDepth+1) - if err == nil { - t.Fatal("expected max depth error") - } - if _, ok := err.(*UnresolvableNameserverError); !ok { - t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err) - } -} - -func TestResolveNSWithAnswer(t *testing.T) { - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0) - if err != nil { - t.Fatalf("ResolveNS with answer: %v", err) - } - if len(addrs) == 0 { - t.Fatal("expected addresses from NS resolution") - } -} - -func TestResolveNSNXDOMAIN(t *testing.T) { - nxMsg := new(dns.Msg) - nxMsg.Rcode = dns.RcodeNameError - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nxMsg.Copy(), nil - }) - - ctx := context.Background() - _, err := tr.ResolveNS(ctx, "nonexistent.invalid.", nil, nil, 0) - if err == nil { - t.Fatal("expected error for NXDOMAIN NS resolution") - } -} - -func TestResolveNSContextCancellation(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - // Keep returning referrals to keep the loop going - refMsg := new(dns.Msg) - refMsg.Rcode = dns.RcodeSuccess - refMsg.Ns = append(refMsg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS, Class: dns.ClassINET}, - Ns: "ns.example.com.", - }) - refMsg.Extra = append(refMsg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }) - return refMsg, nil - }) - - ctx, cancel := context.WithCancel(context.Background()) - cancel() // Cancel immediately - - _, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0) - if err == nil { - t.Fatal("expected error on cancelled context") - } -} - -func TestDiscoverRootsWithRootAddrs(t *testing.T) { - expected := []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")} - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: expected, - }) - - ctx := context.Background() - addrs, err := tr.discoverRoots(ctx) - if err != nil { - t.Fatalf("discoverRoots with RootAddrs: %v", err) - } - if len(addrs) != len(expected) { - t.Errorf("expected %d addresses, got %d", len(expected), len(addrs)) - } -} - -func TestDiscoverRootsFromSystem(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - // No RootAddrs - will call dns.DiscoverRoots - }) - - ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) - defer cancel() - - addrs, err := tr.discoverRoots(ctx) - if err != nil { - t.Logf("discoverRoots without RootAddrs error (may skip): %v", err) - t.Skip() - } - if len(addrs) == 0 { - t.Error("expected at least one root address") - } -} - -func TestTraverserSetHooksAndTraverse(t *testing.T) { - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) - - var events []TraversalEvent - tr.SetHooks(&TraverserHooks{ - OnEvent: func(e TraversalEvent) { - events = append(events, e) - }, - }) - - ctx := context.Background() - _, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) - } - if len(events) == 0 { - t.Error("expected events from hooks") - } -} - -func TestProcessReferralNoAddresses(t *testing.T) { - // Scenario: a referral without addresses. resolveGlueViaSystem fails (expired ctx), - // then ResolveNS is tried via the mock exchange. - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - finalAnswerMsg := new(dns.Msg) - finalAnswerMsg.SetReply(new(dns.Msg)) - finalAnswerMsg.Answer = append(finalAnswerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("5.6.7.8"), - }) - - callCount := 0 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - callCount++ - q := msg.Question[0] - if q.Qtype == dns.TypeA && q.Name == "ns1.example.com." { - return answerMsg.Copy(), nil - } - return finalAnswerMsg.Copy(), nil - }) - - // Create a referral with no addresses (the NS name needs to be resolved) - ref := NewReferral("example.com.", dnsinternal.TypeA, "ns1.example.com.", 1, 1.0, nil) - // Do NOT set addresses - this exercises processReferral's no-address path - - cache := NewInfoCache(nil) - // Use expired context for resolveGlueViaSystem so it returns nil fast - bgCtx := context.Background() - resp := tr.processReferral(bgCtx, ref, cache) - // Result may vary depending on whether 127.0.0.1:53 is available, - // but the function should not panic. - t.Logf("processReferral result type: %v", resp.Type) -} - -func TestReferralResolveAlreadyHasAddresses(t *testing.T) { - ref := NewReferral("example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")} - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, 0) - if err != nil { - t.Fatalf("Resolve with existing addresses: %v", err) - } - if ref.State != StateResolved { - t.Errorf("expected StateResolved, got %v", ref.State) - } -} - -func TestReferralResolveCacheHit(t *testing.T) { - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - cache := NewInfoCache(nil) - cache.StoreGlue("ns1.example.com.", []net.IP{net.ParseIP("1.2.3.4")}) - - ctx := context.Background() - err := ref.Resolve(ctx, tr, cache, nil, 0) - if err != nil { - t.Fatalf("Resolve cache hit: %v", err) - } - if ref.State != StateResolved { - t.Errorf("expected StateResolved, got %v", ref.State) - } -} - -func TestReferralResolveCircular(t *testing.T) { - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - visited := map[string]bool{"ns1.example.com.": true} - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, visited, 0) - if err == nil { - t.Fatal("expected circular referral error") - } - if _, ok := err.(*CircularReferralError); !ok { - t.Errorf("expected CircularReferralError, got %T: %v", err, err) - } -} - -func TestReferralResolveMaxDepth(t *testing.T) { - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, DefaultMaxDepth+1) - if err == nil { - t.Fatal("expected max depth error") - } - if _, ok := err.(*UnresolvableNameserverError); !ok { - t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err) - } -} - -func TestReferralResolveWithAnswer(t *testing.T) { - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, 0) - if err != nil { - t.Fatalf("Resolve: %v", err) - } - if ref.State != StateResolved { - t.Errorf("expected StateResolved, got %v", ref.State) - } - if len(ref.Addresses) == 0 { - t.Error("expected addresses after resolution") - } -} - -func TestReferralResolveNXDOMAIN(t *testing.T) { - nxMsg := new(dns.Msg) - nxMsg.Rcode = dns.RcodeNameError - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nxMsg.Copy(), nil - }) - - ref := NewReferral("nonexistent.invalid.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, 0) - if err == nil { - t.Fatal("expected error for NXDOMAIN") - } - if _, ok := err.(*UnresolvableNameserverError); !ok { - t.Errorf("expected UnresolvableNameserverError, got %T: %v", err, err) - } -} - -func TestReferralResolveContextCancellation(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - refMsg := new(dns.Msg) - refMsg.Rcode = dns.RcodeSuccess - refMsg.Ns = append(refMsg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: ".", Rrtype: dns.TypeNS}, - Ns: "ns.example.com.", - }) - refMsg.Extra = append(refMsg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dns.TypeA}, - A: net.ParseIP("1.2.3.4"), - }) - return refMsg, nil - }) - - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ctx, cancel := context.WithCancel(context.Background()) - cancel() // Cancel immediately - - err := ref.Resolve(ctx, tr, nil, nil, 0) - if err == nil { - t.Fatal("expected error on cancelled context") - } -} - -func TestReferralResolveReferralPath(t *testing.T) { - // Test Resolve when it gets a referral response that pushes to stack - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Ns = append(referralMsg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS, Class: dns.ClassINET}, - Ns: "ns1.example.com.", - }) - referralMsg.Extra = append(referralMsg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET}, - A: net.ParseIP("1.2.3.4"), - }) - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("5.6.7.8"), - }) - - callCount := 0 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - callCount++ - if callCount <= 1 { - return referralMsg.Copy(), nil - } - return answerMsg.Copy(), nil - }) - - ref := NewReferral("ns1.example.com.", dnsinternal.TypeA, ".", 0, 1.0, nil) - ctx := context.Background() - err := ref.Resolve(ctx, tr, nil, nil, 0) - // May succeed or exhaust depending on referral loop - t.Logf("Resolve referral path: err=%v, state=%v", err, ref.State) -} - -func TestResolutionStateStringUnknown(t *testing.T) { - // Cover the default case of ResolutionState.String() - unknown := ResolutionState(99) - s := unknown.String() - if s != "unknown" { - t.Errorf("expected 'unknown' for invalid ResolutionState, got %q", s) - } -} - -func TestTraverserNonFastMode(t *testing.T) { - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: false, // Non-fast mode - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("Traverse non-fast: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results") - } -} - -func TestIterativeQueryWithExchangeUsesConfig(t *testing.T) { - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - QueryConfig: dnsinternal.DefaultQueryConfig(), - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - msg, err := tr.iterativeQueryWithExchange(ctx, net.ParseIP("198.41.0.4"), "example.com.", dnsinternal.TypeA) - if err != nil { - t.Fatalf("iterativeQueryWithExchange with config: %v", err) - } - if msg == nil { - t.Fatal("expected non-nil response") - } -} - -func TestTraverserReferralWithHooks(t *testing.T) { - // Tests that hooks are called with IsResolve=true during ResolveNS sub-traversal - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsinternal.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerMsg.Copy(), nil - }) - - var resolveEvents, progressEvents int - tr.SetHooks(&TraverserHooks{ - OnEvent: func(e TraversalEvent) { - if e.IsResolve { - resolveEvents++ - } else { - progressEvents++ - } - }, - }) - - // Test directly via ResolveNS with hooks - ctx := context.Background() - addrs, err := tr.ResolveNS(ctx, "ns1.example.com.", nil, nil, 0) - if err != nil { - t.Fatalf("ResolveNS: %v", err) - } - _ = addrs - t.Logf("resolveEvents=%d progressEvents=%d", resolveEvents, progressEvents) -} diff --git a/internal/traverse/decoded_query.go b/internal/traverse/decoded_query.go new file mode 100644 index 0000000..9633ba1 --- /dev/null +++ b/internal/traverse/decoded_query.go @@ -0,0 +1,259 @@ +package traverse + +import ( + "fmt" + "strings" + + miekgdns "github.com/miekg/dns" +) + +// Status is the classification of one query outcome. The first eight values +// come from decoded_query.rb / response.rb; noglue and loop are synthesised by +// the resolve phase without sending a query. +type Status string + +const ( + StatusAnswered Status = "answered" + StatusNoData Status = "nodata" + StatusReferral Status = "referral" + StatusRestart Status = "restart" + StatusReferralLame Status = "referral_lame" + StatusError Status = "error" + StatusException Status = "exception" + StatusCNAMELoop Status = "cname_loop" + StatusNoGlue Status = "noglue" + StatusLoop Status = "loop" +) + +// DecodedQuery classifies one DNS response (or network failure) against the +// query that produced it, mirroring decoded_query.rb. Names are canonical +// (lowercase, no trailing dot); Bailiwick "" means root. +type DecodedQuery struct { + Msg *miekgdns.Msg + Err error + + Qname string + Qclass uint16 + Qtype uint16 + IP string + Bailiwick string + + Status Status + Endname string + // ChainTargets lists every CNAME target the in-message chain went + // through (including the final unfollowed target when the chain leaves + // the bailiwick); used for cross-response restart loop detection. + ChainTargets []string + + CacheableGood []miekgdns.RR + CacheableBad []miekgdns.RR + + AuthNS []miekgdns.RR + AuthSOA []miekgdns.RR + AuthOther []miekgdns.RR + + Answers []miekgdns.RR + AuthorityNames []string + + ErrorMessage string + ExceptionMessage string + Warnings []string +} + +// NewDecodedQuery decodes and classifies a response. Pass err non-nil for a +// network-level failure (dnstraverse's "exception"); msg is ignored then. +func NewDecodedQuery(msg *miekgdns.Msg, err error, qname string, qclass, qtype uint16, ip, bailiwick string) *DecodedQuery { + dq := &DecodedQuery{ + Msg: msg, + Err: err, + Qname: canonicalName(qname), + Qclass: qclass, + Qtype: qtype, + IP: ip, + Bailiwick: canonicalName(bailiwick), + } + dq.process() + return dq +} + +func (dq *DecodedQuery) WarningsAdd(warnings ...string) { + dq.Warnings = append(dq.Warnings, warnings...) +} + +// process implements the classification order of decoded_query.rb#process +// exactly (the 7 steps in the design doc). +func (dq *DecodedQuery) process() { + if dq.Err == nil && dq.Msg == nil { + dq.Err = fmt.Errorf("nil DNS response") + } + if dq.Err != nil { + dq.Status = StatusException + dq.ExceptionMessage = dq.Err.Error() + return + } + dq.AuthNS, dq.AuthSOA, dq.AuthOther = msgAuthority(dq.Msg) + dq.CacheableGood, dq.CacheableBad = msgCacheable(dq.Msg, dq.Bailiwick) + endname, targets, ok := msgFollowCNAMEs(dq.Msg, dq.Qname, dq.Qtype, dq.Bailiwick) + if !ok { + dq.Status = StatusCNAMELoop + return + } + dq.Endname = endname + dq.ChainTargets = targets + if dq.Msg.Rcode != miekgdns.RcodeSuccess { + dq.Status = StatusError + dq.ErrorMessage = rcodeErrorMessage(dq.Msg.Rcode) + return + } + if answers := msgAnswers(dq.Msg, dq.Endname, dq.Qclass, dq.Qtype); len(answers) > 0 { + dq.Answers = answers + dq.Status = StatusAnswered + return + } + if dq.Endname != dq.Qname { + dq.Status = StatusRestart + return + } + if len(dq.AuthSOA) > 0 || len(dq.AuthNS) == 0 { + dq.Status = StatusNoData + return + } + dq.Status = StatusReferral + for _, rr := range dq.AuthNS { + if ns, ok := rr.(*miekgdns.NS); ok { + dq.AuthorityNames = append(dq.AuthorityNames, canonicalName(ns.Ns)) + } + } +} + +// rcodeErrorMessage renders the exact error strings of decoded_query.rb +// process_error ("Format error" deliberately fixes the Ruby "Formate" typo — +// documented deviation). +func rcodeErrorMessage(rcode int) string { + switch rcode { + case miekgdns.RcodeFormatError: + return "Format error (FORMERR)" + case miekgdns.RcodeServerFailure: + return "Server failure (SERVFAIL)" + case miekgdns.RcodeNameError: + return "No such domain (NXDOMAIN)" + case miekgdns.RcodeNotImplemented: + return "Not implemented (NOTIMP)" + case miekgdns.RcodeRefused: + return "Refused" + default: + if s, ok := miekgdns.RcodeToString[rcode]; ok { + return s + } + return fmt.Sprintf("RCODE%d", rcode) + } +} + +// insideBailiwick reports whether name is at or below the bailiwick zone: +// bailiwick "" (root), equal fold, or name ends with "."+bailiwick. +func insideBailiwick(name, bailiwick string) bool { + bw := canonicalName(bailiwick) + if bw == "" { + return true + } + n := canonicalName(name) + return n == bw || strings.HasSuffix(n, "."+bw) +} + +// msgAnswers returns the answer-section records matching qname/qclass/qtype +// (message_utility.rb msg_answers?). qtype ANY matches every type. +func msgAnswers(msg *miekgdns.Msg, qname string, qclass, qtype uint16) []miekgdns.RR { + name := canonicalName(qname) + any := qtype == miekgdns.TypeANY + var out []miekgdns.RR + for _, rr := range msg.Answer { + h := rr.Header() + if canonicalName(h.Name) == name && h.Class == qclass && (any || h.Rrtype == qtype) { + out = append(out, rr) + } + } + return out +} + +// msgAuthority partitions the authority section into IN NS, IN SOA and other +// records (message_utility.rb msg_authority). +func msgAuthority(msg *miekgdns.Msg) (ns, soa, other []miekgdns.RR) { + for _, rr := range msg.Ns { + h := rr.Header() + switch { + case h.Rrtype == miekgdns.TypeNS && h.Class == miekgdns.ClassINET: + ns = append(ns, rr) + case h.Rrtype == miekgdns.TypeSOA && h.Class == miekgdns.ClassINET: + soa = append(soa, rr) + default: + other = append(other, rr) + } + } + return ns, soa, other +} + +// msgCacheable partitions ALL sections (answer, authority, additional — in +// that order) into in-bailiwick records worth caching and out-of-bailiwick +// records to discard. OPT pseudo-records are dropped entirely. +func msgCacheable(msg *miekgdns.Msg, bailiwick string) (good, bad []miekgdns.RR) { + for _, section := range [][]miekgdns.RR{msg.Answer, msg.Ns, msg.Extra} { + for _, rr := range section { + if rr.Header().Rrtype == miekgdns.TypeOPT { + continue + } + if insideBailiwick(rr.Header().Name, bailiwick) { + good = append(good, rr) + } else { + bad = append(bad, rr) + } + } + } + return good, bad +} + +// msgFollowCNAMEs follows a CNAME chain within one message and returns the +// final name plus every target passed through (message_utility.rb +// msg_follow_cnames). Following stops — the target is returned unfollowed — +// as soon as the CURRENT owner name is not strictly below the bailiwick +// (Ruby tests `name !~ /\.#{bailiwick}$/i`, so an owner exactly equal to the +// bailiwick also stops the chain). An in-message loop returns ok=false +// (cname_loop). +func msgFollowCNAMEs(msg *miekgdns.Msg, qname string, qtype uint16, bailiwick string) (endname string, targets []string, ok bool) { + name := canonicalName(qname) + bw := canonicalName(bailiwick) + seen := make(map[string]bool) + for { + seen[name] = true + if len(msgAnswers(msg, name, miekgdns.ClassINET, qtype)) > 0 { + return name, targets, true + } + cnames := msgAnswers(msg, name, miekgdns.ClassINET, miekgdns.TypeCNAME) + if len(cnames) == 0 { + return name, targets, true + } + cname, isCNAME := cnames[0].(*miekgdns.CNAME) + if !isCNAME { + return name, targets, true + } + target := canonicalName(cname.Target) + targets = append(targets, target) + if bw != "" && !strings.HasSuffix(name, "."+bw) { + return target, targets, true + } + name = target + if seen[name] { + return "", targets, false + } + } +} + +// isLameReferral implements the response.rb lame rule: a referral is lame +// unless the current bailiwick is root ("") or the new zone is STRICTLY +// deeper than the current bailiwick (equal or sideways zones are lame). +func isLameReferral(bailiwick, newBailiwick string) bool { + bw := canonicalName(bailiwick) + if bw == "" { + return false + } + return !strings.HasSuffix(canonicalName(newBailiwick), "."+bw) +} diff --git a/internal/traverse/decoded_query_test.go b/internal/traverse/decoded_query_test.go new file mode 100644 index 0000000..2bbb696 --- /dev/null +++ b/internal/traverse/decoded_query_test.go @@ -0,0 +1,359 @@ +package traverse + +import ( + "errors" + "testing" + + "github.com/miekg/dns" +) + +func cnameRR(owner, target string) dns.RR { + return &dns.CNAME{ + Hdr: dns.RR_Header{Name: dns.Fqdn(owner), Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, + Target: dns.Fqdn(target), + } +} + +func soaRR(zone string) dns.RR { + return &dns.SOA{ + Hdr: dns.RR_Header{Name: dns.Fqdn(zone), Rrtype: dns.TypeSOA, Class: dns.ClassINET}, + Ns: dns.Fqdn("ns1." + zone), + Mbox: dns.Fqdn("hostmaster." + zone), + Serial: 1, + Refresh: 3600, Retry: 600, Expire: 86400, Minttl: 300, + } +} + +func newMsg(qname string, qtype uint16, rcode int) *dns.Msg { + m := new(dns.Msg) + m.SetQuestion(dns.Fqdn(qname), qtype) + m.Response = true + m.Rcode = rcode + return m +} + +func decode(msg *dns.Msg, qname string, qtype uint16, bailiwick string) *DecodedQuery { + return NewDecodedQuery(msg, nil, qname, dns.ClassINET, qtype, "192.0.2.1", bailiwick) +} + +func TestDecodeException(t *testing.T) { + dq := NewDecodedQuery(nil, errors.New("network timeout"), "example.com", dns.ClassINET, dns.TypeA, "192.0.2.1", "com") + if dq.Status != StatusException { + t.Fatalf("status = %s, want exception", dq.Status) + } + if dq.ExceptionMessage != "network timeout" { + t.Errorf("exception message = %q", dq.ExceptionMessage) + } +} + +func TestDecodeNilMessageIsException(t *testing.T) { + dq := NewDecodedQuery(nil, nil, "example.com", dns.ClassINET, dns.TypeA, "192.0.2.1", "com") + if dq.Status != StatusException { + t.Fatalf("status = %s, want exception", dq.Status) + } +} + +func TestDecodeErrorMessages(t *testing.T) { + tests := []struct { + rcode int + want string + }{ + {dns.RcodeFormatError, "Format error (FORMERR)"}, + {dns.RcodeServerFailure, "Server failure (SERVFAIL)"}, + {dns.RcodeNameError, "No such domain (NXDOMAIN)"}, + {dns.RcodeNotImplemented, "Not implemented (NOTIMP)"}, + {dns.RcodeRefused, "Refused"}, + {dns.RcodeYXDomain, "YXDOMAIN"}, + } + for _, tt := range tests { + msg := newMsg("example.com", dns.TypeA, tt.rcode) + dq := decode(msg, "example.com", dns.TypeA, "com") + if dq.Status != StatusError { + t.Errorf("rcode %d: status = %s, want error", tt.rcode, dq.Status) + } + if dq.ErrorMessage != tt.want { + t.Errorf("rcode %d: message = %q, want %q", tt.rcode, dq.ErrorMessage, tt.want) + } + } +} + +func TestDecodeAnswered(t *testing.T) { + msg := newMsg("example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, aRR("example.com", "93.184.216.34")) + dq := decode(msg, "example.com", dns.TypeA, "example.com") + if dq.Status != StatusAnswered { + t.Fatalf("status = %s, want answered", dq.Status) + } + if len(dq.Answers) != 1 { + t.Errorf("answers = %v", dq.Answers) + } + if dq.Endname != "example.com" { + t.Errorf("endname = %q", dq.Endname) + } +} + +func TestDecodeAnsweredViaCNAMEChain(t *testing.T) { + // In-bailiwick chain ends at a name that has the A answer. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("www.example.com", "web.example.com"), + aRR("web.example.com", "93.184.216.34"), + ) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusAnswered { + t.Fatalf("status = %s, want answered", dq.Status) + } + if dq.Endname != "web.example.com" { + t.Errorf("endname = %q, want web.example.com", dq.Endname) + } +} + +func TestDecodeRestartOnOutOfBailiwickCNAME(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("www.example.com", "cdn.example.org")) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusRestart { + t.Fatalf("status = %s, want restart", dq.Status) + } + if dq.Endname != "cdn.example.org" { + t.Errorf("endname = %q", dq.Endname) + } +} + +func TestDecodeCNAMELoop(t *testing.T) { + msg := newMsg("a.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("a.example.com", "b.example.com"), + cnameRR("b.example.com", "a.example.com"), + ) + dq := decode(msg, "a.example.com", dns.TypeA, "example.com") + if dq.Status != StatusCNAMELoop { + t.Fatalf("status = %s, want cname_loop", dq.Status) + } +} + +func TestDecodeCNAMESelfLoop(t *testing.T) { + msg := newMsg("a.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("a.example.com", "A.EXAMPLE.COM")) + dq := decode(msg, "a.example.com", dns.TypeA, "example.com") + if dq.Status != StatusCNAMELoop { + t.Fatalf("status = %s, want cname_loop (case-insensitive)", dq.Status) + } +} + +func TestDecodeCNAMEChainStopsAtOutOfBailiwickOwner(t *testing.T) { + // Ruby stops following once the CURRENT owner leaves the bailiwick, so a + // two-hop loop through an out-of-bailiwick owner is NOT cname_loop: the + // unfollowed target equals the qname again, leaving endname == qname and + // an empty authority — nodata. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("www.example.com", "a.example.org"), + cnameRR("a.example.org", "www.example.com"), + ) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusNoData { + t.Fatalf("status = %s, want nodata", dq.Status) + } + if dq.Endname != "www.example.com" { + t.Errorf("endname = %q, want www.example.com", dq.Endname) + } +} + +func TestDecodeCNAMEOwnerEqualToBailiwickStopsChain(t *testing.T) { + // Owner exactly equal to the bailiwick is NOT strictly inside it, so the + // chain stops after one hop even though another CNAME exists. + msg := newMsg("example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("example.com", "a.example.com"), + cnameRR("a.example.com", "b.example.com"), + ) + dq := decode(msg, "example.com", dns.TypeA, "example.com") + if dq.Status != StatusRestart { + t.Fatalf("status = %s, want restart", dq.Status) + } + if dq.Endname != "a.example.com" { + t.Errorf("endname = %q, want a.example.com (unfollowed target)", dq.Endname) + } +} + +func TestDecodeQtypeCNAMEIsAnswered(t *testing.T) { + // qtype=CNAME: the CNAME record IS the answer; the chain is never followed. + msg := newMsg("www.example.com", dns.TypeCNAME, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("www.example.com", "web.example.com"), + cnameRR("web.example.com", "www.example.com"), + ) + dq := decode(msg, "www.example.com", dns.TypeCNAME, "example.com") + if dq.Status != StatusAnswered { + t.Fatalf("status = %s, want answered", dq.Status) + } + if dq.Endname != "www.example.com" { + t.Errorf("endname = %q", dq.Endname) + } +} + +func TestDecodeQtypeANYMatchesAnyAnswer(t *testing.T) { + msg := newMsg("example.com", dns.TypeANY, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("example.com", "elsewhere.example.net")) + dq := decode(msg, "example.com", dns.TypeANY, "example.com") + if dq.Status != StatusAnswered { + t.Fatalf("status = %s, want answered (ANY matches CNAME)", dq.Status) + } +} + +func TestDecodeNoDataWithSOA(t *testing.T) { + msg := newMsg("example.com", dns.TypeMX, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, soaRR("example.com")) + dq := decode(msg, "example.com", dns.TypeMX, "example.com") + if dq.Status != StatusNoData { + t.Fatalf("status = %s, want nodata", dq.Status) + } +} + +func TestDecodeNoDataEmptyAuthority(t *testing.T) { + msg := newMsg("example.com", dns.TypeMX, dns.RcodeSuccess) + dq := decode(msg, "example.com", dns.TypeMX, "example.com") + if dq.Status != StatusNoData { + t.Fatalf("status = %s, want nodata", dq.Status) + } +} + +func TestDecodeNoDataSOAWinsOverNS(t *testing.T) { + // SOA + NS in authority is a negative answer, not a referral. + msg := newMsg("example.com", dns.TypeMX, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, soaRR("example.com"), nsRR("example.com", "ns1.example.com")) + dq := decode(msg, "example.com", dns.TypeMX, "example.com") + if dq.Status != StatusNoData { + t.Fatalf("status = %s, want nodata", dq.Status) + } +} + +func TestDecodeReferral(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, + nsRR("example.com", "NS1.Example.COM"), + nsRR("example.com", "ns2.example.net"), + ) + msg.Extra = append(msg.Extra, aRR("ns1.example.com", "1.2.3.4")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + if dq.Status != StatusReferral { + t.Fatalf("status = %s, want referral", dq.Status) + } + if len(dq.AuthorityNames) != 2 || dq.AuthorityNames[0] != "ns1.example.com" || dq.AuthorityNames[1] != "ns2.example.net" { + t.Errorf("authority names = %v", dq.AuthorityNames) + } +} + +func TestDecodeErrorBeatsAnswer(t *testing.T) { + // rcode is checked before answers (step 3 before step 4). + msg := newMsg("example.com", dns.TypeA, dns.RcodeServerFailure) + msg.Answer = append(msg.Answer, aRR("example.com", "1.2.3.4")) + dq := decode(msg, "example.com", dns.TypeA, "com") + if dq.Status != StatusError { + t.Fatalf("status = %s, want error", dq.Status) + } +} + +func TestDecodeCNAMEFollowedIntoNXDOMAIN(t *testing.T) { + // CNAME followed first (step 2), then rcode (step 3): NXDOMAIN after an + // in-message CNAME is still an error, but the loop check ran first. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeNameError) + msg.Answer = append(msg.Answer, cnameRR("www.example.com", "gone.example.com")) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusError { + t.Fatalf("status = %s, want error", dq.Status) + } + if dq.Endname != "gone.example.com" { + t.Errorf("endname = %q", dq.Endname) + } +} + +func TestDecodeCacheablePartition(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("www.example.com", "cdn.example.org")) + msg.Ns = append(msg.Ns, nsRR("example.org", "ns1.example.org")) + msg.Extra = append(msg.Extra, aRR("ns1.example.org", "5.6.7.8")) + opt := new(dns.OPT) + opt.Hdr = dns.RR_Header{Name: ".", Rrtype: dns.TypeOPT} + msg.Extra = append(msg.Extra, opt) + + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if len(dq.CacheableGood) != 1 { + t.Errorf("good = %v, want just the CNAME", dq.CacheableGood) + } + if len(dq.CacheableBad) != 2 { + t.Errorf("bad = %v, want NS+A for example.org", dq.CacheableBad) + } +} + +func TestInsideBailiwick(t *testing.T) { + tests := []struct { + name, bailiwick string + want bool + }{ + {"anything.example.com", "", true}, // root bailiwick + {"anything.example.com", ".", true}, // root as dot + {"example.com", "example.com", true}, // exact + {"Example.COM", "example.com", true}, // exact, case fold + {"www.example.com", "EXAMPLE.com", true}, // suffix, case fold + {"a.b.example.com", "example.com", true}, // deep suffix + {"badexample.com", "example.com", false}, // label boundary + {"example.org", "example.com", false}, // sideways + {"com", "example.com", false}, // shallower + {"www.example.com.", "example.com", true}, // trailing dot + } + for _, tt := range tests { + if got := insideBailiwick(tt.name, tt.bailiwick); got != tt.want { + t.Errorf("insideBailiwick(%q, %q) = %v, want %v", tt.name, tt.bailiwick, got, tt.want) + } + } +} + +func TestIsLameReferral(t *testing.T) { + tests := []struct { + bailiwick, newBailiwick string + want bool + }{ + {"", "com", false}, // root bailiwick never lame + {"", "", false}, // root to root + {"com", "example.com", false}, // strictly deeper + {"com", "a.b.example.com", false}, // much deeper + {"COM", "example.com", false}, // case fold + {"com", "com", true}, // equal zone is lame + {"com", "", true}, // back to root is lame + {"com", "org", true}, // sideways is lame + {"example.com", "com", true}, // shallower is lame + {"example.com", "badexample.com", true}, // label boundary + {"example.com", "www.example.com", false}, // deeper + } + for _, tt := range tests { + if got := isLameReferral(tt.bailiwick, tt.newBailiwick); got != tt.want { + t.Errorf("isLameReferral(%q, %q) = %v, want %v", tt.bailiwick, tt.newBailiwick, got, tt.want) + } + } +} + +func TestMsgFollowCNAMEsNoChain(t *testing.T) { + msg := newMsg("example.com", dns.TypeA, dns.RcodeSuccess) + end, _, ok := msgFollowCNAMEs(msg, "Example.COM.", dns.TypeA, "com") + if !ok || end != "example.com" { + t.Errorf("end = %q ok=%v, want example.com true", end, ok) + } +} + +func TestMsgFollowCNAMEsRootBailiwickFollowsEverything(t *testing.T) { + msg := newMsg("a.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, + cnameRR("a.example.com", "b.example.org"), + cnameRR("b.example.org", "c.example.net"), + aRR("c.example.net", "1.2.3.4"), + ) + end, targets, ok := msgFollowCNAMEs(msg, "a.example.com", dns.TypeA, "") + if !ok || end != "c.example.net" { + t.Errorf("end = %q ok=%v, want c.example.net true", end, ok) + } + if len(targets) != 2 || targets[0] != "b.example.org" || targets[1] != "c.example.net" { + t.Errorf("chain targets = %v", targets) + } +} diff --git a/internal/traverse/hooks.go b/internal/traverse/hooks.go index fdf0d88..3851f35 100644 --- a/internal/traverse/hooks.go +++ b/internal/traverse/hooks.go @@ -1,16 +1,67 @@ package traverse +// EventStage mirrors the :stage symbols reported by traverser.rb's +// report_progress: new/start/answer/resolve plus the fast-mode and +// multi-childset variants. type EventStage int const ( - EventStart EventStage = iota - EventComplete + // StageNew fires when a referral node is created (before processing). + StageNew EventStage = iota + // StageStart fires when a referral is popped for processing. + StageStart + // StageNewReferralSet fires once per extra childset when more than one + // IP of a server produced children. + StageNewReferralSet + // StageNewFast fires instead of StageNew when fast mode already knows + // this referral will be completed from the memo. + StageNewFast + // StageResolve fires after a resolve subtree's statistics are folded + // into the referral (post-order, Ruby's :calc_resolve marker). + StageResolve + // StageAnswer fires after a referral's statistics are calculated + // (post-order, Ruby's :calc_answer marker). + StageAnswer + // StageAnswerFast fires when fast mode replaced the referral with an + // earlier completed one instead of processing it. + StageAnswerFast ) +func (s EventStage) String() string { + switch s { + case StageNew: + return "new" + case StageStart: + return "start" + case StageNewReferralSet: + return "new_referral_set" + case StageNewFast: + return "new_fast" + case StageResolve: + return "resolve" + case StageAnswer: + return "answer" + case StageAnswerFast: + return "answer_fast" + default: + return "unknown" + } +} + +// TraversalEvent is one progress callback. RefID/Status/IsResolve are +// denormalised from Referral so renderers (CLI, web) need not walk the tree. type TraversalEvent struct { - Stage EventStage - Result TraversalResult + Stage EventStage + Referral *Referral + RefID string + // Status summarises the referral's outcome so far (see + // Referral.OverallStatus); empty before any response arrives. + Status Status + // IsResolve is true for nodes inside a glue-resolution subtree. IsResolve bool + // CompletedEarlier carries the refid of the earlier identical referral + // on fast-mode events (StageNewFast / StageAnswerFast). + CompletedEarlier string } type EventHandler func(TraversalEvent) @@ -19,13 +70,16 @@ type TraverserHooks struct { OnEvent EventHandler } -func (h *TraverserHooks) emit(stage EventStage, result TraversalResult, isResolve bool) { - if h == nil || h.OnEvent == nil { +func (h *TraverserHooks) emit(stage EventStage, r *Referral, completedEarlier string) { + if h == nil || h.OnEvent == nil || r == nil { return } h.OnEvent(TraversalEvent{ - Stage: stage, - Result: result, - IsResolve: isResolve, + Stage: stage, + Referral: r, + RefID: r.RefID, + Status: r.OverallStatus(), + IsResolve: r.IsResolve(), + CompletedEarlier: completedEarlier, }) } diff --git a/internal/traverse/hooks_test.go b/internal/traverse/hooks_test.go index 69007e6..7405375 100644 --- a/internal/traverse/hooks_test.go +++ b/internal/traverse/hooks_test.go @@ -1,52 +1,71 @@ package traverse import ( - "context" - "net" "testing" "github.com/miekg/dns" ) -func TestTraverserHooksEmitEvents(t *testing.T) { - answerResp := func() *dns.Msg { - m := new(dns.Msg) - m.SetReply(new(dns.Msg)) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - return m - }() - - var events []TraversalEvent - hooks := &TraverserHooks{ - OnEvent: func(event TraversalEvent) { - events = append(events, event) - }, +func TestEventStageStrings(t *testing.T) { + tests := map[EventStage]string{ + StageNew: "new", + StageStart: "start", + StageNewReferralSet: "new_referral_set", + StageNewFast: "new_fast", + StageResolve: "resolve", + StageAnswer: "answer", + StageAnswerFast: "answer_fast", + EventStage(99): "unknown", } + for stage, want := range tests { + if got := stage.String(); got != want { + t.Errorf("EventStage(%d).String() = %q, want %q", stage, got, want) + } + } +} + +func TestHooksEmitNilSafe(t *testing.T) { + var h *TraverserHooks + h.emit(StageNew, newTestReferral("ns1.example.com", nil), "") // must not panic + (&TraverserHooks{}).emit(StageNew, newTestReferral("ns1.example.com", nil), "") + (&TraverserHooks{OnEvent: func(TraversalEvent) { t.Fatal("emitted for nil referral") }}).emit(StageNew, nil, "") +} + +func TestHooksEventSequenceSimpleAnswer(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "9.9.9.9"))) - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dns.TypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Hooks: hooks, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) + var got []string + cfg := testConfig(false) + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { + got = append(got, ev.Stage.String()+":"+ev.RefID) + }} + runTraversal(t, cfg, m, "example.com") - _, err := tr.Traverse(context.Background(), "example.com") - if err != nil { - t.Fatalf("Traverse: %v", err) + want := []string{"new:", "start:", "new:1", "start:1", "answer:1", "answer:"} + if len(got) != len(want) { + t.Fatalf("events = %v, want %v", got, want) } - if len(events) < 2 { - t.Fatalf("expected start and complete events, got %d", len(events)) + for i := range want { + if got[i] != want[i] { + t.Fatalf("events = %v, want %v", got, want) + } } - if events[0].Stage != EventStart { - t.Fatalf("first event stage = %v, want start", events[0].Stage) - } - if events[1].Stage != EventComplete { - t.Fatalf("second event stage = %v, want complete", events[1].Stage) +} + +func TestHooksEventCarriesStatus(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "example.com", dns.TypeA, answerMsg(aRR("example.com", "9.9.9.9"))) + + var answerStatus Status + cfg := testConfig(false) + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { + if ev.Stage == StageAnswer && ev.RefID == "1" { + answerStatus = ev.Status + } + }} + runTraversal(t, cfg, m, "example.com") + if answerStatus != StatusAnswered { + t.Errorf("answer event status = %q, want answered", answerStatus) } } diff --git a/internal/traverse/referral.go b/internal/traverse/referral.go index 1ed1671..68129cb 100644 --- a/internal/traverse/referral.go +++ b/internal/traverse/referral.go @@ -4,6 +4,7 @@ import ( "context" "fmt" "net" + "sort" "strings" "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" @@ -11,44 +12,87 @@ import ( "golang.org/x/net/idna" ) -type ResolutionState int +// DefaultMaxDepth is the default maximum referral depth (non-zero refid +// components) before a "Maxdepth N exceeded" exception is injected. +const DefaultMaxDepth = 20 + +// ReferralStatus is the resolve-phase status of a Referral node itself +// (referral.rb @status), distinct from the per-IP response statuses. +type ReferralStatus string const ( - StateUnresolved ResolutionState = iota - StateResolving - StateResolved + RefStatusNormal ReferralStatus = "normal" + RefStatusNoGlue ReferralStatus = "noglue" + RefStatusLoop ReferralStatus = "loop" ) -func (s ResolutionState) String() string { - switch s { - case StateUnresolved: - return "unresolved" - case StateResolving: - return "resolving" - case StateResolved: - return "resolved" - default: - return "unknown" - } +// StatsEntry is one aggregated leaf statistic: the probability mass that +// ended in Response's outcome at Referral (referral.rb @stats values). +type StatsEntry struct { + Key string + Prob float64 + Response *ServerResponse + Referral *Referral } +// Referral represents one referral to a specific server for qname/qclass/ +// qtype (referral.rb). The synthetic top node ("rootroot") has Server == "" +// and is never displayed; its children are the root servers. type Referral struct { - Name string - Qtype uint16 - Qclass uint16 - Bailiwick string - - Addresses []net.IP - State ResolutionState - - NSName string + RefID string Parent *Referral - Depth int - Prob float64 + + Qname string + Qclass uint16 + Qtype uint16 + // NSAType is the record type used to resolve nameserver addresses + // (always A; the reference is IPv4-only for transport). + NSAType uint16 + + // Server is the NS hostname this referral queries ("" for rootroot). + Server string + // ServerIPs is nil when the server still needs resolving. After a + // resolve it may also contain "key:..." pseudo entries carrying the + // probability of failed resolutions. + ServerIPs []string + Bailiwick string + ParentIP string + + InfoCache *InfoCache + Status ReferralStatus + + // Responses holds the classified response for each real IP queried. + Responses map[string]*ServerResponse + // Children holds child referrals keyed by the parent IP that produced + // them ("rootroot" for the synthetic top node). + Children map[string][]*Referral + // Resolves is the glue-resolution subtree (refid ".0." components). + Resolves []*Referral + + ServerWeights map[string]float64 + Warnings []string + + // Stats is the post-order aggregation of leaf outcomes below this node. + Stats map[string]*StatsEntry + // StatsResolve aggregates the outcomes of the resolve subtree. + StatsResolve map[string]*StatsEntry + + // ReplacedBy points at the earlier completed referral that fast mode + // substituted for this node. + ReplacedBy *Referral + + // summaryStats memoises SummaryStats() (referral.rb summary_stats). + summaryStats *SummaryStats + + client *dns.Client + maxdepth int + referralResolution bool + processed bool + calculated bool } -// idnaLookup is the IDN lookup profile used to convert internationalised domain -// names (unicode labels) to their ACE/punycode equivalents before querying. +// idnaLookup is the IDN lookup profile used to convert internationalised +// domain names (unicode labels) to their ACE/punycode equivalents. var idnaLookup = idna.New( idna.MapForLookup(), idna.BidiRule(), @@ -56,9 +100,8 @@ var idnaLookup = idna.New( ) // toASCII converts a domain name that may contain unicode labels to its -// punycode (ACE) representation. Pure-ASCII names are returned unchanged. -// On conversion errors the original name is returned so the caller can still -// attempt a query (the server will reject it if truly invalid). +// punycode (ACE) representation. On conversion errors the original name is +// returned so the caller can still attempt a query. func toASCII(name string) string { if name == "" || name == "." { return name @@ -70,210 +113,508 @@ func toASCII(name string) string { return ascii } -func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral { - return &Referral{ - Name: miekgdns.Fqdn(strings.ToLower(toASCII(name))), - Qtype: qtype, - Qclass: miekgdns.ClassINET, - Bailiwick: miekgdns.Fqdn(strings.ToLower(toASCII(bailiwick))), - Depth: depth, - Prob: prob, - Parent: parent, - State: StateUnresolved, +// referralArgs are the per-child overrides for makeReferral; zero values +// inherit from the parent (referral.rb make_referral merge semantics). +type referralArgs struct { + qname string + qtype uint16 + server string + serverIPs []string + bailiwick string + infoCache *InfoCache + refid string + parentIP string + referralResolution bool +} + +func (r *Referral) makeReferral(a referralArgs) *Referral { + child := &Referral{ + RefID: a.refid, + Parent: r, + Qname: r.Qname, + Qclass: r.Qclass, + Qtype: r.Qtype, + NSAType: r.NSAType, + Server: canonicalName(a.server), + ServerIPs: a.serverIPs, + Bailiwick: canonicalName(a.bailiwick), + ParentIP: a.parentIP, + InfoCache: r.InfoCache, + Status: RefStatusNormal, + Responses: make(map[string]*ServerResponse), + Children: make(map[string][]*Referral), + ServerWeights: make(map[string]float64), + client: r.client, + maxdepth: r.maxdepth, + referralResolution: a.referralResolution || r.referralResolution, } -} - -func (r *Referral) InBailiwick(name string) bool { - if r.Bailiwick == "" || r.Bailiwick == "." { - return true + if a.qname != "" { + child.Qname = canonicalName(a.qname) } - fqdn := miekgdns.Fqdn(strings.ToLower(name)) - return miekgdns.IsSubDomain(r.Bailiwick, fqdn) -} - -func (r *Referral) HasAddresses() bool { - return len(r.Addresses) > 0 -} - -func (r *Referral) SetAddresses(addrs []net.IP) { - r.Addresses = addrs - if len(addrs) > 0 { - r.State = StateResolved - } else { - r.State = StateUnresolved + if a.qtype != 0 { + child.Qtype = a.qtype } -} - -type CircularReferralError struct { - Name string - Chain []string -} - -func (e *CircularReferralError) Error() string { - return fmt.Sprintf("circular referral detected for %s: %v", e.Name, e.Chain) -} - -type UnresolvableNameserverError struct { - Name string - Reason string -} - -func (e *UnresolvableNameserverError) Error() string { - return fmt.Sprintf("unresolvable nameserver %s: %s", e.Name, e.Reason) -} - -func (r *Referral) Resolve(ctx context.Context, traverser *Traverser, cache *InfoCache, visited map[string]bool, depth int) error { - if r.HasAddresses() { - r.State = StateResolved - return nil + if a.infoCache != nil { + child.InfoCache = a.infoCache } - - if cache != nil { - if addrs := cache.LookupGlue(r.Name); len(addrs) > 0 { - r.Addresses = addrs - r.State = StateResolved - return nil + // serverweight = 1/len(ips) per IP when the addresses are known. + if child.ServerIPs != nil { + for _, ip := range child.ServerIPs { + child.ServerWeights[ip] = 1.0 / float64(len(child.ServerIPs)) } } + return child +} - if visited != nil { - if visited[r.Name] { - return &CircularReferralError{ - Name: r.Name, - Chain: getVisitedNames(visited), +// IsRootRoot reports whether this is the synthetic top node representing an +// automatic referral to the root servers. +func (r *Referral) IsRootRoot() bool { + return r.Server == "" +} + +// IsResolve reports whether this node is part of a glue-resolution subtree. +func (r *Referral) IsResolve() bool { + return r.referralResolution +} + +// Resolved reports whether the server addresses are known (rootroot is +// always resolved). +func (r *Referral) Resolved() bool { + return r.IsRootRoot() || r.ServerIPs != nil +} + +// Depth counts the non-zero refid components; resolve subtrees (".0.") do +// not count against the depth limit. +func (r *Referral) Depth() int { + return refidDepth(r.RefID) +} + +func refidDepth(refid string) int { + if refid == "" { + return 0 + } + n := 0 + for _, part := range strings.Split(refid, ".") { + if part != "0" { + n++ + } + } + return n +} + +// IPsAsArray returns the real IP addresses known for this referral, +// excluding "key:" pseudo entries. +func (r *Referral) IPsAsArray() []string { + var out []string + for _, ip := range r.ServerIPs { + if !strings.HasPrefix(ip, "key:") { + out = append(out, ip) + } + } + return out +} + +// TxtIPsVerbose renders the per-IP weights, sorted, e.g. +// "50.0%=1.2.3.4,50.0%=noglue:1.2.3.4" (referral.rb txt_ips_verbose). It is +// part of the fast-mode memo key. +func (r *Referral) TxtIPsVerbose() string { + if r.ServerIPs == nil { + return "" + } + parts := make([]string, 0, len(r.ServerIPs)) + for _, ip := range r.ServerIPs { + label := ip + if rest, ok := strings.CutPrefix(ip, "key:"); ok { + // keep the first two colon-separated fields, like Ruby's + // /^key:([^:]+(:[^:]*)?)/ capture. + fields := strings.SplitN(rest, ":", 3) + if len(fields) > 2 { + fields = fields[:2] + } + label = strings.Join(fields, ":") + } + parts = append(parts, fmt.Sprintf("%.1f%%=%s", 100*r.ServerWeights[ip], label)) + } + sort.Strings(parts) + return strings.Join(parts, ",") +} + +// TxtIPs renders the addresses for progress display; failed-resolve pseudo +// entries render as their response description (referral.rb txt_ips). +func (r *Referral) TxtIPs() string { + if r.ServerIPs == nil { + return "" + } + parts := make([]string, 0, len(r.ServerIPs)) + for _, ip := range r.ServerIPs { + if strings.HasPrefix(ip, "key:") { + if e, ok := r.StatsResolve[ip]; ok && e.Response != nil { + parts = append(parts, e.Response.String()) + continue } } - visited[r.Name] = true - } - - if depth > DefaultMaxDepth { - return &UnresolvableNameserverError{ - Name: r.Name, - Reason: "max depth exceeded", - } - } - - roots, err := traverser.discoverRoots(ctx) - if err != nil { - return fmt.Errorf("root discovery: %w", err) - } - - initial := NewReferral(r.Name, dns.TypeA, ".", 0, 1.0, nil) - initial.Addresses = roots - initial.State = StateResolved - - stack := NewStack(DefaultMaxDepth) - stack.Push(initial) - - traversalCache := NewInfoCache(nil) - if visited != nil { - for name := range visited { - traversalCache.StoreGlue(name, []net.IP{}) - } - } - - var lastErr error - for { - select { - case <-ctx.Done(): - return fmt.Errorf("resolution cancelled: %w", ctx.Err()) - default: - } - - ref := stack.Pop() - if ref == nil { - break - } - - cacheForStep := traversalCache - if ref.Parent != nil { - cacheForStep = traversalCache.Child() - } - - resp := traverser.processReferral(ctx, ref, cacheForStep) - - if resp.Type == RespAnswer { - if len(resp.Decoded.Answers) > 0 { - var addrs []net.IP - for _, rr := range resp.Decoded.Answers { - if a, ok := rr.(*miekgdns.A); ok { - addrs = append(addrs, a.A) - } - if aaaa, ok := rr.(*miekgdns.AAAA); ok { - addrs = append(addrs, aaaa.AAAA) - } - } - if len(addrs) > 0 { - r.Addresses = addrs - r.State = StateResolved - if cache != nil { - cache.StoreGlue(r.Name, addrs) - } - return nil - } - } - } - - if resp.Type == RespNXDOMAIN { - lastErr = &UnresolvableNameserverError{ - Name: r.Name, - Reason: "NXDOMAIN", - } - break - } - - if resp.Type == RespSERVFAIL || resp.Type == RespError { - lastErr = fmt.Errorf("server error resolving %s: %s", r.Name, resp.Type) - continue - } - - if resp.Type == RespReferral { - children := resp.ChildReferrals() - for _, child := range children { - // Only skip visited names when they have no addresses; if glue - // was included in the referral response we still need to query - // that child to get the authoritative answer. - if visited != nil && visited[child.Name] && !child.HasAddresses() { - continue - } - if !stack.Push(child) { - lastErr = &UnresolvableNameserverError{ - Name: r.Name, - Reason: "max depth exceeded during resolution", - } - } - } - } - } - - if lastErr != nil { - return lastErr - } - - return &UnresolvableNameserverError{ - Name: r.Name, - Reason: "resolution exhausted without answer", + parts = append(parts, ip) } + sort.Strings(parts) + return strings.Join(parts, ",") } -func getVisitedNames(visited map[string]bool) []string { - var names []string - for name := range visited { - names = append(names, name) - } - return names +func (r *Referral) String() string { + return fmt.Sprintf("%s [%s/%s/%s] server=%s server_ips=%s bailiwick=%s", + r.RefID, r.Qname, ClassToString(r.Qclass), TypeToString(r.Qtype), + r.Server, r.TxtIPs(), r.Bailiwick) } -// IsNameInChain reports whether name appears anywhere in this referral's ancestor -// chain, including this referral itself. Used for CNAME loop detection. -func (r *Referral) IsNameInChain(name string) bool { - n := miekgdns.Fqdn(strings.ToLower(name)) - curr := r - for curr != nil { - if curr.Name == n { +// OverallStatus summarises the node's outcome for event consumers: a resolve +// dead end (noglue/loop), the shared status of every per-IP response, or +// "mixed" when the responses disagree ("" before anything was queried). +func (r *Referral) OverallStatus() Status { + switch r.Status { + case RefStatusNoGlue: + return StatusNoGlue + case RefStatusLoop: + return StatusLoop + } + var s Status + for _, resp := range r.Responses { + if s == "" { + s = resp.Status + } else if s != resp.Status { + return "mixed" + } + } + return s +} + +// StatsList returns the aggregated leaf statistics sorted by stats key. +func (r *Referral) StatsList() []*StatsEntry { + out := make([]*StatsEntry, 0, len(r.Stats)) + for _, e := range r.Stats { + out = append(out, e) + } + sort.Slice(out, func(i, j int) bool { return out[i].Key < out[j].Key }) + return out +} + +// isNoGlue reports a dead end: the server is inside the current bailiwick, +// so its address should have come as glue, but none was provided and no +// deeper zone can name it (referral.rb noglue?). +func (r *Referral) isNoGlue() bool { + return r.ServerIPs == nil && insideBailiwick(r.Server, r.Bailiwick) +} + +// isLoop reports a resolve loop: an ancestor referral asks the same +// qname/qclass/qtype of the same still-unresolved server (referral.rb loop?), +// e.g. b NS c.d while d NS a.b. +func (r *Referral) isLoop() bool { + if r.ServerIPs != nil { + return false + } + for p := r.Parent; p != nil; p = p.Parent { + if p.Qname == r.Qname && p.Qclass == r.Qclass && p.Qtype == r.Qtype && + p.Server == r.Server && p.ServerIPs == nil { return true } - curr = curr.Parent } return false } + +// chainHasQuery reports whether this referral or any ancestor already asks +// qname with the same qclass/qtype. Used to stop cross-response CNAME chains +// (restart loops) that Ruby only catches via the depth limit. +func (r *Referral) chainHasQuery(qname string) bool { + name := canonicalName(qname) + for p := r; p != nil; p = p.Parent { + if p.Qname == name && p.Qclass == r.Qclass && p.Qtype == r.Qtype { + return true + } + } + return false +} + +// resolve turns an address-less referral into either a dead end (noglue/ +// loop) or a resolve subtree querying A from this branch's cache +// (referral.rb resolve). It returns the referrals to process. +func (r *Referral) resolve() ([]*Referral, error) { + if r.isNoGlue() { + r.Status = RefStatusNoGlue + return nil, nil + } + if r.isLoop() { + r.Status = RefStatusLoop + return nil, nil + } + starters, newbailiwick, err := r.InfoCache.GetStartServers(r.Server) + if err != nil { + return nil, err + } + for i, st := range starters { + child := r.makeReferral(referralArgs{ + qname: r.Server, + qtype: r.NSAType, + server: st.Name, + serverIPs: st.IPs, + bailiwick: newbailiwick, + refid: fmt.Sprintf("%s.0.%d", r.RefID, i+1), + referralResolution: true, + }) + r.Resolves = append(r.Resolves, child) + } + return r.Resolves, nil +} + +// resolveCalculate folds the resolve subtree's statistics into per-IP server +// weights (referral.rb resolve_calculate): each answered leaf distributes +// its probability evenly across the A records it returned; every other leaf +// keeps its probability under its "key:" stats key so failures surface in +// the results. +func (r *Referral) resolveCalculate() { + r.StatsResolve = make(map[string]*StatsEntry) + switch r.Status { + case RefStatusNoGlue: + resp := NewNoGlueResponse(r.Qname, r.Qclass, r.Qtype, r.ParentIP, r.Server, r.Bailiwick) + key := resp.StatsKey() + r.StatsResolve[key] = &StatsEntry{Key: key, Prob: 1.0, Response: resp, Referral: r} + case RefStatusLoop: + resp := NewLoopResponse(r.Qname, r.Qclass, r.Qtype, r.ParentIP, r.Server, r.Bailiwick) + key := resp.StatsKey() + r.StatsResolve[key] = &StatsEntry{Key: key, Prob: 1.0, Response: resp, Referral: r} + default: + statsCalculateChildren(r.StatsResolve, r.Resolves, 1.0) + } + + r.ServerWeights = make(map[string]float64) + r.ServerIPs = []string{} + keys := make([]string, 0, len(r.StatsResolve)) + for key := range r.StatsResolve { + keys = append(keys, key) + } + sort.Strings(keys) + addWeight := func(ip string, prob float64) { + if _, ok := r.ServerWeights[ip]; !ok { + r.ServerIPs = append(r.ServerIPs, ip) + } + r.ServerWeights[ip] += prob + } + for _, key := range keys { + data := r.StatsResolve[key] + if data.Response.Status == StatusAnswered { + var addrs []string + for _, rr := range data.Response.DQ.Answers { + if a, ok := rr.(*miekgdns.A); ok { + addrs = append(addrs, a.A.String()) + } + } + for _, addr := range addrs { + addWeight(addr, data.Prob/float64(len(addrs))) + } + if len(addrs) == 0 { + // answered but no A records (e.g. AAAA-only): carry the + // probability as a failure key so mass is not lost. + addWeight(key, data.Prob) + } + } else { + addWeight(key, data.Prob) + } + } +} + +// statsCalculateChildren merges the children's statistics into stats with an +// equal split of weight among them (referral.rb stats_calculate_children). +func statsCalculateChildren(stats map[string]*StatsEntry, children []*Referral, weight float64) { + if len(children) == 0 { + return + } + percent := (1.0 / float64(len(children))) * weight + for _, child := range children { + for key, data := range child.Stats { + if e, ok := stats[key]; ok { + e.Prob += data.Prob * percent + } else { + stats[key] = &StatsEntry{ + Key: key, + Prob: data.Prob * percent, + Response: data.Response, + Referral: data.Referral, + } + } + } + } +} + +// answerCalculate computes this node's aggregated statistics from its +// children, responses and resolve failures (referral.rb answer_calculate). +// Unlike the Ruby source, duplicate stats keys across a referral's IPs merge +// by summing probability (Ruby computes the sum then discards it — a source +// bug that breaks the probabilities-sum-to-1 invariant). +func (r *Referral) answerCalculate() { + r.Stats = make(map[string]*StatsEntry) + if r.IsRootRoot() { + statsCalculateChildren(r.Stats, r.Children["rootroot"], 1.0) + r.calculated = true + return + } + for _, ip := range r.ServerIPs { + serverweight := r.ServerWeights[ip] + if strings.HasPrefix(ip, "key:") { + // resolve failed for some reason - copy the resolve statistics + src := r.StatsResolve[ip] + if e, ok := r.Stats[ip]; ok { + e.Prob += src.Prob + } else { + r.Stats[ip] = &StatsEntry{Key: ip, Prob: src.Prob, Response: src.Response, Referral: src.Referral} + } + continue + } + if children := r.Children[ip]; len(children) > 0 { + statsCalculateChildren(r.Stats, children, serverweight) + continue + } + resp := r.Responses[ip] + if resp == nil { + continue + } + key := resp.StatsKey() + if e, ok := r.Stats[key]; ok { + e.Prob += serverweight + } else { + r.Stats[key] = &StatsEntry{Key: key, Prob: serverweight, Response: resp, Referral: r} + } + } + r.calculated = true +} + +// process queries every real IP of this referral through the packet cache, +// classifies each response, and creates one child per NS name (including +// glueless ones) for referral/restart statuses (referral.rb process/ +// process_normal). It returns one set of children per IP that produced any. +func (r *Referral) process(ctx context.Context) ([][]*Referral, error) { + r.processed = true + if r.IsRootRoot() { + children, err := r.processAddRoots() + if err != nil { + return nil, err + } + return [][]*Referral{children}, nil + } + + // Phase one: query and classify, counting childsets so refids can grow + // an extra childset digit when more than one IP produces children. + childsets := 0 + var order []string + for _, ip := range r.ServerIPs { + if strings.HasPrefix(ip, "key:") { + continue + } + var dq *DecodedQuery + if r.Depth() >= r.maxdepth { + err := fmt.Errorf("Maxdepth %d exceeded", r.maxdepth) + dq = NewDecodedQuery(nil, err, r.Qname, r.Qclass, r.Qtype, ip, r.Bailiwick) + } else { + msg, warnings, err := r.client.Query(ctx, net.ParseIP(ip), r.Qname, r.Qtype) + dq = NewDecodedQuery(msg, err, r.Qname, r.Qclass, r.Qtype, ip, r.Bailiwick) + dq.WarningsAdd(warnings...) + } + resp, err := NewServerResponse(dq, r.Server, r.ParentIP, r.InfoCache) + if err != nil { + return nil, err + } + if resp.Status == StatusRestart { + // Cross-response CNAME loop: any target in the chain that we + // (or an ancestor) are already querying is a dead end. + for _, target := range dq.ChainTargets { + if r.chainHasQuery(target) { + resp.Status = StatusCNAMELoop + break + } + } + } + r.Warnings = append(r.Warnings, dq.Warnings...) + r.Responses[ip] = resp + order = append(order, ip) + if resp.Status == StatusRestart || resp.Status == StatusReferral { + childsets++ + } + } + + // Phase two: create the children. + childset := 0 + var sets [][]*Referral + for _, ip := range order { + resp := r.Responses[ip] + if resp.Status != StatusRestart && resp.Status != StatusReferral { + continue + } + childset++ + refid := r.RefID + if childsets > 1 { + refid = fmt.Sprintf("%s.%d", r.RefID, childset) + } + children := r.makeReferrals(resp, refid, ip) + r.Children[ip] = children + sets = append(sets, children) + } + return sets, nil +} + +// processAddRoots creates one child per root server with equal weight +// (referral.rb process_add_roots); the roots come from the info cache hints. +func (r *Referral) processAddRoots() ([]*Referral, error) { + starters, _, err := r.InfoCache.GetStartServers("") + if err != nil { + return nil, err + } + dot := "" + if r.RefID != "" { + dot = "." + } + var children []*Referral + for i, root := range starters { + child := r.makeReferral(referralArgs{ + server: root.Name, + serverIPs: root.IPs, + refid: fmt.Sprintf("%s%s%d", r.RefID, dot, i+1), + }) + children = append(children, child) + } + r.Children["rootroot"] = children + return children, nil +} + +// makeReferrals creates one child per start server for a referral/restart +// response (referral.rb make_referrals): qname moves to the response's +// endname (the CNAME target on restart), the bailiwick and cache come from +// the response. +func (r *Referral) makeReferrals(resp *ServerResponse, refid, parentIP string) []*Referral { + var children []*Referral + for i, st := range resp.Starters { + children = append(children, r.makeReferral(referralArgs{ + qname: resp.DQ.Endname, + server: st.Name, + serverIPs: st.IPs, + bailiwick: resp.StartersBailiwick, + infoCache: resp.Cache, + refid: fmt.Sprintf("%s.%d", refid, i+1), + parentIP: parentIP, + })) + } + return children +} + +// replaceChild swaps before for after in the children/resolves lists (fast +// mode substitution); before keeps a pointer to its replacement. +func (r *Referral) replaceChild(before, after *Referral) { + before.ReplacedBy = after + for ip := range r.Children { + for i, c := range r.Children[ip] { + if c == before { + r.Children[ip][i] = after + } + } + } + for i, c := range r.Resolves { + if c == before { + r.Resolves[i] = after + } + } +} diff --git a/internal/traverse/referral_test.go b/internal/traverse/referral_test.go index 5e6ecd1..2e95468 100644 --- a/internal/traverse/referral_test.go +++ b/internal/traverse/referral_test.go @@ -1,159 +1,194 @@ package traverse import ( - "net" "testing" - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" + "github.com/miekg/dns" ) -const TypeA = dns.TypeA - -type State = ResolutionState - -func TestNewReferral(t *testing.T) { - ref := NewReferral("example.com", TypeA, ".", 0, 1.0, nil) - if ref.Name != "example.com." { - t.Errorf("Name = %q, want %q", ref.Name, "example.com.") +func newTestReferral(server string, ips []string) *Referral { + r := &Referral{ + RefID: "1", + Qname: "www.example.com", + Qclass: dns.ClassINET, + Qtype: dns.TypeA, + NSAType: dns.TypeA, + Server: server, + ServerIPs: ips, + Bailiwick: "com", + InfoCache: NewInfoCache(nil), + Status: RefStatusNormal, + Responses: make(map[string]*ServerResponse), + Children: make(map[string][]*Referral), + ServerWeights: make(map[string]float64), + maxdepth: DefaultMaxDepth, } - if ref.Qtype != TypeA { - t.Errorf("Qtype = %d, want %d", ref.Qtype, TypeA) - } - if ref.Qclass != 1 { - t.Errorf("Qclass = %d, want 1", ref.Qclass) - } - if ref.State != StateUnresolved { - t.Errorf("State = %d, want %d", ref.State, StateUnresolved) - } - if ref.Depth != 0 { - t.Errorf("Depth = %d, want 0", ref.Depth) - } - if ref.Prob != 1.0 { - t.Errorf("Prob = %f, want 1.0", ref.Prob) - } - if ref.Parent != nil { - t.Error("Parent should be nil") + for _, ip := range ips { + r.ServerWeights[ip] = 1.0 / float64(len(ips)) } + return r } -func TestReferralInBailiwick(t *testing.T) { +func TestRefidDepth(t *testing.T) { tests := []struct { - name string - bailiwick string - testName string - want bool + refid string + want int }{ - {"root bailiwick accepts all", ".", "example.com.", true}, - {"empty bailiwick accepts all", "", "example.com.", true}, - {"subdomain in bailiwick", "com.", "example.com.", true}, - {"deeper subdomain", "com.", "www.example.com.", true}, - {"not in bailiwick", "org.", "example.com.", false}, - {"same zone", "example.com.", "example.com.", true}, - {"sibling zone", "example.com.", "other.com.", false}, + {"", 0}, + {"1", 1}, + {"1.1.2", 3}, + {"1.2.0.1", 3}, + {"1.1.2.0.1.4.2.0.2.0.2", 8}, } - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - ref := &Referral{Bailiwick: tt.bailiwick} - if got := ref.InBailiwick(tt.testName); got != tt.want { - t.Errorf("InBailiwick(%q) = %v, want %v", tt.testName, got, tt.want) - } - }) - } -} - -func TestReferralHasAddresses(t *testing.T) { - ref := &Referral{} - if ref.HasAddresses() { - t.Error("empty referral should not have addresses") - } - - ref.Addresses = []net.IP{net.ParseIP("1.2.3.4")} - if !ref.HasAddresses() { - t.Error("referral with address should have addresses") - } -} - -func TestReferralSetAddresses(t *testing.T) { - ref := &Referral{} - - ref.SetAddresses([]net.IP{net.ParseIP("1.2.3.4")}) - if ref.State != StateResolved { - t.Errorf("State = %d, want %d", ref.State, StateResolved) - } - if !ref.HasAddresses() { - t.Error("should have addresses after SetAddresses") - } - - ref.SetAddresses(nil) - if ref.State != StateUnresolved { - t.Errorf("State = %d, want %d", ref.State, StateUnresolved) - } -} - -func TestResolutionStateString(t *testing.T) { - tests := []struct { - state State - want string - }{ - {StateUnresolved, "unresolved"}, - {StateResolving, "resolving"}, - {StateResolved, "resolved"}, - } - - for _, tt := range tests { - t.Run(tt.want, func(t *testing.T) { - if got := tt.state.String(); got != tt.want { - t.Errorf("String() = %q, want %q", got, tt.want) - } - }) - } -} - -func TestCircularReferralError(t *testing.T) { - err := &CircularReferralError{ - Name: "ns.example.com.", - Chain: []string{"ns1.example.com.", "ns2.example.com."}, - } - - expected := "circular referral detected for ns.example.com.: [ns1.example.com. ns2.example.com.]" - if err.Error() != expected { - t.Errorf("Error() = %q, want %q", err.Error(), expected) - } -} - -func TestUnresolvableNameserverError(t *testing.T) { - err := &UnresolvableNameserverError{ - Name: "ns.example.com.", - Reason: "NXDOMAIN", - } - - expected := "unresolvable nameserver ns.example.com.: NXDOMAIN" - if err.Error() != expected { - t.Errorf("Error() = %q, want %q", err.Error(), expected) - } -} - -func TestGetVisitedNames(t *testing.T) { - visited := map[string]bool{ - "ns1.example.com.": true, - "ns2.example.com.": true, - "ns3.example.com.": true, - } - - names := getVisitedNames(visited) - if len(names) != 3 { - t.Errorf("got %d names, want 3", len(names)) - } - - seen := make(map[string]bool) - for _, name := range names { - if seen[name] { - t.Errorf("duplicate name: %s", name) - } - seen[name] = true - if !visited[name] { - t.Errorf("unexpected name: %s", name) + if got := refidDepth(tt.refid); got != tt.want { + t.Errorf("refidDepth(%q) = %d, want %d", tt.refid, got, tt.want) } } } + +func TestTxtIPsVerbose(t *testing.T) { + r := newTestReferral("ns1.example.com", []string{"2.2.2.2", "1.1.1.1"}) + if got := r.TxtIPsVerbose(); got != "50.0%=1.1.1.1,50.0%=2.2.2.2" { + t.Errorf("TxtIPsVerbose = %q", got) + } + + // key: pseudo entries keep only the first two fields. + r2 := newTestReferral("ns2.example.com", nil) + r2.ServerIPs = []string{"key:noglue:9.9.9.9:www.example.com:IN:A:x:y"} + r2.ServerWeights = map[string]float64{r2.ServerIPs[0]: 1.0} + if got := r2.TxtIPsVerbose(); got != "100.0%=noglue:9.9.9.9" { + t.Errorf("TxtIPsVerbose key entry = %q", got) + } + + var unresolved Referral + if got := unresolved.TxtIPsVerbose(); got != "" { + t.Errorf("unresolved TxtIPsVerbose = %q, want empty", got) + } +} + +func TestFastKeyLowercases(t *testing.T) { + r := newTestReferral("NS1.Example.COM", []string{"1.1.1.1"}) + r.Server = "NS1.Example.COM" // bypass canonicalisation to prove downcasing + key := fastKey(r) + if key != "www.example.com:in:a:ns1.example.com:100.0%=1.1.1.1" { + t.Errorf("fastKey = %q", key) + } +} + +func TestIsNoGlueAndIsLoop(t *testing.T) { + inBailiwick := newTestReferral("ns1.sub.com", nil) + if !inBailiwick.isNoGlue() { + t.Error("in-bailiwick NS without addresses should be noglue") + } + outOfBailiwick := newTestReferral("ns1.other.net", nil) + if outOfBailiwick.isNoGlue() { + t.Error("out-of-bailiwick NS is resolvable, not noglue") + } + resolved := newTestReferral("ns1.sub.com", []string{"1.1.1.1"}) + if resolved.isNoGlue() || resolved.isLoop() { + t.Error("resolved referral is neither noglue nor loop") + } + + parent := newTestReferral("ns.b.net", nil) + parent.Qname = "ns.a.net" + child := newTestReferral("ns.b.net", nil) + child.Qname = "ns.a.net" + child.Parent = parent + if !child.isLoop() { + t.Error("same qname/qclass/qtype/server with unresolved ancestor should loop") + } + child.Server = "ns.c.net" + if child.isLoop() { + t.Error("different server must not loop") + } +} + +func TestChainHasQuery(t *testing.T) { + parent := newTestReferral("ns1.example.com", []string{"1.1.1.1"}) + parent.Qname = "www.a.com" + child := newTestReferral("ns2.example.com", []string{"2.2.2.2"}) + child.Qname = "www.b.net" + child.Parent = parent + if !child.chainHasQuery("www.a.com") { + t.Error("ancestor qname should be found") + } + if !child.chainHasQuery("WWW.B.NET.") { + t.Error("own qname should be found case-insensitively") + } + if child.chainHasQuery("www.c.org") { + t.Error("unknown qname must not match") + } +} + +func TestReplaceChild(t *testing.T) { + parent := newTestReferral("ns1.example.com", []string{"1.1.1.1"}) + before := newTestReferral("ns2.example.com", []string{"2.2.2.2"}) + after := newTestReferral("ns2.example.com", []string{"2.2.2.2"}) + parent.Children["1.1.1.1"] = []*Referral{before} + parent.Resolves = []*Referral{before} + + parent.replaceChild(before, after) + if parent.Children["1.1.1.1"][0] != after || parent.Resolves[0] != after { + t.Error("replaceChild did not swap the node everywhere") + } + if before.ReplacedBy != after { + t.Error("replaced node should point at its replacement") + } +} + +func TestOverallStatus(t *testing.T) { + r := newTestReferral("ns1.example.com", nil) + r.Status = RefStatusNoGlue + if got := r.OverallStatus(); got != StatusNoGlue { + t.Errorf("noglue overall = %q", got) + } + r.Status = RefStatusLoop + if got := r.OverallStatus(); got != StatusLoop { + t.Errorf("loop overall = %q", got) + } + + n := newTestReferral("ns1.example.com", []string{"1.1.1.1", "2.2.2.2"}) + if got := n.OverallStatus(); got != "" { + t.Errorf("no responses overall = %q, want empty", got) + } + n.Responses["1.1.1.1"] = &ServerResponse{Status: StatusAnswered} + if got := n.OverallStatus(); got != StatusAnswered { + t.Errorf("single status overall = %q", got) + } + n.Responses["2.2.2.2"] = &ServerResponse{Status: StatusError} + if got := n.OverallStatus(); got != "mixed" { + t.Errorf("mixed overall = %q", got) + } +} + +func TestIPsAsArraySkipsPseudoKeys(t *testing.T) { + r := newTestReferral("ns1.example.com", []string{"1.1.1.1", "key:noglue:2.2.2.2"}) + got := r.IPsAsArray() + if len(got) != 1 || got[0] != "1.1.1.1" { + t.Errorf("IPsAsArray = %v", got) + } +} + +func TestToASCII(t *testing.T) { + if got := toASCII("bücher.example"); got != "xn--bcher-kva.example" { + t.Errorf("toASCII = %q", got) + } + if got := toASCII("plain.example"); got != "plain.example" { + t.Errorf("ascii name changed: %q", got) + } + if got := toASCII(""); got != "" { + t.Errorf("empty name changed: %q", got) + } +} + +func TestServerResponseString(t *testing.T) { + noglue := NewNoGlueResponse("www.example.com", dns.ClassINET, dns.TypeA, "1.1.1.1", "ns1.example.com", "example.com") + if got := noglue.String(); got != "No glue for ns1.example.com" { + t.Errorf("noglue String = %q", got) + } + loop := NewLoopResponse("www.example.com", dns.ClassINET, dns.TypeA, "1.1.1.1", "ns1.example.com", "example.com") + if got := loop.String(); got != "Loop encountered resolving ns1.example.com" { + t.Errorf("loop String = %q", got) + } +} diff --git a/internal/traverse/response.go b/internal/traverse/response.go deleted file mode 100644 index 7d8371f..0000000 --- a/internal/traverse/response.go +++ /dev/null @@ -1,256 +0,0 @@ -package traverse - -import ( - "net" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" - miekgdns "github.com/miekg/dns" -) - -type ResponseType int - -const ( - RespReferral ResponseType = iota - RespAnswer - RespCNAMEFollow - RespNODATA - RespNXDOMAIN - RespSERVFAIL - RespREFUSED - RespNOTIMPL - RespCNAMELoop - RespError - // RespNSResolutionFailed indicates that the traversal could not resolve the - // IP address of an in-bailiwick nameserver. The domain may still be - // reachable in practice (e.g. via glue records held by the registry), but - // the iterative traversal could not complete that path. - RespNSResolutionFailed -) - -func (rt ResponseType) String() string { - switch rt { - case RespReferral: - return "referral" - case RespAnswer: - return "answer" - case RespCNAMEFollow: - return "cname_follow" - case RespNODATA: - return "nodata" - case RespNXDOMAIN: - return "nxdomain" - case RespSERVFAIL: - return "servfail" - case RespREFUSED: - return "refused" - case RespNOTIMPL: - return "notimp" - case RespCNAMELoop: - return "cname_loop" - case RespError: - return "error" - case RespNSResolutionFailed: - return "ns_error" - default: - return "unknown" - } -} - -type Response struct { - Referral *Referral - Server net.IP - Cache *InfoCache - Decoded *dns.DecodedResponse - Type ResponseType - ErrorMessage string -} - -func NewResponse(ref *Referral, server net.IP, cache *InfoCache) *Response { - return &Response{ - Referral: ref, - Server: server, - Cache: cache, - } -} - -func (r *Response) Process(msg *miekgdns.Msg) *Response { - if msg == nil { - r.Type = RespError - r.ErrorMessage = "nil DNS response" - return r - } - - r.Decoded = dns.DecodeResponse(msg) - if r.Decoded == nil { - r.Type = RespError - r.ErrorMessage = "failed to decode DNS response" - return r - } - - // Synthesize CNAME from DNAME when the server didn't include a synthesized CNAME record. - if len(r.Decoded.CNAMEChain) == 0 && r.Referral != nil && len(r.Decoded.DNAMEMappings) > 0 { - for _, dm := range r.Decoded.DNAMEMappings { - synthesized := dns.SynthesizeCNAMEFromDNAME(r.Referral.Name, dm.Owner, dm.Target) - if synthesized != "" { - r.Decoded.CNAMEChain = append(r.Decoded.CNAMEChain, synthesized) - break - } - } - } - - r.Type = r.classify() - return r -} - -func (r *Response) classify() ResponseType { - switch r.Decoded.Classification { - case dns.ResponseNXDOMAIN: - return RespNXDOMAIN - case dns.ResponseSERVFAIL: - return RespSERVFAIL - case dns.ResponseREFUSED: - return RespREFUSED - case dns.ResponseNOTIMPL: - return RespNOTIMPL - case dns.ResponseAnswer: - if len(r.Decoded.CNAMEChain) > 0 && !r.hasFinalAnswer() { - return RespCNAMEFollow - } - return RespAnswer - case dns.ResponseReferral: - return RespReferral - case dns.ResponseNODATA: - return RespNODATA - default: - return RespError - } -} - -func (r *Response) hasFinalAnswer() bool { - for _, rr := range r.Decoded.Answers { - switch rr.(type) { - case *miekgdns.CNAME, *miekgdns.DNAME, *miekgdns.RRSIG: - // CNAME and DNAME are redirect records, not final answers. - // RRSIG is a DNSSEC signature record — it covers the CNAME/DNAME - // but is not itself the answer to the original question type. - continue - } - return true - } - return false -} - -func (r *Response) ChildReferrals() []*Referral { - if r.Type != RespReferral { - return nil - } - if r.Referral == nil { - return nil - } - - var nameservers []string - for _, rr := range r.Decoded.Authority { - if ns, ok := rr.(*miekgdns.NS); ok { - if r.Referral.InBailiwick(ns.Ns) { - nameservers = append(nameservers, ns.Ns) - } - } - } - - if len(nameservers) == 0 { - for _, rr := range r.Decoded.Authority { - if ns, ok := rr.(*miekgdns.NS); ok { - nameservers = append(nameservers, ns.Ns) - } - } - } - - r.storeAuthority(nameservers) - - prob := r.childProb(len(nameservers)) - var children []*Referral - for _, ns := range nameservers { - child := NewReferral( - r.Referral.Name, - r.Referral.Qtype, - ns, - r.Referral.Depth+1, - prob, - r.Referral, - ) - r.resolveGlue(child) - children = append(children, child) - } - return children -} - -func (r *Response) CNAMEFollowReferral() *Referral { - if r.Type != RespCNAMEFollow || len(r.Decoded.CNAMEChain) == 0 { - return nil - } - target := r.Decoded.CNAMEChain[len(r.Decoded.CNAMEChain)-1] - follow := NewReferral( - target, - r.Referral.Qtype, - r.Referral.Bailiwick, - r.Referral.Depth+1, - r.Referral.Prob, - r.Referral, - ) - if len(r.Referral.Addresses) > 0 { - follow.Addresses = make([]net.IP, len(r.Referral.Addresses)) - copy(follow.Addresses, r.Referral.Addresses) - follow.State = StateResolved - } - return follow -} - -func (r *Response) storeAuthority(nameservers []string) { - if r.Cache == nil { - return - } - zone := r.Referral.Name - r.Cache.StoreNS(zone, nameservers) -} - -func (r *Response) resolveGlue(child *Referral) { - if r.Cache == nil { - return - } - nsName := child.Bailiwick - for _, rr := range r.Decoded.Additional { - switch v := rr.(type) { - case *miekgdns.A: - if normalize(v.Header().Name) == normalize(nsName) { - child.Addresses = append(child.Addresses, v.A) - } - case *miekgdns.AAAA: - if normalize(v.Header().Name) == normalize(nsName) { - child.Addresses = append(child.Addresses, v.AAAA) - } - } - } - if child.HasAddresses() { - child.State = StateResolved - } - r.Cache.StoreGlue(nsName, child.Addresses) -} - -func (r *Response) IsTerminal() bool { - switch r.Type { - case RespAnswer, RespNODATA, RespNXDOMAIN, RespSERVFAIL, RespREFUSED, RespNOTIMPL, RespCNAMELoop, RespError, RespNSResolutionFailed: - return true - default: - return false - } -} - -func (r *Response) childProb(n int) float64 { - if n <= 0 { - return 0 - } - if r.Referral == nil { - return 1.0 / float64(n) - } - return r.Referral.Prob / float64(n) -} diff --git a/internal/traverse/response_test.go b/internal/traverse/response_test.go deleted file mode 100644 index a7bfac1..0000000 --- a/internal/traverse/response_test.go +++ /dev/null @@ -1,340 +0,0 @@ -package traverse - -import ( - "net" - "testing" - - "github.com/miekg/dns" -) - -func TestResponseProcessNil(t *testing.T) { - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(nil) - if r.Type != RespError { - t.Errorf("Type = %d, want %d", r.Type, RespError) - } -} - -func TestResponseClassifyAnswer(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespAnswer { - t.Errorf("Type = %d, want %d", r.Type, RespAnswer) - } - if r.Decoded == nil { - t.Fatal("Decoded should not be nil") - } -} - -func TestResponseClassifyReferral(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = false - msg.Ns = append(msg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS, Class: dns.ClassINET, Ttl: 172800}, - Ns: "a.gtld-servers.net.", - }) - msg.Extra = append(msg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 172800}, - A: net.ParseIP("192.5.6.30"), - }) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - cache := NewInfoCache(nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), cache) - r.Process(msg) - - if r.Type != RespReferral { - t.Errorf("Type = %d, want %d", r.Type, RespReferral) - } -} - -func TestResponseClassifyNXDOMAIN(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeNameError - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespNXDOMAIN { - t.Errorf("Type = %d, want %d", r.Type, RespNXDOMAIN) - } -} - -func TestResponseClassifyNODATA(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = true - msg.Ns = append(msg.Ns, &dns.SOA{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeSOA, Class: dns.ClassINET, Ttl: 3600}, - }) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespNODATA { - t.Errorf("Type = %d, want %d", r.Type, RespNODATA) - } -} - -func TestResponseClassifySERVFAIL(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeServerFailure - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespSERVFAIL { - t.Errorf("Type = %d, want %d", r.Type, RespSERVFAIL) - } -} - -func TestResponseCNAMEFollow(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "example.com.", - }, - ) - - ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespCNAMEFollow { - t.Errorf("Type = %d, want %d", r.Type, RespCNAMEFollow) - } -} - -func TestResponseCNAMEWithFinalAnswer(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "example.com.", - }, - &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }, - ) - - ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if r.Type != RespAnswer { - t.Errorf("Type = %d, want %d (CNAME with final A answer)", r.Type, RespAnswer) - } -} - -func TestResponseChildReferrals(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = false - msg.Ns = append(msg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, Ns: "a.gtld-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeNS}, Ns: "b.gtld-servers.net."}, - ) - msg.Extra = append(msg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("192.5.6.30")}, - &dns.A{Hdr: dns.RR_Header{Name: "b.gtld-servers.net.", Rrtype: dns.TypeA}, A: net.ParseIP("192.33.14.30")}, - ) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - cache := NewInfoCache(nil) - r := NewResponse(ref, net.ParseIP("198.41.0.4"), cache) - r.Process(msg) - - children := r.ChildReferrals() - if len(children) != 2 { - t.Fatalf("expected 2 child referrals, got %d", len(children)) - } - - if children[0].Name != "example.com." { - t.Errorf("child[0] name = %q, want example.com.", children[0].Name) - } - if children[0].Bailiwick != "a.gtld-servers.net." { - t.Errorf("child[0] bailiwick = %q, want a.gtld-servers.net.", children[0].Bailiwick) - } - if children[0].Prob != 0.5 { - t.Errorf("child[0] prob = %f, want 0.5", children[0].Prob) - } - if children[0].Depth != 1 { - t.Errorf("child[0] depth = %d, want 1", children[0].Depth) - } - if !children[0].HasAddresses() { - t.Error("child[0] should have glue addresses") - } - - if !children[1].HasAddresses() { - t.Error("child[1] should have glue addresses") - } - - nsNames := cache.LookupNS("example.com.") - if len(nsNames) != 2 { - t.Errorf("expected 2 NS in cache, got %d", len(nsNames)) - } -} - -func TestResponseChildReferralsNonReferral(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, &dns.A{ - Hdr: dns.RR_Header{Rrtype: dns.TypeA}, - A: net.ParseIP("1.2.3.4"), - }) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - if children := r.ChildReferrals(); children != nil { - t.Error("non-referral should not produce child referrals") - } -} - -func TestResponseCNAMEFollowReferral(t *testing.T) { - msg := new(dns.Msg) - msg.SetReply(new(dns.Msg)) - msg.Answer = append(msg.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME}, - Target: "example.com.", - }, - ) - - ref := NewReferral("www.example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - follow := r.CNAMEFollowReferral() - if follow == nil { - t.Fatal("expected CNAME follow referral") - } - if follow.Name != "example.com." { - t.Errorf("follow name = %q, want %q", follow.Name, "example.com.") - } - if follow.Depth != 1 { - t.Errorf("follow depth = %d, want 1", follow.Depth) - } -} - -func TestResponseIsTerminal(t *testing.T) { - tests := []struct { - respType ResponseType - want bool - }{ - {RespAnswer, true}, - {RespNODATA, true}, - {RespNXDOMAIN, true}, - {RespSERVFAIL, true}, - {RespError, true}, - {RespNSResolutionFailed, true}, - {RespReferral, false}, - {RespCNAMEFollow, false}, - } - - for _, tt := range tests { - t.Run(tt.respType.String(), func(t *testing.T) { - r := &Response{Type: tt.respType} - if got := r.IsTerminal(); got != tt.want { - t.Errorf("IsTerminal() = %v, want %v", got, tt.want) - } - }) - } -} - -func TestResponseTypeString(t *testing.T) { - tests := []struct { - rt ResponseType - want string - }{ - {RespReferral, "referral"}, - {RespAnswer, "answer"}, - {RespCNAMEFollow, "cname_follow"}, - {RespNODATA, "nodata"}, - {RespNXDOMAIN, "nxdomain"}, - {RespSERVFAIL, "servfail"}, - {RespError, "error"}, - {RespNSResolutionFailed, "ns_error"}, - } - - for _, tt := range tests { - t.Run(tt.want, func(t *testing.T) { - if got := tt.rt.String(); got != tt.want { - t.Errorf("String() = %q, want %q", got, tt.want) - } - }) - } -} - -func TestResponseChildReferralsProbabilityInheritance(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = false - msg.Ns = append(msg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "a.gtld-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "b.gtld-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: "com.", Rrtype: dns.TypeNS}, Ns: "c.gtld-servers.net."}, - ) - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 0.5, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - children := r.ChildReferrals() - if len(children) != 3 { - t.Fatalf("expected 3 children, got %d", len(children)) - } - - for _, c := range children { - if c.Prob != 0.5/3.0 { - t.Errorf("child prob = %f, want %f", c.Prob, 0.5/3.0) - } - } -} - -func TestResponseChildReferralsEmptyAuthority(t *testing.T) { - msg := new(dns.Msg) - msg.Rcode = dns.RcodeSuccess - msg.Authoritative = false - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - r := NewResponse(ref, net.ParseIP("1.2.3.4"), nil) - r.Process(msg) - - children := r.ChildReferrals() - if len(children) != 0 { - t.Errorf("expected 0 children with empty authority, got %d", len(children)) - } -} - -func TestResponseNilReferral(t *testing.T) { - r := NewResponse(nil, net.ParseIP("1.2.3.4"), nil) - children := r.ChildReferrals() - if children != nil { - t.Error("nil referral should produce no children") - } - - follow := r.CNAMEFollowReferral() - if follow != nil { - t.Error("nil referral should produce no CNAME follow") - } -} diff --git a/internal/traverse/robustness_test.go b/internal/traverse/robustness_test.go deleted file mode 100644 index 4e1df9a..0000000 --- a/internal/traverse/robustness_test.go +++ /dev/null @@ -1,768 +0,0 @@ -package traverse - -import ( - "context" - "errors" - "net" - "sync/atomic" - "testing" - - "github.com/miekg/dns" -) - -// TestCNAMELoopDetected verifies that a two-step CNAME loop (A → B → A) is -// detected without infinite recursion and produces a RespCNAMELoop result. -func TestCNAMELoopDetected(t *testing.T) { - // www.example.com → CNAME → alias.example.com → CNAME → www.example.com (loop) - cnameToAlias := new(dns.Msg) - cnameToAlias.SetReply(new(dns.Msg)) - cnameToAlias.Answer = append(cnameToAlias.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "alias.example.com.", - }) - - cnameBack := new(dns.Msg) - cnameBack.SetReply(new(dns.Msg)) - cnameBack.Answer = append(cnameBack.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "alias.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "www.example.com.", - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 10, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - switch q.Name { - case "www.example.com.": - return cnameToAlias.Copy(), nil - case "alias.example.com.": - return cnameBack.Copy(), nil - } - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundLoop := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespCNAMELoop { - foundLoop = true - if r.Response.ErrorMessage == "" { - t.Error("expected non-empty ErrorMessage on CNAME loop result") - } - } - } - if !foundLoop { - t.Error("expected RespCNAMELoop result for CNAME loop A → B → A") - } -} - -// TestCNAMEDirectLoop verifies that a direct self-loop (A → A) is handled. -func TestCNAMEDirectLoop(t *testing.T) { - selfLoop := new(dns.Msg) - selfLoop.SetReply(new(dns.Msg)) - selfLoop.Answer = append(selfLoop.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, - Target: "www.example.com.", - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 10, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return selfLoop.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundLoop := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespCNAMELoop { - foundLoop = true - } - } - if !foundLoop { - t.Error("expected RespCNAMELoop for direct self-referencing CNAME") - } -} - -// TestREFUSEDResponse verifies that a REFUSED rcode is classified as RespREFUSED. -func TestREFUSEDResponse(t *testing.T) { - refusedResp := new(dns.Msg) - refusedResp.Rcode = dns.RcodeRefused - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return refusedResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespREFUSED { - t.Errorf("Type = %s, want refused", results[0].Response.Type) - } - if !results[0].Response.IsTerminal() { - t.Error("REFUSED should be a terminal response") - } -} - -// TestNOTIMPLResponse verifies that a NOTIMP rcode is classified as RespNOTIMPL. -func TestNOTIMPLResponse(t *testing.T) { - notImplResp := new(dns.Msg) - notImplResp.Rcode = dns.RcodeNotImplemented - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return notImplResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespNOTIMPL { - t.Errorf("Type = %s, want notimp", results[0].Response.Type) - } - if !results[0].Response.IsTerminal() { - t.Error("NOTIMP should be a terminal response") - } -} - -// TestGracefulDegradationUnreachableServer verifies that when some servers are -// unreachable, traversal continues with the remaining servers and does not panic. -func TestGracefulDegradationUnreachableServer(t *testing.T) { - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - // Referral with two nameservers; first always fails, second provides the answer. - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - referralMsg.Ns = append(referralMsg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, - ) - referralMsg.Extra = append(referralMsg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, - ) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - switch server { - case "198.41.0.4": - return referralMsg.Copy(), nil - case "10.0.0.1": - return nil, errors.New("connection refused") - case "10.0.0.2": - return answerResp.Copy(), nil - } - return nil, errors.New("unexpected server") - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundAnswer := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespAnswer { - foundAnswer = true - } - } - if !foundAnswer { - t.Error("expected an answer result from the reachable server") - } -} - -// TestGracefulDegradationAllUnreachable verifies that when ALL servers fail, -// the traversal returns a SERVFAIL result without panicking. -func TestGracefulDegradationAllUnreachable(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, errors.New("network unreachable") - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("traversal must not return a top-level error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least one result even on total failure") - } - last := results[len(results)-1] - if last.Response == nil { - t.Fatal("last result must have a response") - } - if last.Response.Type != RespSERVFAIL && last.Response.Type != RespError { - t.Errorf("expected SERVFAIL or error when all servers unreachable, got %s", last.Response.Type) - } -} - -// TestDNAMEFollowNoSynthesizedCNAME verifies that a DNAME record in the answer -// section synthesizes a CNAME follow when the server doesn't include one. -func TestDNAMEFollowNoSynthesizedCNAME(t *testing.T) { - // Server returns DNAME only (no synthesized CNAME). - dnameResp := new(dns.Msg) - dnameResp.SetReply(new(dns.Msg)) - dnameResp.Answer = append(dnameResp.Answer, &dns.DNAME{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeDNAME, Class: dns.ClassINET, Ttl: 300}, - Target: "example.net.", - }) - - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "www.example.net.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("203.0.113.1"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "www.example.com." { - return dnameResp.Copy(), nil - } - if q.Name == "www.example.net." { - return answerResp.Copy(), nil - } - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundCNAMEFollow := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespCNAMEFollow { - foundCNAMEFollow = true - } - } - if !foundCNAMEFollow { - t.Error("expected RespCNAMEFollow synthesized from DNAME record") - } -} - -// TestIsNameInChain verifies the ancestor chain lookup. -func TestIsNameInChain(t *testing.T) { - root := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil) - child := NewReferral("www.example.com", dnsTypeA, "example.com.", 1, 1.0, root) - grandchild := NewReferral("sub.www.example.com", dnsTypeA, "www.example.com.", 2, 1.0, child) - - tests := []struct { - ref *Referral - name string - want bool - }{ - {grandchild, "sub.www.example.com", true}, // self - {grandchild, "www.example.com", true}, // parent - {grandchild, "example.com", true}, // grandparent - {grandchild, "other.example.com", false}, // not in chain - {root, "example.com", true}, // root matches itself - {root, "www.example.com", false}, // child not in chain from root - } - - for _, tt := range tests { - got := tt.ref.IsNameInChain(tt.name) - if got != tt.want { - t.Errorf("IsNameInChain(%q) from %q = %v, want %v", tt.name, tt.ref.Name, got, tt.want) - } - } -} - -// TestResponseTypeStrings verifies String() for new response types. -func TestResponseTypeStrings(t *testing.T) { - tests := []struct { - rt ResponseType - want string - }{ - {RespReferral, "referral"}, - {RespAnswer, "answer"}, - {RespCNAMEFollow, "cname_follow"}, - {RespNODATA, "nodata"}, - {RespNXDOMAIN, "nxdomain"}, - {RespSERVFAIL, "servfail"}, - {RespREFUSED, "refused"}, - {RespNOTIMPL, "notimp"}, - {RespCNAMELoop, "cname_loop"}, - {RespError, "error"}, - } - - for _, tt := range tests { - if got := tt.rt.String(); got != tt.want { - t.Errorf("ResponseType(%d).String() = %q, want %q", tt.rt, got, tt.want) - } - } -} - -// TestMalformedResponseNoPanic verifies that a nil response from the exchange -// function does not cause a panic, and produces an error result. -func TestMalformedResponseNoPanic(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil // nil response, no error - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least one result") - } - // Should produce error/servfail, not panic - for _, r := range results { - if r.Response == nil { - t.Error("result has nil response") - } - } -} - -// TestDNSSECRRSIGDoesNotBlockCNAMEFollow verifies that a DNSSEC RRSIG record -// accompanying a CNAME in the answer section is treated as metadata and does -// NOT prevent the traversal from following the CNAME. -func TestDNSSECRRSIGDoesNotBlockCNAMEFollow(t *testing.T) { - // Server returns CNAME + RRSIG (DNSSEC-signed zone response). - cnameWithRRSIG := new(dns.Msg) - cnameWithRRSIG.SetReply(new(dns.Msg)) - cnameWithRRSIG.Answer = append(cnameWithRRSIG.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeCNAME, Class: dns.ClassINET, Ttl: 300}, - Target: "example.com.", - }, - &dns.RRSIG{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dns.TypeRRSIG, Class: dns.ClassINET, Ttl: 300}, - TypeCovered: dns.TypeCNAME, - }, - ) - - finalAnswer := new(dns.Msg) - finalAnswer.SetReply(new(dns.Msg)) - finalAnswer.Answer = append(finalAnswer.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 10, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: true, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "www.example.com." { - return cnameWithRRSIG.Copy(), nil - } - return finalAnswer.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundCNAMEFollow := false - foundAnswer := false - for _, r := range results { - if r.Response != nil { - switch r.Response.Type { - case RespCNAMEFollow: - foundCNAMEFollow = true - case RespAnswer: - foundAnswer = true - } - } - } - if !foundCNAMEFollow { - t.Error("expected RespCNAMEFollow: RRSIG should not block CNAME following") - } - if !foundAnswer { - t.Error("expected final RespAnswer after CNAME follow") - } -} - -// TestFastModeOn verifies that Fast=true uses the shared root cache (default -// behaviour): a child branch can see glue stored by the root referral. -func TestFastModeOn(t *testing.T) { - // Root referral returns two nameservers with glue. Each NS branch returns - // an answer. We verify both branches are queried. - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - referralMsg.Ns = append(referralMsg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, - ) - referralMsg.Extra = append(referralMsg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, - ) - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: true, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - if server == "198.41.0.4" { - return referralMsg.Copy(), nil - } - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - answers := 0 - for _, r := range results { - if r.Response != nil && r.Response.Type == RespAnswer { - answers++ - } - } - if answers == 0 { - t.Error("expected at least one answer with Fast=true") - } -} - -// TestFastModeOff verifies that Fast=false gives each referral its own -// independent cache — no cross-branch glue contamination. -func TestFastModeOff(t *testing.T) { - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - referralMsg.Ns = append(referralMsg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, - ) - referralMsg.Extra = append(referralMsg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, - ) - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: false, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - if server == "198.41.0.4" { - return referralMsg.Copy(), nil - } - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - // Traversal must complete without panic and produce results. - if len(results) == 0 { - t.Fatal("expected at least one result with Fast=false") - } -} - -// TestFastModeDefaultIsTrue verifies that DefaultTraverserConfig has Fast=true. -func TestFastModeDefaultIsTrue(t *testing.T) { - cfg := DefaultTraverserConfig() - if !cfg.Fast { - t.Error("DefaultTraverserConfig().Fast should be true") - } -} - -// TestManyNSRecords verifies that a referral with more than 10 nameservers is -// handled gracefully — no panics, results are produced. -func TestManyNSRecords(t *testing.T) { - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - for i := 1; i <= 12; i++ { - ns := &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, - Ns: net.ParseIP(string(rune('a'+i-1))).String() + ".ns.example.com.", - } - // Use a distinct IP for each NS so glue is resolved. - ip := net.IP{10, 0, 0, byte(i)} - referralMsg.Ns = append(referralMsg.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, - Ns: ns.Ns, - }) - referralMsg.Extra = append(referralMsg.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: ns.Ns, Rrtype: dnsTypeA}, - A: ip, - }) - } - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - var queries int64 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - Fast: true, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - atomic.AddInt64(&queries, 1) - if server == "198.41.0.4" { - return referralMsg.Copy(), nil - } - return answerMsg.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error with 12 NS records: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results with many NS records") - } - foundAnswer := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespAnswer { - foundAnswer = true - } - } - if !foundAnswer { - t.Error("expected at least one answer from the 12-NS referral") - } -} - -// TestIDNPunycodeConversion verifies that a unicode (IDN) domain name is -// converted to its punycode/ACE form before querying. -func TestIDNPunycodeConversion(t *testing.T) { - // "münchen.de" → "xn--mnchen-3ya.de" (after punycode encoding) - ref := NewReferral("münchen.de", dnsTypeA, ".", 0, 1.0, nil) - if ref.Name == "münchen.de." { - t.Errorf("IDN name was not converted to punycode: got %q", ref.Name) - } - // Verify it starts with the expected punycode label. - if ref.Name != "xn--mnchen-3ya.de." { - t.Errorf("unexpected punycode result: got %q, want %q", ref.Name, "xn--mnchen-3ya.de.") - } -} - -// TestASCIIDomainUnchanged verifies that a plain ASCII domain is not mangled -// by the IDN conversion path. -func TestASCIIDomainUnchanged(t *testing.T) { - ref := NewReferral("example.com", dnsTypeA, ".", 0, 1.0, nil) - if ref.Name != "example.com." { - t.Errorf("ASCII domain was mangled: got %q, want %q", ref.Name, "example.com.") - } -} - -// TestWildcardResponse verifies that a wildcard answer (e.g. *.example.com -// returning an A record for sub.example.com) is handled as a regular answer. -func TestWildcardResponse(t *testing.T) { - wildcardAnswer := new(dns.Msg) - wildcardAnswer.SetReply(new(dns.Msg)) - wildcardAnswer.Authoritative = true - wildcardAnswer.Answer = append(wildcardAnswer.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "sub.example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("1.2.3.4"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return wildcardAnswer.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "sub.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results for wildcard response") - } - if results[0].Response.Type != RespAnswer { - t.Errorf("Type = %s, want answer", results[0].Response.Type) - } -} - -// TestLongCNAMEChainDepthLimit verifies that a very long CNAME chain is -// terminated by the MaxDepth limit without infinite recursion or a panic. -func TestLongCNAMEChainDepthLimit(t *testing.T) { - // Every query returns a CNAME to the next label. The MaxDepth setting - // must stop the chain. - counter := 0 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - counter++ - q := msg.Question[0] - resp := new(dns.Msg) - resp.SetReply(msg) - next := "next" + q.Name - resp.Answer = append(resp.Answer, &dns.CNAME{ - Hdr: dns.RR_Header{Name: q.Name, Rrtype: dnsTypeCNAME, Class: dns.ClassINET}, - Target: next, - }) - return resp, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "start.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results") - } - // Traversal must have stopped — counter should not be unbounded. - if counter > 50 { - t.Errorf("too many exchange calls (%d): chain depth limit not enforced", counter) - } -} - -// TestPartialBranchFailureReturnsResults verifies the graceful degradation -// requirement: when some NS branches fail completely, the partial results from -// successful branches are still returned. -func TestPartialBranchFailureReturnsResults(t *testing.T) { - // Three nameservers: first two error, third succeeds. - referralMsg := new(dns.Msg) - referralMsg.Rcode = dns.RcodeSuccess - referralMsg.Authoritative = false - referralMsg.Ns = append(referralMsg.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns1.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns2.example.com."}, - &dns.NS{Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeNS}, Ns: "ns3.example.com."}, - ) - referralMsg.Extra = append(referralMsg.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "ns1.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.1")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns2.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.2")}, - &dns.A{Hdr: dns.RR_Header{Name: "ns3.example.com.", Rrtype: dnsTypeA}, A: net.ParseIP("10.0.0.3")}, - ) - - answerMsg := new(dns.Msg) - answerMsg.SetReply(new(dns.Msg)) - answerMsg.Answer = append(answerMsg.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - switch server { - case "198.41.0.4": - return referralMsg.Copy(), nil - case "10.0.0.1", "10.0.0.2": - return nil, errors.New("server unreachable") - case "10.0.0.3": - return answerMsg.Copy(), nil - } - return nil, errors.New("unexpected server") - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("traversal must not return a top-level error: %v", err) - } - - foundAnswer := false - for _, r := range results { - if r.Response != nil && r.Response.Type == RespAnswer { - foundAnswer = true - } - } - if !foundAnswer { - t.Error("expected an answer from the third (reachable) nameserver despite others failing") - } -} diff --git a/internal/traverse/server_response.go b/internal/traverse/server_response.go new file mode 100644 index 0000000..efd18f8 --- /dev/null +++ b/internal/traverse/server_response.go @@ -0,0 +1,187 @@ +package traverse + +import ( + "fmt" + "sort" + + miekgdns "github.com/miekg/dns" +) + +// ServerResponse wraps a DecodedQuery for one (server, IP) query, mirroring +// response.rb: it owns a child InfoCache seeded with the in-bailiwick records, +// upgrades referral → referral_lame, and computes the start servers for +// referral/restart children. The noglue/loop variants (response_noglue.rb, +// response_loop.rb) are synthetic — no query was sent, DQ is nil. +type ServerResponse struct { + DQ *DecodedQuery + Status Status + + Qname string + Qclass uint16 + Qtype uint16 + IP string + Server string + Bailiwick string + // ParentIP is the address of the referring server; it is part of the + // stats key for referral_lame so lame referrals from different parents + // stay separate. + ParentIP string + + Cache *InfoCache + Starters []StartServer + StartersBailiwick string +} + +// NewServerResponse evaluates a decoded query in the context of parentCache. +// server is the NS hostname that was queried (dq.IP is its address). +func NewServerResponse(dq *DecodedQuery, server, parentIP string, parentCache *InfoCache) (*ServerResponse, error) { + r := &ServerResponse{ + DQ: dq, + Status: dq.Status, + Qname: dq.Qname, + Qclass: dq.Qclass, + Qtype: dq.Qtype, + IP: dq.IP, + Server: canonicalName(server), + Bailiwick: dq.Bailiwick, + ParentIP: parentIP, + Cache: NewInfoCache(parentCache), + } + if err := r.evaluate(); err != nil { + return nil, err + } + return r, nil +} + +// NewNoGlueResponse records a dead end: ip referred us to server inside +// bailiwick without glue and there is no way to resolve it. +func NewNoGlueResponse(qname string, qclass, qtype uint16, ip, server, bailiwick string) *ServerResponse { + return &ServerResponse{ + Status: StatusNoGlue, + Qname: canonicalName(qname), + Qclass: qclass, + Qtype: qtype, + IP: ip, + Server: canonicalName(server), + Bailiwick: canonicalName(bailiwick), + } +} + +// NewLoopResponse records a dead end: resolving server from ip would repeat +// an ancestor referral. +func NewLoopResponse(qname string, qclass, qtype uint16, ip, server, bailiwick string) *ServerResponse { + return &ServerResponse{ + Status: StatusLoop, + Qname: canonicalName(qname), + Qclass: qclass, + Qtype: qtype, + IP: ip, + Server: canonicalName(server), + Bailiwick: canonicalName(bailiwick), + } +} + +// evaluate mirrors response.rb#evaluate: cache the in-bailiwick records, then +// for referral/restart work out the start servers from THIS branch's cache; +// a referral whose new zone is not strictly deeper than the bailiwick is lame. +func (r *ServerResponse) evaluate() error { + if r.Status != StatusException { + r.Cache.Add(r.DQ.CacheableGood) + } + switch r.DQ.Status { + case StatusRestart: + starters, bw, err := r.Cache.GetStartServers(r.DQ.Endname) + if err != nil { + return err + } + r.Starters, r.StartersBailiwick = starters, bw + case StatusReferral: + starters, bw, err := r.Cache.GetStartServers(r.DQ.Endname) + if err != nil { + return err + } + r.Starters, r.StartersBailiwick = starters, bw + if isLameReferral(r.DQ.Bailiwick, bw) { + r.Status = StatusReferralLame + } + starterNames := make([]string, len(starters)) + for i, s := range starters { + starterNames[i] = s.Name + } + if !equalSorted(starterNames, r.DQ.AuthorityNames) { + r.DQ.WarningsAdd("Referred authority names do not match query cache expectations") + } + } + return nil +} + +func equalSorted(a, b []string) bool { + if len(a) != len(b) { + return false + } + as := append([]string(nil), a...) + bs := append([]string(nil), b...) + sort.Strings(as) + sort.Strings(bs) + for i := range as { + if as[i] != bs[i] { + return false + } + } + return true +} + +// StatsKey is the leaf aggregation key (response.rb update_stats_key and the +// noglue/loop variants): identical keys merge by summing probability. +func (r *ServerResponse) StatsKey() string { + qclass := ClassToString(r.Qclass) + qtype := TypeToString(r.Qtype) + switch r.Status { + case StatusNoGlue, StatusLoop: + return fmt.Sprintf("key:%s:%s:%s:%s:%s:%s:%s", + r.Status, r.IP, r.Qname, qclass, qtype, r.Server, r.Bailiwick) + default: + key := fmt.Sprintf("key:%s:%s:%s:%s:%s:%s", + r.Status, r.IP, r.Server, r.Qname, qclass, qtype) + if r.Status == StatusException && r.DQ != nil { + key += ":" + r.DQ.ExceptionMessage + } else if r.Status == StatusReferralLame { + key += ":" + r.ParentIP + } + return key + } +} + +// String renders a short description for progress display (the Ruby +// response to_s variants: "No glue for X" / "Loop encountered resolving X"). +func (r *ServerResponse) String() string { + switch r.Status { + case StatusNoGlue: + return fmt.Sprintf("No glue for %s", r.Server) + case StatusLoop: + return fmt.Sprintf("Loop encountered resolving %s", r.Server) + case StatusException: + if r.DQ != nil { + return r.DQ.ExceptionMessage + } + case StatusError: + if r.DQ != nil { + return r.DQ.ErrorMessage + } + } + return string(r.Status) +} + +func ClassToString(qclass uint16) string { + if s, ok := miekgdns.ClassToString[qclass]; ok { + return s + } + return fmt.Sprintf("CLASS%d", qclass) +} + +func TypeToString(qtype uint16) string { + if s, ok := miekgdns.TypeToString[qtype]; ok { + return s + } + return fmt.Sprintf("TYPE%d", qtype) +} diff --git a/internal/traverse/server_response_test.go b/internal/traverse/server_response_test.go new file mode 100644 index 0000000..959157a --- /dev/null +++ b/internal/traverse/server_response_test.go @@ -0,0 +1,232 @@ +package traverse + +import ( + "testing" + + "github.com/miekg/dns" +) + +// rootedCache returns a cache seeded with root hints, as the traverser will +// always provide. +func rootedCache() *InfoCache { + c := NewInfoCache(nil) + c.AddHints("", []StartServer{{Name: "a.root-servers.net", IPs: []string{"198.41.0.4"}}}) + return c +} + +func TestServerResponseReferralNotLame(t *testing.T) { + // Root server refers com query to the gtld servers: "" → "com" is deeper. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("com", "a.gtld-servers.net")) + msg.Extra = append(msg.Extra, aRR("a.gtld-servers.net", "192.5.6.30")) + dq := decode(msg, "www.example.com", dns.TypeA, "") + + r, err := NewServerResponse(dq, "a.root-servers.net", "", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusReferral { + t.Fatalf("status = %s, want referral", r.Status) + } + if r.StartersBailiwick != "com" { + t.Errorf("starters bailiwick = %q, want com", r.StartersBailiwick) + } + if len(r.Starters) != 1 || r.Starters[0].Name != "a.gtld-servers.net" { + t.Errorf("starters = %v", r.Starters) + } + if len(r.Starters[0].IPs) != 1 || r.Starters[0].IPs[0] != "192.5.6.30" { + t.Errorf("starter IPs = %v", r.Starters[0].IPs) + } + if len(dq.Warnings) != 0 { + t.Errorf("unexpected warnings: %v", dq.Warnings) + } +} + +func TestServerResponseGluelessStarterHasNilIPs(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("example.com", "ns1.example.com")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + + r, err := NewServerResponse(dq, "a.gtld-servers.net", "", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusReferral { + t.Fatalf("status = %s, want referral", r.Status) + } + if r.Starters[0].IPs != nil { + t.Errorf("glueless starter should have nil IPs, got %v", r.Starters[0].IPs) + } +} + +func TestServerResponseLameReferral(t *testing.T) { + // A com server "refers" us to an example.org zone: the NS records are + // out-of-bailiwick so they are discarded, the cache walk falls back to + // the root NS, and "" is not strictly deeper than "com" → lame. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("example.org", "ns1.example.org")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + if dq.Status != StatusReferral { + t.Fatalf("decoded status = %s, want referral", dq.Status) + } + + r, err := NewServerResponse(dq, "a.gtld-servers.net", "192.5.6.30", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusReferralLame { + t.Fatalf("status = %s, want referral_lame", r.Status) + } + if r.StartersBailiwick != "" { + t.Errorf("starters bailiwick = %q, want \"\" (root fallback)", r.StartersBailiwick) + } + found := false + for _, w := range dq.Warnings { + if w == "Referred authority names do not match query cache expectations" { + found = true + } + } + if !found { + t.Errorf("expected mismatch warning, got %v", dq.Warnings) + } + want := "key:referral_lame:192.0.2.1:a.gtld-servers.net:www.example.com:IN:A:192.5.6.30" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +func TestServerResponseEqualZoneReferralIsLame(t *testing.T) { + // Referral back into the SAME zone (com → com) is lame: not strictly deeper. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("com", "b.gtld-servers.net")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + + r, err := NewServerResponse(dq, "a.gtld-servers.net", "192.5.6.30", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusReferralLame { + t.Fatalf("status = %s, want referral_lame", r.Status) + } +} + +func TestServerResponseRestartStarters(t *testing.T) { + // A CNAME out of the bailiwick restarts; starters come from the deepest + // cached zone for the new target (root here). + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, cnameRR("www.example.com", "cdn.example.org")) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + if dq.Status != StatusRestart { + t.Fatalf("decoded status = %s, want restart", dq.Status) + } + + parent := rootedCache() + parent.Add([]dns.RR{nsRR("example.org", "ns1.example.org"), aRR("ns1.example.org", "9.9.9.9")}) + r, err := NewServerResponse(dq, "ns1.example.com", "", parent) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusRestart { + t.Fatalf("status = %s, want restart", r.Status) + } + if r.StartersBailiwick != "example.org" { + t.Errorf("starters bailiwick = %q, want example.org", r.StartersBailiwick) + } + if len(r.Starters) != 1 || r.Starters[0].Name != "ns1.example.org" { + t.Errorf("starters = %v", r.Starters) + } +} + +func TestServerResponseCachesGoodRecordsInChildCache(t *testing.T) { + parent := rootedCache() + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("example.com", "ns1.example.com")) + msg.Extra = append(msg.Extra, + aRR("ns1.example.com", "1.2.3.4"), + aRR("ns1.example.org", "5.6.7.8"), // out of bailiwick — discarded + ) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + + r, err := NewServerResponse(dq, "a.gtld-servers.net", "", parent) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if got := r.Cache.Get("ns1.example.com", dns.ClassINET, dns.TypeA); len(got) != 1 { + t.Errorf("in-bailiwick glue should be cached, got %v", got) + } + if got := r.Cache.Get("ns1.example.org", dns.ClassINET, dns.TypeA); got != nil { + t.Errorf("out-of-bailiwick record must be discarded, got %v", got) + } + // The parent cache stays clean — records live in the response's child. + if got := parent.Get("example.com", dns.ClassINET, dns.TypeNS); got != nil { + t.Errorf("parent cache polluted: %v", got) + } +} + +func TestServerResponseExceptionDoesNotCache(t *testing.T) { + dq := NewDecodedQuery(nil, errTimeout{}, "www.example.com", dns.ClassINET, dns.TypeA, "192.0.2.1", "com") + r, err := NewServerResponse(dq, "a.gtld-servers.net", "", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + if r.Status != StatusException { + t.Fatalf("status = %s, want exception", r.Status) + } + want := "key:exception:192.0.2.1:a.gtld-servers.net:www.example.com:IN:A:query timed out" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +type errTimeout struct{} + +func (errTimeout) Error() string { return "query timed out" } + +func TestServerResponseAnsweredStatsKey(t *testing.T) { + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Answer = append(msg.Answer, aRR("www.example.com", "93.184.216.34")) + dq := decode(msg, "www.example.com", dns.TypeA, "example.com") + + r, err := NewServerResponse(dq, "NS1.Example.Com", "", rootedCache()) + if err != nil { + t.Fatalf("NewServerResponse: %v", err) + } + want := "key:answered:192.0.2.1:ns1.example.com:www.example.com:IN:A" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +func TestNoGlueResponse(t *testing.T) { + r := NewNoGlueResponse("www.example.com", dns.ClassINET, dns.TypeA, "192.5.6.30", "ns1.example.com", "example.com") + if r.Status != StatusNoGlue { + t.Fatalf("status = %s, want noglue", r.Status) + } + // NoGlue/Loop use their own field order: ip, qname, qclass, qtype, server, bailiwick. + want := "key:noglue:192.5.6.30:www.example.com:IN:A:ns1.example.com:example.com" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +func TestLoopResponse(t *testing.T) { + r := NewLoopResponse("www.example.com", dns.ClassINET, dns.TypeA, "192.5.6.30", "ns1.example.com", "example.com") + if r.Status != StatusLoop { + t.Fatalf("status = %s, want loop", r.Status) + } + want := "key:loop:192.5.6.30:www.example.com:IN:A:ns1.example.com:example.com" + if got := r.StatsKey(); got != want { + t.Errorf("stats key = %q, want %q", got, want) + } +} + +func TestServerResponseReferralNoRootHintsErrors(t *testing.T) { + // A lame referral with a completely empty cache chain cannot compute + // starters; the constructor surfaces the "no root hints" error. + msg := newMsg("www.example.com", dns.TypeA, dns.RcodeSuccess) + msg.Ns = append(msg.Ns, nsRR("example.org", "ns1.example.org")) + dq := decode(msg, "www.example.com", dns.TypeA, "com") + if _, err := NewServerResponse(dq, "a.gtld-servers.net", "", NewInfoCache(nil)); err == nil { + t.Fatal("expected error when no NS reachable in cache chain") + } +} diff --git a/internal/traverse/stack.go b/internal/traverse/stack.go deleted file mode 100644 index 7e09cd3..0000000 --- a/internal/traverse/stack.go +++ /dev/null @@ -1,58 +0,0 @@ -package traverse - -const DefaultMaxDepth = 20 - -type Stack struct { - items []*Referral - maxDepth int -} - -func NewStack(maxDepth int) *Stack { - if maxDepth <= 0 { - maxDepth = DefaultMaxDepth - } - return &Stack{ - items: make([]*Referral, 0), - maxDepth: maxDepth, - } -} - -func (s *Stack) Push(r *Referral) bool { - if r == nil { - return false - } - if r.Depth >= s.maxDepth { - return false - } - s.items = append(s.items, r) - return true -} - -func (s *Stack) Pop() *Referral { - if len(s.items) == 0 { - return nil - } - idx := len(s.items) - 1 - item := s.items[idx] - s.items = s.items[:idx] - return item -} - -func (s *Stack) Peek() *Referral { - if len(s.items) == 0 { - return nil - } - return s.items[len(s.items)-1] -} - -func (s *Stack) Len() int { - return len(s.items) -} - -func (s *Stack) MaxDepth() int { - return s.maxDepth -} - -func (s *Stack) IsEmpty() bool { - return len(s.items) == 0 -} diff --git a/internal/traverse/stack_test.go b/internal/traverse/stack_test.go deleted file mode 100644 index d1c9f70..0000000 --- a/internal/traverse/stack_test.go +++ /dev/null @@ -1,143 +0,0 @@ -package traverse - -import ( - "testing" - - "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" -) - -func TestNewStack(t *testing.T) { - s := NewStack(10) - if s.MaxDepth() != 10 { - t.Errorf("MaxDepth = %d, want 10", s.MaxDepth()) - } - if !s.IsEmpty() { - t.Error("new stack should be empty") - } - if s.Len() != 0 { - t.Errorf("Len = %d, want 0", s.Len()) - } -} - -func TestNewStackDefaultDepth(t *testing.T) { - s := NewStack(0) - if s.MaxDepth() != DefaultMaxDepth { - t.Errorf("MaxDepth = %d, want %d", s.MaxDepth(), DefaultMaxDepth) - } - - s = NewStack(-5) - if s.MaxDepth() != DefaultMaxDepth { - t.Errorf("MaxDepth = %d, want %d", s.MaxDepth(), DefaultMaxDepth) - } -} - -func TestStackPushPop(t *testing.T) { - s := NewStack(5) - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - - ok := s.Push(ref) - if !ok { - t.Error("Push should succeed") - } - if s.Len() != 1 { - t.Errorf("Len = %d, want 1", s.Len()) - } - - popped := s.Pop() - if popped != ref { - t.Error("popped referral should match pushed") - } - if s.Len() != 0 { - t.Errorf("Len = %d, want 0", s.Len()) - } -} - -func TestStackLIFO(t *testing.T) { - s := NewStack(5) - r1 := NewReferral("a.com", dns.TypeA, ".", 0, 1.0, nil) - r2 := NewReferral("b.com", dns.TypeA, ".", 1, 1.0, nil) - r3 := NewReferral("c.com", dns.TypeA, ".", 2, 1.0, nil) - - s.Push(r1) - s.Push(r2) - s.Push(r3) - - if popped := s.Pop(); popped != r3 { - t.Error("should pop r3 first (LIFO)") - } - if popped := s.Pop(); popped != r2 { - t.Error("should pop r2 second") - } - if popped := s.Pop(); popped != r1 { - t.Error("should pop r1 third") - } -} - -func TestStackPushNil(t *testing.T) { - s := NewStack(5) - ok := s.Push(nil) - if ok { - t.Error("Push(nil) should return false") - } - if s.Len() != 0 { - t.Errorf("Len = %d, want 0", s.Len()) - } -} - -func TestStackMaxDepth(t *testing.T) { - s := NewStack(3) - - r0 := NewReferral("a.com", dns.TypeA, ".", 0, 1.0, nil) - r1 := NewReferral("b.com", dns.TypeA, ".", 1, 1.0, nil) - r2 := NewReferral("c.com", dns.TypeA, ".", 2, 1.0, nil) - r3 := NewReferral("d.com", dns.TypeA, ".", 3, 1.0, nil) - - if !s.Push(r0) { - t.Error("depth 0 should be accepted") - } - if !s.Push(r1) { - t.Error("depth 1 should be accepted") - } - if !s.Push(r2) { - t.Error("depth 2 should be accepted") - } - if s.Push(r3) { - t.Error("depth 3 should be rejected (maxDepth=3)") - } -} - -func TestStackPopEmpty(t *testing.T) { - s := NewStack(5) - if popped := s.Pop(); popped != nil { - t.Error("Pop on empty stack should return nil") - } -} - -func TestStackPeek(t *testing.T) { - s := NewStack(5) - if peek := s.Peek(); peek != nil { - t.Error("Peek on empty stack should return nil") - } - - ref := NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil) - s.Push(ref) - - if peek := s.Peek(); peek != ref { - t.Error("Peek should return top item") - } - if s.Len() != 1 { - t.Errorf("Peek should not remove item, Len = %d, want 1", s.Len()) - } -} - -func TestStackIsEmpty(t *testing.T) { - s := NewStack(5) - if !s.IsEmpty() { - t.Error("new stack should be empty") - } - - s.Push(NewReferral("example.com", dns.TypeA, ".", 0, 1.0, nil)) - if s.IsEmpty() { - t.Error("stack with item should not be empty") - } -} diff --git a/internal/traverse/stats.go b/internal/traverse/stats.go new file mode 100644 index 0000000..1ad805c --- /dev/null +++ b/internal/traverse/stats.go @@ -0,0 +1,74 @@ +package traverse + +import ( + "sort" + "strings" + + miekgdns "github.com/miekg/dns" +) + +// AnswerStat is one distinct answered RRset with its accumulated probability. +type AnswerStat struct { + // Key groups identical RRset content: sorted rdata strings joined with + // "@@@" (summary_stats.rb get_answer_stats). + Key string + Prob float64 + RRs []miekgdns.RR +} + +// SummaryStats groups the aggregated leaves by status; answered leaves are +// additionally grouped by RRset content, one entry per distinct RRset +// (summary_stats.rb). The ByStatus probabilities sum to 1.0 and the answer +// probabilities sum to ByStatus[StatusAnswered]. +type SummaryStats struct { + ByStatus map[Status]float64 + Answers []AnswerStat +} + +// SummaryStats computes (and memoises) the summary grouping of this node's +// aggregated leaf statistics (referral.rb summary_stats). It returns nil +// until the node's statistics have been calculated. +func (r *Referral) SummaryStats() *SummaryStats { + if r == nil || !r.calculated || len(r.Stats) == 0 { + return nil + } + if r.summaryStats != nil { + return r.summaryStats + } + + stats := &SummaryStats{ByStatus: make(map[Status]float64)} + answers := make(map[string]*AnswerStat) + for _, leaf := range r.StatsList() { + status := leaf.Response.Status + stats.ByStatus[status] += leaf.Prob + if status != StatusAnswered { + continue + } + rdatas := make([]string, 0, len(leaf.Response.DQ.Answers)) + for _, rr := range leaf.Response.DQ.Answers { + rdatas = append(rdatas, rrData(rr)) + } + sort.Strings(rdatas) + key := strings.Join(rdatas, "@@@") + if e, ok := answers[key]; ok { + e.Prob += leaf.Prob + } else { + answers[key] = &AnswerStat{Key: key, Prob: leaf.Prob, RRs: leaf.Response.DQ.Answers} + } + } + + for _, e := range answers { + stats.Answers = append(stats.Answers, *e) + } + sort.Slice(stats.Answers, func(i, j int) bool { return stats.Answers[i].Key < stats.Answers[j].Key }) + r.summaryStats = stats + return stats +} + +// rrData extracts the rdata portion of a record (dnsruby rdata_to_string): +// everything after owner/TTL/class/type in presentation format. +func rrData(rr miekgdns.RR) string { + s := rr.String() + h := rr.Header().String() + return strings.TrimPrefix(s, h) +} diff --git a/internal/traverse/stats_test.go b/internal/traverse/stats_test.go new file mode 100644 index 0000000..743d7c1 --- /dev/null +++ b/internal/traverse/stats_test.go @@ -0,0 +1,302 @@ +package traverse + +import ( + "math" + "net" + "strconv" + "strings" + "testing" + + "github.com/miekg/dns" +) + +// mockCaptureTopology reproduces the delegation behind +// docs/captures/dnstraverse-ruby-www.example.com-A.txt: one root, thirteen +// com gTLD servers, example.com served by two NS with three IPs each, every +// endpoint answering the same two A records (two endpoints return them in +// the opposite order, as in the capture). +func mockCaptureTopology() *mockExchange { + m := newMockExchange() + + gtlds := []string{"a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k", "l", "m"} + var comNS, comGlue []dns.RR + gtldIPs := make([]string, len(gtlds)) + for i, l := range gtlds { + ip := "192.0.2." + strconv.Itoa(i+1) + gtldIPs[i] = ip + comNS = append(comNS, nsRR("com", l+".gtld-servers.net")) + comGlue = append(comGlue, aRR(l+".gtld-servers.net", ip)) + } + m.on("202.12.27.33", "www.example.com", dns.TypeA, referralMsg(comNS, comGlue...)) + + heraIPs := []string{"108.162.192.162", "172.64.32.162", "173.245.58.162"} + elliottIPs := []string{"108.162.195.228", "162.159.44.228", "172.64.35.228"} + exampleReferral := referralMsg( + []dns.RR{ + nsRR("example.com", "hera.ns.cloudflare.com"), + nsRR("example.com", "elliott.ns.cloudflare.com"), + }, + aRR("hera.ns.cloudflare.com", heraIPs[0]), + aRR("hera.ns.cloudflare.com", heraIPs[1]), + aRR("hera.ns.cloudflare.com", heraIPs[2]), + aRR("elliott.ns.cloudflare.com", elliottIPs[0]), + aRR("elliott.ns.cloudflare.com", elliottIPs[1]), + aRR("elliott.ns.cloudflare.com", elliottIPs[2]), + ) + for _, ip := range gtldIPs { + m.on(ip, "www.example.com", dns.TypeA, exampleReferral) + } + + forward := answerMsg( + aRR("www.example.com", "104.20.23.154"), + aRR("www.example.com", "172.66.147.243"), + ) + reversed := answerMsg( + aRR("www.example.com", "172.66.147.243"), + aRR("www.example.com", "104.20.23.154"), + ) + for _, ip := range []string{heraIPs[0], elliottIPs[0], elliottIPs[1], elliottIPs[2]} { + m.on(ip, "www.example.com", dns.TypeA, forward) + } + for _, ip := range []string{heraIPs[1], heraIPs[2]} { + m.on(ip, "www.example.com", dns.TypeA, reversed) + } + return m +} + +// TestCapturePerEndpointFractions asserts the 16.7%-per-endpoint result of +// the www.example.com capture: 13 gTLD paths collapse (fast mode) into six +// endpoint leaves of 1/6 each, and the summary merges the differently +// ordered RRsets into a single 100% answered line. +func TestCapturePerEndpointFractions(t *testing.T) { + m := mockCaptureTopology() + cfg := testConfig(true) + cfg.RootAddrs = []net.IP{net.ParseIP("202.12.27.33")} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 6 { + t.Fatalf("expected 6 answered leaves (one per endpoint), got %d: %v", len(answered), root.StatsList()) + } + for _, leaf := range answered { + if math.Abs(leaf.Prob-1.0/6) > 1e-9 { + t.Errorf("leaf %s prob = %v, want 1/6", leaf.Key, leaf.Prob) + } + } + + stats := root.SummaryStats() + if stats == nil { + t.Fatal("expected summary stats after calculation") + } + if prob := stats.ByStatus[StatusAnswered]; math.Abs(prob-1.0) > 1e-9 { + t.Errorf("answered summary prob = %v, want 1.0", prob) + } + // Both RR orders share the same sorted-rdata key: one summary line. + if len(stats.Answers) != 1 { + t.Fatalf("expected 1 distinct answered RRset, got %d: %v", len(stats.Answers), stats.Answers) + } + if math.Abs(stats.Answers[0].Prob-1.0) > 1e-9 { + t.Errorf("answer group prob = %v, want 1.0", stats.Answers[0].Prob) + } + if !strings.Contains(stats.Answers[0].Key, "@@@") { + t.Errorf("answer key should join rdata with @@@, got %q", stats.Answers[0].Key) + } + if len(stats.Answers[0].RRs) != 2 { + t.Errorf("answer RRs = %v", stats.Answers[0].RRs) + } +} + +// TestDistinctRRsetsSeparateGroups asserts the converse of the capture test: +// two endpoints answering DIFFERENT content produce two separate summary +// groups, each carrying its own share of the answered probability. +func TestDistinctRRsetsSeparateGroups(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + aRR("ns2.example.com", "2.2.2.2"), + )) + // Both RRsets share their first sorted rdata so grouping must consider + // the full content, not just the first record. + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg( + aRR("www.example.com", "1.0.0.1"), + aRR("www.example.com", "9.9.9.9"), + )) + m.on("2.2.2.2", "www.example.com", dns.TypeA, answerMsg( + aRR("www.example.com", "1.0.0.1"), + aRR("www.example.com", "8.8.8.8"), + )) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + stats := root.SummaryStats() + if stats == nil { + t.Fatal("expected summary stats after calculation") + } + if prob := stats.ByStatus[StatusAnswered]; math.Abs(prob-1.0) > 1e-9 { + t.Errorf("answered summary prob = %v, want 1.0", prob) + } + if len(stats.Answers) != 2 { + t.Fatalf("expected 2 distinct answered RRsets, got %d: %v", len(stats.Answers), stats.Answers) + } + if stats.Answers[0].Key == stats.Answers[1].Key { + t.Errorf("answer groups share key %q, want distinct keys", stats.Answers[0].Key) + } + for i, ans := range stats.Answers { + if math.Abs(ans.Prob-0.5) > 1e-9 { + t.Errorf("answer group %d (%q) prob = %v, want 0.5", i, ans.Key, ans.Prob) + } + } + // Answers are sorted by key: "1.0.0.1@@@8.8.8.8" then "1.0.0.1@@@9.9.9.9". + wantRdata := []string{"8.8.8.8", "9.9.9.9"} + for i, ans := range stats.Answers { + if !strings.Contains(ans.Key, "1.0.0.1@@@"+wantRdata[i]) { + t.Errorf("answer group %d key = %q, want it to contain %q", i, ans.Key, "1.0.0.1@@@"+wantRdata[i]) + } + if len(ans.RRs) != 2 { + t.Errorf("answer group %d RRs = %v, want 2 records", i, ans.RRs) + } + } +} + +func TestServfailErrorLeaf(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + aRR("ns2.example.com", "2.2.2.2"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + m.on("2.2.2.2", "www.example.com", dns.TypeA, rcodeMsg(dns.RcodeServerFailure)) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + errs := leavesByStatus(root, StatusError) + if len(errs) != 1 { + t.Fatalf("expected 1 error leaf, got %v", root.StatsList()) + } + if errs[0].Response.DQ.ErrorMessage != "Server failure (SERVFAIL)" { + t.Errorf("error message = %q", errs[0].Response.DQ.ErrorMessage) + } + if math.Abs(errs[0].Prob-0.5) > 1e-9 { + t.Errorf("error prob = %v, want 0.5", errs[0].Prob) + } + + stats := root.SummaryStats() + if math.Abs(stats.ByStatus[StatusError]-0.5) > 1e-9 || + math.Abs(stats.ByStatus[StatusAnswered]-0.5) > 1e-9 { + t.Errorf("summary by status = %v", stats.ByStatus) + } + total := 0.0 + for _, prob := range stats.ByStatus { + total += prob + } + if math.Abs(total-1.0) > 1e-9 { + t.Errorf("summary probabilities sum to %v, want 1.0", total) + } +} + +// TestResolveSubtreeLeavesExcluded asserts that resolve-subtree leaves (the +// A lookups for glueless NS) never reach the main aggregation: +// they surface only through server weights. +func TestResolveSubtreeLeavesExcluded(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns.other.net")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + m.on("198.41.0.4", "ns.other.net", dns.TypeA, answerMsg(aRR("ns.other.net", "4.4.4.4"))) + m.on("4.4.4.4", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + for _, leaf := range root.StatsList() { + if leaf.Response.Qname == "ns.other.net" { + t.Errorf("resolve-subtree leaf leaked into main aggregation: %s", leaf.Key) + } + } + if prob := root.SummaryStats().ByStatus[StatusAnswered]; math.Abs(prob-1.0) > 1e-9 { + t.Errorf("answered summary prob = %v, want 1.0", prob) + } +} + +// TestResolveFailurePseudoIPCarriesMass asserts that a failed glue resolution +// keeps its probability: the failure becomes a "key:" pseudo-IP whose mass +// surfaces in the main aggregation as the failing (resolve) query. +func TestResolveFailurePseudoIPCarriesMass(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns.other.net")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + // The resolve of A ns.other.net fails at the root: SERVFAIL. + m.on("198.41.0.4", "ns.other.net", dns.TypeA, rcodeMsg(dns.RcodeServerFailure)) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + errs := leavesByStatus(root, StatusError) + if len(errs) != 1 { + t.Fatalf("expected 1 error leaf from the failed resolve, got %v", root.StatsList()) + } + leaf := errs[0] + if math.Abs(leaf.Prob-0.5) > 1e-9 { + t.Errorf("failed-resolve prob = %v, want 0.5", leaf.Prob) + } + if leaf.Response.Qname != "ns.other.net" { + t.Errorf("failed-resolve leaf qname = %q, want the resolve target", leaf.Response.Qname) + } + if !strings.HasPrefix(leaf.Key, "key:error:") { + t.Errorf("failed-resolve key = %q", leaf.Key) + } + // The leaf's referral is the resolve-subtree node; its parent is the + // glueless referral, which carries the mass as a pseudo-IP server entry. + glueless := leaf.Referral.Parent + if glueless.Server != "ns.other.net" { + t.Fatalf("glueless referral server = %q", glueless.Server) + } + hasPseudo := false + for ip, weight := range glueless.ServerWeights { + if strings.HasPrefix(ip, "key:") && math.Abs(weight-1.0) <= 1e-9 { + hasPseudo = true + } + } + if !hasPseudo { + t.Errorf("expected a key: pseudo-IP with weight 1.0, got %v", glueless.ServerWeights) + } +} + +func TestSummaryStatsNilAndMemoised(t *testing.T) { + var nilRef *Referral + if nilRef.SummaryStats() != nil { + t.Error("nil referral should produce nil summary") + } + uncalculated := newTestReferral("ns1.example.com", []string{"1.1.1.1"}) + if uncalculated.SummaryStats() != nil { + t.Error("uncalculated referral should produce nil summary") + } + + m := mockSimpleDelegation() + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + first := root.SummaryStats() + if first == nil { + t.Fatal("expected summary stats") + } + if root.SummaryStats() != first { + t.Error("summary stats should be memoised") + } +} diff --git a/internal/traverse/traverse.go b/internal/traverse/traverse.go index 83914e1..284d576 100644 --- a/internal/traverse/traverse.go +++ b/internal/traverse/traverse.go @@ -1,27 +1,29 @@ -// Package traverse implements the core DNS traversal engine for ExploreDNS. +// Package traverse implements the core DNS traversal engine for ExploreDNS, +// a Go port of the Ruby dnstraverse engine (dns.squish.net). // -// The traversal engine starts from the DNS root servers and iteratively -// follows every referral it receives, building a complete picture of the -// delegation path for a domain. Unlike a standard recursive resolver, which -// stops at the first authoritative answer, the traversal engine explores every -// branch so that delegation mismatches, lame delegations, or split authorities -// are all visible in the output. +// The traversal starts from a synthetic "rootroot" node (never displayed) +// with one child per root server, and explores every branch of the +// delegation instead of stopping at the first authoritative answer, so lame +// delegations, missing glue and split authorities are all visible. // // # Architecture // -// A Traverser maintains a stack of Referral objects. Each Referral -// represents a pending query to a specific set of nameservers for a specific -// name and record type. The engine pops referrals one at a time, sends the -// query, classifies the response, and pushes any child referrals back onto the -// stack. -// -// When a referral contains nameserver names but no glue records (IP addresses), -// the engine resolves them via a secondary traversal before continuing. +// A Traverser runs an explicit stack loop over Referral nodes. Each Referral +// queries every IP address of one nameserver for one qname/qclass/qtype, +// classifies each response (DecodedQuery, ServerResponse) and creates one +// child per NS name for referral/restart statuses — including glueless +// nameservers, which get their own resolve subtree (refid ".0." components) +// queried from this branch's cache, never a system resolver. Post-order +// stack markers fold the statistics upwards once all children finished: +// every leaf outcome carries a probability, and the probabilities at the +// root sum to 1.0. // // # Caching // -// An InfoCache stores discovered glue records. In fast mode (default) a -// single root cache is shared across all branches so that glue discovered in -// one branch is immediately available to sibling branches. Disable fast mode -// (TraverserConfig.Fast = false) for fully independent branch resolution. +// Two caches cooperate: the packet-level cache in internal/dns sends each +// (server IP, question, udpsize) at most once per run, and the hierarchical +// per-branch InfoCache holds the in-bailiwick records each response is +// allowed to contribute. Fast mode (default) additionally memoises completed +// referrals so identical subtrees are reported as "completed earlier" +// instead of being walked again. package traverse diff --git a/internal/traverse/traverser.go b/internal/traverse/traverser.go index e3576bd..365fbe5 100644 --- a/internal/traverse/traverser.go +++ b/internal/traverse/traverser.go @@ -4,9 +4,8 @@ import ( "context" "fmt" "net" + "sort" "strings" - "sync" - "time" "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" miekgdns "github.com/miekg/dns" @@ -14,495 +13,329 @@ import ( // TraverserConfig configures the behaviour of a Traverser. type TraverserConfig struct { - // MaxDepth is the maximum referral depth before the traversal gives up. - MaxDepth int + // MaxDepth is the maximum referral depth (non-zero refid components) + // before a "Maxdepth N exceeded" exception is injected. + MaxDepth int // QueryType is the DNS record type to query (e.g. dns.TypeA). - QueryType uint16 + QueryType uint16 // RootConfig controls how root servers are discovered. - RootConfig *dns.RootDiscoveryConfig + RootConfig *dns.RootDiscoveryConfig // QueryConfig controls per-query transport parameters. QueryConfig *dns.QueryConfig // RootAddrs is an optional pre-seeded list of root server IP addresses. - // When non-empty, root discovery via RootConfig is skipped. - RootAddrs []net.IP + // When non-empty, root discovery via RootConfig is skipped and each + // address becomes one root (named by its address). + RootAddrs []net.IP // Hooks provides optional callbacks for traversal events. - Hooks *TraverserHooks - // Fast controls cache sharing across branches. When true (default), child - // branches inherit glue discovered by earlier branches via the shared root - // cache, trading accuracy for speed. When false, each branch gets a - // completely independent cache — slower but results are not contaminated by - // sibling branch observations. + Hooks *TraverserHooks + // Fast enables the completed-referral memo (traverser.rb @answered): + // a referral identical to an earlier completed one (same qname/qclass/ + // qtype/server and per-IP weights) is replaced by it instead of being + // walked again. Non-fast mode re-walks every branch. Fast bool } func DefaultTraverserConfig() *TraverserConfig { return &TraverserConfig{ - MaxDepth: DefaultMaxDepth, - QueryType: dns.TypeA, - RootConfig: nil, - QueryConfig: nil, - RootAddrs: nil, - Fast: true, + MaxDepth: DefaultMaxDepth, + QueryType: dns.TypeA, + Fast: true, } } -// TraversalResult pairs a Referral with the Response received when it was processed. -type TraversalResult struct { - Referral *Referral - Response *Response -} - -// Traverser performs an exhaustive iterative DNS traversal starting from the -// root servers. Create one via NewTraverser and call Traverse to start a run. +// Traverser drives the traversal: it owns the packet-cached query client, +// the fast-mode memo and the explicit stack loop (traverser.rb). type Traverser struct { config *TraverserConfig + client *dns.Client exchange dns.ExchangeFunc - visited map[string]bool - depth int - mu sync.Mutex + // answered is the fast-mode memo of completed referrals. + answered map[string]*Referral + // seen maps every server name encountered to its IP addresses. + seen map[string][]string + // roots memoises root discovery so Roots() and Run() share one lookup. + roots []StartServer } func NewTraverser(cfg *TraverserConfig) *Traverser { if cfg == nil { cfg = DefaultTraverserConfig() } + if cfg.MaxDepth <= 0 { + cfg.MaxDepth = DefaultMaxDepth + } + if cfg.QueryType == 0 { + cfg.QueryType = dns.TypeA + } return &Traverser{ config: cfg, - exchange: nil, - visited: make(map[string]bool), - depth: 0, + client: dns.NewClient(cfg.QueryConfig, nil), + answered: make(map[string]*Referral), + seen: make(map[string][]string), } } +// SetExchange injects a mock wire exchange into the single query path (both +// traversal queries and root discovery); tests use this so no packets leave +// the process. func (t *Traverser) SetExchange(fn dns.ExchangeFunc) { t.exchange = fn + t.client = dns.NewClient(t.config.QueryConfig, fn) } func (t *Traverser) SetHooks(hooks *TraverserHooks) { - if t.config == nil { - t.config = DefaultTraverserConfig() - } t.config.Hooks = hooks } -func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResult, error) { - name = miekgdns.Fqdn(name) - - roots, err := t.discoverRoots(ctx) - if err != nil { - return nil, fmt.Errorf("root discovery: %w", err) - } - - initial := NewReferral(name, t.config.QueryType, ".", 0, 1.0, nil) - initial.Addresses = roots - - stack := NewStack(t.config.MaxDepth) - stack.Push(initial) - - rootCache := NewInfoCache(nil) - var ( - mu sync.Mutex - results []TraversalResult - ) - - for { - select { - case <-ctx.Done(): - return results, fmt.Errorf("traversal cancelled: %w", ctx.Err()) - default: - } - - ref := stack.Pop() - if ref == nil { - break - } - - var cache *InfoCache - if t.config.Fast { - // Fast mode: inherit glue from the shared root cache so earlier - // branch discoveries are visible to later branches. - cache = rootCache - if ref.Parent != nil { - cache = rootCache.Child() - } - } else { - // Non-fast mode: every referral gets its own independent cache so - // no cross-branch glue is reused, ensuring each path is resolved - // from scratch. - cache = NewInfoCache(nil) - } - - if t.config.Hooks != nil { - t.config.Hooks.emit(EventStart, TraversalResult{Referral: ref}, false) - } - - resp := t.processReferral(ctx, ref, cache) - - result := TraversalResult{Referral: ref, Response: resp} - if t.config.Hooks != nil { - t.config.Hooks.emit(EventComplete, result, false) - } - - mu.Lock() - results = append(results, result) - mu.Unlock() - - if resp.IsTerminal() { - continue - } - - if resp.Type == RespReferral { - children := resp.ChildReferrals() - for _, child := range children { - if !stack.Push(child) { - mu.Lock() - results = append(results, TraversalResult{ - Referral: child, - Response: &Response{ - Referral: child, - Type: RespError, - }, - }) - mu.Unlock() - } - } - } - - if resp.Type == RespCNAMEFollow { - follow := resp.CNAMEFollowReferral() - if follow != nil { - // Detect CNAME loop: target name already appears in the ancestor chain. - if follow.Parent != nil && follow.Parent.IsNameInChain(follow.Name) { - mu.Lock() - results = append(results, TraversalResult{ - Referral: follow, - Response: &Response{ - Referral: follow, - Type: RespCNAMELoop, - ErrorMessage: fmt.Sprintf("CNAME loop detected: %s already in traversal chain", follow.Name), - }, - }) - mu.Unlock() - } else if !stack.Push(follow) { - mu.Lock() - results = append(results, TraversalResult{ - Referral: follow, - Response: &Response{ - Referral: follow, - Type: RespError, - }, - }) - mu.Unlock() - } - } - } - } - - return results, nil +// ServersEncountered returns every server name seen during the run mapped to +// its known IP addresses (traverser.rb servers_encountered). +func (t *Traverser) ServersEncountered() map[string][]string { + return t.seen } -func (t *Traverser) discoverRoots(ctx context.Context) ([]net.IP, error) { - if len(t.config.RootAddrs) > 0 { - return t.config.RootAddrs, nil +// Roots performs (and memoises) root discovery, returning the start servers +// the traversal will begin from. Callers may use it before Run to report the +// initial root; Run reuses the memoised result. +func (t *Traverser) Roots(ctx context.Context) ([]StartServer, error) { + if t.roots == nil { + roots, err := t.rootStartServers(ctx) + if err != nil { + return nil, fmt.Errorf("root discovery: %w", err) + } + t.roots = roots } + return t.roots, nil +} - servers, err := dns.DiscoverRoots(ctx, t.config.RootConfig) +// Run traverses the DNS for name and returns the synthetic rootroot node +// (never displayed) whose Stats aggregate every leaf outcome; the per-leaf +// probabilities sum to 1.0. +func (t *Traverser) Run(ctx context.Context, name string) (*Referral, error) { + roots, err := t.Roots(ctx) if err != nil { return nil, err } - var addrs []net.IP - for _, srv := range servers { - addrs = append(addrs, srv.AllIPs(false)...) + cache := NewInfoCache(nil) + cache.AddHints("", roots) + + root := &Referral{ + RefID: "", + Qname: canonicalName(toASCII(name)), + Qclass: miekgdns.ClassINET, + Qtype: t.config.QueryType, + NSAType: dns.TypeA, + Server: "", + Bailiwick: "", + InfoCache: cache, + Status: RefStatusNormal, + Responses: make(map[string]*ServerResponse), + Children: make(map[string][]*Referral), + ServerWeights: make(map[string]float64), + client: t.client, + maxdepth: t.config.MaxDepth, } - return addrs, nil + t.config.Hooks.emit(StageNew, root, "") + + if err := t.run(ctx, root); err != nil { + return root, err + } + return root, nil } -func (t *Traverser) processReferral(ctx context.Context, ref *Referral, cache *InfoCache) *Response { - if !ref.HasAddresses() { - t.mu.Lock() - visitedCopy := make(map[string]bool) - for k, v := range t.visited { - visitedCopy[k] = v - } - t.mu.Unlock() +// stack markers mirroring Ruby's :calc_resolve / :calc_answer placeholders: +// the referral is revisited after its resolves/children finished, giving +// post-order statistics calculation without recursion. +type stackMarker int - // Resolve the nameserver's IP address. The NS hostname is stored in - // Bailiwick; ref.Name is the domain being queried (not the NS name). - nsToResolve := ref.Bailiwick - if nsToResolve == "" || nsToResolve == "." { - nsToResolve = ref.Name - } - nsName := strings.TrimSuffix(nsToResolve, ".") +const ( + markerNone stackMarker = iota + markerCalcResolve + markerCalcAnswer +) - ref.Addresses = t.resolveGlueViaSystem(ctx, nsToResolve, cache) - if len(ref.Addresses) > 0 { - ref.State = StateResolved - } else { - addrs, err := t.ResolveNS(ctx, nsToResolve, cache, visitedCopy, t.depth) - if err != nil { - return &Response{ - Referral: ref, - Type: RespNSResolutionFailed, - ErrorMessage: fmt.Sprintf("nameserver %s could not be resolved", nsName), - } - } - if len(addrs) > 0 { - ref.Addresses = addrs - ref.State = StateResolved - } else { - return &Response{ - Referral: ref, - Type: RespNSResolutionFailed, - ErrorMessage: fmt.Sprintf("nameserver %s could not be resolved", nsName), - } - } - } - } - - for _, addr := range ref.Addresses { - resp := t.queryServer(ctx, ref, addr, cache) - if resp != nil && resp.Type != RespSERVFAIL { - return resp - } - } - - return &Response{ - Referral: ref, - Type: RespSERVFAIL, - } +type stackEntry struct { + ref *Referral + marker stackMarker } -func (t *Traverser) ResolveNS(ctx context.Context, nsName string, cache *InfoCache, visited map[string]bool, depth int) ([]net.IP, error) { - if cache != nil { - if addrs := cache.LookupGlue(nsName); len(addrs) > 0 { - return addrs, nil - } +func (t *Traverser) run(ctx context.Context, root *Referral) error { + stack := []stackEntry{{ref: root}} + pop := func() stackEntry { + e := stack[len(stack)-1] + stack = stack[:len(stack)-1] + return e } - if visited != nil { - if visited[nsName] { - return nil, &CircularReferralError{ - Name: nsName, - Chain: getVisitedNames(visited), - } - } - visited[nsName] = true - } - - if depth > DefaultMaxDepth { - return nil, &UnresolvableNameserverError{ - Name: nsName, - Reason: "max depth exceeded", - } - } - - roots, err := t.discoverRoots(ctx) - if err != nil { - return nil, fmt.Errorf("root discovery: %w", err) - } - - ref := NewReferral(nsName, dns.TypeA, ".", 0, 1.0, nil) - ref.Addresses = roots - ref.State = StateResolved - - traversalCache := NewInfoCache(nil) - if visited != nil { - for name := range visited { - traversalCache.StoreGlue(name, []net.IP{}) - } - } - - var addrs []net.IP - var lastErr error - - stack := NewStack(DefaultMaxDepth) - stack.Push(ref) - - for { + for len(stack) > 0 { select { case <-ctx.Done(): - return nil, fmt.Errorf("resolution cancelled: %w", ctx.Err()) + return fmt.Errorf("traversal cancelled: %w", ctx.Err()) default: } - current := stack.Pop() - if current == nil { - break - } + e := pop() + r := e.ref - cacheForStep := traversalCache - if current.Parent != nil { - cacheForStep = traversalCache.Child() - } - - if t.config.Hooks != nil { - t.config.Hooks.emit(EventStart, TraversalResult{Referral: current}, true) - } - - resp := t.processReferral(ctx, current, cacheForStep) - - if t.config.Hooks != nil { - t.config.Hooks.emit(EventComplete, TraversalResult{Referral: current, Response: resp}, true) - } - - if resp.Type == RespAnswer && len(resp.Decoded.Answers) > 0 { - for _, rr := range resp.Decoded.Answers { - if a, ok := rr.(*miekgdns.A); ok { - addrs = append(addrs, a.A) - } - if aaaa, ok := rr.(*miekgdns.AAAA); ok { - addrs = append(addrs, aaaa.AAAA) - } + switch e.marker { + case markerCalcResolve: + r.resolveCalculate() + t.config.Hooks.emit(StageResolve, r, "") + stack = append(stack, stackEntry{ref: r}) // now needs processing + continue + case markerCalcAnswer: + r.answerCalculate() + t.config.Hooks.emit(StageAnswer, r, "") + if t.config.Fast && r.Status == RefStatusNormal && !hasLameResponse(r) { + t.answered[fastKey(r)] = r } - if len(addrs) > 0 { - if cache != nil { - cache.StoreGlue(nsName, addrs) - } - return addrs, nil + if !r.IsRootRoot() { + t.recordSeen(r) } - } - - if resp.Type == RespNXDOMAIN { - lastErr = &UnresolvableNameserverError{ - Name: nsName, - Reason: "NXDOMAIN", - } - break - } - - if resp.Type == RespSERVFAIL || resp.Type == RespError || resp.Type == RespNSResolutionFailed { - lastErr = fmt.Errorf("server error resolving %s: %s", nsName, resp.Type) continue } - if resp.Type == RespReferral { - children := resp.ChildReferrals() - for _, child := range children { - // Only skip visited names when they have no addresses; if glue - // was included in the referral response we still need to query - // that child to get the authoritative answer. - if visited != nil && visited[child.Name] && !child.HasAddresses() { - continue + // A new item. Fast mode: an identical completed referral replaces + // this one wholesale. noglue/loop nodes are excluded because their + // stats carry node-specific attributes and are cheap to recreate. + if t.config.Fast && r.Parent != nil { + if memo, ok := t.answered[fastKey(r)]; ok && !r.isNoGlue() && !r.isLoop() { + r.Parent.replaceChild(r, memo) + t.config.Hooks.emit(StageAnswerFast, r, memo.RefID) + continue + } + } + + t.config.Hooks.emit(StageStart, r, "") + + if !r.Resolved() { + // Push the resolve subtree with a calc_resolve placeholder so the + // weights are folded in once every resolve leaf completed. + stack = append(stack, stackEntry{ref: r, marker: markerCalcResolve}) + resolves, err := r.resolve() + if err != nil { + return err + } + for _, c := range resolves { + t.config.Hooks.emit(StageNew, c, "") + } + for i := len(resolves) - 1; i >= 0; i-- { + stack = append(stack, stackEntry{ref: resolves[i]}) + } + continue + } + + stack = append(stack, stackEntry{ref: r, marker: markerCalcAnswer}) + childrenSets, err := r.process(ctx) + if err != nil { + return err + } + + seenParentIP := make(map[string]bool) + var flat []*Referral + for _, set := range childrenSets { + for _, c := range set { + if len(childrenSets) > 1 && !seenParentIP[c.ParentIP] { + t.config.Hooks.emit(StageNewReferralSet, c, "") + seenParentIP[c.ParentIP] = true } - if !stack.Push(child) { - lastErr = &UnresolvableNameserverError{ - Name: nsName, - Reason: "max depth exceeded during resolution", + stage, earlier := StageNew, "" + if t.config.Fast { + if memo, ok := t.answered[fastKey(c)]; ok { + stage, earlier = StageNewFast, memo.RefID } } + t.config.Hooks.emit(stage, c, earlier) + flat = append(flat, c) } } - } - - if len(addrs) > 0 { - return addrs, nil - } - - if lastErr != nil { - return nil, lastErr - } - - return nil, &UnresolvableNameserverError{ - Name: nsName, - Reason: "resolution exhausted without answer", - } -} - -func (t *Traverser) queryServer(ctx context.Context, ref *Referral, server net.IP, cache *InfoCache) *Response { - var msg *miekgdns.Msg - var err error - - if t.exchange != nil { - msg, err = t.iterativeQueryWithExchange(ctx, server, ref.Name, ref.Qtype) - } else { - msg, err = dns.Query(ctx, server, ref.Name, ref.Qtype, t.config.QueryConfig) - if err == nil { - msg = t.ensureRDFalse(msg, server, ref.Name, ref.Qtype, t.config.QueryConfig) + for i := len(flat) - 1; i >= 0; i-- { + stack = append(stack, stackEntry{ref: flat[i]}) } } - - if err != nil { - return &Response{ - Referral: ref, - Server: server, - Type: RespError, - } - } - - resp := NewResponse(ref, server, cache) - resp.Process(msg) - return resp -} - -func (t *Traverser) iterativeQueryWithExchange(ctx context.Context, server net.IP, name string, qtype uint16) (*miekgdns.Msg, error) { - if t.config.QueryConfig == nil { - return dns.IterativeQueryWithExchange(ctx, server, name, qtype, nil, t.exchange) - } - return dns.IterativeQueryWithExchange(ctx, server, name, qtype, t.config.QueryConfig, t.exchange) -} - -func (t *Traverser) ensureRDFalse(msg *miekgdns.Msg, server net.IP, name string, qtype uint16, cfg *dns.QueryConfig) *miekgdns.Msg { - if msg != nil && msg.RecursionDesired { - if t.exchange != nil { - ctx := context.Background() - var err error - msg, err = t.iterativeQueryWithExchange(ctx, server, name, qtype) - if err != nil { - return nil - } - return msg - } - msg.RecursionDesired = false - } - return msg -} - -func (t *Traverser) resolveGlueViaSystem(ctx context.Context, name string, cache *InfoCache) []net.IP { - if cache != nil { - if addrs := cache.LookupGlue(name); len(addrs) > 0 { - return addrs - } - } - - c := &miekgdns.Client{ - Net: "udp", - ReadTimeout: 5 * time.Second, - WriteTimeout: 5 * time.Second, - } - if deadline, ok := ctx.Deadline(); ok { - remaining := time.Until(deadline) - if remaining <= 0 { - return nil - } - c.ReadTimeout = remaining - c.WriteTimeout = remaining - } - - fqdn := miekgdns.Fqdn(name) - - aMsg, _, err := c.ExchangeContext(ctx, newAQuery(fqdn), "127.0.0.1:53") - if err == nil { - var addrs []net.IP - for _, rr := range aMsg.Answer { - if a, ok := rr.(*miekgdns.A); ok { - addrs = append(addrs, a.A) - } - } - if len(addrs) > 0 { - if cache != nil { - cache.StoreGlue(name, addrs) - } - return addrs - } - } - return nil } -func newAQuery(name string) *miekgdns.Msg { - m := new(miekgdns.Msg) - m.SetQuestion(name, miekgdns.TypeA) - m.RecursionDesired = true - return m +// fastKey is the fast-mode memo key (traverser.rb): qname/qclass/qtype/ +// server plus the per-IP weights, lowercased. +func fastKey(r *Referral) string { + return strings.ToLower(fmt.Sprintf("%s:%s:%s:%s:%s", + r.Qname, ClassToString(r.Qclass), TypeToString(r.Qtype), r.Server, r.TxtIPsVerbose())) +} + +func hasLameResponse(r *Referral) bool { + for _, resp := range r.Responses { + if resp.Status == StatusReferralLame { + return true + } + } + return false +} + +func (t *Traverser) recordSeen(r *Referral) { + name := strings.ToLower(r.Server) + existing := t.seen[name] + for _, ip := range r.IPsAsArray() { + found := false + for _, have := range existing { + if have == ip { + found = true + break + } + } + if !found { + existing = append(existing, ip) + } + } + t.seen[name] = existing +} + +// rootStartServers returns the root servers as start-server hints: either +// the pre-seeded RootAddrs or the servers found via root discovery (one by +// default, all of them with AllRoots). IPv4 only, like the reference. +func (t *Traverser) rootStartServers(ctx context.Context) ([]StartServer, error) { + if len(t.config.RootAddrs) > 0 { + var out []StartServer + for _, ip := range t.config.RootAddrs { + if v4 := ip.To4(); v4 != nil { + out = append(out, StartServer{Name: v4.String(), IPs: []string{v4.String()}}) + } + } + if len(out) == 0 { + return nil, fmt.Errorf("no usable IPv4 root addresses") + } + return out, nil + } + + rootCfg := t.config.RootConfig + if t.exchange != nil { + var cp dns.RootDiscoveryConfig + if rootCfg != nil { + cp = *rootCfg + } + cp.Exchange = t.exchange + rootCfg = &cp + } + + servers, err := dns.DiscoverRoots(ctx, rootCfg) + if err != nil { + return nil, err + } + + var out []StartServer + for _, srv := range servers { + var ips []string + for _, ip := range srv.IPv4 { + ips = append(ips, ip.String()) + } + if len(ips) == 0 { + continue + } + out = append(out, StartServer{Name: canonicalName(srv.Name), IPs: ips}) + } + if len(out) == 0 { + return nil, fmt.Errorf("no root servers with IPv4 addresses") + } + sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name }) + return out, nil } diff --git a/internal/traverse/traverser_test.go b/internal/traverse/traverser_test.go index e422624..c264e17 100644 --- a/internal/traverse/traverser_test.go +++ b/internal/traverse/traverser_test.go @@ -2,482 +2,774 @@ package traverse import ( "context" + "math" "net" + "strconv" + "strings" + "sync" "testing" + "time" + idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "github.com/miekg/dns" ) -const ( - dnsTypeA = dns.TypeA - dnsTypeNS = dns.TypeNS - dnsTypeCNAME = dns.TypeCNAME - dnsTypeSOA = dns.TypeSOA -) +// --- mock exchange: the single query path used by production and tests --- -func TestDefaultTraverserConfig(t *testing.T) { - cfg := DefaultTraverserConfig() - if cfg.MaxDepth != DefaultMaxDepth { - t.Errorf("MaxDepth = %d, want %d", cfg.MaxDepth, DefaultMaxDepth) - } - if cfg.QueryType != dnsTypeA { - t.Errorf("QueryType = %d, want %d", cfg.QueryType, dnsTypeA) +type mockKey struct { + server string + qname string + qtype uint16 +} + +type mockExchange struct { + mu sync.Mutex + responses map[mockKey]*dns.Msg + errors map[mockKey]error + calls map[mockKey]int +} + +func newMockExchange() *mockExchange { + return &mockExchange{ + responses: make(map[mockKey]*dns.Msg), + errors: make(map[mockKey]error), + calls: make(map[mockKey]int), } } -func TestNewTraverserNilConfig(t *testing.T) { - tr := NewTraverser(nil) - if tr == nil { - t.Fatal("NewTraverser(nil) should not return nil") - } +func (m *mockExchange) on(server, qname string, qtype uint16, msg *dns.Msg) { + m.responses[mockKey{server, dns.Fqdn(qname), qtype}] = msg } -func TestTraverserSimpleTraversal(t *testing.T) { - answerResp := func() *dns.Msg { - m := new(dns.Msg) - m.SetReply(new(dns.Msg)) - m.Answer = append(m.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - return m - }() - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") +func (m *mockExchange) fn(_ context.Context, server string, msg *dns.Msg, _ bool) (*dns.Msg, error) { + host := server + if h, _, err := net.SplitHostPort(server); err == nil { + host = h + } + q := msg.Question[0] + key := mockKey{host, q.Name, q.Qtype} + m.mu.Lock() + m.calls[key]++ + resp, ok := m.responses[key] + err := m.errors[key] + m.mu.Unlock() if err != nil { - t.Fatalf("unexpected error: %v", err) + return nil, err } - if len(results) == 0 { - t.Fatal("expected at least 1 result") + if !ok { + return nil, &net.OpError{Op: "read", Err: &net.DNSError{Err: "no mock response", Name: q.Name}} } + out := resp.Copy() + out.SetReply(msg) + out.Answer = resp.Answer + out.Ns = resp.Ns + out.Extra = resp.Extra + out.Rcode = resp.Rcode + return out, nil +} - found := false - for _, r := range results { - if r.Response.Type == RespAnswer { - found = true - break - } - } - if !found { - t.Error("expected to find an answer response") +func answerMsg(rrs ...dns.RR) *dns.Msg { + m := new(dns.Msg) + m.Answer = rrs + return m +} + +func referralMsg(nsRRs []dns.RR, glue ...dns.RR) *dns.Msg { + m := new(dns.Msg) + m.Ns = nsRRs + m.Extra = glue + return m +} + +func rcodeMsg(rcode int) *dns.Msg { + m := new(dns.Msg) + m.Rcode = rcode + return m +} + +func testConfig(fast bool) *TraverserConfig { + return &TraverserConfig{ + MaxDepth: DefaultMaxDepth, + QueryType: dns.TypeA, + RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, + Fast: fast, + QueryConfig: &idns.QueryConfig{ + Retries: 1, + Timeout: time.Second, + RetryDelay: time.Millisecond, + }, } } -func TestTraverserReferralTraversal(t *testing.T) { - rootAnswer := new(dns.Msg) - rootAnswer.Rcode = dns.RcodeSuccess - rootAnswer.Authoritative = false - rootAnswer.Ns = append(rootAnswer.Ns, &dns.NS{ - Hdr: dns.RR_Header{Name: "com.", Rrtype: dnsTypeNS, Class: dns.ClassINET}, - Ns: "a.gtld-servers.net.", - }) - rootAnswer.Extra = append(rootAnswer.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dnsTypeA}, - A: net.ParseIP("192.5.6.30"), - }) - - tldAnswer := new(dns.Msg) - tldAnswer.SetReply(new(dns.Msg)) - tldAnswer.Answer = append(tldAnswer.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - key := q.Name + "/" + dns.TypeToString[q.Qtype] - if q.Name == "example.com." && server == "198.41.0.4" { - return rootAnswer.Copy(), nil - } - if q.Name == "example.com." { - return tldAnswer.Copy(), nil - } - _ = key - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") +func runTraversal(t *testing.T, cfg *TraverserConfig, m *mockExchange, qname string) (*Traverser, *Referral) { + t.Helper() + tr := NewTraverser(cfg) + tr.SetExchange(m.fn) + root, err := tr.Run(context.Background(), qname) if err != nil { - t.Fatalf("unexpected error: %v", err) + t.Fatalf("Run(%q): %v", qname, err) } - if len(results) < 2 { - t.Fatalf("expected at least 2 results (referral + answer), got %d", len(results)) + return tr, root +} + +func statsSum(root *Referral) float64 { + sum := 0.0 + for _, e := range root.Stats { + sum += e.Prob + } + return sum +} + +func assertSumsToOne(t *testing.T, root *Referral) { + t.Helper() + if sum := statsSum(root); math.Abs(sum-1.0) > 1e-9 { + t.Errorf("aggregated leaf probabilities sum to %v, want 1.0", sum) } } -func TestTraverserMaxDepth(t *testing.T) { - callCount := 0 - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 2, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - callCount++ - m := new(dns.Msg) - m.Rcode = dns.RcodeSuccess - m.Authoritative = false - m.Ns = append(m.Ns, &dns.NS{ - Hdr: dns.RR_Header{Rrtype: dnsTypeNS}, - Ns: "ns.example.com.", - }) - m.Extra = append(m.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dnsTypeA}, - A: net.ParseIP("1.2.3.4"), - }) - return m, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "deep.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - if callCount < 1 { - t.Errorf("expected at least 1 call before max depth, got %d", callCount) - } - - depthExceeded := false - for _, r := range results { - if r.Referral != nil && r.Referral.Depth >= 2 { - depthExceeded = true - } - if r.Response.Type == RespError { - depthExceeded = true +func leavesByStatus(root *Referral, status Status) []*StatsEntry { + var out []*StatsEntry + for _, e := range root.StatsList() { + if e.Response.Status == status { + out = append(out, e) } } - if !depthExceeded { - t.Error("expected to see depth exceeded results") + return out +} + +// --- scenarios --- + +// mockSimpleDelegation wires root → com → example.com with two glued NS that +// both answer. +func mockSimpleDelegation() *mockExchange { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + aRR("ns2.example.com", "2.2.2.2"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + m.on("2.2.2.2", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + return m +} + +func TestReferralFanOut(t *testing.T) { + m := mockSimpleDelegation() + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 2 { + t.Fatalf("expected 2 answered leaves (one per NS), got %d: %v", len(answered), root.StatsList()) + } + for _, leaf := range answered { + if math.Abs(leaf.Prob-0.5) > 1e-9 { + t.Errorf("leaf %s prob = %v, want 0.5", leaf.Key, leaf.Prob) + } + } + + // RefID grammar: rootroot "", root child "1", gtld "1.1", NS "1.1.1"/"1.1.2". + if root.RefID != "" { + t.Errorf("rootroot refid = %q, want empty", root.RefID) + } + top := root.Children["rootroot"] + if len(top) != 1 || top[0].RefID != "1" { + t.Fatalf("top children = %v", top) + } + gtld := top[0].Children["198.41.0.4"] + if len(gtld) != 1 || gtld[0].RefID != "1.1" { + t.Fatalf("gtld children refids wrong: %v", gtld) + } + nsKids := gtld[0].Children["192.5.6.30"] + if len(nsKids) != 2 || nsKids[0].RefID != "1.1.1" || nsKids[1].RefID != "1.1.2" { + t.Fatalf("NS children refids wrong: %v", nsKids) + } + if nsKids[0].Depth() != 3 { + t.Errorf("depth of 1.1.1 = %d, want 3", nsKids[0].Depth()) } } -func TestTraverserContextCancellation(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - m := new(dns.Msg) - m.Rcode = dns.RcodeSuccess - m.Authoritative = false - m.Ns = append(m.Ns, &dns.NS{ - Hdr: dns.RR_Header{Rrtype: dnsTypeNS}, - Ns: "ns.example.com.", - }) - m.Extra = append(m.Extra, &dns.A{ - Hdr: dns.RR_Header{Name: "ns.example.com.", Rrtype: dnsTypeA}, - A: net.ParseIP("1.2.3.4"), - }) - return m, nil - }) +func TestGluelessResolveSubtree(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns.other.net")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + // resolve subtree for A ns.other.net starts back at the root hints + m.on("198.41.0.4", "ns.other.net", dns.TypeA, referralMsg( + []dns.RR{nsRR("net", "d.gtld.net")}, + aRR("d.gtld.net", "3.3.3.3"), + )) + m.on("3.3.3.3", "ns.other.net", dns.TypeA, answerMsg(aRR("ns.other.net", "4.4.4.4"))) + m.on("4.4.4.4", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + var events []TraversalEvent + cfg := testConfig(false) + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { events = append(events, ev) }} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 2 { + t.Fatalf("expected 2 answered leaves, got %v", root.StatsList()) + } + var viaGlueless *StatsEntry + for _, leaf := range answered { + if leaf.Response.IP == "4.4.4.4" { + viaGlueless = leaf + } + } + if viaGlueless == nil { + t.Fatal("no answered leaf via the glueless nameserver") + } + if math.Abs(viaGlueless.Prob-0.5) > 1e-9 { + t.Errorf("glueless leaf prob = %v, want 0.5", viaGlueless.Prob) + } + if got := viaGlueless.Referral.ServerWeights["4.4.4.4"]; math.Abs(got-1.0) > 1e-9 { + t.Errorf("resolved serverweight = %v, want 1.0", got) + } + + // The resolve subtree inserts a .0 refid component and is flagged. + sawResolve := false + for _, ev := range events { + if ev.RefID == "1.1.2.0.1" { + sawResolve = true + if !ev.IsResolve { + t.Error("resolve subtree event not flagged IsResolve") + } + } + } + if !sawResolve { + t.Errorf("no event for resolve refid 1.1.2.0.1; events: %v", refids(events)) + } + // Depth ignores the zero components. + if d := refidDepth("1.1.2.0.1.1"); d != 5 { + t.Errorf("refidDepth(1.1.2.0.1.1) = %d, want 5", d) + } +} + +func TestNoGlueDeadEnd(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + // In-bailiwick NS without glue: dead end. + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + )) + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + noglue := leavesByStatus(root, StatusNoGlue) + if len(noglue) != 1 { + t.Fatalf("expected 1 noglue leaf, got %v", root.StatsList()) + } + leaf := noglue[0] + if math.Abs(leaf.Prob-1.0) > 1e-9 { + t.Errorf("noglue prob = %v, want 1.0", leaf.Prob) + } + if leaf.Response.IP != "192.5.6.30" { + t.Errorf("noglue response IP = %q, want the referring parent IP", leaf.Response.IP) + } + if leaf.Referral.Server != "ns1.example.com" { + t.Errorf("noglue referral server = %q", leaf.Referral.Server) + } + if leaf.Referral.Parent.Server != "a.gtld-servers.net" { + t.Errorf("noglue parent server = %q", leaf.Referral.Parent.Server) + } + if !strings.HasPrefix(leaf.Key, "key:noglue:192.5.6.30:www.example.com:IN:A:ns1.example.com:") { + t.Errorf("noglue stats key = %q", leaf.Key) + } +} + +func TestResolveLoopDeadEnd(t *testing.T) { + m := newMockExchange() + // x.net NS ns.y.net (no glue); y.net NS ns.x.net (no glue): resolving + // either server needs the other, which is a loop. + xReferral := referralMsg([]dns.RR{nsRR("x.net", "ns.y.net")}) + yReferral := referralMsg([]dns.RR{nsRR("y.net", "ns.x.net")}) + m.on("198.41.0.4", "www.x.net", dns.TypeA, xReferral) + m.on("198.41.0.4", "ns.y.net", dns.TypeA, yReferral) + m.on("198.41.0.4", "ns.x.net", dns.TypeA, xReferral) + + _, root := runTraversal(t, testConfig(false), m, "www.x.net") + + assertSumsToOne(t, root) + loops := leavesByStatus(root, StatusLoop) + if len(loops) != 1 { + t.Fatalf("expected 1 loop leaf, got %v", root.StatsList()) + } + if math.Abs(loops[0].Prob-1.0) > 1e-9 { + t.Errorf("loop prob = %v, want 1.0", loops[0].Prob) + } + if loops[0].Referral.Status != RefStatusLoop { + t.Errorf("loop referral status = %q", loops[0].Referral.Status) + } +} + +func TestCNAMERestart(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("a.com", "ns.a.com")}, + aRR("ns.a.com", "5.5.5.5"), + )) + m.on("5.5.5.5", "www.a.com", dns.TypeA, answerMsg(cnameRRT("www.a.com", "www.b.net"))) + // restart resumes from the deepest cached zone — nothing cached for + // b.net, so back to the root. + m.on("198.41.0.4", "www.b.net", dns.TypeA, referralMsg( + []dns.RR{nsRR("net", "e.gtld.net")}, + aRR("e.gtld.net", "6.6.6.6"), + )) + m.on("6.6.6.6", "www.b.net", dns.TypeA, answerMsg(aRR("www.b.net", "7.7.7.7"))) + + _, root := runTraversal(t, testConfig(false), m, "www.a.com") + + assertSumsToOne(t, root) + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 1 { + t.Fatalf("expected 1 answered leaf, got %v", root.StatsList()) + } + leaf := answered[0] + if leaf.Response.Qname != "www.b.net" { + t.Errorf("answered qname = %q, want restart target www.b.net", leaf.Response.Qname) + } + if leaf.Referral.Qname != "www.b.net" { + t.Errorf("restart referral qname = %q", leaf.Referral.Qname) + } + // The restart chain keeps numbering below the restarting node. + if leaf.Referral.RefID != "1.1.1.1.1" { + t.Errorf("answered refid = %q, want 1.1.1.1.1", leaf.Referral.RefID) + } +} + +func TestCNAMERestartLoop(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("a.com", "ns.a.com")}, + aRR("ns.a.com", "5.5.5.5"), + )) + m.on("5.5.5.5", "www.a.com", dns.TypeA, answerMsg(cnameRRT("www.a.com", "www.b.net"))) + // www.b.net points straight back at www.a.com: every chain target is + // checked against the ancestor queries, so this is a CNAME loop. + m.on("198.41.0.4", "www.b.net", dns.TypeA, answerMsg(cnameRRT("www.b.net", "www.a.com"))) + + _, root := runTraversal(t, testConfig(false), m, "www.a.com") + + assertSumsToOne(t, root) + loops := leavesByStatus(root, StatusCNAMELoop) + if len(loops) != 1 { + t.Fatalf("expected 1 cname_loop leaf, got %v", root.StatsList()) + } + if math.Abs(loops[0].Prob-1.0) > 1e-9 { + t.Errorf("cname_loop prob = %v, want 1.0", loops[0].Prob) + } +} + +func TestCNAMERestartLoopIntermediateChainTarget(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.a.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("a.com", "ns.a.com")}, + aRR("ns.a.com", "5.5.5.5"), + )) + m.on("5.5.5.5", "www.a.com", dns.TypeA, answerMsg(cnameRRT("www.a.com", "www.b.net"))) + m.on("198.41.0.4", "www.b.net", dns.TypeA, referralMsg( + []dns.RR{nsRR("b.net", "ns.b.net")}, + aRR("ns.b.net", "6.6.6.6"), + )) + // A three-record chain whose INTERMEDIATE target www.a.com matches an + // ancestor query; the endname other.org does not, so an endname-only + // loop check would miss it. + m.on("6.6.6.6", "www.b.net", dns.TypeA, answerMsg( + cnameRRT("www.b.net", "c.b.net"), + cnameRRT("c.b.net", "www.a.com"), + cnameRRT("www.a.com", "other.org"), + )) + + _, root := runTraversal(t, testConfig(false), m, "www.a.com") + + assertSumsToOne(t, root) + loops := leavesByStatus(root, StatusCNAMELoop) + if len(loops) != 1 { + t.Fatalf("expected 1 cname_loop leaf, got %v", root.StatsList()) + } + leaf := loops[0] + if math.Abs(leaf.Prob-1.0) > 1e-9 { + t.Errorf("cname_loop prob = %v, want 1.0", leaf.Prob) + } + dq := leaf.Response.DQ + if len(dq.ChainTargets) != 3 { + t.Fatalf("ChainTargets = %v, want 3 entries", dq.ChainTargets) + } + if dq.ChainTargets[1] != "www.a.com" { + t.Errorf("ChainTargets[1] = %q, want www.a.com", dq.ChainTargets[1]) + } + if dq.Endname != "other.org" { + t.Errorf("Endname = %q, want other.org", dq.Endname) + } +} + +func TestDepthLimitInjectsException(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.c.b.a", dns.TypeA, referralMsg( + []dns.RR{nsRR("a", "ns.a")}, + aRR("ns.a", "1.1.1.1"), + )) + m.on("1.1.1.1", "www.c.b.a", dns.TypeA, referralMsg( + []dns.RR{nsRR("b.a", "ns.b.a")}, + aRR("ns.b.a", "2.2.2.2"), + )) + // Node 1.1.1 sits at depth 3 == maxdepth: its query is never sent. + cfg := testConfig(false) + cfg.MaxDepth = 3 + _, root := runTraversal(t, cfg, m, "www.c.b.a") + + assertSumsToOne(t, root) + exceptions := leavesByStatus(root, StatusException) + if len(exceptions) != 1 { + t.Fatalf("expected 1 exception leaf, got %v", root.StatsList()) + } + leaf := exceptions[0] + if leaf.Response.DQ.ExceptionMessage != "Maxdepth 3 exceeded" { + t.Errorf("exception message = %q, want %q", leaf.Response.DQ.ExceptionMessage, "Maxdepth 3 exceeded") + } + if math.Abs(leaf.Prob-1.0) > 1e-9 { + t.Errorf("exception prob = %v, want 1.0", leaf.Prob) + } + // No query should have reached depth 3. + m.mu.Lock() + defer m.mu.Unlock() + for key := range m.calls { + if key.server == "2.2.2.2" { + t.Error("query sent beyond the depth limit") + } + } +} + +func TestLameReferral(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + )) + m.on("192.5.6.30", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + // ns1 refers back to the same zone: not strictly deeper, so lame. + m.on("1.1.1.1", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns2.example.com")}, + aRR("ns2.example.com", "2.2.2.2"), + )) + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + lame := leavesByStatus(root, StatusReferralLame) + if len(lame) != 1 { + t.Fatalf("expected 1 referral_lame leaf, got %v", root.StatsList()) + } + leaf := lame[0] + if math.Abs(leaf.Prob-1.0) > 1e-9 { + t.Errorf("lame prob = %v, want 1.0", leaf.Prob) + } + if !strings.HasSuffix(leaf.Key, ":192.5.6.30") { + t.Errorf("lame stats key should end with the parent IP, got %q", leaf.Key) + } + if leaf.Referral.ParentIP != "192.5.6.30" { + t.Errorf("lame referral parent ip = %q", leaf.Referral.ParentIP) + } +} + +func TestChildsetDigitWhenMultipleIPsProduceChildren(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + aRR("a.gtld-servers.net", "192.5.6.31"), + )) + exampleReferral := referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + ) + m.on("192.5.6.30", "www.example.com", dns.TypeA, exampleReferral) + m.on("192.5.6.31", "www.example.com", dns.TypeA, exampleReferral) + m.on("1.1.1.1", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + + var setEvents []TraversalEvent + cfg := testConfig(false) + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { + if ev.Stage == StageNewReferralSet { + setEvents = append(setEvents, ev) + } + }} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + gtld := root.Children["rootroot"][0].Children["198.41.0.4"][0] + set1 := gtld.Children["192.5.6.30"] + set2 := gtld.Children["192.5.6.31"] + if len(set1) != 1 || set1[0].RefID != "1.1.1.1" { + t.Errorf("first childset refid = %v, want 1.1.1.1", refidsOf(set1)) + } + if len(set2) != 1 || set2[0].RefID != "1.1.2.1" { + t.Errorf("second childset refid = %v, want 1.1.2.1", refidsOf(set2)) + } + if len(setEvents) != 2 { + t.Errorf("expected 2 new_referral_set events, got %d", len(setEvents)) + } + // Identical answers from both paths merge into one leaf with prob 1.0. + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 1 || math.Abs(answered[0].Prob-1.0) > 1e-9 { + t.Errorf("answered leaves = %v", root.StatsList()) + } +} + +func TestFastModeReuse(t *testing.T) { + cfg := testConfig(true) + cfg.RootAddrs = []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")} + + m := newMockExchange() + comReferral := referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + ) + m.on("198.41.0.4", "www.example.com", dns.TypeA, comReferral) + m.on("199.9.14.201", "www.example.com", dns.TypeA, comReferral) + m.on("192.5.6.30", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + + var events []TraversalEvent + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { events = append(events, ev) }} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + var fast *TraversalEvent + for i := range events { + if events[i].Stage == StageAnswerFast { + fast = &events[i] + } + } + if fast == nil { + t.Fatal("expected a StageAnswerFast event in fast mode") + } + if fast.RefID != "2.1" || fast.CompletedEarlier != "1.1" { + t.Errorf("fast event refid=%q completedEarlier=%q, want 2.1 / 1.1", fast.RefID, fast.CompletedEarlier) + } + if fast.Referral.ReplacedBy == nil || fast.Referral.ReplacedBy.RefID != "1.1" { + t.Error("fast-replaced referral should point at its replacement") + } + // The second branch's child is the first branch's completed node. + second := root.Children["rootroot"][1] + if got := second.Children["199.9.14.201"][0].RefID; got != "1.1" { + t.Errorf("replaced child refid = %q, want 1.1", got) + } + answered := leavesByStatus(root, StatusAnswered) + if len(answered) != 1 || math.Abs(answered[0].Prob-1.0) > 1e-9 { + t.Errorf("answered leaves = %v", root.StatsList()) + } +} + +func TestNonFastModeReWalks(t *testing.T) { + cfg := testConfig(false) + cfg.RootAddrs = []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")} + + m := newMockExchange() + comReferral := referralMsg( + []dns.RR{nsRR("com", "a.gtld-servers.net")}, + aRR("a.gtld-servers.net", "192.5.6.30"), + ) + m.on("198.41.0.4", "www.example.com", dns.TypeA, comReferral) + m.on("199.9.14.201", "www.example.com", dns.TypeA, comReferral) + m.on("192.5.6.30", "www.example.com", dns.TypeA, answerMsg(aRR("www.example.com", "9.9.9.9"))) + + var events []TraversalEvent + cfg.Hooks = &TraverserHooks{OnEvent: func(ev TraversalEvent) { events = append(events, ev) }} + _, root := runTraversal(t, cfg, m, "www.example.com") + + assertSumsToOne(t, root) + for _, ev := range events { + if ev.Stage == StageAnswerFast || ev.Stage == StageNewFast { + t.Fatalf("unexpected fast-mode event %v in non-fast mode", ev.Stage) + } + } + // Both branches keep their own child node. + second := root.Children["rootroot"][1] + if got := second.Children["199.9.14.201"][0].RefID; got != "2.1" { + t.Errorf("non-fast child refid = %q, want 2.1", got) + } +} + +func TestAllRootsBranching(t *testing.T) { + cfg := testConfig(false) + cfg.RootAddrs = nil + for i := 1; i <= 13; i++ { + cfg.RootAddrs = append(cfg.RootAddrs, net.ParseIP("198.41.0."+strconv.Itoa(i))) + } + + m := newMockExchange() + for i := 1; i <= 13; i++ { + m.on("198.41.0."+strconv.Itoa(i), "example.com", dns.TypeA, answerMsg(aRR("example.com", "9.9.9.9"))) + } + _, root := runTraversal(t, cfg, m, "example.com") + + assertSumsToOne(t, root) + top := root.Children["rootroot"] + if len(top) != 13 { + t.Fatalf("expected 13 top-level children, got %d", len(top)) + } + for i, child := range top { + if child.RefID != strconv.Itoa(i+1) { + t.Errorf("child %d refid = %q, want %q", i, child.RefID, strconv.Itoa(i+1)) + } + } + answered := leavesByStatus(root, StatusAnswered) + // Same answer from 13 different server IPs: 13 distinct leaves of 1/13. + if len(answered) != 13 { + t.Fatalf("expected 13 answered leaves, got %d", len(answered)) + } + for _, leaf := range answered { + if math.Abs(leaf.Prob-1.0/13) > 1e-9 { + t.Errorf("leaf %s prob = %v, want %v", leaf.Key, leaf.Prob, 1.0/13) + } + } +} + +func TestErrorAndNoDataStatuses(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com"), nsRR("example.com", "ns2.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + aRR("ns2.example.com", "2.2.2.2"), + )) + m.on("1.1.1.1", "www.example.com", dns.TypeA, rcodeMsg(dns.RcodeNameError)) + soa := &dns.SOA{ + Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeSOA, Class: dns.ClassINET}, + Ns: "ns1.example.com.", Mbox: "hostmaster.example.com.", + } + nodata := new(dns.Msg) + nodata.Ns = []dns.RR{soa} + m.on("2.2.2.2", "www.example.com", dns.TypeA, nodata) + + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + errs := leavesByStatus(root, StatusError) + if len(errs) != 1 || errs[0].Response.DQ.ErrorMessage != "No such domain (NXDOMAIN)" { + t.Errorf("error leaves = %v", root.StatsList()) + } + nodataLeaves := leavesByStatus(root, StatusNoData) + if len(nodataLeaves) != 1 { + t.Errorf("nodata leaves = %v", root.StatsList()) + } +} + +func TestNetworkExceptionLeaf(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "www.example.com", dns.TypeA, referralMsg( + []dns.RR{nsRR("example.com", "ns1.example.com")}, + aRR("ns1.example.com", "1.1.1.1"), + )) + // no mock for 1.1.1.1 → network error → exception status + _, root := runTraversal(t, testConfig(false), m, "www.example.com") + + assertSumsToOne(t, root) + exceptions := leavesByStatus(root, StatusException) + if len(exceptions) != 1 { + t.Fatalf("expected exception leaf, got %v", root.StatsList()) + } + if math.Abs(exceptions[0].Prob-1.0) > 1e-9 { + t.Errorf("exception prob = %v", exceptions[0].Prob) + } +} + +func TestPacketCacheSingleWireQuery(t *testing.T) { + m := mockSimpleDelegation() + // Both NS answer; querying the same tuple twice must hit the cache. + tr, _ := runTraversal(t, testConfig(false), m, "www.example.com") + + m.mu.Lock() + defer m.mu.Unlock() + for key, n := range m.calls { + if n != 1 { + t.Errorf("query %v sent %d times, want 1", key, n) + } + } + if tr.client == nil { + t.Fatal("traverser has no client") + } +} + +func TestServersEncountered(t *testing.T) { + m := mockSimpleDelegation() + tr, _ := runTraversal(t, testConfig(false), m, "www.example.com") + seen := tr.ServersEncountered() + if len(seen["ns1.example.com"]) != 1 || seen["ns1.example.com"][0] != "1.1.1.1" { + t.Errorf("seen ns1 = %v", seen["ns1.example.com"]) + } + if _, ok := seen["a.gtld-servers.net"]; !ok { + t.Errorf("gtld server missing from seen: %v", seen) + } + if _, ok := seen[""]; ok { + t.Error("rootroot must not be recorded in servers encountered") + } +} + +func TestRunContextCancellation(t *testing.T) { + m := mockSimpleDelegation() + tr := NewTraverser(testConfig(false)) + tr.SetExchange(m.fn) ctx, cancel := context.WithCancel(context.Background()) cancel() - - _, err := tr.Traverse(ctx, "example.com") - if err == nil { - t.Fatal("expected error on cancelled context") + if _, err := tr.Run(ctx, "www.example.com"); err == nil { + t.Fatal("expected cancellation error") } } -func TestTraverserNXDOMAIN(t *testing.T) { - nxdResp := new(dns.Msg) - nxdResp.Rcode = dns.RcodeNameError - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nxdResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "nonexistent.invalid") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespNXDOMAIN { - t.Errorf("Type = %d, want %d", results[0].Response.Type, RespNXDOMAIN) +func TestIDNQnameConvertsToPunycode(t *testing.T) { + m := newMockExchange() + m.on("198.41.0.4", "xn--bcher-kva.example", dns.TypeA, answerMsg(aRR("xn--bcher-kva.example", "9.9.9.9"))) + _, root := runTraversal(t, testConfig(false), m, "bücher.example") + if root.Qname != "xn--bcher-kva.example" { + t.Errorf("qname = %q, want punycode", root.Qname) } + assertSumsToOne(t, root) } -func TestTraverserSERVFAIL(t *testing.T) { - sfResp := new(dns.Msg) - sfResp.Rcode = dns.RcodeServerFailure - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return sfResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespSERVFAIL { - t.Errorf("Type = %d, want %d", results[0].Response.Type, RespSERVFAIL) +func refids(events []TraversalEvent) []string { + out := make([]string, len(events)) + for i, ev := range events { + out[i] = ev.Stage.String() + ":" + ev.RefID } + return out } -func TestTraverserCNAMEFollow(t *testing.T) { - cnameResp := new(dns.Msg) - cnameResp.SetReply(new(dns.Msg)) - cnameResp.Answer = append(cnameResp.Answer, - &dns.CNAME{ - Hdr: dns.RR_Header{Name: "www.example.com.", Rrtype: dnsTypeCNAME, Class: dns.ClassINET}, - Target: "example.com.", - }, - ) - - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "www.example.com." { - return cnameResp.Copy(), nil - } - if q.Name == "example.com." { - return answerResp.Copy(), nil - } - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "www.example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - foundCNAME := false - foundAnswer := false - for _, r := range results { - if r.Response.Type == RespCNAMEFollow { - foundCNAME = true - } - if r.Response.Type == RespAnswer { - foundAnswer = true - } - } - if !foundCNAME { - t.Error("expected CNAME follow response") - } - if !foundAnswer { - t.Error("expected final answer response") +func refidsOf(refs []*Referral) []string { + out := make([]string, len(refs)) + for i, r := range refs { + out[i] = r.RefID } + return out } -func TestTraverserProbabilityCalculation(t *testing.T) { - rootReferral := new(dns.Msg) - rootReferral.Rcode = dns.RcodeSuccess - rootReferral.Authoritative = false - rootReferral.Ns = append(rootReferral.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "a.root-servers.net."}, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "b.root-servers.net."}, - ) - rootReferral.Extra = append(rootReferral.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "a.root-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("198.41.0.4")}, - &dns.A{Hdr: dns.RR_Header{Name: "b.root-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("199.9.14.201")}, - ) - - tldAnswer := new(dns.Msg) - tldAnswer.SetReply(new(dns.Msg)) - tldAnswer.Answer = append(tldAnswer.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("1.2.3.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "example.com." && server == "1.2.3.4" { - return rootReferral.Copy(), nil - } - return tldAnswer.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - for _, r := range results { - if r.Referral != nil && r.Referral.Depth == 1 && r.Referral.Parent != nil { - if r.Referral.Prob != 0.5 { - t.Errorf("child prob = %f, want 0.5", r.Referral.Prob) - } - } - } -} - -func TestTraverserNODATA(t *testing.T) { - nodataResp := new(dns.Msg) - nodataResp.Rcode = dns.RcodeSuccess - nodataResp.Authoritative = true - nodataResp.Ns = append(nodataResp.Ns, &dns.SOA{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeSOA, Class: dns.ClassINET, Ttl: 3600}, - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nodataResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result") - } - if results[0].Response.Type != RespNODATA { - t.Errorf("Type = %d, want %d", results[0].Response.Type, RespNODATA) - } -} - -func TestTraverserMultipleRoots(t *testing.T) { - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return answerResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected results") - } -} - -func TestTraverserNilExchangeResponse(t *testing.T) { - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - return nil, nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - if len(results) == 0 { - t.Fatal("expected at least 1 result even with nil response") - } -} - -func TestTraverserCacheChaining(t *testing.T) { - rootReferral := new(dns.Msg) - rootReferral.Rcode = dns.RcodeSuccess - rootReferral.Authoritative = false - rootReferral.Ns = append(rootReferral.Ns, - &dns.NS{Hdr: dns.RR_Header{Name: ".", Rrtype: dnsTypeNS}, Ns: "a.gtld-servers.net."}, - ) - rootReferral.Extra = append(rootReferral.Extra, - &dns.A{Hdr: dns.RR_Header{Name: "a.gtld-servers.net.", Rrtype: dnsTypeA}, A: net.ParseIP("192.5.6.30")}, - ) - - answerResp := new(dns.Msg) - answerResp.SetReply(new(dns.Msg)) - answerResp.Answer = append(answerResp.Answer, &dns.A{ - Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dnsTypeA, Class: dns.ClassINET, Ttl: 300}, - A: net.ParseIP("93.184.216.34"), - }) - - tr := NewTraverser(&TraverserConfig{ - MaxDepth: 5, - QueryType: dnsTypeA, - RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, - }) - tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { - q := msg.Question[0] - if q.Name == "example.com." && server == "198.41.0.4" { - return rootReferral.Copy(), nil - } - return answerResp.Copy(), nil - }) - - ctx := context.Background() - results, err := tr.Traverse(ctx, "example.com") - if err != nil { - t.Fatalf("unexpected error: %v", err) - } - - cacheHits := 0 - for _, r := range results { - if r.Response != nil && r.Response.Cache != nil { - if r.Response.Cache.NSCount() > 0 { - cacheHits++ - } - } - } - if cacheHits == 0 { - t.Error("expected cache to store NS records from referrals") +func cnameRRT(name, target string) dns.RR { + return &dns.CNAME{ + Hdr: dns.RR_Header{Name: dns.Fqdn(name), Rrtype: dns.TypeCNAME, Class: dns.ClassINET}, + Target: dns.Fqdn(target), } } diff --git a/web/api/handler.go b/web/api/handler.go index 7927057..a3f28ab 100644 --- a/web/api/handler.go +++ b/web/api/handler.go @@ -7,6 +7,8 @@ import ( "fmt" "io/fs" "net/http" + "sort" + "strings" "sync" "time" @@ -40,23 +42,57 @@ type TraverseStartResponse struct { // ProgressEvent carries a single traversal hook event. type ProgressEvent struct { - Stage string `json:"stage"` - Depth int `json:"depth"` - Name string `json:"name"` - QType string `json:"qtype"` - Server string `json:"server,omitempty"` - Bailiwick string `json:"bailiwick,omitempty"` - IsResolve bool `json:"is_resolve,omitempty"` + Stage string `json:"stage"` + RefID string `json:"refid"` + Depth int `json:"depth"` + Name string `json:"name"` + QType string `json:"qtype"` + Server string `json:"server,omitempty"` + IPs string `json:"ips,omitempty"` + Bailiwick string `json:"bailiwick,omitempty"` + Status string `json:"status,omitempty"` + IsResolve bool `json:"is_resolve,omitempty"` + CompletedEarlier string `json:"completed_earlier,omitempty"` } -// ResultItem is a single traversal step result for API consumers. +// ResultItem is one aggregated leaf outcome for API consumers. Parent and +// ParentIP identify the referring server so clients can render the noglue and +// lame-referral wordings; Qname/Qclass/Qtype are the failing query so clients +// can render the "While querying" line when it differs from the original. type ResultItem struct { - Depth int `json:"depth"` - Probability float64 `json:"probability"` - ResponseType string `json:"response_type"` - Server string `json:"server,omitempty"` - Answers []string `json:"answers,omitempty"` - CNAMEChain []string `json:"cname_chain,omitempty"` + RefID string `json:"refid,omitempty"` + Depth int `json:"depth"` + Probability float64 `json:"probability"` + Status string `json:"status"` + Server string `json:"server,omitempty"` + IP string `json:"ip,omitempty"` + Parent string `json:"parent,omitempty"` + ParentIP string `json:"parent_ip,omitempty"` + Qname string `json:"qname,omitempty"` + Qclass string `json:"qclass,omitempty"` + Qtype string `json:"qtype,omitempty"` + Answers []string `json:"answers,omitempty"` + Message string `json:"message,omitempty"` +} + +// SummaryAnswer is one distinct answered RRset with its accumulated +// probability (traverse.SummaryStats). +type SummaryAnswer struct { + Probability float64 `json:"probability"` + Records []string `json:"records"` +} + +// SummaryStatus is the accumulated probability of one non-answered status. +type SummaryStatus struct { + Status string `json:"status"` + Probability float64 `json:"probability"` +} + +// Summary is the grouped view of the aggregated leaves; probabilities across +// Answers plus ByStatus sum to 1.0. +type Summary struct { + Answers []SummaryAnswer `json:"answers,omitempty"` + ByStatus []SummaryStatus `json:"by_status,omitempty"` } // TraversalJob holds all state for a single asynchronous traversal. @@ -66,6 +102,7 @@ type TraversalJob struct { Domain string `json:"domain"` QueryType string `json:"query_type"` Results []ResultItem `json:"results,omitempty"` + Summary *Summary `json:"summary,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"` Error string `json:"error,omitempty"` StartedAt time.Time `json:"started_at"` @@ -76,14 +113,18 @@ type TraversalJob struct { cancel context.CancelFunc } -// subscribe returns a channel that receives future progress events. -// The channel is closed when the job finishes. -func (j *TraversalJob) subscribe() <-chan ProgressEvent { +// subscribeSnapshot atomically registers a subscriber and snapshots the +// progress recorded so far. Publishing appends to Progress and sends to +// subscribers under the same lock, so every event lands either in the +// returned snapshot or on the channel — never both, never neither. +func (j *TraversalJob) subscribeSnapshot() (sub <-chan ProgressEvent, past []ProgressEvent, done bool) { ch := make(chan ProgressEvent, 32) j.mu.Lock() + defer j.mu.Unlock() j.subs = append(j.subs, ch) - j.mu.Unlock() - return ch + past = make([]ProgressEvent, len(j.Progress)) + copy(past, j.Progress) + return ch, past, j.Status != statusRunning } // publishLocked sends ev to all current subscribers. Caller must hold j.mu. @@ -259,6 +300,7 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) { Domain string `json:"domain"` QueryType string `json:"query_type"` Results []ResultItem `json:"results,omitempty"` + Summary *Summary `json:"summary,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"` Error string `json:"error,omitempty"` StartedAt time.Time `json:"started_at"` @@ -269,6 +311,7 @@ func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) { Domain: job.Domain, QueryType: job.QueryType, Results: job.Results, + Summary: job.Summary, Progress: job.Progress, Error: job.Error, StartedAt: job.StartedAt, @@ -300,19 +343,12 @@ func (h *Handler) streamTraversal(w http.ResponseWriter, r *http.Request) { return } - // Subscribe before snapshotting progress so we don't miss events between - // the two operations. Unsubscribe when the client disconnects so stale - // channels don't accumulate. - sub := job.subscribe() + // Subscribe and snapshot atomically so events published in between are + // neither missed nor delivered twice. Unsubscribe when the client + // disconnects so stale channels don't accumulate. + sub, past, alreadyDone := job.subscribeSnapshot() defer job.unsubscribe(sub) - // Replay events already recorded. - job.mu.RLock() - past := make([]ProgressEvent, len(job.Progress)) - copy(past, job.Progress) - alreadyDone := job.Status != statusRunning - job.mu.RUnlock() - sendSSE := func(ev ProgressEvent) bool { b, err := json.Marshal(ev) if err != nil { @@ -375,26 +411,22 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st } cfg.Hooks = &traverse.TraverserHooks{ OnEvent: func(event traverse.TraversalEvent) { - ref := event.Result.Referral - if ref == nil { + ref := event.Referral + if ref == nil || ref.IsRootRoot() { return } - stage := "start" - if event.Stage == traverse.EventComplete { - stage = "complete" - } - server := "" - if event.Result.Response != nil && event.Result.Response.Server != nil { - server = event.Result.Response.Server.String() - } ev := ProgressEvent{ - Stage: stage, - Depth: ref.Depth, - Name: trimFQDN(ref.Name), - QType: idns.QNameType(ref.Qtype), - Server: server, - Bailiwick: trimFQDN(ref.Bailiwick), - IsResolve: event.IsResolve, + Stage: event.Stage.String(), + RefID: event.RefID, + Depth: ref.Depth(), + Name: ref.Qname, + QType: traverse.TypeToString(ref.Qtype), + Server: ref.Server, + IPs: ref.TxtIPs(), + Bailiwick: ref.Bailiwick, + Status: string(event.Status), + IsResolve: event.IsResolve, + CompletedEarlier: event.CompletedEarlier, } job.mu.Lock() @@ -405,7 +437,7 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st } tr := traverse.NewTraverser(cfg) - rawResults, err := tr.Traverse(ctx, domain) + root, err := tr.Run(ctx, domain) now := time.Now() job.mu.Lock() @@ -419,31 +451,87 @@ func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain st return } - items := make([]ResultItem, 0, len(rawResults)) - for _, r := range rawResults { - items = append(items, toResultItem(r)) + var items []ResultItem + if root != nil { + leaves := root.StatsList() + items = make([]ResultItem, 0, len(leaves)) + for _, leaf := range leaves { + items = append(items, toResultItem(leaf)) + } + job.Summary = toSummary(root.SummaryStats()) } job.Results = items job.Status = statusComplete } -// toResultItem converts a TraversalResult to its API representation. -func toResultItem(r traverse.TraversalResult) ResultItem { - item := ResultItem{} - if r.Referral != nil { - item.Depth = r.Referral.Depth - item.Probability = r.Referral.Prob +// toSummary converts the engine's grouped stats to the API representation: +// answers sorted by RRset key (as SummaryStats returns them), remaining +// statuses sorted lexically like the CLI Summary Results section. +func toSummary(stats *traverse.SummaryStats) *Summary { + if stats == nil { + return nil } - if r.Response != nil { - item.ResponseType = r.Response.Type.String() - if r.Response.Server != nil { - item.Server = r.Response.Server.String() + summary := &Summary{} + for _, answer := range stats.Answers { + item := SummaryAnswer{Probability: answer.Prob} + for _, rr := range answer.RRs { + item.Records = append(item.Records, collapseWhitespace(rr.String())) } - if r.Response.Decoded != nil { - for _, rr := range r.Response.Decoded.Answers { - item.Answers = append(item.Answers, idns.FormatRecord(rr)) - } - item.CNAMEChain = append(item.CNAMEChain, r.Response.Decoded.CNAMEChain...) + summary.Answers = append(summary.Answers, item) + } + statuses := make([]traverse.Status, 0, len(stats.ByStatus)) + for status := range stats.ByStatus { + if status != traverse.StatusAnswered { + statuses = append(statuses, status) + } + } + sort.Slice(statuses, func(i, j int) bool { return statuses[i] < statuses[j] }) + for _, status := range statuses { + summary.ByStatus = append(summary.ByStatus, SummaryStatus{ + Status: string(status), + Probability: stats.ByStatus[status], + }) + } + return summary +} + +// collapseWhitespace renders an RR on one line with runs of whitespace +// collapsed to single spaces, matching the CLI summary records. +func collapseWhitespace(s string) string { + return strings.Join(strings.Fields(s), " ") +} + +// toResultItem converts one aggregated leaf to its API representation. +func toResultItem(leaf *traverse.StatsEntry) ResultItem { + resp := leaf.Response + item := ResultItem{ + Probability: leaf.Prob, + Status: string(resp.Status), + IP: resp.IP, + Server: resp.Server, + ParentIP: resp.ParentIP, + Qname: trimFQDN(resp.Qname), + Qclass: traverse.ClassToString(resp.Qclass), + Qtype: traverse.TypeToString(resp.Qtype), + } + if leaf.Referral != nil { + item.RefID = leaf.Referral.RefID + item.Depth = leaf.Referral.Depth() + item.Server = leaf.Referral.Server + item.ParentIP = leaf.Referral.ParentIP + if leaf.Referral.Parent != nil { + item.Parent = leaf.Referral.Parent.Server + } + } + if resp.DQ != nil { + for _, rr := range resp.DQ.Answers { + item.Answers = append(item.Answers, idns.FormatRecord(rr)) + } + switch resp.Status { + case traverse.StatusError: + item.Message = resp.DQ.ErrorMessage + case traverse.StatusException: + item.Message = resp.DQ.ExceptionMessage } } return item @@ -495,4 +583,3 @@ func newUUID() string { b[8] = (b[8] & 0x3f) | 0x80 // variant bits return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:]) } - diff --git a/web/api/handler_internal_test.go b/web/api/handler_internal_test.go new file mode 100644 index 0000000..4324c37 --- /dev/null +++ b/web/api/handler_internal_test.go @@ -0,0 +1,77 @@ +package api + +import ( + "strconv" + "sync" + "testing" +) + +// TestSubscribeSnapshot_NoDuplicates regresses the subscribe/snapshot race: +// subscribers arriving while events are being published must never see the +// same event twice (once from the snapshot replay and once from the channel). +// Events are numbered, so any duplicate breaks strict monotonicity. Slow +// subscribers may legitimately drop events (publishLocked is non-blocking), +// so gaps are not an error. +func TestSubscribeSnapshot_NoDuplicates(t *testing.T) { + job := &TraversalJob{Status: statusRunning} + + const total = 2000 + const subscribers = 8 + + var pub sync.WaitGroup + pub.Add(1) + go func() { + defer pub.Done() + for i := 0; i < total; i++ { + ev := ProgressEvent{RefID: strconv.Itoa(i)} + job.mu.Lock() + job.Progress = append(job.Progress, ev) + job.publishLocked(ev) + job.mu.Unlock() + } + job.mu.Lock() + job.Status = statusComplete + job.mu.Unlock() + job.closeSubscribers() + }() + + var subs sync.WaitGroup + errs := make(chan string, subscribers) + for s := 0; s < subscribers; s++ { + subs.Add(1) + go func() { + defer subs.Done() + sub, past, done := job.subscribeSnapshot() + defer job.unsubscribe(sub) + + last := -1 + check := func(refid string) { + n, err := strconv.Atoi(refid) + if err != nil { + errs <- "bad refid " + refid + return + } + if n <= last { + errs <- "event " + refid + " out of order or duplicated after " + strconv.Itoa(last) + return + } + last = n + } + for _, ev := range past { + check(ev.RefID) + } + if !done { + for ev := range sub { + check(ev.RefID) + } + } + }() + } + + pub.Wait() + subs.Wait() + close(errs) + for msg := range errs { + t.Error(msg) + } +} diff --git a/web/api/handler_test.go b/web/api/handler_test.go index 31a8a32..c0b09ad 100644 --- a/web/api/handler_test.go +++ b/web/api/handler_test.go @@ -5,12 +5,15 @@ import ( "bytes" "encoding/json" "fmt" + "io" "net/http" "net/http/httptest" + "regexp" "strings" "testing" "time" + "gitea.hansenits.com.au/hits/ExploreDNS/internal/config" "gitea.hansenits.com.au/hits/ExploreDNS/web/api" ) @@ -264,6 +267,43 @@ func TestStaticSPA_Index(t *testing.T) { } } +// TestStaticSPA_TypeOptions asserts the SPA type dropdown offers exactly the +// query types config.ParseQueryType accepts. +func TestStaticSPA_TypeOptions(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + resp, err := http.Get("http://" + srv.Addr() + "/") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + body, err := io.ReadAll(resp.Body) + if err != nil { + t.Fatal(err) + } + page := string(body) + + re := regexp.MustCompile(``) + var got []string + for _, m := range re.FindAllStringSubmatch(page, -1) { + got = append(got, m[1]) + } + want := []string{"A", "AAAA", "NS", "CNAME", "MX", "TXT", "SOA", "PTR", "ANY"} + if len(got) != len(want) { + t.Fatalf("type options = %v, want %v", got, want) + } + for i, typ := range want { + if got[i] != typ { + t.Fatalf("type options = %v, want %v", got, want) + } + if _, err := config.ParseQueryType(typ); err != nil { + t.Fatalf("option %s rejected by ParseQueryType: %v", typ, err) + } + } +} + func TestStaticSPA_FallbackToIndex(t *testing.T) { srv := newTestServer(t) defer srv.Shutdown(5 * time.Second) //nolint:errcheck diff --git a/web/api/static/index.html b/web/api/static/index.html index 2212f65..8e101b2 100644 --- a/web/api/static/index.html +++ b/web/api/static/index.html @@ -285,22 +285,28 @@ content: 'Awaiting traversal…'; color: var(--text-dim); } - .pe { display: flex; gap: 0.5rem; } + .pe { display: flex; gap: 0.5rem; align-items: baseline; } .pe-badge { flex-shrink: 0; + min-width: 70px; + text-align: center; font-size: 0.65rem; padding: 1px 5px; border-radius: 3px; font-weight: 600; text-transform: uppercase; letter-spacing: 0.04em; + background: rgba(139,143,168,0.15); + color: var(--text-muted); } - .pe-badge.start { background: rgba(74,124,246,0.18); color: var(--primary); } - .pe-badge.complete { background: rgba(62,207,142,0.18); color: var(--success); } - .pe-badge.resolve { background: rgba(56,189,248,0.18); color: var(--info); } - .pe-depth { color: var(--text-dim); } - .pe-name { color: var(--text); } - .pe-server { color: var(--text-muted); } + .pe-badge.st-working { background: rgba(74,124,246,0.18); color: var(--primary); } + .pe-badge.st-answered { background: rgba(62,207,142,0.18); color: var(--success); } + .pe-badge.st-resolve { background: rgba(56,189,248,0.18); color: var(--info); } + .pe-badge.st-warn { background: rgba(245,166,35,0.18); color: var(--warning); } + .pe-badge.st-err { background: rgba(240,68,56,0.18); color: var(--danger); } + /* Progress lines mirror the CLI: " ()" — keep the + spacing intact. */ + .pe-text { color: var(--text); white-space: pre; } /* ── Results tree ───────────────────────────────────────────────── */ .results-section h2 { @@ -318,57 +324,44 @@ color: var(--text-dim); font-size: 0.875rem; } - .result-item { + /* Aggregated leaves rendered like the CLI Results section: fixed-width + text with significant leading spaces, so white-space must be pre. */ + .result-block { font-family: var(--font-mono); - font-size: 0.8rem; + font-size: 0.78rem; + line-height: 1.5; + white-space: pre; + overflow-x: auto; + margin: 0 0 0.5rem; padding: 0.5rem 0.75rem; - border-radius: var(--radius-sm); border: 1px solid var(--border); - margin-bottom: 0.5rem; + border-left-width: 3px; + border-radius: var(--radius-sm); background: var(--bg-card); - transition: border-color 0.12s; + color: var(--text); } - .result-item:hover { border-color: var(--border-focus); } - .result-header { - display: flex; - align-items: center; - gap: 0.5rem; - flex-wrap: wrap; - } - .result-indent { color: var(--text-dim); flex-shrink: 0; } - .rtype-badge { - font-size: 0.65rem; - padding: 1px 5px; - border-radius: 3px; - font-weight: 700; - text-transform: uppercase; - } - .rtype-answer { background: rgba(62,207,142,0.15); color: var(--success); } - .rtype-referral { background: rgba(74,124,246,0.15); color: var(--primary); } - .rtype-nxdomain, - .rtype-servfail { background: rgba(240,68,56,0.15); color: var(--danger); } - .rtype-timeout { background: rgba(245,166,35,0.15); color: var(--warning); } - .rtype-nodata { background: rgba(56,189,248,0.15); color: var(--info); } - .rtype-cname { background: rgba(168,85,247,0.15); color: #c084fc; } - .rtype-other { background: rgba(139,143,168,0.15);color: var(--text-muted); } + .result-block.st-answered { border-left-color: var(--success); } + .result-block.st-warn { border-left-color: var(--warning); } + .result-block.st-err { border-left-color: var(--danger); } + .result-block.st-other { border-left-color: var(--text-dim); } - .result-prob { - color: var(--text-muted); - font-size: 0.72rem; - } - .result-server { color: var(--text-dim); font-size: 0.72rem; margin-left: auto; } - .result-answers { - margin-top: 0.35rem; - padding-left: 1.2rem; - color: var(--text-muted); + /* ── Summary Results ───────────────────────────────────────────── */ + .summary-pre { + font-family: var(--font-mono); font-size: 0.75rem; + line-height: 1.6; + white-space: pre; + overflow-x: auto; + margin: 0 0 1rem; + padding: 0.6rem 0.75rem; + border: 1px solid var(--border); + border-radius: var(--radius-sm); + background: var(--bg); + color: var(--text); } - .result-answers span { display: block; } - .cname-chain { - color: #c084fc; - font-size: 0.72rem; - margin-top: 0.2rem; - padding-left: 1.2rem; + .summary-pre:empty::before { + content: 'No summary yet.'; + color: var(--text-dim); } /* ── Stats panel ────────────────────────────────────────────────── */ @@ -478,17 +471,16 @@ autocapitalize="off" spellcheck="false" /> +