diff options
Diffstat (limited to 'internal/dns')
| -rw-r--r-- | internal/dns/dns_test.go | 472 | ||||
| -rw-r--r-- | internal/dns/fuzz_test.go | 52 | ||||
| -rw-r--r-- | internal/dns/header.go | 28 | ||||
| -rw-r--r-- | internal/dns/message.go | 9 | ||||
| -rw-r--r-- | internal/dns/name.go | 2 | ||||
| -rw-r--r-- | internal/dns/pack.go | 103 | ||||
| -rw-r--r-- | internal/dns/question.go | 23 | ||||
| -rw-r--r-- | internal/dns/rr.go | 110 | ||||
| -rw-r--r-- | internal/dns/rr_a.go | 34 | ||||
| -rw-r--r-- | internal/dns/rr_aaaa.go | 34 | ||||
| -rw-r--r-- | internal/dns/rr_cname.go | 14 | ||||
| -rw-r--r-- | internal/dns/rr_mx.go | 30 | ||||
| -rw-r--r-- | internal/dns/rr_ns.go | 14 | ||||
| -rw-r--r-- | internal/dns/rr_opt.go | 57 | ||||
| -rw-r--r-- | internal/dns/rr_ptr.go | 14 | ||||
| -rw-r--r-- | internal/dns/rr_soa.go | 65 | ||||
| -rw-r--r-- | internal/dns/rr_txt.go | 43 | ||||
| -rw-r--r-- | internal/dns/rr_unknown.go | 34 | ||||
| -rw-r--r-- | internal/dns/types.go | 12 | ||||
| -rw-r--r-- | internal/dns/unpack.go | 138 |
20 files changed, 1267 insertions, 21 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) + } +} diff --git a/internal/dns/fuzz_test.go b/internal/dns/fuzz_test.go new file mode 100644 index 0000000..8aa1653 --- /dev/null +++ b/internal/dns/fuzz_test.go @@ -0,0 +1,52 @@ +package dns + +import ( + "testing" +) + +func FuzzUnpack(f *testing.F) { + seed := []byte{ + 0x00, 0x01, // ID + 0x01, 0x00, // flags: RD=1 + 0x00, 0x01, // QDCOUNT = 1 + 0x00, 0x00, // ANCOUNT = 0 + 0x00, 0x00, // NSCOUNT = 0 + 0x00, 0x00, // ARCOUNT = 0 + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', + 0x03, 'c', 'o', 'm', + 0x00, // root label + 0x00, 0x01, // QTYPE = A + 0x00, 0x01, // QCLASS = IN + } + f.Add(seed) + + seedResp := []byte{ + 0x00, 0x01, // ID + 0x85, 0x80, // flags: QR=1, AA=1, RD=1, RA=1 + 0x00, 0x01, // QDCOUNT = 1 + 0x00, 0x01, // ANCOUNT = 1 + 0x00, 0x00, // NSCOUNT = 0 + 0x00, 0x00, // ARCOUNT = 0 + 0x03, 'w', 'w', 'w', + 0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', + 0x03, 'c', 'o', 'm', + 0x00, + 0x00, 0x01, // TYPE = A + 0x00, 0x01, // CLASS = IN + 0x00, 0x00, 0x00, 0x3C, // TTL = 60 + 0x00, 0x04, // RDLENGTH = 4 + 0x01, 0x02, 0x03, 0x04, // 1.2.3.4 + } + f.Add(seedResp) + + f.Fuzz(func(t *testing.T, data []byte) { + m, err := Unpack(data) + if err != nil { + return // error is fine; panic is not + } + _, err = Pack(m) + if err != nil { + t.Errorf("Pack after Unpack: %v", err) + } + }) +} diff --git a/internal/dns/header.go b/internal/dns/header.go index 1535500..3b51f67 100644 --- a/internal/dns/header.go +++ b/internal/dns/header.go @@ -36,22 +36,22 @@ func (h *Header) RA() bool { return h.Flags&0x0080 != 0 } // Mutator Flags func (h *Header) SetQR(v bool) { h.set(0x8000, v) } func (h *Header) SetOpCode(v uint8) { - h.Bits = (h.Bits &^ (0xF << 11)) | + h.Flags = (h.Flags &^ (0xF << 11)) | ((uint16(v) & 0xF) << 11) } -func (h *Header) SetAA(v bool) { h.set(0x0400, v) } -func (h *Header) SetTC(v bool) { h.set(0x0200, v) } -func (h *Header) SetRD(v bool) { h.set(0x0100, v) } -func (h *Header) SetRA(v bool) { h.set(0x0080, v) } -func (h *Header) SetZ(v uint8) { h.Bits = (h.Bits &^ 0xF) | (uint16(v) & 0xF) << 4) } -func (h *Header) SetAD(v bool) { h.set(0x0020, v) } -func (h *Header) SetCD(v bool) { h.set(0x0010, v) } -func (h *Header) SetRCode(v uint8) { h.Bits = (h.Bits &^ 0xF) | (uint16(v) & 0xF) } +func (h *Header) SetAA(v bool) { h.set(0x0400, v) } +func (h *Header) SetTC(v bool) { h.set(0x0200, v) } +func (h *Header) SetRD(v bool) { h.set(0x0100, v) } +func (h *Header) SetRA(v bool) { h.set(0x0080, v) } +func (h *Header) SetZ(v uint8) { h.Flags = (h.Flags &^ (0x7 << 4)) | ((uint16(v) & 0x7) << 4) } +func (h *Header) SetAD(v bool) { h.set(0x0020, v) } +func (h *Header) SetCD(v bool) { h.set(0x0010, v) } +func (h *Header) SetRCode(v uint8) { h.Flags = (h.Flags &^ 0xF) | (uint16(v) & 0xF) } func (h *Header) set(mask uint16, v bool) { - if v { - h.Bits |= mask - } else { - h.Bits &^= mask - } + if v { + h.Flags |= mask + } else { + h.Flags &^= mask + } } diff --git a/internal/dns/message.go b/internal/dns/message.go new file mode 100644 index 0000000..5e8c5ba --- /dev/null +++ b/internal/dns/message.go @@ -0,0 +1,9 @@ +package dns + +type Msg struct { + Header Header + Question []Question + Answer []RR + Ns []RR + Extra []RR +} diff --git a/internal/dns/name.go b/internal/dns/name.go index 799e53a..282d72f 100644 --- a/internal/dns/name.go +++ b/internal/dns/name.go @@ -59,7 +59,7 @@ func (n Name) String() string { break } labels = append(labels, string(n.Data[off+1:off+1+l])) - off += 1 + 1 + off += 1 + l } return strings.Join(labels, ".") + "." } diff --git a/internal/dns/pack.go b/internal/dns/pack.go new file mode 100644 index 0000000..dddc2dd --- /dev/null +++ b/internal/dns/pack.go @@ -0,0 +1,103 @@ +package dns + +import ( + "strings" +) + +type packer struct { + buf []byte + names map[string]int +} + +func newPacker() *packer { + return &packer{ + names: make(map[string]int), + } +} + +func (p *packer) writeUint16(v uint16) { + p.buf = append(p.buf, byte(v>>8), byte(v)) +} + +func (p *packer) writeUint32(v uint32) { + p.buf = append(p.buf, byte(v>>24), byte(v>>16), byte(v>>8), byte(v)) +} + +func (p *packer) writeName(n Name) error { + labels := SplitDomainName(n.String()) + for i := 0; i < len(labels); i++ { + suffix := Fqdn(joinLabels(labels[i:])) + if off, ok := p.names[suffix]; ok { + for _, label := range labels[:i] { + p.buf = append(p.buf, byte(len(label))) + p.buf = append(p.buf, label...) + } + p.writeUint16(uint16(off) | 0xC000) + return nil + } + } + + off := len(p.buf) + p.buf = append(p.buf, n.Data...) + + for i := 0; i < len(labels); i++ { + suffix := Fqdn(joinLabels(labels[i:])) + suffixOff := off + for j := 0; j < i; j++ { + suffixOff += 1 + len(labels[j]) + } + if _, exists := p.names[suffix]; !exists { + p.names[suffix] = suffixOff + } + } + return nil +} + +func joinLabels(labels []string) string { + if len(labels) == 0 { + return "." + } + return strings.Join(labels, ".") + "." +} + +func Pack(m *Msg) ([]byte, error) { + p := newPacker() + + headerOff := len(p.buf) + p.buf = append(p.buf, make([]byte, 12)...) + + for _, q := range m.Question { + if err := q.pack(p); err != nil { + return nil, err + } + } + for _, rr := range m.Answer { + if err := rr.pack(p); err != nil { + return nil, err + } + } + for _, rr := range m.Ns { + if err := rr.pack(p); err != nil { + return nil, err + } + } + for _, rr := range m.Extra { + if err := rr.pack(p); err != nil { + return nil, err + } + } + + putUint16(p.buf[headerOff:], m.Header.ID) + putUint16(p.buf[headerOff+2:], m.Header.Flags) + putUint16(p.buf[headerOff+4:], uint16(len(m.Question))) + putUint16(p.buf[headerOff+6:], uint16(len(m.Answer))) + putUint16(p.buf[headerOff+8:], uint16(len(m.Ns))) + putUint16(p.buf[headerOff+10:], uint16(len(m.Extra))) + + return p.buf, nil +} + +func putUint16(b []byte, v uint16) { + b[0] = byte(v >> 8) + b[1] = byte(v) +} diff --git a/internal/dns/question.go b/internal/dns/question.go index 101cb0a..489082c 100644 --- a/internal/dns/question.go +++ b/internal/dns/question.go @@ -1,7 +1,30 @@ package dns +// https://datatracker.ietf.org/doc/html/rfc1035#section-4.1.2 type Question struct { Name Name Type uint16 Class uint16 } + +func (q *Question) pack(p *packer) error { + if err := p.writeName(q.Name); err != nil { + return err + } + p.writeUint16(q.Type) + p.writeUint16(q.Class) + return nil +} +func (q *Question) unpack(u *unpacker) error { + var err error + q.Name, err = u.readName() + if err != nil { + return err + } + q.Type, err = u.readUint16() + if err != nil { + return err + } + q.Class, err = u.readUint16() + return err +} diff --git a/internal/dns/rr.go b/internal/dns/rr.go new file mode 100644 index 0000000..38f467c --- /dev/null +++ b/internal/dns/rr.go @@ -0,0 +1,110 @@ +package dns + +import "fmt" + +type RData interface { + Type() uint16 + Pack(p *packer) error + Unpack(u *unpacker, length uint16) error + String() string +} + +type RR struct { + Name Name + Type uint16 + Class uint16 + TTL uint32 + Data RData +} + +func unpackRR(u *unpacker) (RR, error) { + var rr RR + var err error + rr.Name, err = u.readName() + if err != nil { + return rr, err + } + rr.Type, err = u.readUint16() + if err != nil { + return rr, err + } + rr.Class, err = u.readUint16() + if err != nil { + return rr, err + } + rr.TTL, err = u.readUint32() + if err != nil { + return rr, err + } + rdlength, err := u.readUint16() + if err != nil { + return rr, err + } + + rdStart := u.off + if u.off+int(rdlength) > len(u.buf) { + return rr, fmt.Errorf("dns: short rdata") + } + + switch rr.Type { + case TypeA: + rr.Data = &A{} + case TypeNS: + rr.Data = &NS{} + case TypeCNAME: + rr.Data = &CNAME{} + case TypeSOA: + rr.Data = &SOA{} + case TypePTR: + rr.Data = &PTR{} + case TypeMX: + rr.Data = &MX{} + case TypeTXT: + rr.Data = &TXT{} + case TypeAAAA: + rr.Data = &AAAA{} + case TypeOPT: + opt := &OPT{ + UDPSize: rr.Class, + } + opt.ExtRCode = uint8(rr.TTL >> 24) + opt.Version = uint8(rr.TTL >> 16) + // https://datatracker.ietf.org/doc/html/rfc3225 + // https://datatracker.ietf.org/doc/html/rfc4035 + opt.DO = rr.TTL&0x8000 != 0 + opt.Z = uint16(rr.TTL & 0x7FFF) + rr.Data = opt + default: + rr.Data = &Unknown{RRType: rr.Type} + } + + if err := rr.Data.Unpack(u, rdlength); err != nil { + return rr, err + } + if u.off != rdStart+int(rdlength) { + return rr, fmt.Errorf("dns: rdata length mismatch") + } + return rr, nil +} + +func (rr *RR) pack(p *packer) error { + if err := p.writeName(rr.Name); err != nil { + return err + } + p.writeUint16(rr.Type) + p.writeUint16(rr.Class) + p.writeUint32(rr.TTL) + + rdlenOff := len(p.buf) + p.buf = append(p.buf, 0, 0) + rdStart := len(p.buf) + + if err := rr.Data.Pack(p); err != nil { + return err + } + + rdlen := uint16(len(p.buf) - rdStart) + p.buf[rdlenOff] = byte(rdlen >> 8) + p.buf[rdlenOff+1] = byte(rdlen) + return nil +} diff --git a/internal/dns/rr_a.go b/internal/dns/rr_a.go new file mode 100644 index 0000000..3886ed9 --- /dev/null +++ b/internal/dns/rr_a.go @@ -0,0 +1,34 @@ +package dns + +import ( + "fmt" + "net/netip" +) + +type A struct { + Addr netip.Addr +} + +func (a *A) Type() uint16 { + return TypeA +} + +func (a *A) Pack(p *packer) error { + b := a.Addr.As4() + p.buf = append(p.buf, b[:]...) + return nil +} + +func (a *A) Unpack(u *unpacker, length uint16) error { + if length != 4 { + return fmt.Errorf("dns: A rdata length %d != 4", length) + } + if u.off+4 > len(u.buf) { + return fmt.Errorf("dns: short A rdata") + } + a.Addr = netip.AddrFrom4([4]byte(u.buf[u.off : u.off+4])) + u.off += 4 + return nil +} + +func (a *A) String() string { return a.Addr.String() } diff --git a/internal/dns/rr_aaaa.go b/internal/dns/rr_aaaa.go new file mode 100644 index 0000000..569c0cd --- /dev/null +++ b/internal/dns/rr_aaaa.go @@ -0,0 +1,34 @@ +package dns + +import ( + "fmt" + "net/netip" +) + +type AAAA struct { + Addr netip.Addr +} + +func (a *AAAA) Type() uint16 { + return TypeAAAA +} + +func (a *AAAA) Pack(p *packer) error { + b := a.Addr.As16() + p.buf = append(p.buf, b[:]...) + return nil +} + +func (a *AAAA) Unpack(u *unpacker, length uint16) error { + if length != 16 { + return fmt.Errorf("dns: AAAA rdata length %d != 16", length) + } + if u.off+16 > len(u.buf) { + return fmt.Errorf("dns: short AAAA rdata") + } + a.Addr = netip.AddrFrom16([16]byte(u.buf[u.off : u.off+16])) + u.off += 16 + return nil +} + +func (a *AAAA) String() string { return a.Addr.String() } diff --git a/internal/dns/rr_cname.go b/internal/dns/rr_cname.go new file mode 100644 index 0000000..be19626 --- /dev/null +++ b/internal/dns/rr_cname.go @@ -0,0 +1,14 @@ +package dns + +type CNAME struct { + Cname Name +} + +func (c *CNAME) Type() uint16 { return TypeCNAME } +func (c *CNAME) Pack(p *packer) error { return p.writeName(c.Cname) } +func (c *CNAME) Unpack(u *unpacker, length uint16) error { + var err error + c.Cname, err = u.readName() + return err +} +func (c *CNAME) String() string { return c.Cname.String() } diff --git a/internal/dns/rr_mx.go b/internal/dns/rr_mx.go new file mode 100644 index 0000000..2f49f86 --- /dev/null +++ b/internal/dns/rr_mx.go @@ -0,0 +1,30 @@ +package dns + +import "fmt" + +// https://datatracker.ietf.org/doc/html/rfc1035#section-3.3.9 +type MX struct { + Pref uint16 + Mx Name +} + +func (mx *MX) Type() uint16 { return TypeMX } + +func (mx *MX) Pack(p *packer) error { + p.writeUint16(mx.Pref) + return p.writeName(mx.Mx) +} + +func (mx *MX) Unpack(u *unpacker, length uint16) error { + var err error + mx.Pref, err = u.readUint16() + if err != nil { + return err + } + mx.Mx, err = u.readName() + return err +} + +func (mx *MX) String() string { + return fmt.Sprintf("%d %s", mx.Pref, mx.Mx.String()) +} diff --git a/internal/dns/rr_ns.go b/internal/dns/rr_ns.go new file mode 100644 index 0000000..aa21619 --- /dev/null +++ b/internal/dns/rr_ns.go @@ -0,0 +1,14 @@ +package dns + +type NS struct { + Ns Name +} + +func (ns *NS) Type() uint16 { return TypeNS } +func (ns *NS) Pack(p *packer) error { return p.writeName(ns.Ns) } +func (ns *NS) Unpack(u *unpacker, length uint16) error { + var err error + ns.Ns, err = u.readName() + return err +} +func (ns *NS) String() string { return ns.Ns.String() } diff --git a/internal/dns/rr_opt.go b/internal/dns/rr_opt.go new file mode 100644 index 0000000..77bc1ee --- /dev/null +++ b/internal/dns/rr_opt.go @@ -0,0 +1,57 @@ +package dns + +import "fmt" + +// https://datatracker.ietf.org/doc/html/rfc6891#section-4.1 +type Option struct { + Code uint16 + Data []byte +} + +type OPT struct { + // https://datatracker.ietf.org/doc/html/rfc6891#section-4.3 + UDPSize uint16 + ExtRCode uint8 + Version uint8 + DO bool + Z uint16 + Options []Option +} + +func (o *OPT) Type() uint16 { return TypeOPT } + +func (o *OPT) Pack(p *packer) error { + for _, opt := range o.Options { + p.writeUint16(opt.Code) + p.writeUint16(uint16(len(opt.Data))) + p.buf = append(p.buf, opt.Data...) + } + return nil +} + +func (o *OPT) Unpack(u *unpacker, length uint16) error { + end := u.off + int(length) + for u.off < end { + code, err := u.readUint16() + if err != nil { + return err + } + optLen, err := u.readUint16() + if err != nil { + return err + } + if u.off+int(optLen) > end { + return fmt.Errorf("dns: opt data overflow") + } + o.Options = append(o.Options, Option{ + Code: code, + Data: append([]byte(nil), u.buf[u.off:u.off+int(optLen)]...), + }) + u.off += int(optLen) + } + return nil +} + +func (o *OPT) String() string { + return fmt.Sprintf("OPT udpsize=%d extrcode=%d version=%d do=%t", o.UDPSize, o.ExtRCode, o.Version, o.DO) +} diff --git a/internal/dns/rr_ptr.go b/internal/dns/rr_ptr.go new file mode 100644 index 0000000..660ad82 --- /dev/null +++ b/internal/dns/rr_ptr.go @@ -0,0 +1,14 @@ +package dns + +type PTR struct { + Ptr Name +} + +func (ptr *PTR) Type() uint16 { return TypePTR } +func (ptr *PTR) Pack(p *packer) error { return p.writeName(ptr.Ptr) } +func (ptr *PTR) Unpack(u *unpacker, length uint16) error { + var err error + ptr.Ptr, err = u.readName() + return err +} +func (ptr *PTR) String() string { return ptr.Ptr.String() } diff --git a/internal/dns/rr_soa.go b/internal/dns/rr_soa.go new file mode 100644 index 0000000..dca4baf --- /dev/null +++ b/internal/dns/rr_soa.go @@ -0,0 +1,65 @@ +package dns + +import "fmt" + +// https://datatracker.ietf.org/doc/html/rfc1035#section-3.3.13 +type SOA struct { + Mname Name + Rname Name + Serial uint32 + Refresh uint32 + Retry uint32 + Expire uint32 + Minimum uint32 +} + +func (s *SOA) Type() uint16 { return TypeSOA } + +func (s *SOA) Pack(p *packer) error { + if err := p.writeName(s.Mname); err != nil { + return err + } + if err := p.writeName(s.Rname); err != nil { + return err + } + p.writeUint32(s.Serial) + p.writeUint32(s.Refresh) + p.writeUint32(s.Retry) + p.writeUint32(s.Expire) + p.writeUint32(s.Minimum) + return nil +} + +func (s *SOA) Unpack(u *unpacker, length uint16) error { + var err error + s.Mname, err = u.readName() + if err != nil { + return err + } + s.Rname, err = u.readName() + if err != nil { + return err + } + s.Serial, err = u.readUint32() + if err != nil { + return err + } + s.Refresh, err = u.readUint32() + if err != nil { + return err + } + s.Retry, err = u.readUint32() + if err != nil { + return err + } + s.Expire, err = u.readUint32() + if err != nil { + return err + } + s.Minimum, err = u.readUint32() + return err +} + +func (s *SOA) String() string { + return fmt.Sprintf("%s %s %d %d %d %d %d", s.Mname, s.Rname, s.Serial, s.Refresh, s.Retry, s.Expire, s.Minimum) +} diff --git a/internal/dns/rr_txt.go b/internal/dns/rr_txt.go new file mode 100644 index 0000000..f225afa --- /dev/null +++ b/internal/dns/rr_txt.go @@ -0,0 +1,43 @@ +package dns + +import "fmt" + +type TXT struct { + Txt []string +} + +func (t *TXT) Type() uint16 { + return TypeTXT +} + +func (t *TXT) Pack(p *packer) error { + for _, s := range t.Txt { + if len(s) > 255 { + return fmt.Errorf("dns: txt string > 255 octets") + } + p.buf = append(p.buf, byte(len(s))) + p.buf = append(p.buf, s...) + } + return nil +} + +func (t *TXT) Unpack(u *unpacker, length uint16) error { + end := u.off + int(length) + for u.off < end { + if u.off >= len(u.buf) { + return fmt.Errorf("dns: short txt") + } + l := int(u.buf[u.off]) + u.off++ + if u.off+l > end { + return fmt.Errorf("dns: txt string overflow") + } + t.Txt = append(t.Txt, string(u.buf[u.off:u.off+l])) + u.off += l + } + return nil +} + +func (t *TXT) String() string { + return fmt.Sprintf("%v", t.Txt) +} diff --git a/internal/dns/rr_unknown.go b/internal/dns/rr_unknown.go new file mode 100644 index 0000000..77ae46d --- /dev/null +++ b/internal/dns/rr_unknown.go @@ -0,0 +1,34 @@ +package dns + +// https://datatracker.ietf.org/doc/html/rfc3597 + +import "fmt" + +type Unknown struct { + RRType uint16 + Data []byte +} + +func (u *Unknown) Type() uint16 { + return u.RRType +} + +func (u *Unknown) Pack(p *packer) error { + p.buf = append(p.buf, u.Data...) + return nil +} + +func (u *Unknown) Unpack(ru *unpacker, length uint16) error { + if ru.off+int(length) > len(ru.buf) { + return fmt.Errorf("dns: short unknown rdata") + } + + u.Data = make([]byte, length) + copy(u.Data, ru.buf[ru.off:ru.off+int(length)]) + ru.off += int(length) + return nil +} + +func (u *Unknown) String() string { + return fmt.Sprintf("\\# %d %x", len(u.Data), u.Data) +} diff --git a/internal/dns/types.go b/internal/dns/types.go index 059446c..a71ee8f 100644 --- a/internal/dns/types.go +++ b/internal/dns/types.go @@ -29,12 +29,12 @@ const ( TypeA = 1 TypeNS = 2 TypeCNAME = 5 - typeSOA = 6 - typePTR = 12 - typeMX = 15 - typeTXT = 16 - typeAAAA = 28 // https://www.rfc-editor.org/info/rfc8499/ - typeOPT = 41 + TypeSOA = 6 + TypePTR = 12 + TypeMX = 15 + TypeTXT = 16 + TypeAAAA = 28 + TypeOPT = 41 ) const ( diff --git a/internal/dns/unpack.go b/internal/dns/unpack.go new file mode 100644 index 0000000..48703b1 --- /dev/null +++ b/internal/dns/unpack.go @@ -0,0 +1,138 @@ +package dns + +import "fmt" + +type unpacker struct { + buf []byte + off int +} + +func (u *unpacker) readUint16() (uint16, error) { + if u.off+2 > len(u.buf) { + return 0, fmt.Errorf("dns: short uint16 at %d", u.off) + } + v := uint16(u.buf[u.off])<<8 | uint16(u.buf[u.off+1]) + u.off += 2 + return v, nil +} + +func (u *unpacker) readUint32() (uint32, error) { + if u.off+4 > len(u.buf) { + return 0, fmt.Errorf("dns: short uint32 at %d", u.off) + } + v := uint32(u.buf[u.off])<<24 | uint32(u.buf[u.off+1])<<16 | + uint32(u.buf[u.off+2])<<8 | uint32(u.buf[u.off+3]) + u.off += 4 + return v, nil +} + +func (u *unpacker) readName() (Name, error) { + var name Name + start := u.off + jumped := false + jumps := 0 + + for { + if u.off >= len(u.buf) { + return Name{}, fmt.Errorf("dns: name overflow") + } + c := u.buf[u.off] + + // Compression pointer + if c&0xC0 == 0xC0 { + if u.off+2 > len(u.buf) { + return Name{}, fmt.Errorf("dns: short pointer") + } + ptr := int(uint16(c&0x3F)<<8 | uint16(u.buf[u.off+1])) + if ptr >= len(u.buf) { + return Name{}, fmt.Errorf("dns: pointer out of range") + } + if !jumped { + start = u.off + 2 + } + u.off = ptr + jumped = true + jumps++ + if jumps > 10 { + return Name{}, fmt.Errorf("dns: too many compression jumps") + } + continue + } + + // Reserved label type (0x40, 0x80) + if c&0xC0 != 0 { + return Name{}, fmt.Errorf("dns: reserved label type") + } + + // Root label + if c == 0 { + name.Data = append(name.Data, 0) + u.off++ + break + } + + // Normal label + length := int(c) + if length > 63 { + return Name{}, fmt.Errorf("dns: label too long") + } + if u.off+1+length > len(u.buf) { + return Name{}, fmt.Errorf("dns: short label") + } + name.Data = append(name.Data, u.buf[u.off:u.off+1+length]...) + u.off += 1 + length + } + + if jumped { + u.off = start + } + return name, nil +} + +func Unpack(buf []byte) (*Msg, error) { + if len(buf) < 12 { + return nil, fmt.Errorf("dns: message too short") + } + u := &unpacker{buf: buf} + m := &Msg{} + + m.Header.ID, _ = u.readUint16() + m.Header.Flags, _ = u.readUint16() + qdcount, _ := u.readUint16() + ancount, _ := u.readUint16() + nscount, _ := u.readUint16() + arcount, _ := u.readUint16() + + m.Question = make([]Question, qdcount) + for i := range m.Question { + if err := m.Question[i].unpack(u); err != nil { + return nil, err + } + } + m.Answer = make([]RR, ancount) + for i := range m.Answer { + rr, err := unpackRR(u) + if err != nil { + return nil, err + } + m.Answer[i] = rr + } + m.Ns = make([]RR, nscount) + for i := range m.Ns { + rr, err := unpackRR(u) + if err != nil { + return nil, err + } + m.Ns[i] = rr + } + m.Extra = make([]RR, arcount) + for i := range m.Extra { + rr, err := unpackRR(u) + if err != nil { + return nil, err + } + m.Extra[i] = rr + } + + return m, nil +} |
