package dns import ( "context" "errors" "net" "testing" "time" mdns "github.com/miekg/dns" "internetpressure.io/probe-agent/internal/protocol" ) type timeoutErr struct{} func (timeoutErr) Error() string { return "i/o timeout" } func (timeoutErr) Timeout() bool { return true } func (timeoutErr) Temporary() bool { return true } func TestClassify(t *testing.T) { cases := []struct { err error rcode string want string }{ {nil, "TIMEOUT", protocol.ErrDNSTimeout}, {nil, "SERVFAIL", protocol.ErrDNSServfail}, {nil, "NXDOMAIN", protocol.ErrDNSNxdomain}, {nil, "REFUSED", protocol.ErrDNSFail}, {timeoutErr{}, "ERROR", protocol.ErrDNSTimeout}, {context.DeadlineExceeded, "", protocol.ErrDNSTimeout}, {errors.New("boom"), "ERROR", protocol.ErrDNSFail}, } for _, c := range cases { if got := classify(c.err, c.rcode); got != c.want { t.Errorf("classify(%v,%q)=%q want %q", c.err, c.rcode, got, c.want) } } } // A tiny in-process authoritative server answers example.test with two A records and NXDOMAIN otherwise. func startServer(t *testing.T) string { pc, err := net.ListenPacket("udp", "127.0.0.1:0") if err != nil { t.Fatal(err) } srv := &mdns.Server{PacketConn: pc, Handler: mdns.HandlerFunc(func(w mdns.ResponseWriter, r *mdns.Msg) { m := new(mdns.Msg) m.SetReply(r) switch r.Question[0].Name { case "example.test.": m.Answer = append(m.Answer, &mdns.A{Hdr: mdns.RR_Header{Name: "example.test.", Rrtype: mdns.TypeA, Class: mdns.ClassINET, Ttl: 60}, A: net.ParseIP("10.0.0.2")}, &mdns.A{Hdr: mdns.RR_Header{Name: "example.test.", Rrtype: mdns.TypeA, Class: mdns.ClassINET, Ttl: 60}, A: net.ParseIP("10.0.0.1")}) case "fail.test.": m.Rcode = mdns.RcodeServerFailure case "slow.test.": return // never answers → timeout default: m.Rcode = mdns.RcodeNameError } _ = w.WriteMsg(m) })} go srv.ActivateAndServe() t.Cleanup(func() { srv.Shutdown() }) return pc.LocalAddr().String() } func TestQueryServer(t *testing.T) { addr := startServer(t) c := &Checker{} res := []protocol.Resolver{{ID: "local", Address: addr}} ms := c.Run(context.Background(), protocol.Target{TargetID: "t", Hostname: "example.test"}, res) if len(ms) != 1 { t.Fatalf("expected 1 measurement, got %d", len(ms)) } m := ms[0] if !m.OK || m.Error != "" || m.DNSRcode != "NOERROR" || m.Resolver != "local" || m.Kind != "dns" { t.Fatalf("unexpected: %+v", m) } if len(m.DNSAnswers) != 2 || m.DNSAnswers[0] != "10.0.0.1" || m.DNSAnswers[1] != "10.0.0.2" { t.Fatalf("answers not sorted: %v", m.DNSAnswers) } if m.DNSMs == nil || *m.DNSMs < 0 { t.Fatalf("dns_ms missing") } m = c.Run(context.Background(), protocol.Target{TargetID: "t", Hostname: "nope.test"}, res)[0] if m.OK || m.Error != protocol.ErrDNSNxdomain || m.DNSRcode != "NXDOMAIN" { t.Fatalf("nxdomain: %+v", m) } m = c.Run(context.Background(), protocol.Target{TargetID: "t", Hostname: "fail.test"}, res)[0] if m.OK || m.Error != protocol.ErrDNSServfail || m.DNSRcode != "SERVFAIL" { t.Fatalf("servfail: %+v", m) } ctx, cancel := context.WithTimeout(context.Background(), 700*time.Millisecond) defer cancel() m = c.Run(ctx, protocol.Target{TargetID: "t", Hostname: "slow.test"}, res)[0] if m.OK || m.Error != protocol.ErrDNSTimeout || m.DNSRcode != "TIMEOUT" { t.Fatalf("timeout: %+v", m) } }