From 3e7580b9191528761e10e5c4e6e87ce69ed9050e Mon Sep 17 00:00:00 2001 From: Gary Hansen Date: Mon, 8 Jun 2026 02:47:23 +1000 Subject: [PATCH] feat: implement DNS server fingerprinting (HAN-384) - Add internal/fingerprint package with Fingerprinter type - Sends version.bind CHAOS TXT query to each server - Caches results per IP to avoid redundant queries - Concurrent batch fingerprinting via FingerprintAll - Gracefully handles non-responding servers (returns empty string) - Injectable exchange function for testability - Wire fingerprinting into RunTraversal - Triggered when both ShowVersions and ShowServers are enabled - Collects unique server IPs from traversal results - Stores results in output.Config.Fingerprints before WriteSummary - Display versions in text output (writeServers) - Appends version string after IP when ShowVersions is true - No output change when version is unknown - Include version in JSON output (jsonServer.Version field) - omitempty: field absent when version is unknown - Add comprehensive fingerprint tests (11 test cases) Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: multica-agent --- internal/fingerprint/fingerprint.go | 137 +++++++++++++++ internal/fingerprint/fingerprint_test.go | 208 +++++++++++++++++++++++ internal/output/formatter.go | 4 + internal/output/json.go | 19 ++- internal/output/runner.go | 26 +++ internal/output/text.go | 8 +- 6 files changed, 395 insertions(+), 7 deletions(-) create mode 100644 internal/fingerprint/fingerprint_test.go diff --git a/internal/fingerprint/fingerprint.go b/internal/fingerprint/fingerprint.go index f1b53d1..ad1f890 100644 --- a/internal/fingerprint/fingerprint.go +++ b/internal/fingerprint/fingerprint.go @@ -1 +1,138 @@ package fingerprint + +import ( + "context" + "net" + "sync" + "time" + + miekgdns "github.com/miekg/dns" +) + +const defaultTimeout = 2 * time.Second + +// Fingerprinter queries DNS servers for their software version via the +// version.bind CHAOS TXT query. Results are cached per server IP. +type Fingerprinter struct { + mu sync.Mutex + cache map[string]string + timeout time.Duration + exchange func(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) +} + +// New returns a Fingerprinter with a 2-second per-query timeout. +func New() *Fingerprinter { + return NewWithTimeout(defaultTimeout) +} + +// NewWithTimeout returns a Fingerprinter using the given per-query timeout. +func NewWithTimeout(timeout time.Duration) *Fingerprinter { + return &Fingerprinter{ + cache: make(map[string]string), + timeout: timeout, + } +} + +// 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. +func (f *Fingerprinter) Query(ctx context.Context, ip net.IP) string { + key := ip.String() + + f.mu.Lock() + if v, ok := f.cache[key]; ok { + f.mu.Unlock() + return v + } + f.mu.Unlock() + + version := f.probe(ctx, ip) + + f.mu.Lock() + f.cache[key] = version + f.mu.Unlock() + + return version +} + +// FingerprintAll queries all ips concurrently and returns a map of +// IP string → version string. IPs that don't respond map to "". +// Already-cached IPs are returned from cache without a network round-trip. +func (f *Fingerprinter) FingerprintAll(ctx context.Context, ips []net.IP) map[string]string { + results := make(map[string]string, len(ips)) + + var ( + wg sync.WaitGroup + mu sync.Mutex + toQuery []net.IP + ) + + f.mu.Lock() + for _, ip := range ips { + key := ip.String() + if v, ok := f.cache[key]; ok { + results[key] = v + } else { + toQuery = append(toQuery, ip) + } + } + f.mu.Unlock() + + for _, ip := range toQuery { + wg.Add(1) + go func(ip net.IP) { + defer wg.Done() + version := f.probe(ctx, ip) + key := ip.String() + + f.mu.Lock() + f.cache[key] = version + f.mu.Unlock() + + mu.Lock() + results[key] = version + mu.Unlock() + }(ip) + } + + wg.Wait() + return results +} + +// probe sends a version.bind CHAOS TXT query and returns the version string, +// or "" on any error or non-success response. +func (f *Fingerprinter) probe(ctx context.Context, ip net.IP) string { + m := new(miekgdns.Msg) + m.SetQuestion("version.bind.", miekgdns.TypeTXT) + m.Question[0].Qclass = miekgdns.ClassCHAOS + m.RecursionDesired = false + + target := net.JoinHostPort(ip.String(), "53") + + queryCtx, cancel := context.WithTimeout(ctx, f.timeout) + defer cancel() + + var resp *miekgdns.Msg + var err error + + if f.exchange != nil { + resp, err = f.exchange(queryCtx, target, m) + } else { + client := &miekgdns.Client{ + Net: "udp", + Timeout: f.timeout, + } + resp, _, err = client.ExchangeContext(queryCtx, m, target) + } + + if err != nil || resp == nil || resp.Rcode != miekgdns.RcodeSuccess { + return "" + } + + for _, rr := range resp.Answer { + if txt, ok := rr.(*miekgdns.TXT); ok && len(txt.Txt) > 0 { + return txt.Txt[0] + } + } + return "" +} diff --git a/internal/fingerprint/fingerprint_test.go b/internal/fingerprint/fingerprint_test.go new file mode 100644 index 0000000..fea1d3a --- /dev/null +++ b/internal/fingerprint/fingerprint_test.go @@ -0,0 +1,208 @@ +package fingerprint + +import ( + "context" + "net" + "testing" + "time" + + miekgdns "github.com/miekg/dns" +) + +func mockExchange(version string) func(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) { + return func(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) { + resp := new(miekgdns.Msg) + resp.SetReply(m) + resp.Answer = []miekgdns.RR{ + &miekgdns.TXT{ + Hdr: miekgdns.RR_Header{ + Name: "version.bind.", + Rrtype: miekgdns.TypeTXT, + Class: miekgdns.ClassCHAOS, + Ttl: 0, + }, + Txt: []string{version}, + }, + } + return resp, nil + } +} + +func errorExchange(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) { + return nil, context.DeadlineExceeded +} + +func refusedExchange(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) { + resp := new(miekgdns.Msg) + resp.SetReply(m) + resp.Rcode = miekgdns.RcodeRefused + return resp, nil +} + +func emptyTXTExchange(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) { + resp := new(miekgdns.Msg) + resp.SetReply(m) + // No Answer records. + return resp, nil +} + +func TestQueryReturnsVersion(t *testing.T) { + fp := New() + fp.exchange = mockExchange("BIND 9.18.1") + + version := fp.Query(context.Background(), net.ParseIP("1.2.3.4")) + if version != "BIND 9.18.1" { + t.Fatalf("Query() = %q, want %q", version, "BIND 9.18.1") + } +} + +func TestQueryCachesResult(t *testing.T) { + calls := 0 + fp := New() + fp.exchange = func(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) { + calls++ + return mockExchange("Unbound 1.17.0")(ctx, addr, m) + } + + ip := net.ParseIP("1.2.3.4") + _ = fp.Query(context.Background(), ip) + _ = fp.Query(context.Background(), ip) + + if calls != 1 { + t.Fatalf("expected 1 exchange call (cache hit on second), got %d", calls) + } +} + +func TestQueryReturnsEmptyOnError(t *testing.T) { + fp := New() + fp.exchange = errorExchange + + version := fp.Query(context.Background(), net.ParseIP("1.2.3.4")) + if version != "" { + t.Fatalf("Query() = %q on error, want empty string", version) + } +} + +func TestQueryReturnsEmptyOnRefused(t *testing.T) { + fp := New() + fp.exchange = refusedExchange + + version := fp.Query(context.Background(), net.ParseIP("1.2.3.4")) + if version != "" { + t.Fatalf("Query() = %q on REFUSED, want empty string", version) + } +} + +func TestQueryReturnsEmptyWhenNoTXTRecord(t *testing.T) { + fp := New() + fp.exchange = emptyTXTExchange + + version := fp.Query(context.Background(), net.ParseIP("1.2.3.4")) + if version != "" { + t.Fatalf("Query() = %q with no TXT answer, want empty string", version) + } +} + +func TestFingerprintAllConcurrent(t *testing.T) { + fp := New() + fp.exchange = func(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) { + // Return different version strings per address. + resp := new(miekgdns.Msg) + resp.SetReply(m) + resp.Answer = []miekgdns.RR{ + &miekgdns.TXT{ + Hdr: miekgdns.RR_Header{Name: "version.bind.", Rrtype: miekgdns.TypeTXT, Class: miekgdns.ClassCHAOS}, + Txt: []string{"BIND " + addr}, + }, + } + return resp, nil + } + + ips := []net.IP{ + net.ParseIP("1.1.1.1"), + net.ParseIP("2.2.2.2"), + net.ParseIP("3.3.3.3"), + } + + results := fp.FingerprintAll(context.Background(), ips) + + if len(results) != 3 { + t.Fatalf("FingerprintAll returned %d results, want 3", len(results)) + } + for _, ip := range ips { + if results[ip.String()] == "" { + t.Errorf("FingerprintAll missing version for %s", ip) + } + } +} + +func TestFingerprintAllUsesCache(t *testing.T) { + calls := 0 + fp := New() + fp.exchange = func(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) { + calls++ + return mockExchange("PowerDNS 4.7")(ctx, addr, m) + } + + ip := net.ParseIP("10.0.0.1") + // Prime the cache. + _ = fp.Query(context.Background(), ip) + // FingerprintAll should not re-query cached IPs. + _ = fp.FingerprintAll(context.Background(), []net.IP{ip}) + + if calls != 1 { + t.Fatalf("FingerprintAll re-queried a cached IP: got %d calls, want 1", calls) + } +} + +func TestFingerprintAllHandlesErrors(t *testing.T) { + fp := New() + fp.exchange = errorExchange + + ips := []net.IP{net.ParseIP("1.2.3.4"), net.ParseIP("5.6.7.8")} + results := fp.FingerprintAll(context.Background(), ips) + + for _, ip := range ips { + if v, ok := results[ip.String()]; !ok || v != "" { + t.Errorf("expected empty version for %s on error, got %q (present=%v)", ip, v, ok) + } + } +} + +func TestNewWithTimeout(t *testing.T) { + fp := NewWithTimeout(100 * time.Millisecond) + if fp.timeout != 100*time.Millisecond { + t.Fatalf("timeout = %v, want 100ms", fp.timeout) + } +} + +func TestQueryBuildsCorrectCHAOSQuery(t *testing.T) { + var capturedMsg *miekgdns.Msg + fp := New() + fp.exchange = func(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error) { + capturedMsg = m.Copy() + return emptyTXTExchange(ctx, addr, m) + } + + _ = fp.Query(context.Background(), net.ParseIP("1.2.3.4")) + + if capturedMsg == nil { + t.Fatal("exchange was not called") + } + if len(capturedMsg.Question) != 1 { + t.Fatalf("expected 1 question, got %d", len(capturedMsg.Question)) + } + q := capturedMsg.Question[0] + if q.Name != "version.bind." { + t.Errorf("question name = %q, want %q", q.Name, "version.bind.") + } + if q.Qtype != miekgdns.TypeTXT { + t.Errorf("question type = %d, want TXT (%d)", q.Qtype, miekgdns.TypeTXT) + } + if q.Qclass != miekgdns.ClassCHAOS { + t.Errorf("question class = %d, want CHAOS (%d)", q.Qclass, miekgdns.ClassCHAOS) + } + if capturedMsg.RecursionDesired { + t.Error("RecursionDesired should be false for CHAOS queries") + } +} diff --git a/internal/output/formatter.go b/internal/output/formatter.go index 446eac2..5cd44b7 100644 --- a/internal/output/formatter.go +++ b/internal/output/formatter.go @@ -30,6 +30,10 @@ type Config struct { Quiet bool 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 { diff --git a/internal/output/json.go b/internal/output/json.go index ea96959..29c4200 100644 --- a/internal/output/json.go +++ b/internal/output/json.go @@ -44,8 +44,9 @@ type jsonResult struct { } type jsonServer struct { - Name string `json:"name"` - IPs []string `json:"ips"` + Name string `json:"name"` + IPs []string `json:"ips"` + Version string `json:"version,omitempty"` } type jsonSummary struct { @@ -104,10 +105,16 @@ func (f *jsonFormatter) WriteSummary(results []traverse.TraversalResult) error { if f.cfg.ShowServers { servers := collectServers(results) for name, ips := range servers { - f.payload.Servers = append(f.payload.Servers, jsonServer{ - Name: name, - IPs: ips, - }) + srv := jsonServer{Name: name, IPs: ips} + if f.cfg.ShowVersions && f.cfg.Fingerprints != nil { + for _, ip := range ips { + if v := f.cfg.Fingerprints[ip]; v != "" { + srv.Version = v + break + } + } + } + f.payload.Servers = append(f.payload.Servers, srv) } } diff --git a/internal/output/runner.go b/internal/output/runner.go index cc5354d..67c6b6b 100644 --- a/internal/output/runner.go +++ b/internal/output/runner.go @@ -3,7 +3,9 @@ package output import ( "context" "fmt" + "net" + "github.com/hits/ExploreDNS/internal/fingerprint" "github.com/hits/ExploreDNS/internal/traverse" ) @@ -25,6 +27,13 @@ func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Confi return results, err } + // Fingerprint servers when both ShowVersions and ShowServers are enabled. + // 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)) + } + if err := formatter.WriteSummary(results); err != nil { return results, err } @@ -34,3 +43,20 @@ func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Confi return results, nil } + +// collectUniqueServerIPs returns the set of unique server IPs seen in results. +func collectUniqueServerIPs(results []traverse.TraversalResult) []net.IP { + seen := make(map[string]bool) + var ips []net.IP + for _, r := range results { + if r.Response == nil || r.Response.Server == nil { + continue + } + key := r.Response.Server.String() + if !seen[key] { + seen[key] = true + ips = append(ips, r.Response.Server) + } + } + return ips +} diff --git a/internal/output/text.go b/internal/output/text.go index 48d3bf9..a5f6fae 100644 --- a/internal/output/text.go +++ b/internal/output/text.go @@ -99,7 +99,13 @@ func (f *textFormatter) writeServers(results []traverse.TraversalResult) error { for _, name := range names { for _, ip := range servers[name] { - if _, err := fmt.Fprintf(f.w, "%*s: %-15s\n", width, name, ip); err != nil { + line := fmt.Sprintf("%*s: %-15s", width, name, ip) + if f.cfg.ShowVersions { + if version, ok := f.cfg.Fingerprints[ip]; ok && version != "" { + line += " " + version + } + } + if _, err := fmt.Fprintln(f.w, line); err != nil { return err } }