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

Merged
multica-agent merged 2 commits from agent/go-expert-developer/f3b1b7cf into main 2026-06-07 16:40:02 +00:00
21 changed files with 1239 additions and 86 deletions
+28 -17
View File
@@ -9,6 +9,7 @@ import (
"github.com/hits/ExploreDNS/internal/config" "github.com/hits/ExploreDNS/internal/config"
"github.com/hits/ExploreDNS/internal/dns" "github.com/hits/ExploreDNS/internal/dns"
"github.com/hits/ExploreDNS/internal/output"
"github.com/hits/ExploreDNS/internal/traverse" "github.com/hits/ExploreDNS/internal/traverse"
) )
@@ -26,6 +27,7 @@ func main() {
maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)") maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)")
retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)") retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)")
fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache") fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache")
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
@@ -43,25 +45,18 @@ func main() {
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")
// TODO: ShowResolves is parsed but not yet wired to the traversal/output layer.
showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details") showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details")
noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details") noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details")
// TODO: ShowServers is parsed but not yet wired to the traversal/output layer.
showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried") showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried")
noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried") noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried")
// TODO: ShowVersions is parsed but not yet wired to the traversal/output layer.
showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions") showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions")
noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions") noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions")
// TODO: ShowAllStats is parsed but not yet wired to the traversal/output layer.
showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics") showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics")
noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics") noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics")
// TODO: ShowResults is partially wired; full structured output is pending.
showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results") showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results")
noShowResults := flag.Bool("no-show-results", false, "Hide query results") noShowResults := flag.Bool("no-show-results", false, "Hide query results")
// TODO: ShowSummaryResults is parsed but not yet wired to the traversal/output layer.
showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results") showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results")
noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results") noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results")
@@ -185,25 +180,41 @@ func main() {
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)
} }
outFmt := output.FormatText
if *jsonOutput {
outFmt = output.FormatJSON
}
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") == "",
Debug: cfg.Debug,
}
ctx := context.Background() ctx := context.Background()
results, err := traverser.Traverse(ctx, domain) formatter := output.NewFormatter(outCfg, os.Stdout)
_, err = output.RunTraversal(ctx, traverser, outCfg, formatter, domain)
if err != nil { if err != nil {
fmt.Fprintf(os.Stderr, "Error: traversal failed: %v\n", err) fmt.Fprintf(os.Stderr, "Error: traversal failed: %v\n", err)
os.Exit(1) os.Exit(1)
} }
if cfg.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)
}
}
if cfg.Debug > 0 { if cfg.Debug > 0 {
fmt.Fprintf(os.Stderr, "Debug: Traversal completed with %d results\n", len(results)) fmt.Fprintf(os.Stderr, "Debug: Traversal completed\n")
} }
} }
+7 -7
View File
@@ -36,13 +36,13 @@ type Config struct {
Debug int Debug int
Quiet bool Quiet bool
ShowProgress bool ShowProgress bool
ShowResolves bool ShowResolves bool
ShowServers bool ShowServers bool
ShowVersions bool ShowVersions bool
ShowAllStats bool ShowAllStats bool
ShowResults bool ShowResults bool
ShowSummaryResults bool ShowSummaryResults bool
} }
func ParseQueryType(s string) (uint16, error) { func ParseQueryType(s string) (uint16, error) {
+5 -5
View File
@@ -314,8 +314,8 @@ func TestCNAMEChain(t *testing.T) {
func TestResponseClassificationString(t *testing.T) { func TestResponseClassificationString(t *testing.T) {
tests := []struct { tests := []struct {
rc ResponseClassification rc ResponseClassification
want string want string
}{ }{
{ResponseAnswer, "answer"}, {ResponseAnswer, "answer"},
{ResponseReferral, "referral"}, {ResponseReferral, "referral"},
@@ -344,10 +344,10 @@ func TestFormatRecord(t *testing.T) {
t.Run("A record", func(t *testing.T) { t.Run("A record", func(t *testing.T) {
rr := &dns.A{ rr := &dns.A{
Hdr: dns.RR_Header{ Hdr: dns.RR_Header{
Name: "example.com.", Name: "example.com.",
Rrtype: dns.TypeA, Rrtype: dns.TypeA,
Class: dns.ClassINET, Class: dns.ClassINET,
Ttl: 300, Ttl: 300,
}, },
A: MustParseIP("93.184.216.34"), A: MustParseIP("93.184.216.34"),
} }
-1
View File
@@ -328,4 +328,3 @@ func TestQueryNoTCPFallbackWhenDisabled(t *testing.T) {
t.Error("expected truncated response to be returned as-is") t.Error("expected truncated response to be returned as-is")
} }
} }
+3 -1
View File
@@ -1,6 +1,8 @@
package dns package dns
import "github.com/miekg/dns" import (
"github.com/miekg/dns"
)
const ( const (
TypeA uint16 = dns.TypeA TypeA uint16 = dns.TypeA
+92
View File
@@ -0,0 +1,92 @@
package output
import (
"fmt"
"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
Debug int
}
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
}
logErr := func(context string, err error) {
if err != nil && cfg.Debug > 0 {
fmt.Fprintf(os.Stderr, "Debug: formatter %s: %v\n", context, err)
}
}
return &traverse.TraverserHooks{
OnEvent: func(event traverse.TraversalEvent) {
switch {
case event.IsResolve && cfg.ShowResolves:
logErr("WriteResolve", formatter.WriteResolve(event))
case !event.IsResolve && cfg.ShowProgress:
logErr("WriteProgress", formatter.WriteProgress(event))
}
if event.Stage == traverse.EventComplete && cfg.ShowAllStats {
logErr("WriteResult", formatter.WriteResult(event.Result))
}
},
}
}
+188
View File
@@ -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())
}
}
+191
View File
@@ -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
View File
@@ -1 +0,0 @@
package output
+36
View File
@@ -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
}
+190
View File
@@ -0,0 +1,190 @@
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
})
if len(stats.Answers) == 0 && len(stats.ByType) == 0 {
return nil
}
return stats
}
func answerKey(resp *traverse.Response) (string, []string) {
if resp == nil || resp.Decoded == nil {
return "", nil
}
var rdatas []string
var formatted []string
for _, rr := range resp.Decoded.Answers {
if _, ok := rr.(*miekgdns.CNAME); ok {
continue
}
rdata := rrDataString(rr)
if rdata == "" {
continue
}
rdatas = append(rdatas, rdata)
formatted = append(formatted, dns.FormatRecord(rr))
}
if len(rdatas) == 0 {
return "", nil
}
sort.Strings(rdatas)
return strings.Join(rdatas, " / "), formatted
}
func rrDataString(rr miekgdns.RR) string {
switch v := rr.(type) {
case *miekgdns.A:
return v.A.String()
case *miekgdns.AAAA:
return v.AAAA.String()
case *miekgdns.CNAME:
return v.Target
case *miekgdns.NS:
return v.Ns
case *miekgdns.MX:
return fmt.Sprintf("%d %s", v.Preference, v.Mx)
case *miekgdns.TXT:
return strings.Join(v.Txt, " ")
default:
return rr.String()
}
}
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
}
+271
View File
@@ -0,0 +1,271 @@
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] {
if _, err := fmt.Fprintf(f.w, "%*s: %-15s\n", width, name, ip); 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 ""
}
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"
)
+65
View File
@@ -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")
}
}
+4 -4
View File
@@ -9,10 +9,10 @@ import (
) )
type InfoCache struct { type InfoCache struct {
parent *InfoCache parent *InfoCache
mu sync.RWMutex mu sync.RWMutex
ns map[string][]string ns map[string][]string
glue map[string][]net.IP glue map[string][]net.IP
} }
func NewInfoCache(parent *InfoCache) *InfoCache { func NewInfoCache(parent *InfoCache) *InfoCache {
+31
View File
@@ -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,
})
}
+52
View File
@@ -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)
}
}
+9 -9
View File
@@ -32,18 +32,18 @@ func (s ResolutionState) String() string {
} }
type Referral struct { type Referral struct {
Name string Name string
Qtype uint16 Qtype uint16
Qclass uint16 Qclass uint16
Bailiwick string Bailiwick string
Addresses []net.IP Addresses []net.IP
State ResolutionState State ResolutionState
NSName string NSName string
Parent *Referral Parent *Referral
Depth int Depth int
Prob float64 Prob float64
} }
func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral { func NewReferral(name string, qtype uint16, bailiwick string, depth int, prob float64, parent *Referral) *Referral {
@@ -81,7 +81,7 @@ func (r *Referral) SetAddresses(addrs []net.IP) {
} }
type CircularReferralError struct { type CircularReferralError struct {
Name string Name string
Chain []string Chain []string
} }
@@ -90,7 +90,7 @@ func (e *CircularReferralError) Error() string {
} }
type UnresolvableNameserverError struct { type UnresolvableNameserverError struct {
Name string Name string
Reason string Reason string
} }
+1 -1
View File
@@ -10,7 +10,7 @@ import (
type ResponseType int type ResponseType int
const ( const (
RespReferral ResponseType = iota RespReferral ResponseType = iota
RespAnswer RespAnswer
RespCNAMEFollow RespCNAMEFollow
RespNODATA RespNODATA
+2 -2
View File
@@ -3,7 +3,7 @@ package traverse
const DefaultMaxDepth = 20 const DefaultMaxDepth = 20
type Stack struct { type Stack struct {
items []*Referral items []*Referral
maxDepth int maxDepth int
} }
@@ -12,7 +12,7 @@ func NewStack(maxDepth int) *Stack {
maxDepth = DefaultMaxDepth maxDepth = DefaultMaxDepth
} }
return &Stack{ return &Stack{
items: make([]*Referral, 0), items: make([]*Referral, 0),
maxDepth: maxDepth, maxDepth: maxDepth,
} }
} }
+27 -1
View File
@@ -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 {
+36 -36
View File
@@ -44,9 +44,9 @@ func TestTraverserSimpleTraversal(t *testing.T) {
}() }()
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return answerResp.Copy(), nil return answerResp.Copy(), nil
@@ -94,9 +94,9 @@ func TestTraverserReferralTraversal(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
q := msg.Question[0] q := msg.Question[0]
@@ -124,9 +124,9 @@ func TestTraverserReferralTraversal(t *testing.T) {
func TestTraverserMaxDepth(t *testing.T) { func TestTraverserMaxDepth(t *testing.T) {
callCount := 0 callCount := 0
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 2, MaxDepth: 2,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
callCount++ callCount++
@@ -170,9 +170,9 @@ func TestTraverserMaxDepth(t *testing.T) {
func TestTraverserContextCancellation(t *testing.T) { func TestTraverserContextCancellation(t *testing.T) {
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
m := new(dns.Msg) m := new(dns.Msg)
@@ -203,9 +203,9 @@ func TestTraverserNXDOMAIN(t *testing.T) {
nxdResp.Rcode = dns.RcodeNameError nxdResp.Rcode = dns.RcodeNameError
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return nxdResp.Copy(), nil return nxdResp.Copy(), nil
@@ -229,9 +229,9 @@ func TestTraverserSERVFAIL(t *testing.T) {
sfResp.Rcode = dns.RcodeServerFailure sfResp.Rcode = dns.RcodeServerFailure
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return sfResp.Copy(), nil return sfResp.Copy(), nil
@@ -268,9 +268,9 @@ func TestTraverserCNAMEFollow(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
q := msg.Question[0] q := msg.Question[0]
@@ -328,9 +328,9 @@ func TestTraverserProbabilityCalculation(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("1.2.3.4")}, RootAddrs: []net.IP{net.ParseIP("1.2.3.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
q := msg.Question[0] q := msg.Question[0]
@@ -364,9 +364,9 @@ func TestTraverserNODATA(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return nodataResp.Copy(), nil return nodataResp.Copy(), nil
@@ -394,9 +394,9 @@ func TestTraverserMultipleRoots(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4"), net.ParseIP("199.9.14.201")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return answerResp.Copy(), nil return answerResp.Copy(), nil
@@ -414,9 +414,9 @@ func TestTraverserMultipleRoots(t *testing.T) {
func TestTraverserNilExchangeResponse(t *testing.T) { func TestTraverserNilExchangeResponse(t *testing.T) {
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
return nil, nil return nil, nil
@@ -451,9 +451,9 @@ func TestTraverserCacheChaining(t *testing.T) {
}) })
tr := NewTraverser(&TraverserConfig{ tr := NewTraverser(&TraverserConfig{
MaxDepth: 5, MaxDepth: 5,
QueryType: dnsTypeA, QueryType: dnsTypeA,
RootAddrs: []net.IP{net.ParseIP("198.41.0.4")}, RootAddrs: []net.IP{net.ParseIP("198.41.0.4")},
}) })
tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) { tr.SetExchange(func(ctx context.Context, server string, msg *dns.Msg, useTCP bool) (*dns.Msg, error) {
q := msg.Question[0] q := msg.Question[0]