Files
rdns-go/rdns_test.go
T

305 lines
8.7 KiB
Go

package main
import (
"net"
"os"
"path/filepath"
"testing"
"time"
"github.com/miekg/dns"
"github.com/patrickmn/go-cache"
)
func TestMain(m *testing.M) {
qcache = cache.New(time.Minute, time.Minute)
os.Exit(m.Run())
}
type mockWriter struct {
msg *dns.Msg
addr net.Addr
}
func (m *mockWriter) WriteMsg(msg *dns.Msg) error { m.msg = msg; return nil }
func (m *mockWriter) LocalAddr() net.Addr { return m.addr }
func (m *mockWriter) RemoteAddr() net.Addr { return m.addr }
func (m *mockWriter) Write(b []byte) (int, error) { return len(b), nil }
func (m *mockWriter) Close() error { return nil }
func (m *mockWriter) TsigStatus() error { return nil }
func (m *mockWriter) TsigTimersOnly(bool) {}
func (m *mockWriter) Hijack() {}
func resetMaps() {
clear(hostsMap)
clear(blocklistMap)
clear(zonesMap)
}
func newQuery(name string, qtype uint16, qclass uint16) *dns.Msg {
q := new(dns.Msg)
q.SetQuestion(dns.Fqdn(name), qtype)
q.Question[0].Qclass = qclass
return q
}
func TestGetDomain(t *testing.T) {
cases := map[string]string{
"example.test": "example.test",
"a.example.test": "example.test",
"long.a.example.test": "example.test",
"localhost": "localhost",
}
for in, want := range cases {
if got := getDomain(in); got != want {
t.Errorf("getDomain(%q) = %q, want %q", in, got, want)
}
}
}
func TestDoDot(t *testing.T) {
cases := map[string]string{
"": ".",
".": ".",
"example.test": "example.test.",
}
for in, want := range cases {
if got := doDot(in); got != want {
t.Errorf("doDot(%q) = %q, want %q", in, got, want)
}
}
}
func TestGetRcode(t *testing.T) {
if got := getRcode(dns.RcodeSuccess); got != "NOERROR" {
t.Errorf("getRcode(0) = %q, want NOERROR", got)
}
if got := getRcode(dns.RcodeRefused); got != "REFUSED" {
t.Errorf("getRcode(5) = %q, want REFUSED", got)
}
if got := getRcode(250); got != "SERVFAIL" {
t.Errorf("getRcode(250) = %q, want SERVFAIL", got)
}
}
func TestCacheKeyIncludesClass(t *testing.T) {
a := cacheKey("version.bind", "CH", "TXT")
b := cacheKey("version.bind", "IN", "TXT")
if a == b {
t.Error("cache key must differ between qclasses")
}
}
func TestHandleCacheReturnsCopy(t *testing.T) {
qcache.Flush()
cached := new(dns.Msg)
cached.SetReply(newQuery("example.test", dns.TypeA, dns.ClassINET))
cached.Id = 4321
qcache.SetDefault(cacheKey("example.test", "IN", "A"), cached)
req := newQuery("example.test", dns.TypeA, dns.ClassINET)
req.Id = 1234
msg, rcode, ok := handleCache(req, cacheKey("example.test", "IN", "A"))
if !ok {
t.Fatal("expected cache hit")
}
if rcode != dns.RcodeSuccess {
t.Errorf("rcode = %d, want 0", rcode)
}
if msg.Id != 1234 {
t.Errorf("returned message ID = %d, want 1234 (caller ID)", msg.Id)
}
if cached.Id != 4321 {
t.Errorf("cached entry ID mutated: %d, want 4321", cached.Id)
}
if msg == cached {
t.Error("handleCache must return a copy, not the shared cached entry")
}
if _, _, ok := handleCache(req, cacheKey("other.test", "IN", "A")); ok {
t.Error("expected cache miss for unknown key")
}
}
func TestLoadBlocklist(t *testing.T) {
path := filepath.Join(t.TempDir(), "blocklist.txt")
content := "# comment\n\nExample.TEST\n spaced.test \nblocked.test\n# trailing comment\n"
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
resetMaps()
loadBlocklist(path)
want := map[string]bool{"example.test": true, "spaced.test": true, "blocked.test": true}
for name := range want {
if !blocklistMap[name] {
t.Errorf("blocklist missing %q", name)
}
}
if len(blocklistMap) != len(want) {
t.Errorf("blocklist size = %d, want %d", len(blocklistMap), len(want))
}
}
func TestLoadZones(t *testing.T) {
path := filepath.Join(t.TempDir(), "zones")
content := "# comment\nExample.TEST 192.168.1.102\nplain.test 10.0.0.1 # inline comment\nbroken-line\n"
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
resetMaps()
count := loadZones(path)
if count != 2 {
t.Errorf("count = %d, want 2", count)
}
if ns := getZoneNS("example.test"); ns != "192.168.1.102" {
t.Errorf("getZoneNS(example.test) = %q, want 192.168.1.102", ns)
}
if ns := getZoneNS("unknown.test"); ns != "." {
t.Errorf("getZoneNS(unknown.test) = %q, want \".\"", ns)
}
}
func TestLoadHosts(t *testing.T) {
path := filepath.Join(t.TempDir(), "hosts")
content := "# comment\n192.168.1.10 Example.TEST other.test\nnot-an-ip broken.test\n2001:db8::1 v6.test\n"
if err := os.WriteFile(path, []byte(content), 0o644); err != nil {
t.Fatal(err)
}
resetMaps()
count := loadHosts(path)
if count != 3 {
t.Errorf("count = %d, want 3", count)
}
if ip := getHost("example.test"); ip != "192.168.1.10" {
t.Errorf("getHost(example.test) = %q, want 192.168.1.10", ip)
}
if ip := getHost("broken.test"); ip != "" {
t.Errorf("getHost(broken.test) = %q, want empty (invalid IP must be skipped)", ip)
}
if ip := getHost("v6.test"); ip != "2001:db8::1" {
t.Errorf("getHost(v6.test) = %q, want 2001:db8::1", ip)
}
}
func TestClientIP(t *testing.T) {
w := &mockWriter{addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.10"), Port: 5300}}
if got := clientIP(w); got != "192.0.2.10" {
t.Errorf("clientIP = %q, want 192.0.2.10", got)
}
w = &mockWriter{addr: &net.TCPAddr{IP: net.ParseIP("2001:db8::1"), Port: 5300}}
if got := clientIP(w); got != "2001:db8::1" {
t.Errorf("clientIP = %q, want 2001:db8::1", got)
}
}
func TestHandleQueryHostsAnswer(t *testing.T) {
resetMaps()
qcache.Flush()
hostsMap["example.test"] = "192.168.1.10"
blocklistMap["example.test"] = true // ensures a fallthrough is observable
// A query matches the IPv4 entry
w := &mockWriter{addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.10"), Port: 5300}}
handleQuery(w, newQuery("example.test", dns.TypeA, dns.ClassINET))
if w.msg == nil {
t.Fatal("no response written")
}
if len(w.msg.Answer) != 1 {
t.Fatalf("answers = %d, want 1", len(w.msg.Answer))
}
a, ok := w.msg.Answer[0].(*dns.A)
if !ok {
t.Fatalf("answer type = %T, want *dns.A", w.msg.Answer[0])
}
if !a.A.Equal(net.ParseIP("192.168.1.10")) {
t.Errorf("A = %v, want 192.168.1.10", a.A)
}
// AAAA query does not match an IPv4 entry and must fall through (blocked here)
w = &mockWriter{addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.10"), Port: 5300}}
handleQuery(w, newQuery("example.test", dns.TypeAAAA, dns.ClassINET))
if w.msg == nil {
t.Fatal("no response written")
}
if w.msg.Rcode != dns.RcodeRefused {
t.Errorf("rcode = %d, want REFUSED (fallthrough to blocklist)", w.msg.Rcode)
}
}
func TestHandleQueryBlocked(t *testing.T) {
resetMaps()
qcache.Flush()
blocklistMap["ads.example.test"] = true
w := &mockWriter{addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.10"), Port: 5300}}
handleQuery(w, newQuery("ads.example.test", dns.TypeA, dns.ClassINET))
if w.msg == nil {
t.Fatal("no response written")
}
if w.msg.Rcode != dns.RcodeRefused {
t.Errorf("rcode = %d, want REFUSED", w.msg.Rcode)
}
if !w.msg.RecursionAvailable {
t.Error("RecursionAvailable not set on REFUSED reply")
}
// Second identical query must be served from cache
w2 := &mockWriter{addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.10"), Port: 5300}}
q2 := newQuery("ads.example.test", dns.TypeA, dns.ClassINET)
q2.Id = 7777
handleQuery(w2, q2)
if w2.msg == nil || w2.msg.Rcode != dns.RcodeRefused {
t.Fatalf("cached response rcode = %v, want REFUSED", w2.msg)
}
if w2.msg.Id != 7777 {
t.Errorf("cached response ID = %d, want 7777", w2.msg.Id)
}
}
func TestHandleQueryVersionBind(t *testing.T) {
resetMaps()
qcache.Flush()
w := &mockWriter{addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.10"), Port: 5300}}
handleQuery(w, newQuery("version.bind", dns.TypeTXT, dns.ClassCHAOS))
if w.msg == nil {
t.Fatal("no response written")
}
if len(w.msg.Answer) != 1 {
t.Fatalf("answers = %d, want 1", len(w.msg.Answer))
}
txt, ok := w.msg.Answer[0].(*dns.TXT)
if !ok {
t.Fatalf("answer type = %T, want *dns.TXT", w.msg.Answer[0])
}
if len(txt.Txt) != 1 || txt.Txt[0] != "rdns" {
t.Errorf("TXT = %v, want [rdns]", txt.Txt)
}
// An IN-class version.bind query must not be served from the CH cache entry
if _, _, ok := handleCache(newQuery("version.bind", dns.TypeTXT, dns.ClassINET), cacheKey("version.bind", "IN", "TXT")); ok {
t.Error("IN class query must not hit the CH class cache entry")
}
}
func TestHandleQueryMalformed(t *testing.T) {
w := &mockWriter{addr: &net.UDPAddr{IP: net.ParseIP("192.0.2.10"), Port: 5300}}
handleQuery(w, &dns.Msg{MsgHdr: dns.MsgHdr{Id: 42}})
if w.msg == nil {
t.Fatal("no response written for question-less query")
}
if w.msg.Rcode != dns.RcodeFormatError {
t.Errorf("rcode = %d, want FORMERR", w.msg.Rcode)
}
}