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 }