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