spb/internetpressure
Public
TypeScript 36.3%
Python 31.8%
Go 18%
JavaScript 9.8%
Shell 1.9%
SQL 1.4%
CSS 0.5%
1// Package dns implements the "dns" check: one A query per configured resolver, run sequentially.2//3// Resolver "system" (empty address) uses the OS resolver through net.DefaultResolver; every other resolver is4// queried directly over UDP with github.com/miekg/dns (EDNS0, 3 s timeout), retried once over TCP only when the5// UDP answer is truncated. Exactly one query per resolver per check.6package dns78import (9 "context"10 "errors"11 "net"12 "sort"13 "strings"14 "time"1516 mdns "github.com/miekg/dns"1718 "internetpressure.io/probe-agent/internal/protocol"19)2021// Timeout per query.22const Timeout = 3 * time.Second2324// DefaultResolvers is used by `once` and when the server sends none.25var DefaultResolvers = []protocol.Resolver{26 {ID: "system", Address: ""},27 {ID: "google", Address: "8.8.8.8:53"},28 {ID: "cloudflare", Address: "1.1.1.1:53"},29 {ID: "quad9", Address: "9.9.9.9:53"},30}3132// Checker runs dns checks.33type Checker struct{}3435// Run queries every resolver for the A record of t.Hostname and returns one measurement per resolver.36func (c *Checker) Run(ctx context.Context, t protocol.Target, resolvers []protocol.Resolver) []protocol.Measurement {37 if len(resolvers) == 0 {38 resolvers = DefaultResolvers39 }40 out := make([]protocol.Measurement, 0, len(resolvers))41 for _, r := range resolvers {42 if ctx.Err() != nil {43 break44 }45 out = append(out, c.query(ctx, t, r))46 }47 return out48}4950func (c *Checker) query(ctx context.Context, t protocol.Target, r protocol.Resolver) protocol.Measurement {51 m := protocol.Measurement{TS: protocol.FormatTime(time.Now()), TargetID: t.TargetID, Kind: "dns", Resolver: r.ID}52 ctx, cancel := context.WithTimeout(ctx, Timeout)53 defer cancel()5455 start := time.Now()56 var answers []string57 var rcode string58 var err error59 if r.Address == "" || r.ID == "system" {60 answers, rcode, err = querySystem(ctx, t.Hostname)61 } else {62 answers, rcode, err = queryServer(ctx, t.Hostname, r.Address)63 }64 m.DNSMs = protocol.F(protocol.Ms(time.Since(start)))65 m.DNSRcode = rcode66 sort.Strings(answers)67 m.DNSAnswers = answers6869 switch {70 case err != nil:71 m.Error = classify(err, rcode)72 case rcode == "NOERROR" && len(answers) > 0:73 m.OK = true74 case rcode == "NOERROR":75 m.Error = protocol.ErrDNSFail // NOERROR but no A record (e.g. AAAA-only or CNAME chain without A)76 default:77 m.Error = classify(nil, rcode)78 }79 return m80}8182// querySystem resolves through the OS stub resolver. The rcode is inferred from the error class.83func querySystem(ctx context.Context, host string) ([]string, string, error) {84 ips, err := net.DefaultResolver.LookupIP(ctx, "ip4", host)85 if err != nil {86 var dnsErr *net.DNSError87 if errors.As(err, &dnsErr) {88 switch {89 case dnsErr.IsNotFound:90 return nil, "NXDOMAIN", err91 case dnsErr.IsTimeout || errors.Is(err, context.DeadlineExceeded):92 return nil, "TIMEOUT", err93 default:94 return nil, "SERVFAIL", err95 }96 }97 if errors.Is(err, context.DeadlineExceeded) {98 return nil, "TIMEOUT", err99 }100 return nil, "ERROR", err101 }102 out := make([]string, 0, len(ips))103 for _, ip := range ips {104 if v4 := ip.To4(); v4 != nil {105 out = append(out, v4.String())106 }107 }108 return out, "NOERROR", nil109}110111// queryServer sends one A query over UDP (EDNS0 1232) and retries once over TCP if truncated.112func queryServer(ctx context.Context, host, server string) ([]string, string, error) {113 msg := new(mdns.Msg)114 msg.SetQuestion(mdns.Fqdn(host), mdns.TypeA)115 msg.RecursionDesired = true116 msg.SetEdns0(1232, false)117118 client := &mdns.Client{Net: "udp", Timeout: Timeout, UDPSize: 1232}119 resp, _, err := client.ExchangeContext(ctx, msg, server)120 if err == nil && resp != nil && resp.Truncated {121 tcp := &mdns.Client{Net: "tcp", Timeout: Timeout}122 resp, _, err = tcp.ExchangeContext(ctx, msg, server)123 }124 if err != nil {125 var ne net.Error126 if errors.Is(err, context.DeadlineExceeded) || (errors.As(err, &ne) && ne.Timeout()) ||127 strings.Contains(err.Error(), "i/o timeout") {128 return nil, "TIMEOUT", err129 }130 return nil, "ERROR", err131 }132 rcode := mdns.RcodeToString[resp.Rcode]133 if rcode == "" {134 rcode = "RCODE" + itoa(resp.Rcode)135 }136 var out []string137 for _, rr := range resp.Answer {138 if a, ok := rr.(*mdns.A); ok {139 out = append(out, a.A.String())140 }141 }142 return out, rcode, nil143}144145// classify maps (transport error, rcode) to a protocol error code.146func classify(err error, rcode string) string {147 switch rcode {148 case "TIMEOUT":149 return protocol.ErrDNSTimeout150 case "SERVFAIL":151 return protocol.ErrDNSServfail152 case "NXDOMAIN":153 return protocol.ErrDNSNxdomain154 }155 if err != nil {156 var ne net.Error157 if errors.Is(err, context.DeadlineExceeded) || (errors.As(err, &ne) && ne.Timeout()) {158 return protocol.ErrDNSTimeout159 }160 }161 return protocol.ErrDNSFail162}163164func itoa(n int) string {165 if n == 0 {166 return "0"167 }168 var b [20]byte169 i := len(b)170 for n > 0 {171 i--172 b[i] = byte('0' + n%10)173 n /= 10174 }175 return string(b[i:])176}177