diff --git a/cmd/exploredns/main.go b/cmd/exploredns/main.go index 33528a7..01a59e7 100644 --- a/cmd/exploredns/main.go +++ b/cmd/exploredns/main.go @@ -1,209 +1,219 @@ package main import ( - "context" - "flag" - "fmt" - "os" - "time" +"context" +"flag" +"fmt" +"os" +"time" - "github.com/hits/ExploreDNS/internal/config" - "github.com/hits/ExploreDNS/internal/dns" - "github.com/hits/ExploreDNS/internal/traverse" +"github.com/hits/ExploreDNS/internal/config" +"github.com/hits/ExploreDNS/internal/dns" +"github.com/hits/ExploreDNS/internal/output" +"github.com/hits/ExploreDNS/internal/traverse" ) 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)") - rootServer := flag.String("root-server", cfg.RootServer, "Override root server") - allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Use all 13 root servers") - rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses") - followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals") - udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)") - allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation") - alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP") - maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)") - retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)") - fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache") +queryType := flag.String("type", cfg.QueryType, "Record type (A, AAAA, NS, CNAME, MX, TXT, SOA, PTR, ANY)") +rootServer := flag.String("root-server", cfg.RootServer, "Override root server") +allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Use all 13 root servers") +rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses") +followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals") +udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)") +allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation") +alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP") +maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)") +retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)") +fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache") +jsonOutput := flag.Bool("json", false, "Output results as JSON") - // Verbose: long and short form share the same variable. - var verboseVal bool - flag.BoolVar(&verboseVal, "verbose", cfg.Verbose, "Verbose output") - flag.BoolVar(&verboseVal, "v", cfg.Verbose, "Verbose output (shorthand)") +// Verbose: long and short form share the same variable. +var verboseVal bool +flag.BoolVar(&verboseVal, "verbose", cfg.Verbose, "Verbose output") +flag.BoolVar(&verboseVal, "v", cfg.Verbose, "Verbose output (shorthand)") - // Debug: -d sets level 1, -dd sets level 2 (library debug). - var dFlag, ddFlag bool - flag.BoolVar(&dFlag, "d", false, "Debug mode (stackable: -dd for library debug)") - flag.BoolVar(&dFlag, "debug", false, "Debug mode") - flag.BoolVar(&ddFlag, "dd", false, "Library debug mode (equivalent to -d -d)") +// Debug: -d sets level 1, -dd sets level 2 (library debug). +var dFlag, ddFlag bool +flag.BoolVar(&dFlag, "d", false, "Debug mode (stackable: -dd for library debug)") +flag.BoolVar(&dFlag, "debug", false, "Debug mode") +flag.BoolVar(&ddFlag, "dd", false, "Library debug mode (equivalent to -d -d)") - // Quiet: long and short form share the same variable. - var quietVal bool - flag.BoolVar(&quietVal, "quiet", cfg.Quiet, "Suppress supplementary info") - flag.BoolVar(&quietVal, "q", cfg.Quiet, "Suppress supplementary info (shorthand)") +// Quiet: long and short form share the same variable. +var quietVal bool +flag.BoolVar(&quietVal, "quiet", cfg.Quiet, "Suppress supplementary info") +flag.BoolVar(&quietVal, "q", cfg.Quiet, "Suppress supplementary info (shorthand)") - // TODO: ShowProgress is parsed but not yet wired to the traversal/output layer. - showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress") - noShowProgress := flag.Bool("no-show-progress", false, "Hide traversal progress") - // TODO: ShowResolves is parsed but not yet wired to the traversal/output layer. - showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details") - noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details") - // TODO: ShowServers is parsed but not yet wired to the traversal/output layer. - showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried") - noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried") - // TODO: ShowVersions is parsed but not yet wired to the traversal/output layer. - showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions") - noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions") - // TODO: ShowAllStats is parsed but not yet wired to the traversal/output layer. - showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics") - noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics") - // TODO: ShowResults is partially wired; full structured output is pending. - showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results") - noShowResults := flag.Bool("no-show-results", false, "Hide query results") - // TODO: ShowSummaryResults is parsed but not yet wired to the traversal/output layer. - showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results") - noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results") +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") +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") +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.RootServer = *rootServer - cfg.AllRootServers = *allRootServers - cfg.RootAAAA = *rootAAAA - cfg.FollowAAAA = *followAAAA - cfg.UDPSize = *udpSize - cfg.AllowTCP = *allowTCP - cfg.AlwaysTCP = *alwaysTCP - cfg.MaxDepth = *maxDepth - cfg.Retries = *retries - cfg.Fast = *fast - cfg.Verbose = verboseVal - cfg.Debug = config.ParseDebugLevel(dFlag, ddFlag) - cfg.Quiet = quietVal +cfg.QueryType = *queryType +cfg.RootServer = *rootServer +cfg.AllRootServers = *allRootServers +cfg.RootAAAA = *rootAAAA +cfg.FollowAAAA = *followAAAA +cfg.UDPSize = *udpSize +cfg.AllowTCP = *allowTCP +cfg.AlwaysTCP = *alwaysTCP +cfg.MaxDepth = *maxDepth +cfg.Retries = *retries +cfg.Fast = *fast +cfg.Verbose = verboseVal +cfg.Debug = config.ParseDebugLevel(dFlag, ddFlag) +cfg.Quiet = quietVal - if *noShowProgress { - cfg.ShowProgress = false - } else if *showProgress { - cfg.ShowProgress = true - } - if *noShowResolves { - cfg.ShowResolves = false - } else if *showResolves { - cfg.ShowResolves = true - } - if *noShowServers { - cfg.ShowServers = false - } else if *showServers { - cfg.ShowServers = true - } - if *noShowVersions { - cfg.ShowVersions = false - } else if *showVersions { - cfg.ShowVersions = true - } - if *noShowAllStats { - cfg.ShowAllStats = false - } else if *showAllStats { - cfg.ShowAllStats = true - } - if *noShowResults { - cfg.ShowResults = false - } else if *showResults { - cfg.ShowResults = true - } - if *noShowSummaryResults { - cfg.ShowSummaryResults = false - } else if *showSummaryResults { - cfg.ShowSummaryResults = true - } +if *noShowProgress { +cfg.ShowProgress = false +} else if *showProgress { +cfg.ShowProgress = true +} +if *noShowResolves { +cfg.ShowResolves = false +} else if *showResolves { +cfg.ShowResolves = true +} +if *noShowServers { +cfg.ShowServers = false +} else if *showServers { +cfg.ShowServers = true +} +if *noShowVersions { +cfg.ShowVersions = false +} else if *showVersions { +cfg.ShowVersions = true +} +if *noShowAllStats { +cfg.ShowAllStats = false +} else if *showAllStats { +cfg.ShowAllStats = true +} +if *noShowResults { +cfg.ShowResults = false +} else if *showResults { +cfg.ShowResults = true +} +if *noShowSummaryResults { +cfg.ShowSummaryResults = false +} else if *showSummaryResults { +cfg.ShowSummaryResults = true +} - if err := cfg.Validate(); err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) - } +if err := cfg.Validate(); err != nil { +fmt.Fprintf(os.Stderr, "Error: %v\n", err) +os.Exit(1) +} - domain, err := cfg.GetDomain(args) - if err != nil { - fmt.Fprintf(os.Stderr, "Error: %v\n", err) - os.Exit(1) - } +domain, err := cfg.GetDomain(args) +if err != nil { +fmt.Fprintf(os.Stderr, "Error: %v\n", err) +os.Exit(1) +} - rootIP, err := cfg.ParseRootServer() - if err != nil { - fmt.Fprintf(os.Stderr, "Error: invalid root server: %v\n", err) - os.Exit(1) - } +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) - } +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() - } +var rootServerAddr string +if rootIP != nil { +rootServerAddr = rootIP.String() +} - queryConfig := &dns.QueryConfig{ - UDPSize: cfg.UDPSize, - Timeout: 5 * time.Second, - Retries: cfg.Retries, - UseTCP: cfg.AlwaysTCP, - AllowTCP: cfg.AllowTCP, - } +queryConfig := &dns.QueryConfig{ +UDPSize: cfg.UDPSize, +Timeout: 5 * time.Second, +Retries: cfg.Retries, +UseTCP: cfg.AlwaysTCP, +AllowTCP: cfg.AllowTCP, +} - rootConfig := &dns.RootDiscoveryConfig{ - IncludeAAAA: cfg.RootAAAA, - Server: rootServerAddr, - AllRoots: cfg.AllRootServers, - } +rootConfig := &dns.RootDiscoveryConfig{ +IncludeAAAA: cfg.RootAAAA, +Server: rootServerAddr, +AllRoots: cfg.AllRootServers, +} - traverserConfig := &traverse.TraverserConfig{ - MaxDepth: cfg.MaxDepth, - QueryType: queryTypeValue, - RootConfig: rootConfig, - QueryConfig: queryConfig, - } +traverserConfig := &traverse.TraverserConfig{ +MaxDepth: cfg.MaxDepth, +QueryType: queryTypeValue, +RootConfig: rootConfig, +QueryConfig: queryConfig, +} - traverser := traverse.NewTraverser(traverserConfig) +traverser := traverse.NewTraverser(traverserConfig) - if cfg.Debug > 0 { - fmt.Fprintf(os.Stderr, "Debug: Config loaded\n") - fmt.Fprintf(os.Stderr, " Domain: %s\n", domain) - fmt.Fprintf(os.Stderr, " Query Type: %s\n", cfg.QueryType) - fmt.Fprintf(os.Stderr, " Max Depth: %d\n", cfg.MaxDepth) - fmt.Fprintf(os.Stderr, " UDP Size: %d\n", cfg.UDPSize) - fmt.Fprintf(os.Stderr, " Retries: %d\n", cfg.Retries) - fmt.Fprintf(os.Stderr, " Allow TCP: %v\n", cfg.AllowTCP) - fmt.Fprintf(os.Stderr, " Always TCP: %v\n", cfg.AlwaysTCP) - fmt.Fprintf(os.Stderr, " Fast Mode: %v\n", cfg.Fast) - } +if cfg.Debug > 0 { +fmt.Fprintf(os.Stderr, "Debug: Config loaded\n") +fmt.Fprintf(os.Stderr, " Domain: %s\n", domain) +fmt.Fprintf(os.Stderr, " Query Type: %s\n", cfg.QueryType) +fmt.Fprintf(os.Stderr, " Max Depth: %d\n", cfg.MaxDepth) +fmt.Fprintf(os.Stderr, " UDP Size: %d\n", cfg.UDPSize) +fmt.Fprintf(os.Stderr, " Retries: %d\n", cfg.Retries) +fmt.Fprintf(os.Stderr, " Allow TCP: %v\n", cfg.AllowTCP) +fmt.Fprintf(os.Stderr, " Always TCP: %v\n", cfg.AlwaysTCP) +fmt.Fprintf(os.Stderr, " Fast Mode: %v\n", cfg.Fast) +} - if !cfg.Quiet { - fmt.Printf("ExploreDNS - exploring: %s (type: %s)\n", domain, cfg.QueryType) - } +if !cfg.Quiet && !*jsonOutput { +fmt.Printf("ExploreDNS - exploring: %s (type: %s)\n", domain, cfg.QueryType) +} - ctx := context.Background() - results, err := traverser.Traverse(ctx, domain) - if err != nil { - fmt.Fprintf(os.Stderr, "Error: traversal failed: %v\n", err) - os.Exit(1) - } +outFmt := output.FormatText +if *jsonOutput { +outFmt = output.FormatJSON +} - if cfg.ShowResults && len(results) > 0 { - fmt.Printf("\nResults:\n") - for i, result := range results { - fmt.Printf(" [%d] %s -> %s\n", i+1, result.Referral.Name, result.Response.Type) - } - } +outCfg := &output.Config{ +Format: outFmt, +Domain: domain, +QueryType: cfg.QueryType, +ShowProgress: cfg.ShowProgress, +ShowResolves: cfg.ShowResolves, +ShowServers: cfg.ShowServers, +ShowVersions: cfg.ShowVersions, +ShowAllStats: cfg.ShowAllStats, +ShowResults: cfg.ShowResults, +ShowSummaryResults: cfg.ShowSummaryResults, +Verbose: cfg.Verbose, +Quiet: cfg.Quiet, +Color: os.Getenv("NO_COLOR") == "", +} - if cfg.Debug > 0 { - fmt.Fprintf(os.Stderr, "Debug: Traversal completed with %d results\n", len(results)) - } +ctx := context.Background() +formatter := output.NewFormatter(outCfg, os.Stdout) +_, err = output.RunTraversal(ctx, traverser, outCfg, formatter, domain) +if err != nil { +fmt.Fprintf(os.Stderr, "Error: traversal failed: %v\n", err) +os.Exit(1) +} + +if cfg.Debug > 0 { +fmt.Fprintf(os.Stderr, "Debug: Traversal completed\n") +} } diff --git a/internal/dns/types.go b/internal/dns/types.go index f17064a..f082ecb 100644 --- a/internal/dns/types.go +++ b/internal/dns/types.go @@ -1,6 +1,11 @@ package dns -import "github.com/miekg/dns" +import ( + "fmt" + "strings" + + "github.com/miekg/dns" +) const ( TypeA uint16 = dns.TypeA @@ -35,6 +40,32 @@ func QNameType(qtype uint16) string { 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 { return 2048 } diff --git a/internal/output/formatter.go b/internal/output/formatter.go new file mode 100644 index 0000000..62266f2 --- /dev/null +++ b/internal/output/formatter.go @@ -0,0 +1,85 @@ +package output + +import ( + "io" + "os" + + "github.com/hits/ExploreDNS/internal/traverse" +) + +type Format int + +const ( + FormatText Format = iota + FormatJSON +) + +type Config struct { + Format Format + Domain string + QueryType string + ShowProgress bool + ShowResolves bool + ShowServers bool + ShowVersions bool + ShowAllStats bool + ShowResults bool + ShowSummaryResults bool + Verbose bool + Quiet bool + Color bool +} + +func DefaultConfig() *Config { + return &Config{ + Format: FormatText, + ShowProgress: true, + ShowResolves: true, + ShowServers: true, + ShowVersions: true, + ShowAllStats: true, + ShowResults: true, + ShowSummaryResults: true, + Color: os.Getenv("NO_COLOR") == "", + } +} + +type Formatter interface { + WriteProgress(event traverse.TraversalEvent) error + WriteResolve(event traverse.TraversalEvent) error + WriteResult(result traverse.TraversalResult) error + WriteSummary(results []traverse.TraversalResult) error + Flush() error +} + +func NewFormatter(cfg *Config, w io.Writer) Formatter { + if cfg == nil { + cfg = DefaultConfig() + } + if w == nil { + w = os.Stdout + } + if cfg.Format == FormatJSON { + return newJSONFormatter(cfg, w) + } + return newTextFormatter(cfg, w) +} + +func AttachHooks(cfg *Config, formatter Formatter) *traverse.TraverserHooks { + if cfg == nil || formatter == nil { + return nil + } + return &traverse.TraverserHooks{ + OnEvent: func(event traverse.TraversalEvent) { + switch { + case event.IsResolve && cfg.ShowResolves: + _ = formatter.WriteResolve(event) + case !event.IsResolve && cfg.ShowProgress: + _ = formatter.WriteProgress(event) + } + if event.Stage == traverse.EventComplete && cfg.ShowAllStats { + _ = formatter.WriteResult(event.Result) + } + }, + } +} diff --git a/internal/output/formatter_test.go b/internal/output/formatter_test.go new file mode 100644 index 0000000..1f5e774 --- /dev/null +++ b/internal/output/formatter_test.go @@ -0,0 +1,188 @@ +package output + +import ( + "bytes" + "context" + "encoding/json" + "net" + "strings" + "testing" + + "github.com/hits/ExploreDNS/internal/dns" + "github.com/hits/ExploreDNS/internal/traverse" + miekgdns "github.com/miekg/dns" +) + +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) + } +} + +func TestNewFormatterSelectsImplementation(t *testing.T) { + text := NewFormatter(DefaultConfig(), &bytes.Buffer{}) + if _, ok := text.(*textFormatter); !ok { + t.Fatalf("expected text formatter, got %T", text) + } + + jsonCfg := DefaultConfig() + jsonCfg.Format = FormatJSON + jsonFmt := NewFormatter(jsonCfg, &bytes.Buffer{}) + if _, ok := jsonFmt.(*jsonFormatter); !ok { + t.Fatalf("expected json formatter, got %T", jsonFmt) + } +} + +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"), + }, + }, + }, + } + + var buf bytes.Buffer + cfg := DefaultConfig() + 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") + if err != nil { + t.Fatalf("RunTraversal: %v", err) + } + if !strings.Contains(buf.String(), "Summary:") { + t.Fatalf("expected formatted summary output, got %q", buf.String()) + } +} diff --git a/internal/output/json.go b/internal/output/json.go new file mode 100644 index 0000000..e1d8ba3 --- /dev/null +++ b/internal/output/json.go @@ -0,0 +1,191 @@ +package output + +import ( + "encoding/json" + "io" + + "github.com/hits/ExploreDNS/internal/dns" + "github.com/hits/ExploreDNS/internal/traverse" +) + +type jsonFormatter struct { + cfg *Config + w io.Writer + payload jsonDocument +} + +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"` +} + +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"` +} + +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"` +} + +type jsonServer struct { + Name string `json:"name"` + IPs []string `json:"ips"` +} + +type jsonSummary struct { + ByType map[string]float64 `json:"by_type,omitempty"` + Answers []jsonAnswerStat `json:"answers,omitempty"` +} + +type jsonAnswerStat struct { + RData string `json:"rdata"` + Probability float64 `json:"probability"` + Records []string `json:"records,omitempty"` +} + +func newJSONFormatter(cfg *Config, w io.Writer) *jsonFormatter { + return &jsonFormatter{ + cfg: cfg, + w: w, + payload: jsonDocument{ + Domain: cfg.Domain, + QueryType: cfg.QueryType, + }, + } +} + +func (f *jsonFormatter) WriteProgress(event traverse.TraversalEvent) error { + if !f.cfg.ShowProgress { + return nil + } + f.payload.Progress = append(f.payload.Progress, f.eventToJSON(event)) + 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)) + 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)) + 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)) + } + } + + if f.cfg.ShowServers { + servers := collectServers(results) + for name, ips := range servers { + f.payload.Servers = append(f.payload.Servers, jsonServer{ + Name: name, + IPs: ips, + }) + } + } + + if f.cfg.ShowSummaryResults { + stats := ComputeSummary(results) + if stats != nil { + f.payload.Summary = jsonSummary{ + ByType: stats.ByType, + } + for _, answer := range stats.Answers { + f.payload.Summary.Answers = append(f.payload.Summary.Answers, jsonAnswerStat{ + RData: answer.RData, + Probability: answer.Prob, + Records: answer.RRs, + }) + } + } + } + + return nil +} + +func (f *jsonFormatter) Flush() error { + enc := json.NewEncoder(f.w) + enc.SetIndent("", " ") + return enc.Encode(f.payload) +} + +func (f *jsonFormatter) eventToJSON(event traverse.TraversalEvent) jsonProgressEvent { + ref := event.Result.Referral + if ref == nil { + return jsonProgressEvent{} + } + + 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 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 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...) + } + } + 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/output.go b/internal/output/output.go deleted file mode 100644 index ad89311..0000000 --- a/internal/output/output.go +++ /dev/null @@ -1 +0,0 @@ -package output diff --git a/internal/output/runner.go b/internal/output/runner.go new file mode 100644 index 0000000..cc5354d --- /dev/null +++ b/internal/output/runner.go @@ -0,0 +1,36 @@ +package output + +import ( + "context" + "fmt" + + "github.com/hits/ExploreDNS/internal/traverse" +) + +func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Config, formatter Formatter, domain string) ([]traverse.TraversalResult, error) { + if traverser == nil { + return nil, fmt.Errorf("traverser is required") + } + if cfg == nil { + cfg = DefaultConfig() + } + if formatter == nil { + formatter = NewFormatter(cfg, nil) + } + + traverser.SetHooks(AttachHooks(cfg, formatter)) + + results, err := traverser.Traverse(ctx, domain) + if err != nil { + return results, err + } + + if err := formatter.WriteSummary(results); err != nil { + return results, err + } + if err := formatter.Flush(); err != nil { + return results, err + } + + return results, nil +} diff --git a/internal/output/stats.go b/internal/output/stats.go new file mode 100644 index 0000000..50fb188 --- /dev/null +++ b/internal/output/stats.go @@ -0,0 +1,187 @@ +package output + +import ( + "fmt" + "sort" + "strings" + + "github.com/hits/ExploreDNS/internal/dns" + "github.com/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 + }) + + 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() + } +} + +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": + return "found no such record" + case "nxdomain": + return "name does not exist" + case "servfail": + return "resulted in SERVFAIL" + case "error": + return "resulted in an error" + case "referral": + return "resulted in a referral" + default: + return respType + } +} + +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 +} diff --git a/internal/output/text.go b/internal/output/text.go new file mode 100644 index 0000000..ae6183d --- /dev/null +++ b/internal/output/text.go @@ -0,0 +1,278 @@ +package output + +import ( + "fmt" + "io" + "sort" + "strings" + + "github.com/hits/ExploreDNS/internal/dns" + "github.com/hits/ExploreDNS/internal/traverse" +) + +type textFormatter struct { + cfg *Config + w io.Writer +} + +func newTextFormatter(cfg *Config, w io.Writer) *textFormatter { + return &textFormatter{cfg: cfg, w: w} +} + +func (f *textFormatter) WriteProgress(event traverse.TraversalEvent) error { + if event.Stage != traverse.EventStart { + return nil + } + line := f.formatReferralLine(event.Result, false) + if !event.Result.Referral.HasAddresses() { + line += " -- resolving" + } + return f.writeLine(line) +} + +func (f *textFormatter) WriteResolve(event traverse.TraversalEvent) error { + if event.Stage != traverse.EventStart { + 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 + } + prefix := strings.Repeat(" ", result.Referral.Depth+1) + line := prefix + f.formatResultLine(result) + return f.writeLine(line) +} + +func (f *textFormatter) WriteSummary(results []traverse.TraversalResult) error { + if f.cfg.ShowServers { + if err := f.writeServers(results); err != nil { + return err + } + } + + if f.cfg.ShowResults { + if err := f.writeResults(results); err != nil { + return err + } + } + + if f.cfg.ShowSummaryResults { + if err := f.writeSummaryResults(results); err != nil { + return err + } + } + + return nil +} + +func (f *textFormatter) Flush() error { + return nil +} + +func (f *textFormatter) writeServers(results []traverse.TraversalResult) error { + servers := collectServers(results) + if len(servers) == 0 { + return nil + } + + if _, err := fmt.Fprintln(f.w, "The following servers were encountered:"); err != nil { + return err + } + + names := make([]string, 0, len(servers)) + 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) + } + } + + for _, name := range names { + for _, ip := range servers[name] { + version := "" + 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 + } + } + } + _, err := fmt.Fprintln(f.w) + return err +} + +func (f *textFormatter) writeResults(results []traverse.TraversalResult) error { + if _, err := fmt.Fprintln(f.w, "Results:"); err != nil { + return err + } + + terminal := terminalResults(results) + for _, result := range terminal { + prefix := strings.Repeat(" ", result.Referral.Depth+1) + line := prefix + f.formatResultLine(result) + if _, err := fmt.Fprintln(f.w, line); err != nil { + return err + } + } + _, err := fmt.Fprintln(f.w) + return err +} + +func (f *textFormatter) writeSummaryResults(results []traverse.TraversalResult) error { + stats := ComputeSummary(results) + if stats == nil { + return nil + } + + if _, err := fmt.Fprintln(f.w, "Summary:"); err != nil { + return err + } + + prefix := " " + for _, answer := range stats.Answers { + line := fmt.Sprintf("%s%s answered with %s", prefix, formatProbability(answer.Prob), answer.RData) + 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) + } + sort.Strings(types) + + for _, respType := range types { + line := fmt.Sprintf("%s%s %s", prefix, formatProbability(stats.ByType[respType]), summaryTypeLabel(respType)) + if _, err := fmt.Fprintln(f.w, line); err != nil { + return err + } + } + + _, err := fmt.Fprintln(f.w) + return err +} + +func (f *textFormatter) formatReferralLine(result traverse.TraversalResult, isResolve bool) string { + ref := result.Referral + if ref == nil { + return "" + } + + 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) + } + 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, rrs := answerKey(result.Response) + if key == "" { + return fmt.Sprintf("%s resulted in answer", prob) + } + if len(rrs) == 1 { + return f.colorize(fmt.Sprintf("%s answered with %s", prob, rrs[0]), colorGreen) + } + return f.colorize(fmt.Sprintf("%s answered with %s", prob, strings.Join(rrs, " / ")), 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.RespError: + return f.colorize(fmt.Sprintf("%s resulted in an error", prob), colorRed) + default: + return fmt.Sprintf("%s %s", prob, result.Response.Type) + } +} + +func (f *textFormatter) writeLine(line string) error { + if line == "" { + return nil + } + _, err := fmt.Fprintln(f.w, line) + return err +} + +func (f *textFormatter) colorize(text, color string) string { + if !f.cfg.Color || color == "" { + return text + } + return color + text + colorReset +} + +func referralID(ref *traverse.Referral) string { + if ref == nil { + return "" + } + if ref.Depth == 0 { + return "1" + } + 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" + colorYellow = "\033[33m" + colorRed = "\033[31m" +) diff --git a/internal/output/text_test.go b/internal/output/text_test.go new file mode 100644 index 0000000..dc1ebdc --- /dev/null +++ b/internal/output/text_test.go @@ -0,0 +1,65 @@ +package output + +import ( + "bytes" + "strings" + "testing" + + "github.com/hits/ExploreDNS/internal/dns" + "github.com/hits/ExploreDNS/internal/traverse" +) + +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") + } +} diff --git a/internal/traverse/hooks.go b/internal/traverse/hooks.go new file mode 100644 index 0000000..fdf0d88 --- /dev/null +++ b/internal/traverse/hooks.go @@ -0,0 +1,31 @@ +package traverse + +type EventStage int + +const ( + EventStart EventStage = iota + EventComplete +) + +type TraversalEvent struct { + Stage EventStage + Result TraversalResult + IsResolve bool +} + +type EventHandler func(TraversalEvent) + +type TraverserHooks struct { + OnEvent EventHandler +} + +func (h *TraverserHooks) emit(stage EventStage, result TraversalResult, isResolve bool) { + if h == nil || h.OnEvent == nil { + return + } + h.OnEvent(TraversalEvent{ + Stage: stage, + Result: result, + IsResolve: isResolve, + }) +} diff --git a/internal/traverse/hooks_test.go b/internal/traverse/hooks_test.go new file mode 100644 index 0000000..69007e6 --- /dev/null +++ b/internal/traverse/hooks_test.go @@ -0,0 +1,52 @@ +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) + }, + } + + 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 + }) + + _, err := tr.Traverse(context.Background(), "example.com") + if err != nil { + t.Fatalf("Traverse: %v", err) + } + if len(events) < 2 { + t.Fatalf("expected start and complete events, got %d", len(events)) + } + 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) + } +} diff --git a/internal/traverse/traverser.go b/internal/traverse/traverser.go index b262f69..3b2c259 100644 --- a/internal/traverse/traverser.go +++ b/internal/traverse/traverser.go @@ -16,6 +16,7 @@ type TraverserConfig struct { RootConfig *dns.RootDiscoveryConfig QueryConfig *dns.QueryConfig RootAddrs []net.IP + Hooks *TraverserHooks } func DefaultTraverserConfig() *TraverserConfig { @@ -57,6 +58,13 @@ func (t *Traverser) SetExchange(fn dns.ExchangeFunc) { t.exchange = 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) @@ -94,9 +102,19 @@ func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResul cache = rootCache.Child() } + 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, TraversalResult{Referral: ref, Response: resp}) + results = append(results, result) mu.Unlock() if resp.IsTerminal() { @@ -266,8 +284,16 @@ func (t *Traverser) ResolveNS(ctx context.Context, nsName string, cache *InfoCac 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 {