From 263b1ecf11eaa3d93fe90413fb847578dd07ceb8 Mon Sep 17 00:00:00 2001 From: Yoshio HANAWA Date: Wed, 19 Aug 2026 03:55:33 +0900 Subject: [PATCH] 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. --- dns/client.go | 9 ++- dns/dns_test.go | 152 ++++++++++++++++++++++++------------------ x/xnet/stack-async.go | 3 +- 3 files changed, 98 insertions(+), 66 deletions(-) diff --git a/dns/client.go b/dns/client.go index 8897066..0a62847 100644 --- a/dns/client.go +++ b/dns/client.go @@ -24,6 +24,9 @@ type ResolveConfig struct { Questions []Question Additional []Resource EnableRecursion bool + // MaxResponseAnswers limits how many answer records are decoded from the + // DNS response. If zero it defaults to the number of Questions. + MaxResponseAnswers uint16 } func (sudp *Client) Protocol() uint64 { return uint64(lneto.IPProtoUDP) } @@ -37,8 +40,12 @@ func (c *Client) StartResolve(localPort, txid uint16, cfg ResolveConfig) error { if nd > math.MaxUint16 { return lneto.ErrInvalidConfig } + maxAns := cfg.MaxResponseAnswers + if maxAns == 0 { + maxAns = uint16(nd) + } c.reset(localPort, txid, CQueryPending, cfg.EnableRecursion) - c.msg.LimitResourceDecoding(uint16(nd), uint16(nd), 0, 0) + c.msg.LimitResourceDecoding(uint16(nd), maxAns, 0, 0) c.msg.AddQuestions(cfg.Questions) c.msg.AddAdditionals(cfg.Additional) return nil diff --git a/dns/dns_test.go b/dns/dns_test.go index 64bfc1b..75d4a6e 100644 --- a/dns/dns_test.go +++ b/dns/dns_test.go @@ -243,76 +243,100 @@ 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[:]), - }, + 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}, } - - // 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) + 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][:]) + } - // Set up the DNS client. - var client Client - client.StartResolve(clientPort, txid, ResolveConfig{ - Questions: []Question{{ - Name: MustNewName(hostname), - Type: TypeA, - Class: ClassINET, - }}, - EnableRecursion: true, - }) + // 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) + } - // Simulate sending by calling Encapsulate (changes state to AwaitResponse). - var dummy [512]byte - client.Encapsulate(dummy[:], 0, 0) + 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) + } - // Call Demux with DNS payload. - err = client.Demux(dnsPayload, 0) - if err != nil { - t.Fatal("Client Demux error:", 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) + } - // Check the client received the answer. - var addrs [4]netip.Addr - answers, err := client.ResponseAnswerLookup(addrs[:], hostname) - if answers != 1 { - t.Fatalf("expected 1 answer, got %d", answers) - } - addr := addrs[0] - if !addr.Is4() { - t.Fatalf("expected 4 bytes in answer, got %d", addr.BitLen()/8) - } - if addr.As4() != wantIP { - t.Errorf("expected IP %v, got %v", wantIP, addr.String()) - } + 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) + } + } - // Test MessageCopyTo as well. - var lookup Message - lookup.LimitResourceDecoding(1, 1, 0, 0) - done, err := client.ResponseCopyTo(&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)) + 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)) + } + }) } } diff --git a/x/xnet/stack-async.go b/x/xnet/stack-async.go index 2e974c5..62e24fe 100644 --- a/x/xnet/stack-async.go +++ b/x/xnet/stack-async.go @@ -629,7 +629,8 @@ func (s *StackAsync) StartLookupIPType(host string, qtype dns.Type) error { Additional: []dns.Resource{ s.ednsopt, }, - EnableRecursion: true, + EnableRecursion: true, + MaxResponseAnswers: uint16(len(s.addrbufnip)), }) if err != nil { return err