summaryrefslogtreecommitdiff
path: root/internal
diff options
context:
space:
mode:
authorradhitya <alif@radhitya.org>2026-07-04 14:29:02 +0700
committerradhitya <alif@radhitya.org>2026-07-04 14:29:02 +0700
commit52a0d3845ead07f840e9e99cb4a8c3507c84ab30 (patch)
treeb9a08117bdcb35a1a0e61543b0859db8f4aa6bfd /internal
parente9d3659ef41b9b5fd5e7cc7bd04ae3fc07850044 (diff)
fiuh
Diffstat (limited to 'internal')
-rw-r--r--internal/dns/dns_test.go472
-rw-r--r--internal/dns/fuzz_test.go52
-rw-r--r--internal/dns/header.go28
-rw-r--r--internal/dns/message.go9
-rw-r--r--internal/dns/name.go2
-rw-r--r--internal/dns/pack.go103
-rw-r--r--internal/dns/question.go23
-rw-r--r--internal/dns/rr.go110
-rw-r--r--internal/dns/rr_a.go34
-rw-r--r--internal/dns/rr_aaaa.go34
-rw-r--r--internal/dns/rr_cname.go14
-rw-r--r--internal/dns/rr_mx.go30
-rw-r--r--internal/dns/rr_ns.go14
-rw-r--r--internal/dns/rr_opt.go57
-rw-r--r--internal/dns/rr_ptr.go14
-rw-r--r--internal/dns/rr_soa.go65
-rw-r--r--internal/dns/rr_txt.go43
-rw-r--r--internal/dns/rr_unknown.go34
-rw-r--r--internal/dns/types.go12
-rw-r--r--internal/dns/unpack.go138
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
+}