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:
2026-03-03 18:37:02 -05:00
parent 5e28d0bd8c
commit f0ce5a4042
30 changed files with 2847 additions and 473 deletions
+175
View File
@@ -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
}
+346
View File
@@ -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)
}
}