diff options
Diffstat (limited to 'internal/dns/dns_test.go')
| -rw-r--r-- | internal/dns/dns_test.go | 472 |
1 files changed, 472 insertions, 0 deletions
diff --git a/internal/dns/dns_test.go b/internal/dns/dns_test.go new file mode 100644 index 0000000..c930c41 --- /dev/null +++ b/internal/dns/dns_test.go @@ -0,0 +1,472 @@ +package dns + +import ( + "net/netip" + "strings" + "testing" +) + +func TestFqdn(t *testing.T) { + tests := []struct { + in, out string + }{ + {"", "."}, + {".", "."}, + {"example.com", "example.com."}, + {"example.com.", "example.com."}, + } + for _, tt := range tests { + if got := Fqdn(tt.in); got != tt.out { + t.Errorf("Fqdn(%q) = %q, want %q", tt.in, got, tt.out) + } + } +} + +func TestSplitDomainName(t *testing.T) { + tests := []struct { + in string + out []string + }{ + {".", nil}, + {"example.com", []string{"example", "com"}}, + {"www.example.com.", []string{"www", "example", "com"}}, + } + for _, tt := range tests { + if got := SplitDomainName(tt.in); !stringSliceEqual(got, tt.out) { + t.Errorf("SplitDomainName(%q) = %v, want %v", tt.in, got, tt.out) + } + } +} + +func stringSliceEqual(a, b []string) bool { + if len(a) != len(b) { + return false + } + for i := range a { + if a[i] != b[i] { + return false + } + } + return true +} + +func TestNewName(t *testing.T) { + n, err := NewName(".") + if err != nil { + t.Fatal(err) + } + if n.String() != "." { + t.Errorf("root name String() = %q, want %q", n.String(), ".") + } + + n, err = NewName("www.example.com") + if err != nil { + t.Fatal(err) + } + if n.String() != "www.example.com." { + t.Errorf("String() = %q, want %q", n.String(), "www.example.com.") + } + + long := strings.Repeat("a", 64) + _, err = NewName(long + ".com") + if err == nil { + t.Error("expected error for label > 63 octets") + } + + var parts []string + for i := 0; i < 32; i++ { + parts = append(parts, strings.Repeat("a", 7)) + } + _, err = NewName(strings.Join(parts, ".")) + if err == nil { + t.Error("expected error for name > 255 octets") + } +} + +func TestNewNameRoundTrip(t *testing.T) { + names := []string{ + ".", + "com.", + "example.com.", + "www.example.com.", + "_xmpp._tcp.example.com.", + "a.b.c.d.e.f.g.", + } + for _, want := range names { + n, err := NewName(want) + if err != nil { + t.Errorf("NewName(%q): %v", want, err) + continue + } + if got := n.String(); got != want { + t.Errorf("NewName(%q).String() = %q", want, got) + } + } +} + +func TestNewNameWireRoundTrip(t *testing.T) { + names := []string{ + ".", + "com.", + "example.com.", + "www.example.com.", + } + for _, s := range names { + n, err := NewName(s) + if err != nil { + t.Fatal(err) + } + p := newPacker() + if err := p.writeName(n); err != nil { + t.Fatal(err) + } + u := &unpacker{buf: p.buf} + n2, err := u.readName() + if err != nil { + t.Fatalf("round-trip unpack %q: %v", s, err) + } + if n2.String() != s { + t.Errorf("round-trip %q: got %q", s, n2.String()) + } + } +} + +func TestHeader(t *testing.T) { + var h Header + h.SetQR(true) + if !h.QR() { + t.Error("QR should be set") + } + h.SetQR(false) + if h.QR() { + t.Error("QR should be clear") + } + h.SetOpCode(4) + if h.OpCode() != 4 { + t.Errorf("OpCode = %d, want 4", h.OpCode()) + } + h.SetAA(true) + if !h.AA() { + t.Error("AA should be set") + } + h.SetTC(true) + if !h.TC() { + t.Error("TC should be set") + } + h.SetRD(true) + if !h.RD() { + t.Error("RD should be set") + } + h.SetRA(true) + if !h.RA() { + t.Error("RA should be set") + } + h.SetAD(true) + h.SetCD(true) + if h.Flags&0x0030 != 0x0030 { + t.Error("AD+CD bits not set") + } + h.SetRCode(5) + if h.Flags&0xF != 5 { + t.Errorf("RCode = %d, want 5", h.Flags&0xF) + } +} + +func TestARoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 1, Flags: 0x0100}, + Question: []Question{{Name: mustName("example.com."), Type: TypeA, Class: ClassIN}}, + Answer: []RR{{Name: mustName("example.com."), Type: TypeA, Class: ClassIN, TTL: 300, Data: &A{Addr: netip.MustParseAddr("1.2.3.4")}}}, + } + roundTrip(t, m) +} + +func TestAAAARoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 2, Flags: 0x8000}, + Question: []Question{{Name: mustName("example.com."), Type: TypeAAAA, Class: ClassIN}}, + Answer: []RR{{Name: mustName("example.com."), Type: TypeAAAA, Class: ClassIN, TTL: 300, Data: &AAAA{Addr: netip.MustParseAddr("::1")}}}, + } + roundTrip(t, m) +} + +func TestNSRoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 3}, + Question: []Question{{Name: mustName("example.com."), Type: TypeNS, Class: ClassIN}}, + Answer: []RR{{Name: mustName("example.com."), Type: TypeNS, Class: ClassIN, TTL: 300, Data: &NS{Ns: mustName("ns1.example.com.")}}}, + } + roundTrip(t, m) +} + +func TestCNAMERoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 4}, + Question: []Question{{Name: mustName("www.example.com."), Type: TypeCNAME, Class: ClassIN}}, + Answer: []RR{{Name: mustName("www.example.com."), Type: TypeCNAME, Class: ClassIN, TTL: 300, Data: &CNAME{Cname: mustName("example.com.")}}}, + } + roundTrip(t, m) +} + +func TestPTRRoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 5}, + Question: []Question{{Name: mustName("4.3.2.1.in-addr.arpa."), Type: TypePTR, Class: ClassIN}}, + Answer: []RR{{Name: mustName("4.3.2.1.in-addr.arpa."), Type: TypePTR, Class: ClassIN, TTL: 300, Data: &PTR{Ptr: mustName("example.com.")}}}, + } + roundTrip(t, m) +} + +func TestMXRoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 6}, + Question: []Question{{Name: mustName("example.com."), Type: TypeMX, Class: ClassIN}}, + Answer: []RR{{Name: mustName("example.com."), Type: TypeMX, Class: ClassIN, TTL: 300, Data: &MX{Pref: 10, Mx: mustName("mail.example.com.")}}}, + } + roundTrip(t, m) +} + +func TestTXTRoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 7}, + Question: []Question{{Name: mustName("example.com."), Type: TypeTXT, Class: ClassIN}}, + Answer: []RR{{Name: mustName("example.com."), Type: TypeTXT, Class: ClassIN, TTL: 300, Data: &TXT{Txt: []string{"hello", "world"}}}}, + } + roundTrip(t, m) +} + +func TestSOARoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 8}, + Question: []Question{{Name: mustName("example.com."), Type: TypeSOA, Class: ClassIN}}, + Answer: []RR{ + { + Name: mustName("example.com."), Type: TypeSOA, Class: ClassIN, TTL: 300, + Data: &SOA{ + Mname: mustName("ns1.example.com."), Rname: mustName("admin.example.com."), + Serial: 2026062701, Refresh: 3600, Retry: 900, Expire: 86400, Minimum: 300, + }, + }, + }, + } + roundTrip(t, m) +} + +func TestOPTRoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 9}, + Question: []Question{{Name: mustName("."), Type: TypeA, Class: ClassIN}}, + Extra: []RR{ + { + Name: mustName("."), Type: TypeOPT, Class: 4096, TTL: 0x00_00_80_00, + Data: &OPT{ + UDPSize: 4096, DO: true, + Options: []Option{{Code: 10, Data: []byte{0x00, 0x04}}}, + }, + }, + }, + } + roundTrip(t, m) +} + +func TestUnknownRoundTrip(t *testing.T) { + m := &Msg{ + Header: Header{ID: 10}, + Question: []Question{{Name: mustName("example.com."), Type: 100, Class: ClassIN}}, + Answer: []RR{{Name: mustName("example.com."), Type: 100, Class: ClassIN, TTL: 300, Data: &Unknown{RRType: 100, Data: []byte{0x01, 0x02, 0x03}}}}, + } + packed, err := Pack(m) + if err != nil { + t.Fatal(err) + } + m2, err := Unpack(packed) + if err != nil { + t.Fatal(err) + } + if len(m2.Answer) != 1 { + t.Fatalf("expected 1 answer, got %d", len(m2.Answer)) + } + u, ok := m2.Answer[0].Data.(*Unknown) + if !ok { + t.Fatalf("Data type = %T, want *Unknown", m2.Answer[0].Data) + } + if u.RRType != 100 { + t.Errorf("RRType = %d, want 100", u.RRType) + } + if len(u.Data) != 3 || u.Data[0] != 1 || u.Data[1] != 2 || u.Data[2] != 3 { + t.Errorf("Data = %v, want [1 2 3]", u.Data) + } +} + +func TestMultiSectionRoundTrip(t *testing.T) { + var h Header + h.ID = 42 + h.SetQR(true) + h.SetAA(true) + h.SetRD(true) + h.SetRA(true) + m := &Msg{ + Header: h, + Question: []Question{{Name: mustName("example.com."), Type: TypeA, Class: ClassIN}}, + Answer: []RR{{Name: mustName("example.com."), Type: TypeA, Class: ClassIN, TTL: 300, Data: &A{Addr: netip.MustParseAddr("1.2.3.4")}}}, + Ns: []RR{{Name: mustName("example.com."), Type: TypeNS, Class: ClassIN, TTL: 600, Data: &NS{Ns: mustName("ns1.example.com.")}}}, + Extra: []RR{{Name: mustName("ns1.example.com."), Type: TypeA, Class: ClassIN, TTL: 300, Data: &A{Addr: netip.MustParseAddr("4.3.2.1")}}}, + } + wire, err := Pack(m) + if err != nil { + t.Fatal(err) + } + m2, err := Unpack(wire) + if err != nil { + t.Fatal(err) + } + if m2.Header.ID != 42 { + t.Errorf("ID = %d, want 42", m2.Header.ID) + } + if !m2.Header.QR() { + t.Error("QR not set") + } + if !m2.Header.AA() { + t.Error("AA not set") + } + if !m2.Header.RD() { + t.Error("RD not set") + } + if !m2.Header.RA() { + t.Error("RA not set") + } + if len(m2.Question) != 1 || len(m2.Answer) != 1 || len(m2.Ns) != 1 || len(m2.Extra) != 1 { + t.Errorf("section lengths: Q=%d A=%d NS=%d Extra=%d", len(m2.Question), len(m2.Answer), len(m2.Ns), len(m2.Extra)) + } +} + +func TestNameCompression(t *testing.T) { + m := &Msg{ + Header: Header{ID: 11}, + Question: []Question{{Name: mustName("example.com."), Type: TypeA, Class: ClassIN}}, + Answer: []RR{ + {Name: mustName("example.com."), Type: TypeA, Class: ClassIN, TTL: 300, Data: &A{Addr: netip.MustParseAddr("1.2.3.4")}}, + {Name: mustName("example.com."), Type: TypeAAAA, Class: ClassIN, TTL: 300, Data: &AAAA{Addr: netip.MustParseAddr("::1")}}, + }, + } + wire, err := Pack(m) + if err != nil { + t.Fatal(err) + } + m2, err := Unpack(wire) + if err != nil { + t.Fatal(err) + } + if len(m2.Answer) != 2 { + t.Fatalf("expected 2 answers, got %d", len(m2.Answer)) + } + if m2.Answer[1].Name.String() != "example.com." { + t.Errorf("second answer name = %q, want %q", m2.Answer[1].Name.String(), "example.com.") + } +} + +func TestUnpackShortBuffer(t *testing.T) { + _, err := Unpack([]byte{0, 0, 0, 0, 0, 0}) + if err == nil { + t.Error("expected error for short buffer") + } +} + +func TestUnpackTruncatedMessage(t *testing.T) { + wire := make([]byte, 12) + wire[4] = 0 + wire[5] = 1 // QDCOUNT = 1 + _, err := Unpack(wire) + if err == nil { + t.Error("expected error for truncated message with missing question") + } +} + +func TestNameCompressionLoop(t *testing.T) { + wire := []byte{0xC0, 0x00} + u := &unpacker{buf: wire} + _, err := u.readName() + if err == nil { + t.Error("expected error for compression loop") + } +} + +func TestReservedLabelType(t *testing.T) { + wire := []byte{0x40, 0x00} + u := &unpacker{buf: wire} + _, err := u.readName() + if err == nil { + t.Error("expected error for reserved label type 0x40") + } +} + +func mustName(s string) Name { + n, err := NewName(s) + if err != nil { + panic(err) + } + return n +} + +func roundTrip(t *testing.T, m *Msg) { + t.Helper() + wire, err := Pack(m) + if err != nil { + t.Fatal(err) + } + m2, err := Unpack(wire) + if err != nil { + t.Fatalf("Unpack: %v\nwire: %x", err, wire) + } + if len(m2.Question) != len(m.Question) { + t.Errorf("Question count: %d vs %d", len(m2.Question), len(m.Question)) + } + if len(m2.Answer) != len(m.Answer) { + t.Errorf("Answer count: %d vs %d", len(m2.Answer), len(m.Answer)) + } + if len(m2.Ns) != len(m.Ns) { + t.Errorf("Ns count: %d vs %d", len(m2.Ns), len(m.Ns)) + } + if len(m2.Extra) != len(m.Extra) { + t.Errorf("Extra count: %d vs %d", len(m2.Extra), len(m.Extra)) + } + for i := range m.Answer { + if m2.Answer[i].Name.String() != m.Answer[i].Name.String() { + t.Errorf("Answer[%d] name: %q vs %q", i, m2.Answer[i].Name.String(), m.Answer[i].Name.String()) + } + if m2.Answer[i].Type != m.Answer[i].Type { + t.Errorf("Answer[%d] type: %d vs %d", i, m2.Answer[i].Type, m.Answer[i].Type) + } + if m2.Answer[i].Class != m.Answer[i].Class { + t.Errorf("Answer[%d] class: %d vs %d", i, m2.Answer[i].Class, m.Answer[i].Class) + } + if m2.Answer[i].TTL != m.Answer[i].TTL { + t.Errorf("Answer[%d] TTL: %d vs %d", i, m2.Answer[i].TTL, m.Answer[i].TTL) + } + } +} + +func BenchmarkPack(b *testing.B) { + m := &Msg{ + Header: Header{ID: 1}, + Question: []Question{{Name: mustName("example.com."), Type: TypeA, Class: ClassIN}}, + Answer: []RR{{Name: mustName("example.com."), Type: TypeA, Class: ClassIN, TTL: 300, Data: &A{Addr: netip.MustParseAddr("1.2.3.4")}}}, + } + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, _ = Pack(m) + } +} + +func BenchmarkUnpack(b *testing.B) { + m := &Msg{ + Header: Header{ID: 1}, + Question: []Question{{Name: mustName("example.com."), Type: TypeA, Class: ClassIN}}, + Answer: []RR{{Name: mustName("example.com."), Type: TypeA, Class: ClassIN, TTL: 300, Data: &A{Addr: netip.MustParseAddr("1.2.3.4")}}}, + } + wire, _ := Pack(m) + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, _ = Unpack(wire) + } +} |
