fix dns Message.CopyFrom bug

This commit is contained in:
Patricio Whittingslow
2026-01-12 09:14:10 -03:00
parent c4ab7cd012
commit 0ca2eae4d2
3 changed files with 82 additions and 20 deletions
+4
View File
@@ -638,6 +638,10 @@ func (dst *Message) CopyFrom(m Message) {
internal.SliceReuse(&dst.Answers, len(m.Answers))
internal.SliceReuse(&dst.Authorities, len(m.Authorities))
internal.SliceReuse(&dst.Additionals, len(m.Additionals))
dst.Questions = dst.Questions[:len(m.Questions)]
dst.Answers = dst.Answers[:len(m.Answers)]
dst.Authorities = dst.Authorities[:len(m.Authorities)]
dst.Additionals = dst.Additionals[:len(m.Additionals)]
for i := range dst.Questions {
dst.Questions[i].CopyFrom(m.Questions[i])
}
+78
View File
@@ -237,3 +237,81 @@ func TestDecodeMessage(t *testing.T) {
t.Fatal(incomplete, err, off)
}
}
func TestClient_ReceivesDNSResponse(t *testing.T) {
const hostname = "example.com"
const txid = uint16(12345)
const clientPort = uint16(54321)
wantIP := [4]byte{93, 184, 216, 34}
// Build a DNS response message.
name := MustNewName(hostname)
responseMsg := Message{
Questions: []Question{{
Name: name,
Type: TypeA,
Class: ClassINET,
}},
Answers: []Resource{
NewResource(name, TypeA, ClassINET, 300, wantIP[:]),
},
}
// Response flags: QR=1 (response), RD=1, RA=1.
responseFlags := HeaderFlags(1<<15 | 1<<8 | 1<<7)
var buf [512]byte
dnsPayload, err := responseMsg.AppendTo(buf[:0], txid, responseFlags)
if err != nil {
t.Fatal("failed to build DNS response:", err)
}
// Set up the DNS client.
var client Client
client.StartResolve(clientPort, txid, ResolveConfig{
Questions: []Question{{
Name: MustNewName(hostname),
Type: TypeA,
Class: ClassINET,
}},
EnableRecursion: true,
})
// Simulate sending by calling Encapsulate (changes state to AwaitResponse).
var dummy [512]byte
client.Encapsulate(dummy[:], 0, 0)
// Call Demux with DNS payload.
err = client.Demux(dnsPayload, 0)
if err != nil {
t.Fatal("Client Demux error:", err)
}
// Check the client received the answer.
answers := client.Answers()
if len(answers) != 1 {
t.Fatalf("expected 1 answer, got %d", len(answers))
}
data := answers[0].RawData()
if len(data) != 4 {
t.Fatalf("expected 4 bytes in answer, got %d", len(data))
}
if [4]byte(data) != wantIP {
t.Errorf("expected IP %v, got %v", wantIP, data)
}
// Test MessageCopyTo as well.
var lookup Message
lookup.LimitResourceDecoding(1, 1, 0, 0)
done, err := client.MessageCopyTo(&lookup)
if err != nil {
t.Fatal("MessageCopyTo error:", err)
}
if !done {
t.Fatal("expected done=true")
}
if len(lookup.Answers) != 1 {
t.Fatalf("MessageCopyTo: expected 1 answer, got %d", len(lookup.Answers))
}
}
-20
View File
@@ -74,26 +74,6 @@ func TestDNS_QueryReceivesAnswer(t *testing.T) {
t.Fatal("client Demux failed:", err)
}
// Verify our DNS response is valid by decoding it separately.
var testMsg dns.Message
testMsg.LimitResourceDecoding(1, 1, 0, 0)
// Find DNS payload offset in response packet.
const ethLen, ipLen, udpLen = 14, 20, 8
dnsOffset := ethLen + ipLen + udpLen
_, _, decodeErr := testMsg.Decode(responsePkt[dnsOffset:])
if decodeErr != nil {
t.Logf("DNS decode error: %v", decodeErr)
}
t.Logf("Test decode: questions=%d, answers=%d", len(testMsg.Questions), len(testMsg.Answers))
if len(testMsg.Answers) > 0 {
data := testMsg.Answers[0].RawData()
t.Logf("Answer data: %v (len=%d)", data, len(data))
}
// Also verify DNS header flags in the packet.
dnsFrame, _ := dns.NewFrame(responsePkt[dnsOffset:])
t.Logf("DNS Frame: txid=%d, flags=%s, QD=%d, AN=%d", dnsFrame.TxID(), dnsFrame.Flags(), dnsFrame.QDCount(), dnsFrame.ANCount())
// Check the result.
addrs, done, err := client.ResultLookupIP(hostname)
if err != nil {