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