Files
lneto/dns/dns_test.go
T
Yoshio HANAWA 263b1ecf11 fix(dns): decode up to four answer records (#186)
* fix(dns): decode up to four answer records

A single DNS question can return multiple answer records. Use an answer
decode limit of four in Client.StartResolve and add a regression test
covering a single-question response with multiple A records.

* fix(dns): add MaxResponseAnswers to ResolveConfig

This allows callers to explicitly declare the maximum number of answer
records to retain, removing the hardcoded limit in Client.StartResolve.

xnet.StackAsync is updated to set this limit to match the length of its
lookup-result buffer. This maintains the zero-allocation design while
fixing the issue where multiple A records were ignored.
2026-08-18 11:55:33 -07:00

343 lines
8.8 KiB
Go

package dns
import (
"fmt"
"net/netip"
"strings"
"testing"
)
var defaultMessageFlags = NewClientHeaderFlags(OpCodeQuery, true)
func TestNameString(t *testing.T) {
var name Name
domain := "foo.bar.org"
domainSplit := strings.Split(domain, ".")
for i, label := range domainSplit {
name.AddLabel(label)
s := name.String()
if s != strings.Join(domainSplit[:i+1], ".")+"." {
t.Fatalf("unexpected name string %q", s)
}
}
}
func TestNameAppendDecode(t *testing.T) {
const domain = "foo.bar.org"
name, err := NewName(domain)
if err != nil {
t.Fatal(err)
} else if name.String() != domain+"." {
t.Fatalf("unexpected name string %q", name.String())
}
var buf [512]byte
b, err := name.AppendTo(buf[:0])
if err != nil {
t.Fatal(err)
}
if uint16(len(b)) != name.Len() {
t.Fatalf("unexpected name length %d", len(b))
}
if b[len(b)-1] != 0 {
t.Fatalf("unexpected name terminator byte after construction: %q", b[len(b)-1])
}
var name2 Name
n, err := name2.Decode(b, 0)
if err != nil {
t.Fatal(err)
}
if n != name.Len() {
t.Errorf("unexpected name parsed length %q (%d), want %q (%d)", name.data, n, b, name.Len())
}
if name2.String() != name.String() {
t.Errorf("unexpected name string %q, want %q", name2.String(), name.String())
}
// Re-decode.
const okvalidName = "\x03www\x02go\x03dev\x00"
_, err = name.Decode([]byte(okvalidName), 0)
if err != nil {
t.Error("got error decoding valid name", err)
} else if name.String() != "www.go.dev." {
t.Error("unexpected name string", name.String())
}
b, err = name.AppendTo(buf[:0])
if err != nil {
t.Fatal(err)
}
if b[len(b)-1] != 0 {
t.Fatalf("unexpected name terminator byte after decoding: %q", b[len(b)-1])
}
if string(b) != okvalidName {
t.Errorf("unexpected name bytes after decode %q, want %q", b, okvalidName)
}
// Decode invalid name.
const invalidName = "\x03w.w\x02go\x03dev\x00"
_, err = name.Decode([]byte(invalidName), 0)
if err == nil {
t.Error("expected error for invalid name")
} else if err != errInvalidName {
t.Errorf("unexpected error %v, want %v", err, errInvalidName)
}
}
func TestMessageAppendEncode(t *testing.T) {
var tests = []struct {
Message Message
error error
}{
{
Message: Message{
Questions: []Question{
{
Name: MustNewName("."),
Type: TypeA,
Class: ClassINET,
},
},
Answers: []Resource{
{
header: ResourceHeader{
Name: MustNewName("."),
Type: TypeA,
Class: ClassINET,
TTL: 256,
Length: 3,
},
data: []byte{1, 2, 3},
},
},
},
},
}
var buf [512]byte
for _, tt := range tests {
b, err := tt.Message.AppendTo(buf[:0], 123, defaultMessageFlags)
if err != nil {
t.Fatal(err)
}
var msg Message
msg.LimitResourceDecoding(uint16(len(tt.Message.Questions)), uint16(len(tt.Message.Answers)), uint16(len(tt.Message.Authorities)), uint16(len(tt.Message.Additionals)))
_, incomplete, err := msg.Decode(b)
if err != nil {
t.Fatal(err)
} else if incomplete {
t.Fatal("incomplete parse")
}
if msg.String() != tt.Message.String() {
t.Errorf("mismatch message strings after append/decode:\n%s\n%s", tt.Message.String(), msg.String())
}
}
}
func TestMessageAppendEncodeIncompleteOK(t *testing.T) {
var tests = []struct {
Message Message
error error
}{
{
Message: Message{
Questions: []Question{
{
Name: MustNewName("."),
Type: TypeA,
Class: ClassINET,
},
},
Answers: []Resource{
{
header: ResourceHeader{
Name: MustNewName("."),
Type: TypeA,
Class: ClassINET,
TTL: 256,
Length: 3,
},
data: []byte{1, 2, 3},
},
{
header: ResourceHeader{
Name: MustNewName("."),
Type: TypeA,
Class: ClassINET,
TTL: 256,
Length: 3,
},
data: []byte{1, 2, 3},
},
},
},
},
}
var buf [512]byte
for _, tt := range tests {
b, err := tt.Message.AppendTo(buf[:0], 123, defaultMessageFlags)
if err != nil {
t.Fatal(err)
}
var msg Message
// Limit answers to 1 to test incomplete parsing (message has 2 answers).
msg.LimitResourceDecoding(uint16(len(tt.Message.Questions)), 1, uint16(len(tt.Message.Authorities)), uint16(len(tt.Message.Additionals)))
_, incomplete, err := msg.Decode(b)
if err != nil && !incomplete {
t.Fatal(err)
} else if !incomplete {
t.Fatal("expected incomplete parse")
}
tt.Message.Answers = tt.Message.Answers[:1] // Trim to match the limited decode.
if msg.String() != tt.Message.String() {
t.Errorf("mismatch message strings after append/decode:\n%s\n%s", tt.Message.String(), msg.String())
}
}
}
func (m *Message) String() string {
// s := fmt.Sprintf("Message: %#v\n", &m.Header)
var s strings.Builder
if len(m.Questions) > 0 {
s.WriteString("-- Questions\n")
for _, q := range m.Questions {
s.WriteString(fmt.Sprintf("%#v\n", q))
}
}
if len(m.Answers) > 0 {
s.WriteString("-- Answers\n")
for _, a := range m.Answers {
s.WriteString(fmt.Sprintf("%#v\n", a))
}
}
if len(m.Authorities) > 0 {
s.WriteString("-- Authorities\n")
for _, ns := range m.Authorities {
s.WriteString(fmt.Sprintf("%#v\n", ns))
}
}
if len(m.Additionals) > 0 {
s.WriteString("-- Additionals\n")
for _, e := range m.Additionals {
s.WriteString(fmt.Sprintf("%#v\n", e))
}
}
return s.String()
}
func TestDecodeMessage(t *testing.T) {
var data = []byte{
0x84, 0x05, 0x81, 0x80, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x01, 0x0b, 0x77, 0x68, 0x69,
0x74, 0x74, 0x69, 0x6c, 0x65, 0x61, 0x6b, 0x73, 0x03, 0x63, 0x6f, 0x6d, 0x00, 0x00, 0x01, 0x00,
0x01, 0xc0, 0x0c, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x1e, 0xaf, 0x00, 0x04, 0xc6, 0x31, 0x17,
0x91, 0x00, 0x00, 0x29, 0x10, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
}
var msg Message
msg.LimitResourceDecoding(5, 5, 5, 5)
off, incomplete, err := msg.Decode(data)
if incomplete || err != nil {
t.Fatal(incomplete, err, off)
}
}
func TestClient_ReceivesDNSResponse(t *testing.T) {
const hostname = "example.com"
const txid = uint16(12345)
const clientPort = uint16(54321)
const maxAnswers = 4
allIPs := [5][4]byte{
{192, 0, 2, 1},
{192, 0, 2, 2},
{192, 0, 2, 3},
{192, 0, 2, 4},
{192, 0, 2, 5},
}
tests := []struct {
name string
responseIPs [][4]byte
wantAnswers int
}{
{name: "single_answer", responseIPs: allIPs[:1], wantAnswers: 1},
{name: "multiple_answers", responseIPs: allIPs[:4], wantAnswers: 4},
{name: "answer_limit", responseIPs: allIPs[:5], wantAnswers: maxAnswers},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
name := MustNewName(hostname)
responseMsg := Message{
Questions: []Question{{
Name: name,
Type: TypeA,
Class: ClassINET,
}},
Answers: make([]Resource, len(tt.responseIPs)),
}
for i := range tt.responseIPs {
responseMsg.Answers[i] = NewResource(name, TypeA, ClassINET, 300, tt.responseIPs[i][:])
}
// Response flags: QR=1 (response), RD=1, RA=1.
responseFlags := HeaderFlags(1<<15 | 1<<8 | 1<<7)
var responseBuf [512]byte
dnsPayload, err := responseMsg.AppendTo(responseBuf[:0], txid, responseFlags)
if err != nil {
t.Fatal("failed to build DNS response:", err)
}
var client Client
err = client.StartResolve(clientPort, txid, ResolveConfig{
Questions: []Question{{
Name: name,
Type: TypeA,
Class: ClassINET,
}},
EnableRecursion: true,
MaxResponseAnswers: maxAnswers,
})
if err != nil {
t.Fatal("failed to start DNS resolve:", err)
}
// Encapsulate the query to move the client into the outstanding state.
var queryBuf [512]byte
_, err = client.Encapsulate(queryBuf[:], 0, 0)
if err != nil {
t.Fatal("failed to encapsulate DNS query:", err)
}
if err := client.Demux(dnsPayload, 0); err != nil {
t.Fatal("failed to demux DNS response:", err)
}
var addrs [maxAnswers]netip.Addr
answers, err := client.ResponseAnswerLookup(addrs[:], hostname)
if err != nil {
t.Fatal("failed to look up DNS response answers:", err)
}
if answers != uint16(tt.wantAnswers) {
t.Fatalf("expected %d answers, got %d", tt.wantAnswers, answers)
}
for i := 0; i < tt.wantAnswers; i++ {
addr := addrs[i]
if !addr.Is4() {
t.Errorf("answer %d: expected IPv4 address, got %v", i, addr)
continue
}
if addr.As4() != tt.responseIPs[i] {
t.Errorf("answer %d: expected IP %v, got %v", i, tt.responseIPs[i], addr)
}
}
var lookup Message
done, err := client.ResponseCopyTo(&lookup)
if err != nil {
t.Fatal("failed to copy DNS response:", err)
}
if !done {
t.Fatal("expected done=true")
}
if len(lookup.Answers) != tt.wantAnswers {
t.Fatalf("expected %d copied answers, got %d", tt.wantAnswers, len(lookup.Answers))
}
})
}
}