// Package dns implements the "dns" check: one A query per configured resolver, run sequentially. // // Resolver "system" (empty address) uses the OS resolver through net.DefaultResolver; every other resolver is // queried directly over UDP with github.com/miekg/dns (EDNS0, 3 s timeout), retried once over TCP only when the // UDP answer is truncated. Exactly one query per resolver per check. package dns import ( "context" "errors" "net" "sort" "strings" "time" mdns "github.com/miekg/dns" "internetpressure.io/probe-agent/internal/protocol" ) // Timeout per query. const Timeout = 3 * time.Second // DefaultResolvers is used by `once` and when the server sends none. var DefaultResolvers = []protocol.Resolver{ {ID: "system", Address: ""}, {ID: "google", Address: "8.8.8.8:53"}, {ID: "cloudflare", Address: "1.1.1.1:53"}, {ID: "quad9", Address: "9.9.9.9:53"}, } // Checker runs dns checks. type Checker struct{} // Run queries every resolver for the A record of t.Hostname and returns one measurement per resolver. func (c *Checker) Run(ctx context.Context, t protocol.Target, resolvers []protocol.Resolver) []protocol.Measurement { if len(resolvers) == 0 { resolvers = DefaultResolvers } out := make([]protocol.Measurement, 0, len(resolvers)) for _, r := range resolvers { if ctx.Err() != nil { break } out = append(out, c.query(ctx, t, r)) } return out } func (c *Checker) query(ctx context.Context, t protocol.Target, r protocol.Resolver) protocol.Measurement { m := protocol.Measurement{TS: protocol.FormatTime(time.Now()), TargetID: t.TargetID, Kind: "dns", Resolver: r.ID} ctx, cancel := context.WithTimeout(ctx, Timeout) defer cancel() start := time.Now() var answers []string var rcode string var err error if r.Address == "" || r.ID == "system" { answers, rcode, err = querySystem(ctx, t.Hostname) } else { answers, rcode, err = queryServer(ctx, t.Hostname, r.Address) } m.DNSMs = protocol.F(protocol.Ms(time.Since(start))) m.DNSRcode = rcode sort.Strings(answers) m.DNSAnswers = answers switch { case err != nil: m.Error = classify(err, rcode) case rcode == "NOERROR" && len(answers) > 0: m.OK = true case rcode == "NOERROR": m.Error = protocol.ErrDNSFail // NOERROR but no A record (e.g. AAAA-only or CNAME chain without A) default: m.Error = classify(nil, rcode) } return m } // querySystem resolves through the OS stub resolver. The rcode is inferred from the error class. func querySystem(ctx context.Context, host string) ([]string, string, error) { ips, err := net.DefaultResolver.LookupIP(ctx, "ip4", host) if err != nil { var dnsErr *net.DNSError if errors.As(err, &dnsErr) { switch { case dnsErr.IsNotFound: return nil, "NXDOMAIN", err case dnsErr.IsTimeout || errors.Is(err, context.DeadlineExceeded): return nil, "TIMEOUT", err default: return nil, "SERVFAIL", err } } if errors.Is(err, context.DeadlineExceeded) { return nil, "TIMEOUT", err } return nil, "ERROR", err } out := make([]string, 0, len(ips)) for _, ip := range ips { if v4 := ip.To4(); v4 != nil { out = append(out, v4.String()) } } return out, "NOERROR", nil } // queryServer sends one A query over UDP (EDNS0 1232) and retries once over TCP if truncated. func queryServer(ctx context.Context, host, server string) ([]string, string, error) { msg := new(mdns.Msg) msg.SetQuestion(mdns.Fqdn(host), mdns.TypeA) msg.RecursionDesired = true msg.SetEdns0(1232, false) client := &mdns.Client{Net: "udp", Timeout: Timeout, UDPSize: 1232} resp, _, err := client.ExchangeContext(ctx, msg, server) if err == nil && resp != nil && resp.Truncated { tcp := &mdns.Client{Net: "tcp", Timeout: Timeout} resp, _, err = tcp.ExchangeContext(ctx, msg, server) } if err != nil { var ne net.Error if errors.Is(err, context.DeadlineExceeded) || (errors.As(err, &ne) && ne.Timeout()) || strings.Contains(err.Error(), "i/o timeout") { return nil, "TIMEOUT", err } return nil, "ERROR", err } rcode := mdns.RcodeToString[resp.Rcode] if rcode == "" { rcode = "RCODE" + itoa(resp.Rcode) } var out []string for _, rr := range resp.Answer { if a, ok := rr.(*mdns.A); ok { out = append(out, a.A.String()) } } return out, rcode, nil } // classify maps (transport error, rcode) to a protocol error code. func classify(err error, rcode string) string { switch rcode { case "TIMEOUT": return protocol.ErrDNSTimeout case "SERVFAIL": return protocol.ErrDNSServfail case "NXDOMAIN": return protocol.ErrDNSNxdomain } if err != nil { var ne net.Error if errors.Is(err, context.DeadlineExceeded) || (errors.As(err, &ne) && ne.Timeout()) { return protocol.ErrDNSTimeout } } return protocol.ErrDNSFail } func itoa(n int) string { if n == 0 { return "0" } var b [20]byte i := len(b) for n > 0 { i-- b[i] = byte('0' + n%10) n /= 10 } return string(b[i:]) }