Files
dnshelper/resolver/resolver.go
T
dyoder f0ce5a4042 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.
2026-03-03 18:37:02 -05:00

176 lines
4.7 KiB
Go

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
}