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) } }