feat: Implement platform-specific hosts file path retrieval
- Created a new `platform` package to handle OS-specific hosts file paths. - Added implementations for macOS, Linux, and Windows to retrieve the hosts file path. - Introduced tests for the `GetHostsFilePath` function to ensure correct behavior across platforms. feat: Add DNS resolver for IP lookups - Implemented a `resolver` package for performing DNS lookups. - Created a `DNSResolver` struct with a `LookupIP` method to handle A-record lookups and CNAME resolution. - Added comprehensive tests for various DNS scenarios, including CNAME chains and error handling. chore: Refactor project structure and update dependencies - Restructured the project to follow Go's standard package layout. - Updated `go.mod` to Go 1.20 and added `golang.org/x/net` dependency. - Removed old source files and ensured a clean build with all tests passing.
This commit is contained in:
@@ -0,0 +1,175 @@
|
||||
package resolver
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/net/dns/dnsmessage"
|
||||
)
|
||||
|
||||
const (
|
||||
defaultTimeout = 5 * time.Second
|
||||
maxCNAMEDepth = 10
|
||||
udpBufSize = 1232
|
||||
)
|
||||
|
||||
// Resolver performs DNS lookups against a specific server.
|
||||
type Resolver interface {
|
||||
LookupIP(hostname, server string) ([]string, error)
|
||||
}
|
||||
|
||||
// DNSResolver implements Resolver using golang.org/x/net/dns/dnsmessage.
|
||||
type DNSResolver struct{}
|
||||
|
||||
// New returns a new DNSResolver.
|
||||
func New() *DNSResolver {
|
||||
return &DNSResolver{}
|
||||
}
|
||||
|
||||
// LookupIP performs a DNS A-record lookup for hostname using the given DNS server.
|
||||
// server may include a port (e.g. "8.8.8.8:53") or just an IP/host (":53" appended).
|
||||
func (d *DNSResolver) LookupIP(hostname, server string) ([]string, error) {
|
||||
return d.lookupIPWithDepth(hostname, server, 0)
|
||||
}
|
||||
|
||||
func (d *DNSResolver) lookupIPWithDepth(hostname, server string, depth int) ([]string, error) {
|
||||
if depth > maxCNAMEDepth {
|
||||
return nil, fmt.Errorf("CNAME chain depth exceeded for %s (max %d)", hostname, maxCNAMEDepth)
|
||||
}
|
||||
|
||||
// Ensure FQDN (trailing dot required by dnsmessage.NewName).
|
||||
fqdn := hostname
|
||||
if !strings.HasSuffix(fqdn, ".") {
|
||||
fqdn += "."
|
||||
}
|
||||
name, err := dnsmessage.NewName(fqdn)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid hostname %q: %w", hostname, err)
|
||||
}
|
||||
|
||||
// Build A query with random ID.
|
||||
id := uint16(rand.Uint32()) //nolint:gosec
|
||||
msg := dnsmessage.Message{
|
||||
Header: dnsmessage.Header{
|
||||
ID: id,
|
||||
RecursionDesired: true,
|
||||
},
|
||||
Questions: []dnsmessage.Question{{
|
||||
Name: name,
|
||||
Type: dnsmessage.TypeA,
|
||||
Class: dnsmessage.ClassINET,
|
||||
}},
|
||||
}
|
||||
packed, err := msg.Pack()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("packing DNS query: %w", err)
|
||||
}
|
||||
|
||||
// Connect over UDP.
|
||||
addr := server
|
||||
if !strings.Contains(server, ":") {
|
||||
addr = server + ":53"
|
||||
}
|
||||
conn, err := net.DialTimeout("udp", addr, defaultTimeout)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("connecting to DNS server %s: %w", server, err)
|
||||
}
|
||||
defer conn.Close()
|
||||
conn.SetDeadline(time.Now().Add(defaultTimeout)) //nolint:errcheck
|
||||
|
||||
if _, err := conn.Write(packed); err != nil {
|
||||
return nil, fmt.Errorf("sending DNS query to %s: %w", server, err)
|
||||
}
|
||||
|
||||
buf := make([]byte, udpBufSize)
|
||||
n, err := conn.Read(buf)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reading DNS response from %s: %w", server, err)
|
||||
}
|
||||
|
||||
// Parse response using the streaming Parser API (idiomatic dnsmessage approach).
|
||||
var parser dnsmessage.Parser
|
||||
respHeader, err := parser.Start(buf[:n])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parsing DNS response: %w", err)
|
||||
}
|
||||
|
||||
// Validate response ID to guard against spoofing / mismatched replies.
|
||||
if respHeader.ID != id {
|
||||
return nil, fmt.Errorf("DNS response ID mismatch (expected %d, got %d)", id, respHeader.ID)
|
||||
}
|
||||
|
||||
// Reject truncated responses — no TCP fallback per contract.
|
||||
if respHeader.Truncated {
|
||||
return nil, fmt.Errorf("DNS response truncated for %s; TCP fallback not supported", hostname)
|
||||
}
|
||||
|
||||
// Non-success rcodes.
|
||||
if respHeader.RCode != dnsmessage.RCodeSuccess {
|
||||
return nil, fmt.Errorf("DNS query for %s failed: %s", hostname, respHeader.RCode.String())
|
||||
}
|
||||
|
||||
// Skip the questions section.
|
||||
if err := parser.SkipAllQuestions(); err != nil {
|
||||
return nil, fmt.Errorf("parsing DNS response questions: %w", err)
|
||||
}
|
||||
|
||||
// Collect A records and CNAME targets from the answer section.
|
||||
var ips []string
|
||||
var cnameTargets []string
|
||||
for {
|
||||
hdr, err := parser.AnswerHeader()
|
||||
if err == dnsmessage.ErrSectionDone {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parsing DNS answer header: %w", err)
|
||||
}
|
||||
switch hdr.Type {
|
||||
case dnsmessage.TypeA:
|
||||
aRec, err := parser.AResource()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parsing A record: %w", err)
|
||||
}
|
||||
ips = append(ips, net.IP(aRec.A[:]).String())
|
||||
case dnsmessage.TypeCNAME:
|
||||
cnameRec, err := parser.CNAMEResource()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parsing CNAME record: %w", err)
|
||||
}
|
||||
cnameTargets = append(cnameTargets, cnameRec.CNAME.String())
|
||||
default:
|
||||
if err := parser.SkipAnswer(); err != nil {
|
||||
return nil, fmt.Errorf("skipping DNS answer: %w", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// If no A records but CNAME targets exist, follow the CNAME chain.
|
||||
// Only follow when there are no A records — some servers return both.
|
||||
if len(ips) == 0 && len(cnameTargets) > 0 {
|
||||
for _, target := range cnameTargets {
|
||||
// Strip trailing dot: dnsmessage CNAME targets are FQDNs with dot.
|
||||
targetHost := strings.TrimSuffix(target, ".")
|
||||
cnameIPs, err := d.lookupIPWithDepth(targetHost, server, depth+1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ips = append(ips, cnameIPs...)
|
||||
}
|
||||
}
|
||||
|
||||
// Deduplicate IPs preserving order.
|
||||
seen := make(map[string]bool)
|
||||
var unique []string
|
||||
for _, ip := range ips {
|
||||
if !seen[ip] {
|
||||
seen[ip] = true
|
||||
unique = append(unique, ip)
|
||||
}
|
||||
}
|
||||
return unique, nil
|
||||
}
|
||||
@@ -0,0 +1,346 @@
|
||||
package resolver_test
|
||||
|
||||
import (
|
||||
"encoding/binary"
|
||||
"net"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"dns-helper/resolver"
|
||||
|
||||
"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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user