fix: address code review issues on PR #10 (HAN-384)
- runner.go: store fingerprints in local var, call formatter.SetFingerprints() instead of mutating caller's *Config; removes Fingerprints field from Config - fingerprint.go: use singleflight.Group in Query to coalesce concurrent probes for the same IP, eliminating the double-probe race - output layer: add SetFingerprints(map[string]string) to Formatter interface; textFormatter and jsonFormatter each hold their own fingerprints field - formatter_test.go: add TestTextFormatterWriteSummaryShowsVersions and TestJSONFormatterWriteSummaryShowsVersions verifying version output Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Co-authored-by: multica-agent <github@multica.ai>
This commit is contained in:
co-authored by
Copilot
multica-agent
parent
3e7580b919
commit
c96d18859f
@@ -2,12 +2,14 @@ module github.com/hits/ExploreDNS
|
|||||||
|
|
||||||
go 1.25.6
|
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 (
|
require (
|
||||||
golang.org/x/mod v0.31.0 // indirect
|
golang.org/x/mod v0.31.0 // indirect
|
||||||
golang.org/x/net v0.48.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/sys v0.39.0 // indirect
|
||||||
golang.org/x/tools v0.40.0 // indirect
|
golang.org/x/tools v0.40.0 // indirect
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
|
|
||||||
miekgdns "github.com/miekg/dns"
|
miekgdns "github.com/miekg/dns"
|
||||||
|
"golang.org/x/sync/singleflight"
|
||||||
)
|
)
|
||||||
|
|
||||||
const defaultTimeout = 2 * time.Second
|
const defaultTimeout = 2 * time.Second
|
||||||
@@ -17,6 +18,7 @@ type Fingerprinter struct {
|
|||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
cache map[string]string
|
cache map[string]string
|
||||||
timeout time.Duration
|
timeout time.Duration
|
||||||
|
group singleflight.Group
|
||||||
exchange func(ctx context.Context, addr string, m *miekgdns.Msg) (*miekgdns.Msg, error)
|
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
|
// Query returns the version string for ip, or "" if the server doesn't
|
||||||
// respond or doesn't support the version.bind CHAOS query.
|
// respond or doesn't support the version.bind CHAOS query.
|
||||||
// Results are cached: subsequent calls for the same IP return immediately.
|
// 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 {
|
func (f *Fingerprinter) Query(ctx context.Context, ip net.IP) string {
|
||||||
key := ip.String()
|
key := ip.String()
|
||||||
|
|
||||||
@@ -46,13 +50,15 @@ func (f *Fingerprinter) Query(ctx context.Context, ip net.IP) string {
|
|||||||
}
|
}
|
||||||
f.mu.Unlock()
|
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()
|
return v.(string)
|
||||||
f.cache[key] = version
|
|
||||||
f.mu.Unlock()
|
|
||||||
|
|
||||||
return version
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// FingerprintAll queries all ips concurrently and returns a map of
|
// FingerprintAll queries all ips concurrently and returns a map of
|
||||||
|
|||||||
@@ -31,9 +31,6 @@ type Config struct {
|
|||||||
Color bool
|
Color bool
|
||||||
Debug int
|
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 {
|
func DefaultConfig() *Config {
|
||||||
@@ -56,6 +53,9 @@ type Formatter interface {
|
|||||||
WriteResult(result traverse.TraversalResult) error
|
WriteResult(result traverse.TraversalResult) error
|
||||||
WriteSummary(results []traverse.TraversalResult) error
|
WriteSummary(results []traverse.TraversalResult) error
|
||||||
Flush() 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 {
|
func NewFormatter(cfg *Config, w io.Writer) Formatter {
|
||||||
|
|||||||
@@ -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) {
|
func TestRunTraversalUsesHooks(t *testing.T) {
|
||||||
answerResp := func() *miekgdns.Msg {
|
answerResp := func() *miekgdns.Msg {
|
||||||
m := new(miekgdns.Msg)
|
m := new(miekgdns.Msg)
|
||||||
|
|||||||
+10
-5
@@ -9,9 +9,10 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type jsonFormatter struct {
|
type jsonFormatter struct {
|
||||||
cfg *Config
|
cfg *Config
|
||||||
w io.Writer
|
w io.Writer
|
||||||
payload jsonDocument
|
payload jsonDocument
|
||||||
|
fingerprints map[string]string
|
||||||
}
|
}
|
||||||
|
|
||||||
type jsonDocument struct {
|
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 {
|
func (f *jsonFormatter) WriteProgress(event traverse.TraversalEvent) error {
|
||||||
if !f.cfg.ShowProgress {
|
if !f.cfg.ShowProgress {
|
||||||
return nil
|
return nil
|
||||||
@@ -106,9 +111,9 @@ func (f *jsonFormatter) WriteSummary(results []traverse.TraversalResult) error {
|
|||||||
servers := collectServers(results)
|
servers := collectServers(results)
|
||||||
for name, ips := range servers {
|
for name, ips := range servers {
|
||||||
srv := jsonServer{Name: name, IPs: ips}
|
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 {
|
for _, ip := range ips {
|
||||||
if v := f.cfg.Fingerprints[ip]; v != "" {
|
if v := f.fingerprints[ip]; v != "" {
|
||||||
srv.Version = v
|
srv.Version = v
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -31,7 +31,8 @@ func RunTraversal(ctx context.Context, traverser *traverse.Traverser, cfg *Confi
|
|||||||
// Gating on ShowServers avoids unnecessary network calls when versions
|
// Gating on ShowServers avoids unnecessary network calls when versions
|
||||||
// would not be displayed anyway.
|
// would not be displayed anyway.
|
||||||
if cfg.ShowVersions && cfg.ShowServers {
|
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 {
|
if err := formatter.WriteSummary(results); err != nil {
|
||||||
|
|||||||
@@ -11,14 +11,19 @@ import (
|
|||||||
)
|
)
|
||||||
|
|
||||||
type textFormatter struct {
|
type textFormatter struct {
|
||||||
cfg *Config
|
cfg *Config
|
||||||
w io.Writer
|
w io.Writer
|
||||||
|
fingerprints map[string]string
|
||||||
}
|
}
|
||||||
|
|
||||||
func newTextFormatter(cfg *Config, w io.Writer) *textFormatter {
|
func newTextFormatter(cfg *Config, w io.Writer) *textFormatter {
|
||||||
return &textFormatter{cfg: cfg, w: w}
|
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 {
|
func (f *textFormatter) WriteProgress(event traverse.TraversalEvent) error {
|
||||||
if event.Stage != traverse.EventStart {
|
if event.Stage != traverse.EventStart {
|
||||||
return nil
|
return nil
|
||||||
@@ -101,7 +106,7 @@ func (f *textFormatter) writeServers(results []traverse.TraversalResult) error {
|
|||||||
for _, ip := range servers[name] {
|
for _, ip := range servers[name] {
|
||||||
line := fmt.Sprintf("%*s: %-15s", width, name, ip)
|
line := fmt.Sprintf("%*s: %-15s", width, name, ip)
|
||||||
if f.cfg.ShowVersions {
|
if f.cfg.ShowVersions {
|
||||||
if version, ok := f.cfg.Fingerprints[ip]; ok && version != "" {
|
if version, ok := f.fingerprints[ip]; ok && version != "" {
|
||||||
line += " " + version
|
line += " " + version
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user