- 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.
176 lines
4.7 KiB
Go
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
|
|
}
|