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") } }