Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 38 additions & 14 deletions src/cares_wrap.cc
Original file line number Diff line number Diff line change
Expand Up @@ -249,6 +249,23 @@ std::vector<std::pair<std::string, int>> ParseServersCsv(const char* csv) {
return servers;
}

int GetAnswerCountForTTLBuffer(const unsigned char* buf, int len) {
static constexpr int kDNSAnswerCountOffset = 6;
static constexpr int kAresDefaultTTLBufferLength = 256;
if (len <= kDNSAnswerCountOffset + 1) {
return kAresDefaultTTLBufferLength;
}

const int answer_count = (static_cast<int>(buf[kDNSAnswerCountOffset]) << 8) |
static_cast<int>(buf[kDNSAnswerCountOffset + 1]);
return answer_count == 0 ? 1 : answer_count;
}

template <typename T>
std::vector<T> MakeAddrTTLBuffer(const unsigned char* buf, int len) {
return std::vector<T>(GetAnswerCountForTTLBuffer(buf, len));
}

Maybe<int> ParseGeneralReply(Environment* env,
const unsigned char* buf,
int len,
Expand Down Expand Up @@ -1217,11 +1234,12 @@ Maybe<int> AnyTraits::Parse(QueryAnyWrap* wrap,
int type, status, old_count;

/* Parse A records or CNAME records */
ares_addrttl addrttls[256];
int naddrttls = arraysize(addrttls);
std::vector<ares_addrttl> addrttls =
MakeAddrTTLBuffer<ares_addrttl>(buf, len);
int naddrttls = static_cast<int>(addrttls.size());

type = ns_t_cname_or_a;
if (!ParseGeneralReply(env, buf, len, &type, ret, addrttls, &naddrttls)
if (!ParseGeneralReply(env, buf, len, &type, ret, addrttls.data(), &naddrttls)
.To(&status)) {
return Nothing<int>();
}
Expand Down Expand Up @@ -1294,11 +1312,13 @@ Maybe<int> AnyTraits::Parse(QueryAnyWrap* wrap,
}

/* Parse AAAA records */
ares_addr6ttl addr6ttls[256];
int naddr6ttls = arraysize(addr6ttls);
std::vector<ares_addr6ttl> addr6ttls =
MakeAddrTTLBuffer<ares_addr6ttl>(buf, len);
int naddr6ttls = static_cast<int>(addr6ttls.size());

type = ns_t_aaaa;
if (!ParseGeneralReply(env, buf, len, &type, ret, addr6ttls, &naddr6ttls)
if (!ParseGeneralReply(
env, buf, len, &type, ret, addr6ttls.data(), &naddr6ttls)
.To(&status)) {
return Nothing<int>();
}
Expand Down Expand Up @@ -1486,20 +1506,22 @@ Maybe<int> ATraits::Parse(QueryAWrap* wrap,
HandleScope handle_scope(env->isolate());
Context::Scope context_scope(env->context());

ares_addrttl addrttls[256];
int naddrttls = arraysize(addrttls), status;
std::vector<ares_addrttl> addrttls =
MakeAddrTTLBuffer<ares_addrttl>(buf, len);
int naddrttls = static_cast<int>(addrttls.size()), status;
Local<Array> ret = Array::New(env->isolate());

int type = ns_t_a;
if (!ParseGeneralReply(env, buf, len, &type, ret, addrttls, &naddrttls)
if (!ParseGeneralReply(env, buf, len, &type, ret, addrttls.data(), &naddrttls)
.To(&status)) {
return Nothing<int>();
}
if (status != ARES_SUCCESS) {
return Just<int>(status);
}

Local<Array> ttls = AddrTTLToArray<ares_addrttl>(env, addrttls, naddrttls);
Local<Array> ttls =
AddrTTLToArray<ares_addrttl>(env, addrttls.data(), naddrttls);

wrap->CallOnComplete(ret, ttls);
return Just<int>(ARES_SUCCESS);
Expand All @@ -1518,20 +1540,22 @@ Maybe<int> AaaaTraits::Parse(QueryAaaaWrap* wrap,
HandleScope handle_scope(env->isolate());
Context::Scope context_scope(env->context());

ares_addr6ttl addrttls[256];
int naddrttls = arraysize(addrttls), status;
std::vector<ares_addr6ttl> addrttls =
MakeAddrTTLBuffer<ares_addr6ttl>(buf, len);
int naddrttls = static_cast<int>(addrttls.size()), status;
Local<Array> ret = Array::New(env->isolate());

int type = ns_t_aaaa;
if (!ParseGeneralReply(env, buf, len, &type, ret, addrttls, &naddrttls)
if (!ParseGeneralReply(env, buf, len, &type, ret, addrttls.data(), &naddrttls)
.To(&status)) {
return Nothing<int>();
}
if (status != ARES_SUCCESS) {
return Just<int>(status);
}

Local<Array> ttls = AddrTTLToArray<ares_addr6ttl>(env, addrttls, naddrttls);
Local<Array> ttls =
AddrTTLToArray<ares_addr6ttl>(env, addrttls.data(), naddrttls);

wrap->CallOnComplete(ret, ttls);
return Just<int>(ARES_SUCCESS);
Expand Down
69 changes: 69 additions & 0 deletions test/parallel/test-dns-resolveany-ttl-overflow.js
Original file line number Diff line number Diff line change
@@ -0,0 +1,69 @@
'use strict';
const common = require('../common');
const dnstools = require('../common/dns');
const assert = require('assert');
const dgram = require('dgram');
const dns = require('dns');

const dnsPromises = dns.promises;

const kRecordCount = 257;
const kADomain = 'many-a.example.org';

const server = dgram.createSocket('udp4');

server.on('message', common.mustCall((msg, { address, port }) => {
const parsed = dnstools.parseDNSPacket(msg);
const question = parsed.questions[0];
const { domain } = question;

assert.strictEqual(question.type, 'ANY');
assert.strictEqual(domain, kADomain);

server.send(dnstools.writeDNSPacket({
id: parsed.id,
questions: parsed.questions,
answers: createARecords(domain),
}), port, address);
}, 2));

server.bind(0, common.mustCall(async () => {
const { port } = server.address();
const callbackResolver = new dns.Resolver({ timeout: 1000, tries: 1 });
const promiseResolver = new dnsPromises.Resolver({ timeout: 1000, tries: 1 });
callbackResolver.setServers([`127.0.0.1:${port}`]);
promiseResolver.setServers([`127.0.0.1:${port}`]);

validateRecords(await promiseResolver.resolveAny(kADomain), 'A');
validateRecords(await resolveAny(callbackResolver, kADomain), 'A');

server.close();
}));

function createARecords(domain) {
return Array.from({ length: kRecordCount }, (_, i) => ({
type: 'A',
address: `10.0.${i >> 8}.${i & 0xff}`,
ttl: 60 + i,
domain,
}));
}

function resolveAny(resolver, domain) {
return new Promise((resolve) => {
resolver.resolveAny(domain, common.mustSucceed(resolve));
});
}

function validateRecords(records, type) {
assert.strictEqual(records.length, kRecordCount);
for (const record of records) {
assert.strictEqual(record.type, type);
}

assert.strictEqual(records[0].ttl, 60);
assert.strictEqual(records[255].ttl, 315);
assert.strictEqual(records[256].ttl, 316);
assert.strictEqual(records[0].address, '10.0.0.0');
assert.strictEqual(records[256].address, '10.0.1.0');
}
Loading