From 0ca2eae4d295b6c3d57821c41603ff25b1bdf069 Mon Sep 17 00:00:00 2001 From: Patricio Whittingslow Date: Mon, 12 Jan 2026 09:14:10 -0300 Subject: [PATCH] fix dns Message.CopyFrom bug --- dns/dns.go | 4 +++ dns/dns_test.go | 78 +++++++++++++++++++++++++++++++++++++++++ x/xnet/xnet_dns_test.go | 20 ----------- 3 files changed, 82 insertions(+), 20 deletions(-) diff --git a/dns/dns.go b/dns/dns.go index c3abd1f..1521e4b 100644 --- a/dns/dns.go +++ b/dns/dns.go @@ -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]) } diff --git a/dns/dns_test.go b/dns/dns_test.go index 5380a7c..e2effdd 100644 --- a/dns/dns_test.go +++ b/dns/dns_test.go @@ -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)) + } +} diff --git a/x/xnet/xnet_dns_test.go b/x/xnet/xnet_dns_test.go index 1b91a0a..0822366 100644 --- a/x/xnet/xnet_dns_test.go +++ b/x/xnet/xnet_dns_test.go @@ -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 {