package dns import ( "testing" "time" "github.com/miekg/dns" ) func TestDecodeResponseREFUSED(t *testing.T) { msg := newTestMsg(dns.RcodeRefused) d := DecodeResponse(msg) if d.Classification != ResponseREFUSED { t.Errorf("classification = %v, want ResponseREFUSED", d.Classification) } if d.RcodeName != "REFUSED" { t.Errorf("RcodeName = %q, want REFUSED", d.RcodeName) } } func TestDecodeResponseNOTIMPL(t *testing.T) { msg := newTestMsg(dns.RcodeNotImplemented) d := DecodeResponse(msg) if d.Classification != ResponseNOTIMPL { t.Errorf("classification = %v, want ResponseNOTIMPL", d.Classification) } if d.RcodeName != "NOTIMP" { t.Errorf("RcodeName = %q, want NOTIMP", d.RcodeName) } } func TestDecodeResponseREFUSEDString(t *testing.T) { if got := ResponseREFUSED.String(); got != "refused" { t.Errorf("ResponseREFUSED.String() = %q, want \"refused\"", got) } if got := ResponseNOTIMPL.String(); got != "notimp" { t.Errorf("ResponseNOTIMPL.String() = %q, want \"notimp\"", got) } } func TestExtractDNAMEMappings(t *testing.T) { t.Run("no DNAME", func(t *testing.T) { msg := new(dns.Msg) msg.Answer = append(msg.Answer, &dns.A{ Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeA}, A: MustParseIP("1.2.3.4"), }) d := DecodeResponse(msg) if len(d.DNAMEMappings) != 0 { t.Errorf("expected 0 DNAME mappings, got %d", len(d.DNAMEMappings)) } }) t.Run("DNAME in answer", func(t *testing.T) { msg := new(dns.Msg) msg.SetReply(new(dns.Msg)) msg.Answer = append(msg.Answer, &dns.DNAME{ Hdr: dns.RR_Header{Name: "example.com.", Rrtype: dns.TypeDNAME, Class: dns.ClassINET, Ttl: 300}, Target: "example.net.", }) d := DecodeResponse(msg) if len(d.DNAMEMappings) != 1 { t.Fatalf("expected 1 DNAME mapping, got %d", len(d.DNAMEMappings)) } if d.DNAMEMappings[0].Owner != "example.com." { t.Errorf("Owner = %q, want %q", d.DNAMEMappings[0].Owner, "example.com.") } if d.DNAMEMappings[0].Target != "example.net." { t.Errorf("Target = %q, want %q", d.DNAMEMappings[0].Target, "example.net.") } }) } func TestSynthesizeCNAMEFromDNAME(t *testing.T) { tests := []struct { queryName string dnameOwner string dnameTarget string want string }{ { queryName: "foo.example.com.", dnameOwner: "example.com.", dnameTarget: "example.net.", want: "foo.example.net.", }, { queryName: "bar.foo.example.com.", dnameOwner: "example.com.", dnameTarget: "example.net.", want: "bar.foo.example.net.", }, { // Owner itself is not redirected queryName: "example.com.", dnameOwner: "example.com.", dnameTarget: "example.net.", want: "", }, { // Not a subdomain queryName: "other.com.", dnameOwner: "example.com.", dnameTarget: "example.net.", want: "", }, } for _, tt := range tests { got := SynthesizeCNAMEFromDNAME(tt.queryName, tt.dnameOwner, tt.dnameTarget) if got != tt.want { t.Errorf("SynthesizeCNAMEFromDNAME(%q, %q, %q) = %q, want %q", tt.queryName, tt.dnameOwner, tt.dnameTarget, got, tt.want) } } } func TestBackoffDelay(t *testing.T) { tests := []struct { attempt int want time.Duration }{ {0, 0}, {1, 100 * time.Millisecond}, {2, 200 * time.Millisecond}, {3, 400 * time.Millisecond}, {4, 800 * time.Millisecond}, {5, 1600 * time.Millisecond}, {6, 2000 * time.Millisecond}, // capped at 2s {10, 2000 * time.Millisecond}, // still capped } for _, tt := range tests { got := backoffDelay(tt.attempt) if got != tt.want { t.Errorf("backoffDelay(%d) = %v, want %v", tt.attempt, got, tt.want) } } }