feat: implement output formatting and display (HAN-383) #9
+187
-177
@@ -1,209 +1,219 @@
|
|||||||
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/traverse"
|
"github.com/hits/ExploreDNS/internal/output"
|
||||||
|
"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")
|
||||||
|
|
||||||
// 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)")
|
||||||
|
|
||||||
// TODO: ShowProgress is parsed but not yet wired to the traversal/output layer.
|
showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress")
|
||||||
showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress")
|
noShowProgress := flag.Bool("no-show-progress", false, "Hide traversal progress")
|
||||||
noShowProgress := flag.Bool("no-show-progress", false, "Hide traversal progress")
|
showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details")
|
||||||
// TODO: ShowResolves is parsed but not yet wired to the traversal/output layer.
|
noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details")
|
||||||
showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details")
|
showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried")
|
||||||
noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details")
|
noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried")
|
||||||
// TODO: ShowServers is parsed but not yet wired to the traversal/output layer.
|
showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions")
|
||||||
showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried")
|
noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions")
|
||||||
noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried")
|
showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics")
|
||||||
// TODO: ShowVersions is parsed but not yet wired to the traversal/output layer.
|
noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics")
|
||||||
showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions")
|
showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results")
|
||||||
noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions")
|
noShowResults := flag.Bool("no-show-results", false, "Hide query results")
|
||||||
// TODO: ShowAllStats is parsed but not yet wired to the traversal/output layer.
|
showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results")
|
||||||
showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics")
|
noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results")
|
||||||
noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics")
|
|
||||||
// TODO: ShowResults is partially wired; full structured output is pending.
|
|
||||||
showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results")
|
|
||||||
noShowResults := flag.Bool("no-show-results", false, "Hide query results")
|
|
||||||
// TODO: ShowSummaryResults is parsed but not yet wired to the traversal/output layer.
|
|
||||||
showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results")
|
|
||||||
noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results")
|
|
||||||
|
|
||||||
flag.Usage = config.PrintUsage
|
flag.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 {
|
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)
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx := context.Background()
|
outFmt := output.FormatText
|
||||||
results, err := traverser.Traverse(ctx, domain)
|
if *jsonOutput {
|
||||||
if err != nil {
|
outFmt = output.FormatJSON
|
||||||
fmt.Fprintf(os.Stderr, "Error: traversal failed: %v\n", err)
|
}
|
||||||
os.Exit(1)
|
|
||||||
}
|
|
||||||
|
|
||||||
if cfg.ShowResults && len(results) > 0 {
|
outCfg := &output.Config{
|
||||||
fmt.Printf("\nResults:\n")
|
Format: outFmt,
|
||||||
for i, result := range results {
|
Domain: domain,
|
||||||
fmt.Printf(" [%d] %s -> %s\n", i+1, result.Referral.Name, result.Response.Type)
|
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 {
|
ctx := context.Background()
|
||||||
fmt.Fprintf(os.Stderr, "Debug: Traversal completed with %d results\n", len(results))
|
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")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
+32
-1
@@ -1,6 +1,11 @@
|
|||||||
package dns
|
package dns
|
||||||
|
|
||||||
import "github.com/miekg/dns"
|
import (
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/miekg/dns"
|
||||||
|
)
|
||||||
|
|
||||||
const (
|
const (
|
||||||
TypeA uint16 = dns.TypeA
|
TypeA uint16 = dns.TypeA
|
||||||
@@ -35,6 +40,32 @@ 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
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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)
|
||||||
|
}
|
||||||
|
},
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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())
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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"
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1 +0,0 @@
|
|||||||
package output
|
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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
|
||||||
|
}
|
||||||
@@ -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"
|
||||||
|
)
|
||||||
@@ -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")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -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,
|
||||||
|
})
|
||||||
|
}
|
||||||
@@ -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)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -16,6 +16,7 @@ type TraverserConfig struct {
|
|||||||
RootConfig *dns.RootDiscoveryConfig
|
RootConfig *dns.RootDiscoveryConfig
|
||||||
QueryConfig *dns.QueryConfig
|
QueryConfig *dns.QueryConfig
|
||||||
RootAddrs []net.IP
|
RootAddrs []net.IP
|
||||||
|
Hooks *TraverserHooks
|
||||||
}
|
}
|
||||||
|
|
||||||
func DefaultTraverserConfig() *TraverserConfig {
|
func DefaultTraverserConfig() *TraverserConfig {
|
||||||
@@ -57,6 +58,13 @@ func (t *Traverser) SetExchange(fn dns.ExchangeFunc) {
|
|||||||
t.exchange = fn
|
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) {
|
func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResult, error) {
|
||||||
name = miekgdns.Fqdn(name)
|
name = miekgdns.Fqdn(name)
|
||||||
|
|
||||||
@@ -94,9 +102,19 @@ func (t *Traverser) Traverse(ctx context.Context, name string) ([]TraversalResul
|
|||||||
cache = rootCache.Child()
|
cache = rootCache.Child()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if t.config.Hooks != nil {
|
||||||
|
t.config.Hooks.emit(EventStart, TraversalResult{Referral: ref}, false)
|
||||||
|
}
|
||||||
|
|
||||||
resp := t.processReferral(ctx, ref, cache)
|
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()
|
mu.Lock()
|
||||||
results = append(results, TraversalResult{Referral: ref, Response: resp})
|
results = append(results, result)
|
||||||
mu.Unlock()
|
mu.Unlock()
|
||||||
|
|
||||||
if resp.IsTerminal() {
|
if resp.IsTerminal() {
|
||||||
@@ -266,8 +284,16 @@ func (t *Traverser) ResolveNS(ctx context.Context, nsName string, cache *InfoCac
|
|||||||
cacheForStep = traversalCache.Child()
|
cacheForStep = traversalCache.Child()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if t.config.Hooks != nil {
|
||||||
|
t.config.Hooks.emit(EventStart, TraversalResult{Referral: current}, true)
|
||||||
|
}
|
||||||
|
|
||||||
resp := t.processReferral(ctx, current, cacheForStep)
|
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 {
|
if resp.Type == RespAnswer && len(resp.Decoded.Answers) > 0 {
|
||||||
for _, rr := range resp.Decoded.Answers {
|
for _, rr := range resp.Decoded.Answers {
|
||||||
if a, ok := rr.(*miekgdns.A); ok {
|
if a, ok := rr.(*miekgdns.A); ok {
|
||||||
|
|||||||
Reference in New Issue
Block a user