mirror of
https://github.com/soypat/lneto.git
synced 2026-08-12 19:03:42 +00:00
fix dns Message.CopyFrom bug
This commit is contained in:
@@ -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])
|
||||
}
|
||||
|
||||
@@ -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))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user