diff --git a/go.mod b/go.mod index 29ccbfa..d1e67bf 100644 --- a/go.mod +++ b/go.mod @@ -2,12 +2,14 @@ module github.com/hits/ExploreDNS go 1.25.6 -require github.com/miekg/dns v1.1.72 +require ( + github.com/miekg/dns v1.1.72 + golang.org/x/sync v0.19.0 +) require ( golang.org/x/mod v0.31.0 // indirect golang.org/x/net v0.48.0 // indirect - golang.org/x/sync v0.19.0 // indirect golang.org/x/sys v0.39.0 // indirect golang.org/x/tools v0.40.0 // indirect ) diff --git a/internal/fingerprint/fingerprint.go b/internal/fingerprint/fingerprint.go index ad1f890..d467822 100644 --- a/internal/fingerprint/fingerprint.go +++ b/internal/fingerprint/fingerprint.go @@ -7,6 +7,7 @@ import ( "time" miekgdns "github.com/miekg/dns" + "golang.org/x/sync/singleflight" ) const defaultTimeout = 2 * time.Second @@ -17,6 +18,7 @@ type Fingerprinter struct { mu sync.Mutex cache map[string]string timeout time.Duration + group singleflight.Group exchange func(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) } @@ -36,6 +38,8 @@ func NewWithTimeout(timeout time.Duration) *Fingerprinter { // Query returns the version string for ip, or "" if the server doesn't // respond or doesn't support the version.bind CHAOS query. // Results are cached: subsequent calls for the same IP return immediately. +// Concurrent calls for the same IP are coalesced via singleflight so only +// one network probe is issued per IP at a time. func (f *Fingerprinter) Query(ctx context.Context, ip net.IP) string { key := ip.String() @@ -46,13 +50,15 @@ func (f *Fingerprinter) Query(ctx context.Context, ip net.IP) string { } f.mu.Unlock() - version := f.probe(ctx, ip) + v, _, _ := f.group.Do(key, func() (interface{}, error) { + version := f.probe(ctx, ip) + f.mu.Lock() + f.cache[key] = version + f.mu.Unlock() + return version, nil + }) - f.mu.Lock() - f.cache[key] = version - f.mu.Unlock() - - return version + return v.(string) } // FingerprintAll queries all ips concurrently and returns a map of diff --git a/internal/output/formatter.go b/internal/output/formatter.go index 5cd44b7..85d8532 100644 --- a/internal/output/formatter.go +++ b/internal/output/formatter.go @@ -31,9 +31,6 @@ type Config struct { Color bool Debug int - // Fingerprints maps server IP strings to their version.bind version strings. - // Populated by RunTraversal when ShowVersions and ShowServers are both true. - Fingerprints map[string]string } func DefaultConfig() *Config { @@ -56,6 +53,9 @@ type Formatter interface { WriteResult(result traverse.TraversalResult) error WriteSummary(results []traverse.TraversalResult) error Flush() error + // SetFingerprints supplies server-version data to the formatter. + // Call before WriteSummary when ShowVersions is true. + SetFingerprints(fps map[string]string) } func NewFormatter(cfg *Config, w io.Writer) Formatter { diff --git a/internal/output/formatter_test.go b/internal/output/formatter_test.go index 1f5e774..48348bc 100644 --- a/internal/output/formatter_test.go +++ b/internal/output/formatter_test.go @@ -150,6 +150,80 @@ func TestJSONFormatterProducesValidOutput(t *testing.T) { } } +func TestTextFormatterWriteSummaryShowsVersions(t *testing.T) { + ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 0, 1.0, nil) + serverIP := net.ParseIP("198.41.0.4") + resp := &traverse.Response{ + Referral: ref, + Server: serverIP, + Type: traverse.RespAnswer, + } + + var buf bytes.Buffer + cfg := DefaultConfig() + cfg.Color = false + cfg.ShowVersions = true + cfg.ShowServers = true + cfg.ShowResults = false + cfg.ShowSummaryResults = false + formatter := NewFormatter(cfg, &buf) + formatter.SetFingerprints(map[string]string{serverIP.String(): "BIND 9.18.1"}) + + if err := formatter.WriteSummary([]traverse.TraversalResult{{Referral: ref, Response: resp}}); err != nil { + t.Fatalf("WriteSummary: %v", err) + } + + out := buf.String() + if !strings.Contains(out, "BIND 9.18.1") { + t.Fatalf("expected version string in text output, got %q", out) + } +} + +func TestJSONFormatterWriteSummaryShowsVersions(t *testing.T) { + ref := traverse.NewReferral("example.com.", dns.TypeA, "com.", 0, 1.0, nil) + serverIP := net.ParseIP("198.41.0.4") + resp := &traverse.Response{ + Referral: ref, + Server: serverIP, + Type: traverse.RespAnswer, + } + + var buf bytes.Buffer + cfg := DefaultConfig() + cfg.Format = FormatJSON + cfg.Domain = "example.com" + cfg.QueryType = "A" + cfg.ShowVersions = true + cfg.ShowServers = true + cfg.ShowResults = false + cfg.ShowSummaryResults = false + formatter := NewFormatter(cfg, &buf) + formatter.SetFingerprints(map[string]string{serverIP.String(): "Unbound 1.17.0"}) + + 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()) + } + servers, ok := payload["servers"].([]any) + if !ok || len(servers) == 0 { + t.Fatalf("expected servers in json output, got %v", payload) + } + srv, ok := servers[0].(map[string]any) + if !ok { + t.Fatalf("expected server object, got %T", servers[0]) + } + if srv["version"] != "Unbound 1.17.0" { + t.Fatalf("expected version = %q, got %v", "Unbound 1.17.0", srv["version"]) + } +} + func TestRunTraversalUsesHooks(t *testing.T) { answerResp := func() *miekgdns.Msg { m := new(miekgdns.Msg) diff --git a/internal/output/json.go b/internal/output/json.go index 29c4200..407b6c8 100644 --- a/internal/output/json.go +++ b/internal/output/json.go @@ -9,9 +9,10 @@ import ( ) type jsonFormatter struct { - cfg *Config - w io.Writer - payload jsonDocument + cfg *Config + w io.Writer + payload jsonDocument + fingerprints map[string]string } type jsonDocument struct { @@ -71,6 +72,10 @@ func newJSONFormatter(cfg *Config, w io.Writer) *jsonFormatter { } } +func (f *jsonFormatter) SetFingerprints(fps map[string]string) { + f.fingerprints = fps +} + func (f *jsonFormatter) WriteProgress(event traverse.TraversalEvent) error { if !f.cfg.ShowProgress { return nil @@ -106,9 +111,9 @@ func (f *jsonFormatter) WriteSummary(results []traverse.TraversalResult) error { servers := collectServers(results) for name, ips := range servers { srv := jsonServer{Name: name, IPs: ips} - if f.cfg.ShowVersions && f.cfg.Fingerprints != nil { + if f.cfg.ShowVersions && f.fingerprints != nil { for _, ip := range ips { - if v := f.cfg.Fingerprints[ip]; v != "" { + if v := f.fingerprints[ip]; v != "" { srv.Version = v break } diff --git a/internal/output/runner.go b/internal/output/runner.go index 67c6b6b..4e3139b 100644 --- a/internal/output/runner.go +++ b/internal/output/runner.go @@ -31,7 +31,8 @@ func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Confi // Gating on ShowServers avoids unnecessary network calls when versions // would not be displayed anyway. if cfg.ShowVersions && cfg.ShowServers { - cfg.Fingerprints = fingerprint.New().FingerprintAll(ctx, collectUniqueServerIPs(results)) + fps := fingerprint.New().FingerprintAll(ctx, collectUniqueServerIPs(results)) + formatter.SetFingerprints(fps) } if err := formatter.WriteSummary(results); err != nil { diff --git a/internal/output/text.go b/internal/output/text.go index a5f6fae..5933b62 100644 --- a/internal/output/text.go +++ b/internal/output/text.go @@ -11,14 +11,19 @@ import ( ) type textFormatter struct { - cfg *Config - w io.Writer + cfg *Config + w io.Writer + fingerprints map[string]string } func newTextFormatter(cfg *Config, w io.Writer) *textFormatter { return &textFormatter{cfg: cfg, w: w} } +func (f *textFormatter) SetFingerprints(fps map[string]string) { + f.fingerprints = fps +} + func (f *textFormatter) WriteProgress(event traverse.TraversalEvent) error { if event.Stage != traverse.EventStart { return nil @@ -101,7 +106,7 @@ func (f *textFormatter) writeServers(results []traverse.TraversalResult) error { for _, ip := range servers[name] { line := fmt.Sprintf("%*s: %-15s", width, name, ip) if f.cfg.ShowVersions { - if version, ok := f.cfg.Fingerprints[ip]; ok && version != "" { + if version, ok := f.fingerprints[ip]; ok && version != "" { line += " " + version } }