305 lines
8.7 KiB
Go
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)
|
|
}
|
|
}
|