feat: implement output formatting and display (HAN-383) #9

Merged
multica-agent merged 2 commits from agent/go-expert-developer/f3b1b7cf into main 2026-06-07 16:40:02 +00:00
14 changed files with 268 additions and 294 deletions
Showing only changes of commit 8e7beacc22 - Show all commits
+187 -186
View File
@@ -1,219 +1,220 @@
package main package main
import ( import (
"context" "context"
"flag" "flag"
"fmt" "fmt"
"os" "os"
"time" "time"
"github.com/hits/ExploreDNS/internal/config" "github.com/hits/ExploreDNS/internal/config"
"github.com/hits/ExploreDNS/internal/dns" "github.com/hits/ExploreDNS/internal/dns"
"github.com/hits/ExploreDNS/internal/output" "github.com/hits/ExploreDNS/internal/output"
"github.com/hits/ExploreDNS/internal/traverse" "github.com/hits/ExploreDNS/internal/traverse"
) )
func main() { func main() {
cfg := config.DefaultConfig() cfg := config.DefaultConfig()
queryType := flag.String("type", cfg.QueryType, "Record type (A, AAAA, NS, CNAME, MX, TXT, SOA, PTR, ANY)") 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") rootServer := flag.String("root-server", cfg.RootServer, "Override root server")
allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Use all 13 root servers") allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Use all 13 root servers")
rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses") rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses")
followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals") followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals")
udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)") udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)")
allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation") allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation")
alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP") alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP")
maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)") maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)")
retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)") retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)")
fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache") fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache")
jsonOutput := flag.Bool("json", false, "Output results as JSON") jsonOutput := flag.Bool("json", false, "Output results as JSON")
// Verbose: long and short form share the same variable. // Verbose: long and short form share the same variable.
var verboseVal bool var verboseVal bool
flag.BoolVar(&verboseVal, "verbose", cfg.Verbose, "Verbose output") flag.BoolVar(&verboseVal, "verbose", cfg.Verbose, "Verbose output")
flag.BoolVar(&verboseVal, "v", cfg.Verbose, "Verbose output (shorthand)") flag.BoolVar(&verboseVal, "v", cfg.Verbose, "Verbose output (shorthand)")
// Debug: -d sets level 1, -dd sets level 2 (library debug). // Debug: -d sets level 1, -dd sets level 2 (library debug).
var dFlag, ddFlag bool var dFlag, ddFlag bool
flag.BoolVar(&dFlag, "d", false, "Debug mode (stackable: -dd for library debug)") flag.BoolVar(&dFlag, "d", false, "Debug mode (stackable: -dd for library debug)")
flag.BoolVar(&dFlag, "debug", false, "Debug mode") flag.BoolVar(&dFlag, "debug", false, "Debug mode")
flag.BoolVar(&ddFlag, "dd", false, "Library debug mode (equivalent to -d -d)") flag.BoolVar(&ddFlag, "dd", false, "Library debug mode (equivalent to -d -d)")
// Quiet: long and short form share the same variable. // Quiet: long and short form share the same variable.
var quietVal bool var quietVal bool
flag.BoolVar(&quietVal, "quiet", cfg.Quiet, "Suppress supplementary info") flag.BoolVar(&quietVal, "quiet", cfg.Quiet, "Suppress supplementary info")
flag.BoolVar(&quietVal, "q", cfg.Quiet, "Suppress supplementary info (shorthand)") flag.BoolVar(&quietVal, "q", cfg.Quiet, "Suppress supplementary info (shorthand)")
showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress") showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress")
noShowProgress := flag.Bool("no-show-progress", false, "Hide traversal progress") noShowProgress := flag.Bool("no-show-progress", false, "Hide traversal progress")
showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details") showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details")
noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details") noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details")
showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried") showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried")
noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried") noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried")
showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions") showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions")
noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions") noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions")
showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics") showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics")
noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics") noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics")
showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results") showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results")
noShowResults := flag.Bool("no-show-results", false, "Hide query results") noShowResults := flag.Bool("no-show-results", false, "Hide query results")
showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results") showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results")
noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results") noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results")
flag.Usage = config.PrintUsage flag.Usage = config.PrintUsage
flag.Parse() flag.Parse()
args := flag.Args() args := flag.Args()
cfg.QueryType = *queryType cfg.QueryType = *queryType
cfg.RootServer = *rootServer cfg.RootServer = *rootServer
cfg.AllRootServers = *allRootServers cfg.AllRootServers = *allRootServers
cfg.RootAAAA = *rootAAAA cfg.RootAAAA = *rootAAAA
cfg.FollowAAAA = *followAAAA cfg.FollowAAAA = *followAAAA
cfg.UDPSize = *udpSize cfg.UDPSize = *udpSize
cfg.AllowTCP = *allowTCP cfg.AllowTCP = *allowTCP
cfg.AlwaysTCP = *alwaysTCP cfg.AlwaysTCP = *alwaysTCP
cfg.MaxDepth = *maxDepth cfg.MaxDepth = *maxDepth
cfg.Retries = *retries cfg.Retries = *retries
cfg.Fast = *fast cfg.Fast = *fast
cfg.Verbose = verboseVal cfg.Verbose = verboseVal
cfg.Debug = config.ParseDebugLevel(dFlag, ddFlag) cfg.Debug = config.ParseDebugLevel(dFlag, ddFlag)
cfg.Quiet = quietVal cfg.Quiet = quietVal
if *noShowProgress { if *noShowProgress {
cfg.ShowProgress = false cfg.ShowProgress = false
} else if *showProgress { } else if *showProgress {
cfg.ShowProgress = true cfg.ShowProgress = true
} }
if *noShowResolves { if *noShowResolves {
cfg.ShowResolves = false cfg.ShowResolves = false
} else if *showResolves { } else if *showResolves {
cfg.ShowResolves = true cfg.ShowResolves = true
} }
if *noShowServers { if *noShowServers {
cfg.ShowServers = false cfg.ShowServers = false
} else if *showServers { } else if *showServers {
cfg.ShowServers = true cfg.ShowServers = true
} }
if *noShowVersions { if *noShowVersions {
cfg.ShowVersions = false cfg.ShowVersions = false
} else if *showVersions { } else if *showVersions {
cfg.ShowVersions = true cfg.ShowVersions = true
} }
if *noShowAllStats { if *noShowAllStats {
cfg.ShowAllStats = false cfg.ShowAllStats = false
} else if *showAllStats { } else if *showAllStats {
cfg.ShowAllStats = true cfg.ShowAllStats = true
} }
if *noShowResults { if *noShowResults {
cfg.ShowResults = false cfg.ShowResults = false
} else if *showResults { } else if *showResults {
cfg.ShowResults = true cfg.ShowResults = true
} }
if *noShowSummaryResults { if *noShowSummaryResults {
cfg.ShowSummaryResults = false cfg.ShowSummaryResults = false
} else if *showSummaryResults { } else if *showSummaryResults {
cfg.ShowSummaryResults = true cfg.ShowSummaryResults = true
} }
if err := cfg.Validate(); err != nil { if err := cfg.Validate(); err != nil {
fmt.Fprintf(os.Stderr, "Error: %v\n", err) fmt.Fprintf(os.Stderr, "Error: %v\n", err)
os.Exit(1) os.Exit(1)
} }
domain, err := cfg.GetDomain(args) domain, err := cfg.GetDomain(args)
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "Error: %v\n", err) fmt.Fprintf(os.Stderr, "Error: %v\n", err)
os.Exit(1) os.Exit(1)
} }
rootIP, err := cfg.ParseRootServer() rootIP, err := cfg.ParseRootServer()
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "Error: invalid root server: %v\n", err) fmt.Fprintf(os.Stderr, "Error: invalid root server: %v\n", err)
os.Exit(1) os.Exit(1)
} }
queryTypeValue, err := config.ParseQueryType(cfg.QueryType) queryTypeValue, err := config.ParseQueryType(cfg.QueryType)
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "Error: %v\n", err) fmt.Fprintf(os.Stderr, "Error: %v\n", err)
os.Exit(1) os.Exit(1)
} }
var rootServerAddr string var rootServerAddr string
if rootIP != nil { if rootIP != nil {
rootServerAddr = rootIP.String() rootServerAddr = rootIP.String()
} }
queryConfig := &dns.QueryConfig{ queryConfig := &dns.QueryConfig{
UDPSize: cfg.UDPSize, UDPSize: cfg.UDPSize,
Timeout: 5 * time.Second, Timeout: 5 * time.Second,
Retries: cfg.Retries, Retries: cfg.Retries,
UseTCP: cfg.AlwaysTCP, UseTCP: cfg.AlwaysTCP,
AllowTCP: cfg.AllowTCP, AllowTCP: cfg.AllowTCP,
} }
rootConfig := &dns.RootDiscoveryConfig{ rootConfig := &dns.RootDiscoveryConfig{
IncludeAAAA: cfg.RootAAAA, IncludeAAAA: cfg.RootAAAA,
Server: rootServerAddr, Server: rootServerAddr,
AllRoots: cfg.AllRootServers, AllRoots: cfg.AllRootServers,
} }
traverserConfig := &traverse.TraverserConfig{ traverserConfig := &traverse.TraverserConfig{
MaxDepth: cfg.MaxDepth, MaxDepth: cfg.MaxDepth,
QueryType: queryTypeValue, QueryType: queryTypeValue,
RootConfig: rootConfig, RootConfig: rootConfig,
QueryConfig: queryConfig, QueryConfig: queryConfig,
} }
traverser := traverse.NewTraverser(traverserConfig) traverser := traverse.NewTraverser(traverserConfig)
if cfg.Debug > 0 { if cfg.Debug > 0 {
fmt.Fprintf(os.Stderr, "Debug: Config loaded\n") fmt.Fprintf(os.Stderr, "Debug: Config loaded\n")
fmt.Fprintf(os.Stderr, " Domain: %s\n", domain) fmt.Fprintf(os.Stderr, " Domain: %s\n", domain)
fmt.Fprintf(os.Stderr, " Query Type: %s\n", cfg.QueryType) fmt.Fprintf(os.Stderr, " Query Type: %s\n", cfg.QueryType)
fmt.Fprintf(os.Stderr, " Max Depth: %d\n", cfg.MaxDepth) fmt.Fprintf(os.Stderr, " Max Depth: %d\n", cfg.MaxDepth)
fmt.Fprintf(os.Stderr, " UDP Size: %d\n", cfg.UDPSize) fmt.Fprintf(os.Stderr, " UDP Size: %d\n", cfg.UDPSize)
fmt.Fprintf(os.Stderr, " Retries: %d\n", cfg.Retries) fmt.Fprintf(os.Stderr, " Retries: %d\n", cfg.Retries)
fmt.Fprintf(os.Stderr, " Allow TCP: %v\n", cfg.AllowTCP) fmt.Fprintf(os.Stderr, " Allow TCP: %v\n", cfg.AllowTCP)
fmt.Fprintf(os.Stderr, " Always TCP: %v\n", cfg.AlwaysTCP) fmt.Fprintf(os.Stderr, " Always TCP: %v\n", cfg.AlwaysTCP)
fmt.Fprintf(os.Stderr, " Fast Mode: %v\n", cfg.Fast) fmt.Fprintf(os.Stderr, " Fast Mode: %v\n", cfg.Fast)
} }
if !cfg.Quiet && !*jsonOutput { if !cfg.Quiet && !*jsonOutput {
fmt.Printf("ExploreDNS - exploring: %s (type: %s)\n", domain, cfg.QueryType) fmt.Printf("ExploreDNS - exploring: %s (type: %s)\n", domain, cfg.QueryType)
} }
outFmt := output.FormatText outFmt := output.FormatText
if *jsonOutput { if *jsonOutput {
outFmt = output.FormatJSON outFmt = output.FormatJSON
} }
outCfg := &output.Config{ outCfg := &output.Config{
Format: outFmt, Format: outFmt,
Domain: domain, Domain: domain,
QueryType: cfg.QueryType, QueryType: cfg.QueryType,
ShowProgress: cfg.ShowProgress, ShowProgress: cfg.ShowProgress,
ShowResolves: cfg.ShowResolves, ShowResolves: cfg.ShowResolves,
ShowServers: cfg.ShowServers, ShowServers: cfg.ShowServers,
ShowVersions: cfg.ShowVersions, ShowVersions: cfg.ShowVersions,
ShowAllStats: cfg.ShowAllStats, ShowAllStats: cfg.ShowAllStats,
ShowResults: cfg.ShowResults, ShowResults: cfg.ShowResults,
ShowSummaryResults: cfg.ShowSummaryResults, ShowSummaryResults: cfg.ShowSummaryResults,
Verbose: cfg.Verbose, Verbose: cfg.Verbose,
Quiet: cfg.Quiet, Quiet: cfg.Quiet,
Color: os.Getenv("NO_COLOR") == "", Color: os.Getenv("NO_COLOR") == "",
} Debug: cfg.Debug,
}
ctx := context.Background() ctx := context.Background()
formatter := output.NewFormatter(outCfg, os.Stdout) formatter := output.NewFormatter(outCfg, os.Stdout)
_, err = output.RunTraversal(ctx, traverser, outCfg, formatter, domain) _, err = output.RunTraversal(ctx, traverser, outCfg, formatter, domain)
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "Error: traversal failed: %v\n", err) fmt.Fprintf(os.Stderr, "Error: traversal failed: %v\n", err)
os.Exit(1) os.Exit(1)
} }
if cfg.Debug > 0 { if cfg.Debug > 0 {
fmt.Fprintf(os.Stderr, "Debug: Traversal completed\n") fmt.Fprintf(os.Stderr, "Debug: Traversal completed\n")
} }
} }
+7 -7
View File
@@ -36,13 +36,13 @@ type Config struct {
Debug int Debug int
Quiet bool Quiet bool
ShowProgress bool ShowProgress bool
ShowResolves bool ShowResolves bool
ShowServers bool ShowServers bool
ShowVersions bool ShowVersions bool
ShowAllStats bool ShowAllStats bool
ShowResults bool ShowResults bool
ShowSummaryResults bool ShowSummaryResults bool
} }
func ParseQueryType(s string) (uint16, error) { func ParseQueryType(s string) (uint16, error) {
+5 -5
View File
@@ -314,8 +314,8 @@ func TestCNAMEChain(t *testing.T) {
func TestResponseClassificationString(t *testing.T) { func TestResponseClassificationString(t *testing.T) {
tests := []struct { tests := []struct {
rc ResponseClassification rc ResponseClassification
want string want string
}{ }{
{ResponseAnswer, "answer"}, {ResponseAnswer, "answer"},
{ResponseReferral, "referral"}, {ResponseReferral, "referral"},
@@ -344,10 +344,10 @@ func TestFormatRecord(t *testing.T) {
t.Run("A record", func(t *testing.T) { t.Run("A record", func(t *testing.T) {
rr := &dns.A{ rr := &dns.A{
Hdr: dns.RR_Header{ Hdr: dns.RR_Header{
Name: "example.com.", Name: "example.com.",
Rrtype: dns.TypeA, Rrtype: dns.TypeA,
Class: dns.ClassINET, Class: dns.ClassINET,
Ttl: 300, Ttl: 300,
}, },
A: MustParseIP("93.184.216.34"), A: MustParseIP("93.184.216.34"),
} }
-1
View File
@@ -328,4 +328,3 @@ func TestQueryNoTCPFallbackWhenDisabled(t *testing.T) {
t.Error("expected truncated response to be returned as-is") t.Error("expected truncated response to be returned as-is")
} }
} }
-29
View File
@@ -1,9 +1,6 @@
package dns package dns
import ( import (
"fmt"
"strings"
"github.com/miekg/dns" "github.com/miekg/dns"
) )
@@ -40,32 +37,6 @@ func QNameType(qtype uint16) string {
return dns.TypeToString[qtype] return dns.TypeToString[qtype]
} }
func ParseQueryType(s string) (uint16, error) {
s = strings.ToUpper(strings.TrimSpace(s))
switch s {
case "A":
return TypeA, nil
case "AAAA":
return TypeAAAA, nil
case "NS":
return TypeNS, nil
case "CNAME":
return TypeCNAME, nil
case "MX":
return TypeMX, nil
case "TXT":
return TypeTXT, nil
case "SOA":
return TypeSOA, nil
case "PTR":
return TypePTR, nil
case "ANY":
return TypeANY, nil
default:
return 0, fmt.Errorf("invalid query type: %s", s)
}
}
func DefaultEDNS0UDPSize() int { func DefaultEDNS0UDPSize() int {
return 2048 return 2048
} }
+10 -3
View File
@@ -1,6 +1,7 @@
package output package output
import ( import (
"fmt"
"io" "io"
"os" "os"
@@ -28,6 +29,7 @@ type Config struct {
Verbose bool Verbose bool
Quiet bool Quiet bool
Color bool Color bool
Debug int
} }
func DefaultConfig() *Config { func DefaultConfig() *Config {
@@ -69,16 +71,21 @@ func AttachHooks(cfg *Config, formatter Formatter) *traverse.TraverserHooks {
if cfg == nil || formatter == nil { if cfg == nil || formatter == nil {
return nil 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)
}
}
return &traverse.TraverserHooks{ return &traverse.TraverserHooks{
OnEvent: func(event traverse.TraversalEvent) { OnEvent: func(event traverse.TraversalEvent) {
switch { switch {
case event.IsResolve && cfg.ShowResolves: case event.IsResolve && cfg.ShowResolves:
_ = formatter.WriteResolve(event) logErr("WriteResolve", formatter.WriteResolve(event))
case !event.IsResolve && cfg.ShowProgress: case !event.IsResolve && cfg.ShowProgress:
_ = formatter.WriteProgress(event) logErr("WriteProgress", formatter.WriteProgress(event))
} }
if event.Stage == traverse.EventComplete && cfg.ShowAllStats { if event.Stage == traverse.EventComplete && cfg.ShowAllStats {
_ = formatter.WriteResult(event.Result) logErr("WriteResult", formatter.WriteResult(event.Result))
} }
}, },
} }
+2 -2
View File
@@ -54,8 +54,8 @@ type jsonSummary struct {
} }
type jsonAnswerStat struct { type jsonAnswerStat struct {
RData string `json:"rdata"` RData string `json:"rdata"`
Probability float64 `json:"probability"` Probability float64 `json:"probability"`
Records []string `json:"records,omitempty"` Records []string `json:"records,omitempty"`
} }
+3
View File
@@ -73,6 +73,9 @@ func ComputeSummary(results []traverse.TraversalResult) *SummaryStats {
return stats.Answers[i].RData < stats.Answers[j].RData return stats.Answers[i].RData < stats.Answers[j].RData
}) })
if len(stats.Answers) == 0 && len(stats.ByType) == 0 {
return nil
}
return stats return stats
} }
+1 -8
View File
@@ -99,11 +99,7 @@ func (f *textFormatter) writeServers(results []traverse.TraversalResult) error {
for _, name := range names { for _, name := range names {
for _, ip := range servers[name] { for _, ip := range servers[name] {
version := "" if _, err := fmt.Fprintf(f.w, "%*s: %-15s\n", width, name, ip); err != nil {
if f.cfg.ShowVersions {
version = " (version lookup pending)"
}
if _, err := fmt.Fprintf(f.w, "%*s: %-15s%s\n", width, name, ip, version); err != nil {
return err return err
} }
} }
@@ -234,9 +230,6 @@ func referralID(ref *traverse.Referral) string {
if ref == nil { if ref == nil {
return "" return ""
} }
if ref.Depth == 0 {
return "1"
}
return fmt.Sprintf("%d", ref.Depth+1) return fmt.Sprintf("%d", ref.Depth+1)
} }
+4 -4
View File
@@ -9,10 +9,10 @@ import (
) )
type InfoCache struct { type InfoCache struct {
parent *InfoCache parent *InfoCache
mu sync.RWMutex mu sync.RWMutex
ns map[string][]string ns map[string][]string
glue map[string][]net.IP glue map[string][]net.IP
} }
func NewInfoCache(parent *InfoCache) *InfoCache { func NewInfoCache(parent *InfoCache) *InfoCache {
+9 -9
View File
@@ -32,18 +32,18 @@ func (s ResolutionState) String() string {
} }
type Referral struct { type Referral struct {
Name string Name string
Qtype uint16 Qtype uint16
Qclass uint16 Qclass uint16
Bailiwick string Bailiwick string
Addresses []net.IP Addresses []net.IP
State ResolutionState State ResolutionState
NSName string NSName string
Parent *Referral Parent *Referral
Depth int Depth int
Prob float64 Prob float64
} }
func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral { func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral {
@@ -81,7 +81,7 @@ func (r *Referral) SetAddresses(addrs []net.IP) {
} }
type CircularReferralError struct { type CircularReferralError struct {
Name string Name string
Chain []string Chain []string
} }
@@ -90,7 +90,7 @@ func (e *CircularReferralError) Error() string {
} }
type UnresolvableNameserverError struct { type UnresolvableNameserverError struct {
Name string Name string
Reason string Reason string
} }
+1 -1
View File
@@ -10,7 +10,7 @@ import (
type ResponseType int type ResponseType int
const ( const (
RespReferral ResponseType = iota RespReferral ResponseType = iota
RespAnswer RespAnswer
RespCNAMEFollow RespCNAMEFollow
RespNODATA RespNODATA
+2 -2
View File
@@ -3,7 +3,7 @@ package traverse
const DefaultMaxDepth = 20 const DefaultMaxDepth = 20
type Stack struct { type Stack struct {
items []*Referral items []*Referral
maxDepth int maxDepth int
} }
@@ -12,7 +12,7 @@ func NewStack(maxDepth int) *Stack {
maxDepth = DefaultMaxDepth maxDepth = DefaultMaxDepth
} }
return &Stack{ return &Stack{
items: make([]*Referral, 0), items: make([]*Referral, 0),
maxDepth: maxDepth, maxDepth: maxDepth,
} }
} }
+36 -36
View File
@@ -44,9 +44,9 @@ func TestTraverserSimpleTraversal(t *testing.T) {
}() }()
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return answerResp.Copy(), nil return answerResp.Copy(), nil
@@ -94,9 +94,9 @@ func TestTraverserReferralTraversal(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
q := msg.Question[0] q := msg.Question[0]
@@ -124,9 +124,9 @@ func TestTraverserReferralTraversal(t *testing.T) {
func TestTraverserMaxDepth(t *testing.T) { func TestTraverserMaxDepth(t *testing.T) {
callCount := 0 callCount := 0
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 2, MaxDepth: 2,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
callCount++ callCount++
@@ -170,9 +170,9 @@ func TestTraverserMaxDepth(t *testing.T) {
func TestTraverserContextCancellation(t *testing.T) { func TestTraverserContextCancellation(t *testing.T) {
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
m := new(dns.Msg) m := new(dns.Msg)
@@ -203,9 +203,9 @@ func TestTraverserNXDOMAIN(t *testing.T) {
nxdResp.Rcode = dns.RcodeNameError nxdResp.Rcode = dns.RcodeNameError
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return nxdResp.Copy(), nil return nxdResp.Copy(), nil
@@ -229,9 +229,9 @@ func TestTraverserSERVFAIL(t *testing.T) {
sfResp.Rcode = dns.RcodeServerFailure sfResp.Rcode = dns.RcodeServerFailure
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return sfResp.Copy(), nil return sfResp.Copy(), nil
@@ -268,9 +268,9 @@ func TestTraverserCNAMEFollow(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
q := msg.Question[0] q := msg.Question[0]
@@ -328,9 +328,9 @@ func TestTraverserProbabilityCalculation(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("1.2.3.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
q := msg.Question[0] q := msg.Question[0]
@@ -364,9 +364,9 @@ func TestTraverserNODATA(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return nodataResp.Copy(), nil return nodataResp.Copy(), nil
@@ -394,9 +394,9 @@ func TestTraverserMultipleRoots(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return answerResp.Copy(), nil return answerResp.Copy(), nil
@@ -414,9 +414,9 @@ func TestTraverserMultipleRoots(t *testing.T) {
func TestTraverserNilExchangeResponse(t *testing.T) { func TestTraverserNilExchangeResponse(t *testing.T) {
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return nil, nil return nil, nil
@@ -451,9 +451,9 @@ func TestTraverserCacheChaining(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, 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) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
q := msg.Question[0] q := msg.Question[0]