Files

346 lines
9.0 KiB
Go

package resolver_test
import (
"ekdns/resolver"
"encoding/binary"
"net"
"strings"
"testing"
"time"
"golang.org/x/net/dns/dnsmessage"
)
// ---------------------------------------------------------------------------
// Fake DNS server helpers
// ---------------------------------------------------------------------------
// startFakeDNS starts a UDP listener that handles a single request and responds
// using the provided handler. Returns the server address "host:port".
func startFakeDNS(t *testing.T, handler func(query []byte) []byte) string {
t.Helper()
conn, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatalf("startFakeDNS: %v", err)
}
t.Cleanup(func() { conn.Close() })
go func() {
buf := make([]byte, 1232)
n, addr, err := conn.ReadFrom(buf)
if err != nil {
return
}
resp := handler(buf[:n])
if resp != nil {
conn.WriteTo(resp, addr) //nolint:errcheck
}
}()
return conn.LocalAddr().String()
}
// startFakeDNSMulti starts a UDP listener that handles multiple requests.
func startFakeDNSMulti(t *testing.T, handler func(query []byte) []byte) string {
t.Helper()
conn, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatalf("startFakeDNSMulti: %v", err)
}
t.Cleanup(func() { conn.Close() })
go func() {
buf := make([]byte, 1232)
for {
n, addr, err := conn.ReadFrom(buf)
if err != nil {
return
}
resp := handler(buf[:n])
if resp != nil {
conn.WriteTo(resp, addr) //nolint:errcheck
}
}
}()
return conn.LocalAddr().String()
}
// queryID extracts the DNS message ID from the first 2 bytes.
func queryID(msg []byte) uint16 {
if len(msg) < 2 {
return 0
}
return binary.BigEndian.Uint16(msg[:2])
}
// buildAResponse builds a DNS response with the given A records.
func buildAResponse(id uint16, name dnsmessage.Name, ips [][4]byte) []byte {
answers := make([]dnsmessage.Resource, len(ips))
for i, ip := range ips {
answers[i] = dnsmessage.Resource{
Header: dnsmessage.ResourceHeader{
Name: name,
Type: dnsmessage.TypeA,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.AResource{A: ip},
}
}
msg := dnsmessage.Message{
Header: dnsmessage.Header{
ID: id,
Response: true,
RCode: dnsmessage.RCodeSuccess,
},
Questions: []dnsmessage.Question{{
Name: name,
Type: dnsmessage.TypeA,
Class: dnsmessage.ClassINET,
}},
Answers: answers,
}
packed, err := msg.Pack()
if err != nil {
panic("buildAResponse: pack failed: " + err.Error())
}
return packed
}
// buildCNAMEResponse builds a DNS response containing a single CNAME record.
func buildCNAMEResponse(id uint16, queryName, cnameTarget dnsmessage.Name) []byte {
msg := dnsmessage.Message{
Header: dnsmessage.Header{
ID: id,
Response: true,
RCode: dnsmessage.RCodeSuccess,
},
Questions: []dnsmessage.Question{{
Name: queryName,
Type: dnsmessage.TypeA,
Class: dnsmessage.ClassINET,
}},
Answers: []dnsmessage.Resource{{
Header: dnsmessage.ResourceHeader{
Name: queryName,
Type: dnsmessage.TypeCNAME,
Class: dnsmessage.ClassINET,
TTL: 60,
},
Body: &dnsmessage.CNAMEResource{CNAME: cnameTarget},
}},
}
packed, err := msg.Pack()
if err != nil {
panic("buildCNAMEResponse: pack failed: " + err.Error())
}
return packed
}
// buildNXDOMAINResponse builds a DNS response with NXDOMAIN rcode.
func buildNXDOMAINResponse(id uint16, name dnsmessage.Name) []byte {
msg := dnsmessage.Message{
Header: dnsmessage.Header{
ID: id,
Response: true,
RCode: dnsmessage.RCodeNameError, // NXDOMAIN
},
Questions: []dnsmessage.Question{{
Name: name,
Type: dnsmessage.TypeA,
Class: dnsmessage.ClassINET,
}},
}
packed, err := msg.Pack()
if err != nil {
panic("buildNXDOMAINResponse: pack failed: " + err.Error())
}
return packed
}
// mustNewName creates a dnsmessage.Name from an FQDN, panicking on error.
func mustNewName(fqdn string) dnsmessage.Name {
n, err := dnsmessage.NewName(fqdn)
if err != nil {
panic("mustNewName: " + err.Error())
}
return n
}
// ---------------------------------------------------------------------------
// Tests
// ---------------------------------------------------------------------------
func TestLookupIP_ARecordSuccess(t *testing.T) {
name := mustNewName("example.com.")
serverAddr := startFakeDNS(t, func(query []byte) []byte {
id := queryID(query)
return buildAResponse(id, name, [][4]byte{
{1, 1, 1, 1},
{2, 2, 2, 2},
})
})
r := resolver.New()
ips, err := r.LookupIP("example.com", serverAddr)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(ips) != 2 {
t.Fatalf("expected 2 IPs, got %d: %v", len(ips), ips)
}
want := map[string]bool{"1.1.1.1": true, "2.2.2.2": true}
for _, ip := range ips {
if !want[ip] {
t.Errorf("unexpected IP %q", ip)
}
}
}
func TestLookupIP_CNAMEChainToARecords(t *testing.T) {
queryName := mustNewName("alias.example.com.")
targetName := mustNewName("real.example.com.")
callCount := 0
serverAddr := startFakeDNSMulti(t, func(query []byte) []byte {
id := queryID(query)
callCount++
if callCount == 1 {
// First query: return CNAME
return buildCNAMEResponse(id, queryName, targetName)
}
// Second query (for CNAME target): return A record
return buildAResponse(id, targetName, [][4]byte{{3, 3, 3, 3}})
})
r := resolver.New()
ips, err := r.LookupIP("alias.example.com", serverAddr)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(ips) != 1 || ips[0] != "3.3.3.3" {
t.Errorf("expected [3.3.3.3], got %v", ips)
}
}
func TestLookupIP_CNAMEDepthLimit(t *testing.T) {
// Always return a CNAME chain; the resolver should fail at depth > 10.
counter := 0
serverAddr := startFakeDNSMulti(t, func(query []byte) []byte {
id := queryID(query)
counter++
// Build a CNAME pointing to a unique target to avoid caching
srcName := mustNewName("host.example.com.")
targetName := mustNewName("host.example.com.")
return buildCNAMEResponse(id, srcName, targetName)
})
r := resolver.New()
_, err := r.LookupIP("host.example.com", serverAddr)
if err == nil {
t.Fatal("expected error for CNAME depth limit")
}
if !strings.Contains(err.Error(), "CNAME chain depth") {
t.Errorf("expected 'CNAME chain depth' in error, got: %v", err)
}
}
func TestLookupIP_NXDOMAIN(t *testing.T) {
name := mustNewName("notexist.example.com.")
serverAddr := startFakeDNS(t, func(query []byte) []byte {
id := queryID(query)
return buildNXDOMAINResponse(id, name)
})
r := resolver.New()
_, err := r.LookupIP("notexist.example.com", serverAddr)
if err == nil {
t.Fatal("expected error for NXDOMAIN")
}
if !strings.Contains(err.Error(), "notexist.example.com") {
t.Errorf("error should mention hostname, got: %v", err)
}
}
func TestLookupIP_ServerTimeout(t *testing.T) {
// Use a real listener but never reply — causes a read timeout.
conn, err := net.ListenPacket("udp", "127.0.0.1:0")
if err != nil {
t.Fatalf("ListenPacket: %v", err)
}
serverAddr := conn.LocalAddr().String()
conn.Close() // close immediately — resolver can't connect
r := resolver.New()
done := make(chan error, 1)
go func() {
_, err := r.LookupIP("example.com", serverAddr)
done <- err
}()
select {
case err := <-done:
if err == nil {
t.Fatal("expected error for unreachable server")
}
case <-time.After(10 * time.Second):
t.Fatal("LookupIP did not return within 10s")
}
}
func TestLookupIP_Deduplication(t *testing.T) {
name := mustNewName("example.com.")
serverAddr := startFakeDNS(t, func(query []byte) []byte {
id := queryID(query)
// Return the same IP twice
return buildAResponse(id, name, [][4]byte{
{1, 1, 1, 1},
{1, 1, 1, 1},
})
})
r := resolver.New()
ips, err := r.LookupIP("example.com", serverAddr)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(ips) != 1 {
t.Errorf("expected 1 unique IP after dedup, got %d: %v", len(ips), ips)
}
}
func TestLookupIP_FQDNNormalization(t *testing.T) {
// Hostname without trailing dot — resolver must append it internally.
name := mustNewName("example.com.")
serverAddr := startFakeDNS(t, func(query []byte) []byte {
id := queryID(query)
return buildAResponse(id, name, [][4]byte{{5, 5, 5, 5}})
})
r := resolver.New()
// Pass hostname WITHOUT trailing dot
ips, err := r.LookupIP("example.com", serverAddr)
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if len(ips) != 1 || ips[0] != "5.5.5.5" {
t.Errorf("expected [5.5.5.5], got %v", ips)
}
}
func TestLookupIP_ResponseIDMismatch(t *testing.T) {
name := mustNewName("example.com.")
serverAddr := startFakeDNS(t, func(query []byte) []byte {
id := queryID(query)
// Return response with wrong ID (XOR with 0xFFFF)
wrongID := id ^ 0xFFFF
return buildAResponse(wrongID, name, [][4]byte{{1, 1, 1, 1}})
})
r := resolver.New()
_, err := r.LookupIP("example.com", serverAddr)
if err == nil {
t.Fatal("expected error for response ID mismatch")
}
if !strings.Contains(strings.ToLower(err.Error()), "mismatch") {
t.Errorf("error should mention 'mismatch', got: %v", err)
}
}