346 lines
9.0 KiB
Go
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)
|
|
}
|
|
}
|