fix(dns): resolve CNAME chains and prevent invalid IP parsing

- Add automated in-band CNAME chain resolution to extract final A/AAAA IP addresses.
- Replace `MaxResponseAnswers` with `MaxIPs` and `MaxCNAMEs` to explicitly bound resource decoding and prevent memory exhaustion.
- Bound CNAME chain traversal to prevent infinite loops from cyclic records.
This commit is contained in:
Yoshio HANAWA
2026-08-20 22:38:28 +09:00
parent 263b1ecf11
commit a33e9c5fab
5 changed files with 328 additions and 35 deletions
+165 -12
View File
@@ -239,26 +239,178 @@ func TestDecodeMessage(t *testing.T) {
}
}
// Regression test for CNAME-following: a response for www.yahoo.co.jp
// contains a CNAME record to edge12.g.yimg.jp (with compressed labels in its
// RDATA) followed by the A record for the canonical name. The CNAME RDATA
// must not be interpreted as an IP address and the A record must be returned.
func TestClient_CNAMEResponse(t *testing.T) {
const hostname = "www.yahoo.co.jp"
const txid = uint16(0x1234)
const clientPort = uint16(54321)
response := []byte{
// Header: txid 0x1234, QR|RD|RA, QD=1 AN=2 NS=0 AR=0.
0x12, 0x34, 0x81, 0x80, 0x00, 0x01, 0x00, 0x02, 0x00, 0x00, 0x00, 0x00,
// Question: www.yahoo.co.jp A IN.
0x03, 'w', 'w', 'w', 0x05, 'y', 'a', 'h', 'o', 'o', 0x02, 'c', 'o', 0x02, 'j', 'p', 0x00,
0x00, 0x01, 0x00, 0x01,
// Answer 1: (ptr to question) CNAME IN ttl=842 rdlen=16
// rdata: edge12.g.yimg.jp with "jp" as compression pointer to offset 0x19.
0xc0, 0x0c, 0x00, 0x05, 0x00, 0x01, 0x00, 0x00, 0x03, 0x4a, 0x00, 0x10,
0x06, 'e', 'd', 'g', 'e', '1', '2', 0x01, 'g', 0x04, 'y', 'i', 'm', 'g', 0xc0, 0x19,
// Answer 2: (ptr into CNAME rdata) A IN ttl=36 rdlen=4 182.22.23.124.
0xc0, 0x2d, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x24, 0x00, 0x04, 0xb6, 0x16, 0x17, 0x7c,
}
name := MustNewName(hostname)
var client Client
err := client.StartResolve(clientPort, txid, ResolveConfig{
Questions: []Question{{
Name: name,
Type: TypeA,
Class: ClassINET,
}},
EnableRecursion: true,
MaxIPs: 4,
MaxCNAMEs: 2,
})
if err != nil {
t.Fatal("failed to start DNS resolve:", err)
}
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(response, 0); err != nil {
t.Fatal("failed to demux DNS response:", err)
}
var addrs [4]netip.Addr
n, err := client.ResponseAnswerLookup(addrs[:], hostname)
if err != nil {
t.Fatal("failed to look up DNS response answers:", err)
}
if n != 1 {
t.Fatalf("expected 1 answer, got %d: %v", n, addrs[:n])
}
if addrs[0] != (netip.AddrFrom4([4]byte{182, 22, 23, 124})) {
t.Fatalf("expected 182.22.23.124, got %v", addrs[0])
}
}
// Table-driven tests for Message.WriteAnswers covering answer reordering,
// non-address record types and cyclic CNAME aliases.
func TestMessage_WriteAnswers(t *testing.T) {
tests := []struct {
name string
host string
response []byte
want []netip.Addr
}{
{
name: "A record before its CNAME",
host: "www.yahoo.co.jp",
// Answer 1 is the A record for edge12.g.yimg.jp, spelled out with
// a trailing compression pointer to "jp" in the question. Answer 2
// is the CNAME from www.yahoo.co.jp whose RDATA is a single
// backward compression pointer to answer 1's owner name.
response: []byte{
// Header: txid 0x1234, QR|RD|RA, QD=1 AN=2 NS=0 AR=0.
0x12, 0x34, 0x81, 0x80, 0x00, 0x01, 0x00, 0x02, 0x00, 0x00, 0x00, 0x00,
// Question: www.yahoo.co.jp A IN.
0x03, 'w', 'w', 'w', 0x05, 'y', 'a', 'h', 'o', 'o', 0x02, 'c', 'o', 0x02, 'j', 'p', 0x00,
0x00, 0x01, 0x00, 0x01,
// Answer 1: edge12.g.yimg.jp A IN ttl=36 rdlen=4 182.22.23.124.
0x06, 'e', 'd', 'g', 'e', '1', '2', 0x01, 'g', 0x04, 'y', 'i', 'm', 'g', 0xc0, 0x19,
0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x00, 0x24, 0x00, 0x04, 0xb6, 0x16, 0x17, 0x7c,
// Answer 2: (ptr to question) CNAME IN ttl=842 rdlen=2, target
// is a pointer to answer 1's owner name at offset 0x21.
0xc0, 0x0c, 0x00, 0x05, 0x00, 0x01, 0x00, 0x00, 0x03, 0x4a, 0x00, 0x02, 0xc0, 0x21,
},
want: []netip.Addr{netip.AddrFrom4([4]byte{182, 22, 23, 124})},
},
{
name: "ignores non-address records",
host: "example.com",
response: []byte{
// Header: txid 0x5678, QR|RD|RA, QD=1 AN=2 NS=0 AR=0.
0x56, 0x78, 0x81, 0x80, 0x00, 0x01, 0x00, 0x02, 0x00, 0x00, 0x00, 0x00,
// Question: example.com A IN.
0x07, 'e', 'x', 'a', 'm', 'p', 'l', 'e', 0x03, 'c', 'o', 'm', 0x00,
0x00, 0x01, 0x00, 0x01,
// Answer 1: (ptr to question) TXT IN ttl=60 rdlen=6 "hello".
0xc0, 0x0c, 0x00, 0x10, 0x00, 0x01, 0x00, 0x00, 0x00, 0x3c, 0x00, 0x06, 0x05, 'h', 'e', 'l', 'l', 'o',
// Answer 2: (ptr to question) A IN ttl=300 rdlen=4 192.0.2.1.
0xc0, 0x0c, 0x00, 0x01, 0x00, 0x01, 0x00, 0x00, 0x01, 0x2c, 0x00, 0x04, 0xc0, 0x00, 0x02, 0x01,
},
want: []netip.Addr{netip.AddrFrom4([4]byte{192, 0, 2, 1})},
},
{
name: "CNAME cycle terminates",
host: "a.com",
response: []byte{
// Header: txid 0xabcd, QR|RD|RA, QD=1 AN=2 NS=0 AR=0.
0xab, 0xcd, 0x81, 0x80, 0x00, 0x01, 0x00, 0x02, 0x00, 0x00, 0x00, 0x00,
// Question: a.com A IN.
0x01, 'a', 0x03, 'c', 'o', 'm', 0x00, 0x00, 0x01, 0x00, 0x01,
// Answer 1: a.com CNAME b.com.
0x01, 'a', 0x03, 'c', 'o', 'm', 0x00, 0x00, 0x05, 0x00, 0x01, 0x00, 0x00, 0x00, 0x0a, 0x00, 0x07, 0x01, 'b', 0x03, 'c', 'o', 'm', 0x00,
// Answer 2: b.com CNAME a.com.
0x01, 'b', 0x03, 'c', 'o', 'm', 0x00, 0x00, 0x05, 0x00, 0x01, 0x00, 0x00, 0x00, 0x0a, 0x00, 0x07, 0x01, 'a', 0x03, 'c', 'o', 'm', 0x00,
},
want: nil,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
var msg Message
msg.LimitResourceDecoding(1, 4, 0, 0)
_, incomplete, err := msg.Decode(tt.response)
if incomplete || err != nil {
t.Fatal("decode:", incomplete, err)
}
var addrs [4]netip.Addr
n, err := msg.WriteAnswers(addrs[:], tt.host)
if err != nil {
t.Fatal("write answers:", err)
}
if n != uint16(len(tt.want)) {
t.Fatalf("expected %d addresses, got %d: %v", len(tt.want), n, addrs[:n])
}
for i, want := range tt.want {
if addrs[i] != want {
t.Errorf("address %d: expected %v, got %v", i, want, addrs[i])
}
}
})
}
}
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{
const maxIPs = 4
const maxCNAMEs = 2
const maxDecoded = maxIPs + maxCNAMEs
allIPs := [8][4]byte{
{192, 0, 2, 1},
{192, 0, 2, 2},
{192, 0, 2, 3},
{192, 0, 2, 4},
{192, 0, 2, 5},
{192, 0, 2, 6},
{192, 0, 2, 7},
{192, 0, 2, 8},
}
tests := []struct {
name string
responseIPs [][4]byte
wantAnswers int
wantDecoded int // Answers decoded (and copied by ResponseCopyTo).
wantAnswers int // Addresses returned by ResponseAnswerLookup.
}{
{name: "single_answer", responseIPs: allIPs[:1], wantAnswers: 1},
{name: "multiple_answers", responseIPs: allIPs[:4], wantAnswers: 4},
{name: "answer_limit", responseIPs: allIPs[:5], wantAnswers: maxAnswers},
{name: "single_answer", responseIPs: allIPs[:1], wantDecoded: 1, wantAnswers: 1},
{name: "multiple_answers", responseIPs: allIPs[:4], wantDecoded: 4, wantAnswers: 4},
{name: "answer_limit", responseIPs: allIPs[:5], wantDecoded: 5, wantAnswers: maxIPs},
{name: "decode_limit", responseIPs: allIPs[:8], wantDecoded: maxDecoded, wantAnswers: maxIPs},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
@@ -290,8 +442,9 @@ func TestClient_ReceivesDNSResponse(t *testing.T) {
Type: TypeA,
Class: ClassINET,
}},
EnableRecursion: true,
MaxResponseAnswers: maxAnswers,
EnableRecursion: true,
MaxIPs: maxIPs,
MaxCNAMEs: maxCNAMEs,
})
if err != nil {
t.Fatal("failed to start DNS resolve:", err)
@@ -307,12 +460,12 @@ func TestClient_ReceivesDNSResponse(t *testing.T) {
t.Fatal("failed to demux DNS response:", err)
}
var addrs [maxAnswers]netip.Addr
var addrs [maxIPs]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) {
if int(answers) != tt.wantAnswers {
t.Fatalf("expected %d answers, got %d", tt.wantAnswers, answers)
}
for i := 0; i < tt.wantAnswers; i++ {
@@ -334,8 +487,8 @@ func TestClient_ReceivesDNSResponse(t *testing.T) {
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))
if len(lookup.Answers) != tt.wantDecoded {
t.Fatalf("expected %d copied answers, got %d", tt.wantDecoded, len(lookup.Answers))
}
})
}