Coverage Report

Created: 2026-09-13 07:08

next uncovered line (L), next uncovered region (R), next uncovered branch (B)
/src/pdns/pdns/dnsparser.cc
Line
Count
Source
1
/*
2
 * This file is part of PowerDNS or dnsdist.
3
 * Copyright -- PowerDNS.COM B.V. and its contributors
4
 *
5
 * This program is free software; you can redistribute it and/or modify
6
 * it under the terms of version 2 of the GNU General Public License as
7
 * published by the Free Software Foundation.
8
 *
9
 * In addition, for the avoidance of any doubt, permission is granted to
10
 * link this program with OpenSSL and to (re)distribute the binaries
11
 * produced as the result of such linking.
12
 *
13
 * This program is distributed in the hope that it will be useful,
14
 * but WITHOUT ANY WARRANTY; without even the implied warranty of
15
 * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE.  See the
16
 * GNU General Public License for more details.
17
 *
18
 * You should have received a copy of the GNU General Public License
19
 * along with this program; if not, write to the Free Software
20
 * Foundation, Inc., 51 Franklin Street, Fifth Floor, Boston, MA 02110-1301 USA.
21
 */
22
#include "dnsparser.hh"
23
#include "dnswriter.hh"
24
#include <boost/algorithm/string.hpp>
25
#include <boost/format.hpp>
26
#include <cstdint>
27
#include <stdexcept>
28
#include <string>
29
30
#include "dns_random.hh"
31
#include "namespaces.hh"
32
#include "noinitvector.hh"
33
34
std::atomic<bool> DNSRecordContent::d_locked{false};
35
36
UnknownRecordContent::UnknownRecordContent(const string& zone)
37
712
{
38
  // The expected format is '\#', followed by the length in decimal, and as
39
  // many pairs of hex digits as the length, which may be separated by
40
  // whitespace.
41
  // Because of this, using stringtok() might be horribly suboptimal for large
42
  // data with every byte separated by spaces.
43
44
  // The following is equivalent to strintok(parts, zone), but stops after
45
  // filling two parts, and stores the beginning of the actual payload for
46
  // further consumption.
47
712
  constexpr const char *delimiters = " \t\n";
48
712
  std::vector<std::string> parts;
49
712
  parts.reserve(2);
50
712
  std::string::size_type pos{0};
51
712
  {
52
712
    const auto len = zone.length();
53
54
1.31k
    while (pos<len) {
55
      // eat leading whitespace
56
1.30k
      pos = zone.find_first_not_of (delimiters, pos);
57
1.30k
      if (pos == string::npos) {
58
2
        break;   // nothing left but white space
59
2
      }
60
61
      // find the end of the token
62
1.29k
      std::string::size_type epos = zone.find_first_of (delimiters, pos);
63
64
      // push token
65
1.29k
      if (epos == std::string::npos) {
66
359
        parts.push_back (zone.substr(pos));
67
359
        pos = epos;
68
359
        break;
69
359
      }
70
940
      parts.push_back (zone.substr(pos, epos-pos));
71
      // set up for next loop
72
940
      pos = epos + 1;
73
940
      if (parts.size() == 2) {
74
338
        break;
75
338
      }
76
940
    }
77
712
  }
78
79
712
  if (parts.empty() || parts.at(0) != "\\#") {
80
173
    throw MOADNSException("Unknown record was stored incorrectly, should start with '\\#'");
81
173
  }
82
83
539
  if (parts.size() < 2) {
84
4
    throw MOADNSException("Unknown record was stored incorrectly, missing size field");
85
4
  }
86
535
  auto total = pdns::checked_stoi<unsigned long>(parts.at(1));
87
535
  if (total == 0) {
88
76
    if (pos != std::string::npos) {
89
20
      throw MOADNSException("Unknown record was stored incorrectly, spurious data after zero size field");
90
20
    }
91
56
    return;
92
76
  }
93
94
459
  if (total > std::numeric_limits<uint16_t>::max()) {
95
130
    throw MOADNSException((boost::format("invalid unknown record length size (%d)") % total).str());
96
130
  }
97
98
329
  std::string out;
99
329
  out.reserve(total);
100
101
  // This loops mimics stringtok() again
102
329
  unsigned int byte = 0;
103
1.25k
  while (byte < total) {
104
    // eat leading whitespace
105
1.07k
    pos = zone.find_first_not_of (delimiters, pos);
106
1.07k
    if (pos == std::string::npos) { // nothing left but white space
107
78
      throw MOADNSException("Unknown record was stored incorrectly, truncated after byte " + std::to_string(byte) + " of " + std::to_string(total));
108
78
    }
109
110
    // find the end of the token
111
997
    std::string::size_type epos = zone.find_first_of (delimiters, pos);
112
113
    // extract token
114
997
    std::string_view chunk{};
115
997
    if (epos == std::string::npos) {
116
      // TODO: replace with zone.subview(pos) once we can use C++26
117
211
      chunk = std::string_view(&zone.at(pos));
118
211
      pos = epos;
119
786
    } else {
120
      // TODO: replace with zone.subview(pos, epos-pos) once we can use C++26
121
786
      chunk = std::string_view(&zone.at(pos), epos-pos);
122
786
    }
123
124
    // process token
125
997
    if ((chunk.size() % 2) != 0) {
126
33
      throw MOADNSException("Unknown record was stored incorrectly, sequence of digits for byte " + std::to_string(byte) + " of " + std::to_string(total) + " onward at offset " + std::to_string(pos) + " has uneven length");
127
33
    }
128
964
    if ((chunk.size() / 2) > total - byte) {
129
22
      throw MOADNSException("Unknown record was stored incorrectly, sequence of digits for byte " + std::to_string(byte) + " of " + std::to_string(total) + " onward at offset " + std::to_string(pos) + " is too long");
130
22
    }
131
4.08k
    for (std::string_view::size_type subpos = 0; subpos < chunk.size(); subpos += 2) {
132
3.16k
      int chr{0};
133
3.16k
      if (sscanf(&chunk.at(subpos), "%02x", &chr) != 1) {
134
19
        throw MOADNSException("unable to read data for byte " + std::to_string(byte) + " of " + std::to_string(total) + " at offset " + std::to_string(pos + subpos) + " from unknown record");
135
19
      }
136
3.14k
      out.append(1, static_cast<char>(chr));
137
3.14k
      ++byte;
138
3.14k
    }
139
140
    // set up for next loop
141
923
    if (epos == std::string::npos) {
142
167
      pos = epos;
143
167
    }
144
756
    else {
145
756
      pos = epos + 1;
146
756
    }
147
923
  }
148
149
177
  d_record.insert(d_record.end(), out.begin(), out.end());
150
177
}
151
152
string UnknownRecordContent::getZoneRepresentation(bool /* noDot */) const
153
480
{
154
480
  ostringstream str;
155
480
  str<<"\\# "<<(unsigned int)d_record.size();
156
480
  if (!d_record.empty()) {
157
473
    std::array<char,4> hex{};
158
473
    str << " ";
159
1.09M
    for (auto byte : d_record) {
160
1.09M
      snprintf(hex.data(), hex.size(), "%02x", byte);
161
1.09M
      str << hex.data();
162
1.09M
    }
163
473
  }
164
480
  return str.str();
165
480
}
166
167
void UnknownRecordContent::toPacket(DNSPacketWriter& pw) const
168
31
{
169
31
  pw.xfrBlob(string(d_record.begin(),d_record.end()));
170
31
}
171
172
shared_ptr<DNSRecordContent> DNSRecordContent::deserialize(const DNSName& qname, uint16_t qtype, const string& serialized, uint16_t qclass, bool internalRepresentation)
173
13.5k
{
174
13.5k
  dnsheader dnsheader;
175
13.5k
  memset(&dnsheader, 0, sizeof(dnsheader));
176
13.5k
  dnsheader.qdcount=htons(1);
177
13.5k
  dnsheader.ancount=htons(1);
178
179
13.5k
  PacketBuffer packet; // build pseudo packet
180
  /* will look like: dnsheader, 5 bytes, encoded qname, dns record header, serialized data */
181
13.5k
  const auto& encoded = qname.getStorage();
182
13.5k
  packet.resize(sizeof(dnsheader) + 5 + encoded.size() + sizeof(struct dnsrecordheader) + serialized.size());
183
184
13.5k
  uint16_t pos=0;
185
13.5k
  memcpy(&packet[0], &dnsheader, sizeof(dnsheader)); pos+=sizeof(dnsheader);
186
187
13.5k
  constexpr std::array<uint8_t, 5> tmp= {'\x0', '\x0', '\x1', '\x0', '\x1' }; // root question for ns_t_a
188
13.5k
  memcpy(&packet[pos], tmp.data(), tmp.size()); pos += tmp.size();
189
190
13.5k
  memcpy(&packet[pos], encoded.c_str(), encoded.size()); pos+=(uint16_t)encoded.size();
191
192
13.5k
  struct dnsrecordheader drh;
193
13.5k
  drh.d_type=htons(qtype);
194
13.5k
  drh.d_class=htons(qclass);
195
13.5k
  drh.d_ttl=0;
196
13.5k
  drh.d_clen=htons(serialized.size());
197
198
13.5k
  memcpy(&packet[pos], &drh, sizeof(drh)); pos+=sizeof(drh);
199
13.5k
  if (!serialized.empty()) {
200
13.5k
    memcpy(&packet[pos], serialized.c_str(), serialized.size());
201
13.5k
    pos += (uint16_t) serialized.size();
202
13.5k
    (void) pos;
203
13.5k
  }
204
205
13.5k
  DNSRecord dr;
206
13.5k
  dr.d_class = qclass;
207
13.5k
  dr.d_type = qtype;
208
13.5k
  dr.d_name = qname;
209
13.5k
  dr.d_clen = serialized.size();
210
  // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast): packet.data() is uint8_t *
211
13.5k
  PacketReader reader(std::string_view(reinterpret_cast<const char*>(packet.data()), packet.size()), packet.size() - serialized.size() - sizeof(dnsrecordheader), internalRepresentation);
212
  /* needed to get the record boundaries right */
213
13.5k
  reader.getDnsrecordheader(drh);
214
13.5k
  auto content = DNSRecordContent::make(dr, reader, Opcode::Query);
215
13.5k
  return content;
216
13.5k
}
217
218
std::shared_ptr<DNSRecordContent> DNSRecordContent::make(const DNSRecord& dr,
219
                                                         PacketReader& pr)
220
192
{
221
192
  uint16_t searchclass = (dr.d_type == QType::OPT) ? 1 : dr.d_class; // class is invalid for OPT
222
223
192
  auto i = getTypemap().find(pair(searchclass, dr.d_type));
224
192
  if(i==getTypemap().end() || !i->second) {
225
0
    return std::make_shared<UnknownRecordContent>(dr, pr);
226
0
  }
227
228
192
  return i->second(dr, pr);
229
192
}
230
231
std::shared_ptr<DNSRecordContent> DNSRecordContent::make(uint16_t qtype, uint16_t qclass,
232
                                                         const string& content)
233
9.81k
{
234
9.81k
  auto i = getZmakermap().find(pair(qclass, qtype));
235
9.81k
  if(i==getZmakermap().end()) {
236
422
    return std::make_shared<UnknownRecordContent>(content);
237
422
  }
238
239
  // If the record type is known, but is provided as raw data, we need to first
240
  // parse it as an UnknownRecordContent, and then pretend the data is coming
241
  // from the wire.
242
9.39k
  if (content.length() >= 3 && content.at(0) == '\\' && content.at(1) == '#' && isspace(content.at(2)) != 0) {
243
290
    UnknownRecordContent urc(content);
244
290
    const auto& rawdata = urc.getRawContent();
245
    // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast): rawdata.data() is uint8_t *
246
290
    PacketReader reader(std::string_view(reinterpret_cast<const char *>(rawdata.data()), rawdata.size()), 0, false, true /* standalone */);
247
290
    DNSRecord rec;
248
290
    rec.d_class = qclass;
249
290
    rec.d_type = qtype;
250
290
    rec.d_ttl = 0;
251
290
    rec.d_clen = rawdata.size();
252
290
    return make(rec, reader);
253
290
  }
254
255
9.10k
  return i->second(content);
256
9.39k
}
257
258
std::shared_ptr<DNSRecordContent> DNSRecordContent::make(const DNSRecord& dr, PacketReader& pr, uint16_t oc)
259
196k
{
260
  // For opcode UPDATE and where the DNSRecord is an answer record, we don't care about content, because this is
261
  // not used within the prerequisite section of RFC2136, so - we can simply use unknownrecordcontent.
262
  // For section 3.2.3, we do need content so we need to get it properly. But only for the correct QClasses.
263
196k
  if (oc == Opcode::Update && dr.d_place == DNSResourceRecord::ANSWER && dr.d_class != 1)
264
5.25k
    return std::make_shared<UnknownRecordContent>(dr, pr);
265
266
191k
  uint16_t searchclass = (dr.d_type == QType::OPT) ? 1 : dr.d_class; // class is invalid for OPT
267
268
191k
  auto i = getTypemap().find(pair(searchclass, dr.d_type));
269
191k
  if(i==getTypemap().end() || !i->second) {
270
60.9k
    return std::make_shared<UnknownRecordContent>(dr, pr);
271
60.9k
  }
272
273
130k
  return i->second(dr, pr);
274
191k
}
275
276
0
string DNSRecordContent::upgradeContent(const DNSName& qname, const QType& qtype, const string& content) {
277
  // seamless upgrade for previously unsupported but now implemented types.
278
0
  UnknownRecordContent unknown_content(content);
279
0
  shared_ptr<DNSRecordContent> rc = DNSRecordContent::deserialize(qname, qtype.getCode(), unknown_content.serialize(qname));
280
0
  return rc->getZoneRepresentation();
281
0
}
282
283
DNSRecordContent::typemap_t& DNSRecordContent::getTypemap()
284
383k
{
285
383k
  static DNSRecordContent::typemap_t typemap;
286
383k
  return typemap;
287
383k
}
288
289
DNSRecordContent::n2typemap_t& DNSRecordContent::getN2Typemap()
290
295k
{
291
295k
  static DNSRecordContent::n2typemap_t n2typemap;
292
295k
  return n2typemap;
293
295k
}
294
295
DNSRecordContent::t2namemap_t& DNSRecordContent::getT2Namemap()
296
36.8k
{
297
36.8k
  static DNSRecordContent::t2namemap_t t2namemap;
298
36.8k
  return t2namemap;
299
36.8k
}
300
301
DNSRecordContent::zmakermap_t& DNSRecordContent::getZmakermap()
302
20.0k
{
303
20.0k
  static DNSRecordContent::zmakermap_t zmakermap;
304
20.0k
  return zmakermap;
305
20.0k
}
306
307
bool DNSRecordContent::isRegisteredType(uint16_t rtype, uint16_t rclass)
308
0
{
309
0
  return getTypemap().count(pair(rclass, rtype)) != 0;
310
0
}
311
312
0
DNSRecord::DNSRecord(const DNSResourceRecord& rr): d_name(rr.qname)
313
0
{
314
0
  d_type = rr.qtype.getCode();
315
0
  d_ttl = rr.ttl;
316
0
  d_class = rr.qclass;
317
0
  d_place = DNSResourceRecord::ANSWER;
318
0
  d_clen = 0;
319
0
  d_content = DNSRecordContent::make(d_type, rr.qclass, rr.content);
320
0
}
321
322
// If you call this and you are not parsing a packet coming from a socket, you are doing it wrong.
323
DNSResourceRecord DNSResourceRecord::fromWire(const DNSRecord& wire)
324
0
{
325
0
  DNSResourceRecord resourceRecord;
326
0
  resourceRecord.qname = wire.d_name;
327
0
  resourceRecord.qtype = QType(wire.d_type);
328
0
  resourceRecord.ttl = wire.d_ttl;
329
0
  resourceRecord.content = wire.getContent()->getZoneRepresentation(true);
330
0
  resourceRecord.auth = false;
331
0
  resourceRecord.qclass = wire.d_class;
332
0
  return resourceRecord;
333
0
}
334
335
void MOADNSParser::init(bool query, const std::string_view& packet)
336
15.7k
{
337
15.7k
  if (packet.size() < sizeof(dnsheader))
338
16
    throw MOADNSException("Packet shorter than minimal header");
339
340
15.7k
  memcpy(&d_header, packet.data(), sizeof(dnsheader));
341
342
15.7k
  if(d_header.opcode != Opcode::Query && d_header.opcode != Opcode::Notify && d_header.opcode != Opcode::Update)
343
12
    throw MOADNSException("Can't parse non-query packet with opcode="+ std::to_string(d_header.opcode));
344
345
15.7k
  d_header.qdcount=ntohs(d_header.qdcount);
346
15.7k
  d_header.ancount=ntohs(d_header.ancount);
347
15.7k
  d_header.nscount=ntohs(d_header.nscount);
348
15.7k
  d_header.arcount=ntohs(d_header.arcount);
349
350
15.7k
  if (query && (d_header.qdcount > 1))
351
248
    throw MOADNSException("Query with QD > 1 ("+std::to_string(d_header.qdcount)+")");
352
353
15.5k
  unsigned int n=0;
354
355
15.5k
  PacketReader pr(packet);
356
15.5k
  bool validPacket=false;
357
15.5k
  try {
358
15.5k
    d_qtype = d_qclass = 0; // sometimes replies come in with no question, don't present garbage then
359
360
64.0k
    for(n=0;n < d_header.qdcount; ++n) {
361
48.5k
      d_qname=pr.getName();
362
48.5k
      d_qtype=pr.get16BitInt();
363
48.5k
      d_qclass=pr.get16BitInt();
364
48.5k
    }
365
366
15.5k
    struct dnsrecordheader ah;
367
15.5k
    vector<unsigned char> record;
368
15.5k
    bool seenTSIG = false;
369
15.5k
    validPacket=true;
370
15.5k
    unsigned int supposedRecordCount = d_header.ancount + d_header.nscount + d_header.arcount;
371
    // No need to reserve more memory for more records than the request can
372
    // contain. We could try to be smarter and actually count the records
373
    // by doing getDnsrecordheader and skip the payload in a loop, but all
374
    // we really want here is to avoid reserving too much memory in case of
375
    // maliciously high record counts.
376
15.5k
    auto reserveRecordCount = std::min(1 + (packet.size() - sizeof(dnsheader)) / sizeof(dnsrecordheader), static_cast<size_t>(supposedRecordCount));
377
15.5k
    d_answers.reserve(reserveRecordCount);
378
259k
    for (n = 0; n < supposedRecordCount; ++n) {
379
244k
      DNSRecord dr;
380
381
244k
      if(n < d_header.ancount)
382
160k
        dr.d_place=DNSResourceRecord::ANSWER;
383
83.6k
      else if(n < d_header.ancount + d_header.nscount)
384
47.2k
        dr.d_place=DNSResourceRecord::AUTHORITY;
385
36.4k
      else
386
36.4k
        dr.d_place=DNSResourceRecord::ADDITIONAL;
387
388
244k
      unsigned int recordStartPos=pr.getPosition();
389
390
244k
      DNSName name=pr.getName();
391
392
244k
      pr.getDnsrecordheader(ah);
393
244k
      dr.d_ttl=ah.d_ttl;
394
244k
      dr.d_type=ah.d_type;
395
244k
      dr.d_class=ah.d_class;
396
397
244k
      dr.d_name = std::move(name);
398
244k
      dr.d_clen = ah.d_clen;
399
400
244k
      if (query &&
401
55.6k
          !(d_qtype == QType::IXFR && dr.d_place == DNSResourceRecord::AUTHORITY && dr.d_type == QType::SOA) && // IXFR queries have a SOA in their AUTHORITY section
402
54.7k
          (dr.d_place == DNSResourceRecord::ANSWER || dr.d_place == DNSResourceRecord::AUTHORITY || (dr.d_type != QType::OPT && dr.d_type != QType::TSIG && dr.d_type != QType::SIG && dr.d_type != QType::TKEY) || ((dr.d_type == QType::TSIG || dr.d_type == QType::SIG || dr.d_type == QType::TKEY) && dr.d_class != QClass::ANY))) {
403
//        cerr<<"discarding RR, query is "<<query<<", place is "<<dr.d_place<<", type is "<<dr.d_type<<", class is "<<dr.d_class<<endl;
404
52.8k
        dr.setContent(std::make_shared<UnknownRecordContent>(dr, pr));
405
52.8k
      }
406
191k
      else {
407
//        cerr<<"parsing RR, query is "<<query<<", place is "<<dr.d_place<<", type is "<<dr.d_type<<", class is "<<dr.d_class<<endl;
408
191k
        dr.setContent(DNSRecordContent::make(dr, pr, d_header.opcode));
409
191k
      }
410
411
244k
      if (dr.d_place == DNSResourceRecord::ADDITIONAL && seenTSIG) {
412
192
        throw MOADNSException("Packet ("+d_qname.toString()+"|#"+std::to_string(d_qtype)+") has an unexpected record ("+std::to_string(dr.d_type)+") after a TSIG one.");
413
192
      }
414
415
243k
      if(dr.d_type == QType::TSIG && dr.d_class == QClass::ANY) {
416
552
        if(seenTSIG || dr.d_place != DNSResourceRecord::ADDITIONAL) {
417
297
          throw MOADNSException("Packet ("+d_qname.toLogString()+"|#"+std::to_string(d_qtype)+") has a TSIG record in an invalid position.");
418
297
        }
419
255
        seenTSIG = true;
420
255
        d_tsigPos = recordStartPos;
421
255
      }
422
423
243k
      d_answers.emplace_back(std::move(dr));
424
243k
    }
425
426
#if 0
427
    if(pr.getPosition()!=packet.size()) {
428
      throw MOADNSException("Packet ("+d_qname+"|#"+std::to_string(d_qtype)+") has trailing garbage ("+ std::to_string(pr.getPosition()) + " < " +
429
                            std::to_string(packet.size()) + ")");
430
    }
431
#endif
432
15.5k
  }
433
15.5k
  catch(const std::out_of_range &re) {
434
11.0k
    if(validPacket && d_header.tc) { // don't sweat it over truncated packets, but do adjust an, ns and arcount
435
2.29k
      if(n < d_header.ancount) {
436
1.43k
        d_header.ancount=n; d_header.nscount = d_header.arcount = 0;
437
1.43k
      }
438
861
      else if(n < d_header.ancount + d_header.nscount) {
439
456
        d_header.nscount = n - d_header.ancount; d_header.arcount=0;
440
456
      }
441
405
      else {
442
405
        d_header.arcount = n - d_header.ancount - d_header.nscount;
443
405
      }
444
2.29k
    }
445
8.78k
    else {
446
8.78k
      throw MOADNSException("Error parsing packet of "+std::to_string(packet.size())+" bytes (rd="+
447
8.78k
                            std::to_string(d_header.rd)+
448
8.78k
                            "), out of bounds: "+string(re.what()));
449
8.78k
    }
450
11.0k
  }
451
15.5k
}
452
453
bool MOADNSParser::hasEDNS() const
454
0
{
455
0
  if (d_header.arcount == 0 || d_answers.empty()) {
456
0
    return false;
457
0
  }
458
459
0
  for (const auto& record : d_answers) {
460
0
    if (record.d_place == DNSResourceRecord::ADDITIONAL && record.d_type == QType::OPT) {
461
0
      return true;
462
0
    }
463
0
  }
464
465
0
  return false;
466
0
}
467
468
void PacketReader::getDnsrecordheader(struct dnsrecordheader &ah)
469
251k
{
470
251k
  unsigned char *p = reinterpret_cast<unsigned char*>(&ah);
471
472
2.76M
  for(unsigned int n = 0; n < sizeof(dnsrecordheader); ++n) {
473
2.51M
    p[n] = d_content.at(d_pos++);
474
2.51M
  }
475
476
251k
  ah.d_type = ntohs(ah.d_type);
477
251k
  ah.d_class = ntohs(ah.d_class);
478
251k
  ah.d_clen = ntohs(ah.d_clen);
479
251k
  ah.d_ttl = ntohl(ah.d_ttl);
480
481
251k
  d_startrecordpos = d_pos; // needed for getBlob later on
482
251k
  d_recordlen = ah.d_clen;
483
251k
  if (d_pos > d_content.size() || (d_content.size() - d_pos) < (d_recordlen)) {
484
1.16k
    throw std::out_of_range("DNS record length (" + std::to_string(d_recordlen) + " starting at " + std::to_string(d_pos) + ") goes beyond the packet's content (" + std::to_string(d_content.size()) + ")");
485
1.16k
  }
486
251k
}
487
488
489
void PacketReader::copyRecord(vector<unsigned char>& dest, uint16_t len)
490
119k
{
491
119k
  if (len == 0) {
492
73.2k
    return;
493
73.2k
  }
494
45.8k
  if ((d_pos + len) > d_content.size()) {
495
0
    throw std::out_of_range("Attempt to copy outside of packet");
496
0
  }
497
498
45.8k
  dest.resize(len);
499
500
13.9M
  for (uint16_t n = 0; n < len; ++n) {
501
13.8M
    dest.at(n) = d_content.at(d_pos++);
502
13.8M
  }
503
45.8k
}
504
505
void PacketReader::copyRecord(unsigned char* dest, uint16_t len)
506
1.65k
{
507
1.65k
  if (d_pos + len > d_content.size()) {
508
0
    throw std::out_of_range("Attempt to copy outside of packet");
509
0
  }
510
511
1.65k
  memcpy(dest, &d_content.at(d_pos), len);
512
1.65k
  d_pos += len;
513
1.65k
}
514
515
void PacketReader::xfrNodeOrLocatorID(NodeOrLocatorID& ret)
516
3.09k
{
517
3.09k
  if (d_pos + sizeof(ret) > d_content.size()) {
518
55
    throw std::out_of_range("Attempt to read 64 bit value outside of packet");
519
55
  }
520
3.04k
  memcpy(&ret.content, &d_content.at(d_pos), sizeof(ret.content));
521
3.04k
  d_pos += sizeof(ret);
522
3.04k
}
523
524
void PacketReader::xfr48BitInt(uint64_t& ret)
525
719
{
526
719
  ret=0;
527
719
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
528
719
  ret<<=8;
529
719
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
530
719
  ret<<=8;
531
719
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
532
719
  ret<<=8;
533
719
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
534
719
  ret<<=8;
535
719
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
536
719
  ret<<=8;
537
719
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
538
719
}
539
540
uint32_t PacketReader::get32BitInt()
541
36.9k
{
542
36.9k
  uint32_t ret=0;
543
36.9k
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
544
36.9k
  ret<<=8;
545
36.9k
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
546
36.9k
  ret<<=8;
547
36.9k
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
548
36.9k
  ret<<=8;
549
36.9k
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
550
551
36.9k
  return ret;
552
36.9k
}
553
554
555
uint16_t PacketReader::get16BitInt()
556
1.02M
{
557
1.02M
  uint16_t ret=0;
558
1.02M
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
559
1.02M
  ret<<=8;
560
1.02M
  ret+=static_cast<uint8_t>(d_content.at(d_pos++));
561
562
1.02M
  return ret;
563
1.02M
}
564
565
uint8_t PacketReader::get8BitInt()
566
512k
{
567
512k
  return d_content.at(d_pos++);
568
512k
}
569
570
DNSName PacketReader::getName()
571
362k
{
572
362k
  unsigned int consumed;
573
362k
  try {
574
362k
    uint16_t minOffset = d_standalone ? 0 : sizeof(dnsheader);
575
362k
    bool uncompress = !d_standalone;
576
362k
    DNSName name((const char*) d_content.data(), d_content.size(), d_pos, uncompress, nullptr /* qtype */, nullptr /* qclass */, &consumed, minOffset);
577
578
362k
    d_pos+=consumed;
579
362k
    return name;
580
362k
  }
581
362k
  catch(const std::range_error& re) {
582
9.35k
    throw std::out_of_range(string("dnsname issue: ")+re.what());
583
9.35k
  }
584
362k
  catch(...) {
585
0
    throw std::out_of_range("dnsname issue");
586
0
  }
587
0
  throw PDNSException("PacketReader::getName(): name is empty");
588
362k
}
589
590
// FIXME see #6010 and #3503 if you want a proper solution
591
string txtEscape(const string &name)
592
619k
{
593
619k
  string ret;
594
619k
  std::array<char, 5> ebuf{};
595
596
10.5M
  for (char letter : name) {
597
10.5M
    const unsigned uch = static_cast<unsigned char>(letter);
598
10.5M
    if (uch >= 127 || uch < 32) {
599
4.79M
      snprintf(ebuf.data(), ebuf.size(), "\\%03u", uch);
600
4.79M
      ret += ebuf.data();
601
4.79M
    }
602
5.74M
    else if (letter == '"' || letter == '\\'){
603
311k
      ret += '\\';
604
311k
      ret += letter;
605
311k
    }
606
5.42M
    else {
607
5.42M
      ret += letter;
608
5.42M
    }
609
10.5M
  }
610
619k
  return ret;
611
619k
}
612
613
// exceptions thrown here do not result in logging in the main pdns auth server - just so you know!
614
string PacketReader::getText(bool multi, bool lenField)
615
20.9k
{
616
20.9k
  string ret;
617
20.9k
  ret.reserve(40);
618
423k
  while (d_pos < d_startrecordpos + d_recordlen ) {
619
404k
    if (!ret.empty()) {
620
395k
      ret.append(1,' ');
621
395k
    }
622
404k
    uint16_t labellen;
623
404k
    if (lenField) {
624
401k
      labellen = static_cast<uint8_t>(d_content.at(d_pos++));
625
401k
    }
626
3.79k
    else {
627
3.79k
      labellen = d_recordlen - (d_pos - d_startrecordpos);
628
3.79k
    }
629
630
404k
    const uint16_t remaining = (d_startrecordpos + d_recordlen) - d_pos;
631
404k
    if (labellen > remaining) {
632
605
      throw std::out_of_range("label length in text record exceeds record boundary");
633
605
    }
634
635
404k
    ret.append(1, '"');
636
404k
    if (labellen) { // no need to do anything for an empty string
637
106k
      string val(&d_content.at(d_pos), &d_content.at(d_pos + labellen - 1) + 1);
638
106k
      ret.append(txtEscape(val)); // the end is one beyond the packet
639
106k
    }
640
404k
    ret.append(1, '"');
641
404k
    d_pos += labellen;
642
404k
    if (!multi) {
643
1.75k
      break;
644
1.75k
    }
645
404k
  }
646
647
20.3k
  if (ret.empty() && !lenField) {
648
    // all lenField == false cases (CAA and URI at the time of this writing) want that emptiness to be explicit
649
2.43k
    return "\"\"";
650
2.43k
  }
651
17.8k
  return ret;
652
20.3k
}
653
654
string PacketReader::getUnquotedText(bool lenField)
655
4.34k
{
656
4.34k
  uint16_t stop_at{};
657
4.34k
  if (lenField) {
658
4.34k
    stop_at = static_cast<uint8_t>(d_content.at(d_pos)) + d_pos + 1;
659
4.34k
  }
660
0
  else {
661
0
    stop_at = d_recordlen;
662
0
  }
663
664
  /* think unsigned overflow */
665
4.34k
  if (stop_at < d_pos) {
666
2
    throw std::out_of_range("getUnquotedText out of record range");
667
2
  }
668
669
  /* Validate against record boundary */
670
4.34k
  const uint16_t recordEnd = d_startrecordpos + d_recordlen;
671
4.34k
  if (stop_at > recordEnd) {
672
141
    throw std::out_of_range("getUnquotedText: length exceeds record boundary");
673
141
  }
674
675
4.20k
  if (stop_at == d_pos) {
676
0
    return "";
677
0
  }
678
679
4.20k
  d_pos++;
680
4.20k
  string ret(d_content.substr(d_pos, stop_at-d_pos));
681
4.20k
  d_pos = stop_at;
682
4.20k
  return ret;
683
4.20k
}
684
685
void PacketReader::xfrBlob(string& blob)
686
45.4k
{
687
45.4k
  try {
688
45.4k
    if(d_recordlen && !(d_pos == (d_startrecordpos + d_recordlen))) {
689
15.5k
      if (d_pos > (d_startrecordpos + d_recordlen)) {
690
21
        throw std::out_of_range("xfrBlob out of record range");
691
21
      }
692
15.5k
      blob.assign(&d_content.at(d_pos), &d_content.at(d_startrecordpos + d_recordlen - 1 ) + 1);
693
15.5k
    }
694
29.8k
    else {
695
29.8k
      blob.clear();
696
29.8k
    }
697
698
45.4k
    d_pos = d_startrecordpos + d_recordlen;
699
45.4k
  }
700
45.4k
  catch(...)
701
45.4k
  {
702
21
    throw std::out_of_range("xfrBlob out of range");
703
21
  }
704
45.4k
}
705
706
2.69k
void PacketReader::xfrBlobNoSpaces(string& blob, int length) {
707
2.69k
  xfrBlob(blob, length);
708
2.69k
}
709
710
void PacketReader::xfrBlob(string& blob, int length)
711
596k
{
712
596k
  if(length) {
713
591k
    if (length < 0) {
714
0
      throw std::out_of_range("xfrBlob out of range (negative length)");
715
0
    }
716
591k
    auto available = (d_startrecordpos + d_recordlen) - d_pos;
717
591k
    if (available < length) {
718
584
      throw std::out_of_range("xfrBlob out of range (excessive length)");
719
584
    }
720
721
591k
    blob.assign(&d_content.at(d_pos), &d_content.at(d_pos + length - 1 ) + 1 );
722
723
591k
    d_pos += length;
724
591k
  }
725
5.05k
  else {
726
5.05k
    blob.clear();
727
5.05k
  }
728
596k
}
729
730
12.0k
void PacketReader::xfrSvcParamKeyVals(set<SvcParam> &kvs) {
731
12.0k
  int32_t lastKey{-1}; // Keep track of the last key, as ordering should be strict
732
733
25.6k
  while (d_pos < (d_startrecordpos + d_recordlen)) {
734
14.6k
    if (d_pos + 2 > (d_startrecordpos + d_recordlen)) {
735
75
      throw std::out_of_range("incomplete key");
736
75
    }
737
14.5k
    uint16_t keyInt;
738
14.5k
    xfr16BitInt(keyInt);
739
740
14.5k
    if (keyInt <= lastKey) {
741
235
      throw std::out_of_range("Found SVCParamKey " + std::to_string(keyInt) + " after SVCParamKey " + std::to_string(lastKey));
742
235
    }
743
14.2k
    lastKey = keyInt;
744
745
14.2k
    auto key = static_cast<SvcParam::SvcParamKey>(keyInt);
746
747
14.2k
    uint16_t len;
748
14.2k
    xfr16BitInt(len);
749
750
14.2k
    if (d_pos + len > (d_startrecordpos + d_recordlen)) {
751
351
      throw std::out_of_range("record is shorter than SVCB lengthfield implies");
752
351
    }
753
754
13.9k
    switch (key)
755
13.9k
    {
756
1.66k
    case SvcParam::mandatory: {
757
1.66k
      if (len % 2 != 0) {
758
6
        throw std::out_of_range("mandatory SvcParam has invalid length");
759
6
      }
760
1.65k
      if (len == 0) {
761
53
        throw std::out_of_range("empty 'mandatory' values");
762
53
      }
763
1.60k
      std::set<SvcParam::SvcParamKey> paramKeys;
764
1.60k
      size_t stop = d_pos + len;
765
470k
      while (d_pos < stop) {
766
468k
        uint16_t keyval;
767
468k
        xfr16BitInt(keyval);
768
468k
        paramKeys.insert(static_cast<SvcParam::SvcParamKey>(keyval));
769
468k
      }
770
1.60k
      kvs.insert(SvcParam(key, std::move(paramKeys)));
771
1.60k
      break;
772
1.65k
    }
773
2.03k
    case SvcParam::alpn: {
774
2.03k
      size_t stop = d_pos + len;
775
2.03k
      std::vector<string> alpns;
776
277k
      while (d_pos < stop) {
777
275k
        string alpn;
778
275k
        uint8_t alpnLen = 0;
779
275k
        xfr8BitInt(alpnLen);
780
275k
        if (alpnLen == 0) {
781
50
          throw std::out_of_range("alpn length of 0");
782
50
        }
783
275k
        if (d_pos + alpnLen > stop) {
784
106
          throw std::out_of_range("alpn length is larger than rest of alpn SVC Param");
785
106
        }
786
275k
        xfrBlob(alpn, alpnLen);
787
275k
        alpns.push_back(std::move(alpn));
788
275k
      }
789
1.87k
      kvs.insert(SvcParam(key, std::move(alpns)));
790
1.87k
      break;
791
2.03k
    }
792
298
    case SvcParam::ohttp:
793
617
    case SvcParam::no_default_alpn: {
794
617
      if (len != 0) {
795
8
        throw std::out_of_range("invalid length for " + SvcParam::keyToString(key));
796
8
      }
797
609
      kvs.insert(SvcParam(key));
798
609
      break;
799
617
    }
800
992
    case SvcParam::port: {
801
992
      if (len != 2) {
802
4
        throw std::out_of_range("invalid length for port");
803
4
      }
804
988
      uint16_t port;
805
988
      xfr16BitInt(port);
806
988
      kvs.insert(SvcParam(key, port));
807
988
      break;
808
992
    }
809
768
    case SvcParam::ipv4hint:
810
1.39k
    case SvcParam::ipv6hint: {
811
1.39k
      size_t addrLen = (key == SvcParam::ipv4hint ? 4 : 16);
812
1.39k
      if (len % addrLen != 0) {
813
11
        throw std::out_of_range("invalid length for " + SvcParam::keyToString(key));
814
11
      }
815
1.38k
      vector<ComboAddress> addresses;
816
1.38k
      auto stop = d_pos + len;
817
309k
      while (d_pos < stop)
818
308k
      {
819
308k
        ComboAddress addr;
820
308k
        xfrCAWithoutPort(key, addr);
821
308k
        addresses.push_back(addr);
822
308k
      }
823
      // If there were no addresses, and the input comes from internal
824
      // representation, we can reasonably assume this is the serialization
825
      // of "auto".
826
1.38k
      bool doAuto{d_internal && len == 0};
827
1.38k
      auto param = SvcParam(key, std::move(addresses));
828
1.38k
      param.setAutoHint(doAuto);
829
1.38k
      kvs.insert(std::move(param));
830
1.38k
      break;
831
1.39k
    }
832
1.27k
    case SvcParam::ech: {
833
1.27k
      std::string blob;
834
1.27k
      blob.reserve(len);
835
1.27k
      xfrBlobNoSpaces(blob, len);
836
1.27k
      kvs.insert(SvcParam(key, blob));
837
1.27k
      break;
838
1.39k
    }
839
1.28k
    case SvcParam::tls_supported_groups: {
840
1.28k
      if (len % 2 != 0) {
841
4
        throw std::out_of_range("invalid length for " + SvcParam::keyToString(key));
842
4
      }
843
1.28k
      vector<uint16_t> groups;
844
1.28k
      groups.reserve(len / 2);
845
1.28k
      auto stop = d_pos + len;
846
331k
      while (d_pos < stop)
847
330k
      {
848
330k
        uint16_t group = 0;
849
330k
        xfr16BitInt(group);
850
330k
        groups.push_back(group);
851
330k
      }
852
1.28k
      auto param = SvcParam(key, std::move(groups));
853
1.28k
      kvs.insert(std::move(param));
854
1.28k
      break;
855
1.28k
    }
856
4.57k
    default: {
857
4.57k
      std::string blob;
858
4.57k
      blob.reserve(len);
859
4.57k
      xfrBlob(blob, len);
860
4.57k
      kvs.insert(SvcParam(key, blob));
861
4.57k
      break;
862
1.28k
    }
863
13.9k
    }
864
13.9k
  }
865
12.0k
}
866
867
868
void PacketReader::xfrHexBlob(string& blob, bool /* keepReading */)
869
13.8k
{
870
13.8k
  xfrBlob(blob);
871
13.8k
}
872
873
//FIXME400 remove this method completely
874
string simpleCompress(const string& elabel, const string& root)
875
0
{
876
0
  string label=elabel;
877
  // FIXME400: this relies on the semi-canonical escaped output from getName
878
0
  if(strchr(label.c_str(), '\\')) {
879
0
    boost::replace_all(label, "\\.", ".");
880
0
    boost::replace_all(label, "\\032", " ");
881
0
    boost::replace_all(label, "\\\\", "\\");
882
0
  }
883
0
  typedef vector<pair<unsigned int, unsigned int> > parts_t;
884
0
  parts_t parts;
885
0
  vstringtok(parts, label, ".");
886
0
  string ret;
887
0
  ret.reserve(label.size()+4);
888
0
  for(const auto & part : parts) {
889
    // NOLINTNEXTLINE(cppcoreguidelines-pro-bounds-pointer-arithmetic)
890
0
    auto label_part = std::string_view(label.c_str() + part.first, 1 + label.length() - part.first); // also match trailing 0, hence '1 +'
891
0
    if(!root.empty() && pdns_ilexicographical_compare_three_way(root, label_part) == 0) {
892
0
      const unsigned char rootptr[2]={0xc0,0x11};
893
0
      ret.append((const char *) rootptr, 2);
894
0
      return ret;
895
0
    }
896
0
    ret.append(1, (char)(part.second - part.first));
897
0
    ret.append(label.c_str() + part.first, part.second - part.first);
898
0
  }
899
0
  ret.append(1, (char)0);
900
0
  return ret;
901
0
}
902
903
// method of operation: silently fail if it doesn't work - we're only trying to be nice, don't fall over on it
904
void editDNSPacketTTL(char* packet, size_t length, const std::function<uint32_t(uint8_t, uint16_t, uint16_t, uint32_t)>& visitor)
905
0
{
906
0
  if(length < sizeof(dnsheader))
907
0
    return;
908
0
  try
909
0
  {
910
0
    dnsheader dh;
911
0
    memcpy((void*)&dh, (const dnsheader*)packet, sizeof(dh));
912
0
    uint64_t numrecords = ntohs(dh.ancount) + ntohs(dh.nscount) + ntohs(dh.arcount);
913
0
    DNSPacketMangler dpm(packet, length);
914
915
0
    uint64_t n;
916
0
    for(n=0; n < ntohs(dh.qdcount) ; ++n) {
917
0
      dpm.skipDomainName();
918
      /* type and class */
919
0
      dpm.skipBytes(4);
920
0
    }
921
922
0
    for(n=0; n < numrecords; ++n) {
923
0
      dpm.skipDomainName();
924
925
0
      uint8_t section = n < ntohs(dh.ancount) ? 1 : (n < (ntohs(dh.ancount) + ntohs(dh.nscount)) ? 2 : 3);
926
0
      uint16_t dnstype = dpm.get16BitInt();
927
0
      uint16_t dnsclass = dpm.get16BitInt();
928
929
0
      if(dnstype == QType::OPT) // not getting near that one with a stick
930
0
        break;
931
932
0
      uint32_t dnsttl = dpm.get32BitInt();
933
0
      uint32_t newttl = visitor(section, dnsclass, dnstype, dnsttl);
934
0
      if (newttl) {
935
0
        dpm.rewindBytes(sizeof(newttl));
936
0
        dpm.setAndSkip32BitInt(newttl);
937
0
      }
938
0
      dpm.skipRData();
939
0
    }
940
0
  }
941
0
  catch(...)
942
0
  {
943
0
    return;
944
0
  }
945
0
}
946
947
static bool checkIfPacketContainsRecords(const PacketBuffer& packet, const std::unordered_set<QType>& qtypes)
948
0
{
949
0
  auto length = packet.size();
950
0
  if (length < sizeof(dnsheader)) {
951
0
    return false;
952
0
  }
953
954
0
  try {
955
0
    const dnsheader_aligned dh(packet.data());
956
0
    DNSPacketMangler dpm(const_cast<char*>(reinterpret_cast<const char*>(packet.data())), length);
957
958
0
    const uint16_t qdcount = ntohs(dh->qdcount);
959
0
    for (size_t n = 0; n < qdcount; ++n) {
960
0
      dpm.skipDomainName();
961
      /* type and class */
962
0
      dpm.skipBytes(4);
963
0
    }
964
0
    const size_t recordsCount = static_cast<size_t>(ntohs(dh->ancount)) + ntohs(dh->nscount) + ntohs(dh->arcount);
965
0
    for (size_t n = 0; n < recordsCount; ++n) {
966
0
      dpm.skipDomainName();
967
0
      uint16_t dnstype = dpm.get16BitInt();
968
0
      uint16_t dnsclass = dpm.get16BitInt();
969
0
      if (dnsclass == QClass::IN && qtypes.count(dnstype) > 0) {
970
0
        return true;
971
0
      }
972
      /* ttl */
973
0
      dpm.skipBytes(4);
974
0
      dpm.skipRData();
975
0
    }
976
0
  }
977
0
  catch (...) {
978
0
  }
979
980
0
  return false;
981
0
}
982
983
static int rewritePacketWithoutRecordTypes(const PacketBuffer& initialPacket, PacketBuffer& newContent, const std::unordered_set<QType>& qtypes)
984
0
{
985
0
  static const std::unordered_set<QType>& safeTypes{QType::A, QType::AAAA, QType::DHCID, QType::TXT, QType::OPT, QType::HINFO, QType::DNSKEY, QType::CDNSKEY, QType::DS, QType::CDS, QType::DLV, QType::SSHFP, QType::KEY, QType::CERT, QType::TLSA, QType::SMIMEA, QType::OPENPGPKEY, QType::SVCB, QType::HTTPS, QType::NSEC3, QType::CSYNC, QType::NSEC3PARAM, QType::LOC, QType::NID, QType::L32, QType::L64, QType::EUI48, QType::EUI64, QType::URI, QType::CAA};
986
987
0
  if (initialPacket.size() < sizeof(dnsheader)) {
988
0
    return EINVAL;
989
0
  }
990
0
  try {
991
0
    const dnsheader_aligned dh(initialPacket.data());
992
993
0
    if (ntohs(dh->qdcount) == 0)
994
0
      return ENOENT;
995
0
    auto packetView = std::string_view(reinterpret_cast<const char*>(initialPacket.data()), initialPacket.size());
996
997
0
    PacketReader pr(packetView);
998
999
0
    size_t idx = 0;
1000
0
    DNSName rrname;
1001
0
    uint16_t qdcount = ntohs(dh->qdcount);
1002
0
    uint16_t ancount = ntohs(dh->ancount);
1003
0
    uint16_t nscount = ntohs(dh->nscount);
1004
0
    uint16_t arcount = ntohs(dh->arcount);
1005
0
    uint16_t rrtype;
1006
0
    uint16_t rrclass;
1007
0
    string blob;
1008
0
    struct dnsrecordheader ah;
1009
1010
0
    rrname = pr.getName();
1011
0
    rrtype = pr.get16BitInt();
1012
0
    rrclass = pr.get16BitInt();
1013
1014
0
    GenericDNSPacketWriter<PacketBuffer> pw(newContent, rrname, rrtype, rrclass, dh->opcode);
1015
0
    pw.getHeader()->id=dh->id;
1016
0
    pw.getHeader()->qr=dh->qr;
1017
0
    pw.getHeader()->aa=dh->aa;
1018
0
    pw.getHeader()->tc=dh->tc;
1019
0
    pw.getHeader()->rd=dh->rd;
1020
0
    pw.getHeader()->ra=dh->ra;
1021
0
    pw.getHeader()->ad=dh->ad;
1022
0
    pw.getHeader()->cd=dh->cd;
1023
0
    pw.getHeader()->rcode=dh->rcode;
1024
1025
    /* consume remaining qd if any */
1026
0
    if (qdcount > 1) {
1027
0
      for(idx = 1; idx < qdcount; idx++) {
1028
0
        rrname = pr.getName();
1029
0
        rrtype = pr.get16BitInt();
1030
0
        rrclass = pr.get16BitInt();
1031
0
        (void) rrtype;
1032
0
        (void) rrclass;
1033
0
      }
1034
0
    }
1035
1036
    /* copy AN */
1037
0
    for (idx = 0; idx < ancount; idx++) {
1038
0
      rrname = pr.getName();
1039
0
      pr.getDnsrecordheader(ah);
1040
0
      pr.xfrBlob(blob);
1041
1042
0
      if (qtypes.find(ah.d_type) == qtypes.end()) {
1043
        // if this is not a safe type
1044
0
        if (safeTypes.find(ah.d_type) == safeTypes.end()) {
1045
          // "unsafe" types might contain compressed data, so cancel rewrite
1046
0
          newContent.clear();
1047
0
          return EIO;
1048
0
        }
1049
0
        pw.startRecord(rrname, ah.d_type, ah.d_ttl, ah.d_class, DNSResourceRecord::ANSWER, true);
1050
0
        pw.xfrBlob(blob);
1051
0
      }
1052
0
    }
1053
1054
    /* copy NS */
1055
0
    for (idx = 0; idx < nscount; idx++) {
1056
0
      rrname = pr.getName();
1057
0
      pr.getDnsrecordheader(ah);
1058
0
      pr.xfrBlob(blob);
1059
1060
0
      if (qtypes.find(ah.d_type) == qtypes.end()) {
1061
0
        if (safeTypes.find(ah.d_type) == safeTypes.end()) {
1062
          // "unsafe" types might contain compressed data, so cancel rewrite
1063
0
          newContent.clear();
1064
0
          return EIO;
1065
0
        }
1066
0
        pw.startRecord(rrname, ah.d_type, ah.d_ttl, ah.d_class, DNSResourceRecord::AUTHORITY, true);
1067
0
        pw.xfrBlob(blob);
1068
0
      }
1069
0
    }
1070
    /* copy AR */
1071
0
    for (idx = 0; idx < arcount; idx++) {
1072
0
      rrname = pr.getName();
1073
0
      pr.getDnsrecordheader(ah);
1074
0
      pr.xfrBlob(blob);
1075
1076
0
      if (qtypes.find(ah.d_type) == qtypes.end()) {
1077
0
        if (safeTypes.find(ah.d_type) == safeTypes.end()) {
1078
          // "unsafe" types might contain compressed data, so cancel rewrite
1079
0
          newContent.clear();
1080
0
          return EIO;
1081
0
        }
1082
0
        pw.startRecord(rrname, ah.d_type, ah.d_ttl, ah.d_class, DNSResourceRecord::ADDITIONAL, true);
1083
0
        pw.xfrBlob(blob);
1084
0
      }
1085
0
    }
1086
0
    pw.commit();
1087
1088
0
  }
1089
0
  catch (...)
1090
0
  {
1091
0
    newContent.clear();
1092
0
    return EIO;
1093
0
  }
1094
0
  return 0;
1095
0
}
1096
1097
void clearDNSPacketRecordTypes(vector<uint8_t>& packet, const std::unordered_set<QType>& qtypes)
1098
0
{
1099
0
  return clearDNSPacketRecordTypes(reinterpret_cast<PacketBuffer&>(packet), qtypes);
1100
0
}
1101
1102
void clearDNSPacketRecordTypes(PacketBuffer& packet, const std::unordered_set<QType>& qtypes)
1103
0
{
1104
0
  if (!checkIfPacketContainsRecords(packet, qtypes)) {
1105
0
    return;
1106
0
  }
1107
1108
0
  PacketBuffer newContent;
1109
1110
0
  auto result = rewritePacketWithoutRecordTypes(packet, newContent, qtypes);
1111
0
  if (!result) {
1112
0
    packet = std::move(newContent);
1113
0
  }
1114
0
}
1115
1116
// method of operation: silently fail if it doesn't work - we're only trying to be nice, don't fall over on it
1117
void ageDNSPacket(char* packet, size_t length, uint32_t seconds, const dnsheader_aligned& aligned_dh)
1118
0
{
1119
0
  if (length < sizeof(dnsheader)) {
1120
0
    return;
1121
0
  }
1122
0
  try {
1123
0
    const dnsheader* dhp = aligned_dh.get();
1124
0
    const uint64_t dqcount = ntohs(dhp->qdcount);
1125
0
    const uint64_t numrecords = ntohs(dhp->ancount) + ntohs(dhp->nscount) + ntohs(dhp->arcount);
1126
0
    DNSPacketMangler dpm(packet, length);
1127
1128
0
    for (uint64_t rec = 0; rec < dqcount; ++rec) {
1129
0
      dpm.skipDomainName();
1130
      /* type and class */
1131
0
      dpm.skipBytes(4);
1132
0
    }
1133
1134
0
    for(uint64_t rec = 0; rec < numrecords; ++rec) {
1135
0
      dpm.skipDomainName();
1136
1137
0
      uint16_t dnstype = dpm.get16BitInt();
1138
      /* class */
1139
0
      dpm.skipBytes(2);
1140
1141
0
      if (dnstype != QType::OPT) { // not aging that one with a stick
1142
0
        dpm.decreaseAndSkip32BitInt(seconds);
1143
0
      } else {
1144
0
        dpm.skipBytes(4);
1145
0
      }
1146
0
      dpm.skipRData();
1147
0
    }
1148
0
  }
1149
0
  catch(...) {
1150
0
  }
1151
0
}
1152
1153
void ageDNSPacket(std::string& packet, uint32_t seconds, const dnsheader_aligned& aligned_dh)
1154
0
{
1155
0
  ageDNSPacket(packet.data(), packet.length(), seconds, aligned_dh);
1156
0
}
1157
1158
void shuffleDNSPacket(char* packet, size_t length, const dnsheader_aligned& aligned_dh)
1159
0
{
1160
0
  if (length < sizeof(dnsheader)) {
1161
0
    return;
1162
0
  }
1163
0
  try {
1164
0
    const dnsheader* dhp = aligned_dh.get();
1165
0
    const uint16_t ancount = ntohs(dhp->ancount);
1166
0
    if (ancount == 1) {
1167
      // quick exit, nothing to shuffle
1168
0
      return;
1169
0
    }
1170
1171
0
    DNSPacketMangler dpm(packet, length);
1172
1173
0
    const uint16_t qdcount = ntohs(dhp->qdcount);
1174
1175
0
    for(size_t iter = 0; iter < qdcount; ++iter) {
1176
0
      dpm.skipDomainName();
1177
      /* type and class */
1178
0
      dpm.skipBytes(4);
1179
0
    }
1180
1181
    // for now shuffle only first rrset, only As and AAAAs
1182
0
    uint16_t rrset_type = 0;
1183
0
    DNSName rrset_dnsname{};
1184
0
    std::vector<std::pair<uint32_t, uint32_t>> rrdata_indexes;
1185
0
    rrdata_indexes.reserve(ancount);
1186
1187
0
    for(size_t iter = 0; iter < ancount; ++iter) {
1188
0
      auto domain_start = dpm.getOffset();
1189
0
      dpm.skipDomainName();
1190
0
      const uint16_t dnstype = dpm.get16BitInt();
1191
0
      if (dnstype == QType::A || dnstype == QType::AAAA) {
1192
0
        if (rrdata_indexes.empty()) {
1193
0
          rrset_type = dnstype;
1194
0
          rrset_dnsname = DNSName(packet, length, domain_start, true);
1195
0
        } else {
1196
0
          if (dnstype != rrset_type) {
1197
0
            break;
1198
0
          }
1199
0
          if (DNSName(packet, length, domain_start, true) != rrset_dnsname) {
1200
0
            break;
1201
0
          }
1202
0
        }
1203
        /* class */
1204
0
        dpm.skipBytes(2);
1205
1206
        /* ttl */
1207
0
        dpm.skipBytes(4);
1208
0
        rrdata_indexes.push_back(dpm.skipRDataAndReturnOffsets());
1209
0
      } else {
1210
0
        if (!rrdata_indexes.empty()) {
1211
0
          break;
1212
0
        }
1213
        /* class */
1214
0
        dpm.skipBytes(2);
1215
1216
        /* ttl */
1217
0
        dpm.skipBytes(4);
1218
0
        dpm.skipRData();
1219
0
      }
1220
0
    }
1221
1222
0
    if (rrdata_indexes.size() >= 2) {
1223
0
      using uid = std::uniform_int_distribution<std::vector<std::pair<uint32_t, uint32_t>>::size_type>;
1224
0
      uid dist;
1225
1226
0
      pdns::dns_random_engine randomEngine;
1227
0
      for (auto swapped = rrdata_indexes.size() - 1; swapped > 0; --swapped) {
1228
0
        auto swapped_with = dist(randomEngine, uid::param_type(0, swapped));
1229
0
        if (swapped != swapped_with) {
1230
0
          dpm.swapInPlace(rrdata_indexes.at(swapped), rrdata_indexes.at(swapped_with));
1231
0
        }
1232
0
      }
1233
0
    }
1234
0
  }
1235
0
  catch(...) {
1236
0
  }
1237
0
}
1238
1239
uint32_t getDNSPacketMinTTL(const char* packet, size_t length, bool* seenAuthSOA)
1240
0
{
1241
0
  uint32_t result = std::numeric_limits<uint32_t>::max();
1242
0
  if(length < sizeof(dnsheader)) {
1243
0
    return result;
1244
0
  }
1245
0
  try
1246
0
  {
1247
0
    const dnsheader_aligned dh(packet);
1248
0
    DNSPacketMangler dpm(const_cast<char*>(packet), length);
1249
1250
0
    const uint16_t qdcount = ntohs(dh->qdcount);
1251
0
    for(size_t n = 0; n < qdcount; ++n) {
1252
0
      dpm.skipDomainName();
1253
      /* type and class */
1254
0
      dpm.skipBytes(4);
1255
0
    }
1256
0
    const size_t numrecords = ntohs(dh->ancount) + ntohs(dh->nscount) + ntohs(dh->arcount);
1257
0
    for(size_t n = 0; n < numrecords; ++n) {
1258
0
      dpm.skipDomainName();
1259
0
      const uint16_t dnstype = dpm.get16BitInt();
1260
      /* class */
1261
0
      const uint16_t dnsclass = dpm.get16BitInt();
1262
1263
0
      if(dnstype == QType::OPT) {
1264
0
        break;
1265
0
      }
1266
1267
      /* report it if we see a SOA record in the AUTHORITY section */
1268
0
      if(dnstype == QType::SOA && dnsclass == QClass::IN && seenAuthSOA != nullptr && n >= ntohs(dh->ancount) && n < (ntohs(dh->ancount) + ntohs(dh->nscount))) {
1269
0
        *seenAuthSOA = true;
1270
0
      }
1271
1272
0
      const uint32_t ttl = dpm.get32BitInt();
1273
0
      result = std::min(result, ttl);
1274
1275
0
      dpm.skipRData();
1276
0
    }
1277
0
  }
1278
0
  catch(...)
1279
0
  {
1280
0
  }
1281
0
  return result;
1282
0
}
1283
1284
uint32_t getDNSPacketLength(const char* packet, size_t length)
1285
0
{
1286
0
  uint32_t result = length;
1287
0
  if(length < sizeof(dnsheader)) {
1288
0
    return result;
1289
0
  }
1290
0
  try
1291
0
  {
1292
0
    const dnsheader_aligned dh(packet);
1293
0
    DNSPacketMangler dpm(const_cast<char*>(packet), length);
1294
1295
0
    const uint16_t qdcount = ntohs(dh->qdcount);
1296
0
    for(size_t n = 0; n < qdcount; ++n) {
1297
0
      dpm.skipDomainName();
1298
      /* type and class */
1299
0
      dpm.skipBytes(4);
1300
0
    }
1301
0
    const size_t numrecords = ntohs(dh->ancount) + ntohs(dh->nscount) + ntohs(dh->arcount);
1302
0
    for(size_t n = 0; n < numrecords; ++n) {
1303
0
      dpm.skipDomainName();
1304
      /* type (2), class (2) and ttl (4) */
1305
0
      dpm.skipBytes(8);
1306
0
      dpm.skipRData();
1307
0
    }
1308
0
    result = dpm.getOffset();
1309
0
  }
1310
0
  catch(...)
1311
0
  {
1312
0
  }
1313
0
  return result;
1314
0
}
1315
1316
uint16_t getRecordsOfTypeCount(const char* packet, size_t length, uint8_t section, uint16_t type)
1317
0
{
1318
0
  uint16_t result = 0;
1319
0
  if(length < sizeof(dnsheader)) {
1320
0
    return result;
1321
0
  }
1322
0
  try
1323
0
  {
1324
0
    const dnsheader_aligned dh(packet);
1325
0
    DNSPacketMangler dpm(const_cast<char*>(packet), length);
1326
1327
0
    const uint16_t qdcount = ntohs(dh->qdcount);
1328
0
    for(size_t n = 0; n < qdcount; ++n) {
1329
0
      dpm.skipDomainName();
1330
0
      if (section == 0) {
1331
0
        uint16_t dnstype = dpm.get16BitInt();
1332
0
        if (dnstype == type) {
1333
0
          result++;
1334
0
        }
1335
        /* class */
1336
0
        dpm.skipBytes(2);
1337
0
      } else {
1338
        /* type and class */
1339
0
        dpm.skipBytes(4);
1340
0
      }
1341
0
    }
1342
0
    const uint16_t ancount = ntohs(dh->ancount);
1343
0
    for(size_t n = 0; n < ancount; ++n) {
1344
0
      dpm.skipDomainName();
1345
0
      if (section == 1) {
1346
0
        uint16_t dnstype = dpm.get16BitInt();
1347
0
        if (dnstype == type) {
1348
0
          result++;
1349
0
        }
1350
        /* class */
1351
0
        dpm.skipBytes(2);
1352
0
      } else {
1353
        /* type and class */
1354
0
        dpm.skipBytes(4);
1355
0
      }
1356
      /* ttl */
1357
0
      dpm.skipBytes(4);
1358
0
      dpm.skipRData();
1359
0
    }
1360
0
    const uint16_t nscount = ntohs(dh->nscount);
1361
0
    for(size_t n = 0; n < nscount; ++n) {
1362
0
      dpm.skipDomainName();
1363
0
      if (section == 2) {
1364
0
        uint16_t dnstype = dpm.get16BitInt();
1365
0
        if (dnstype == type) {
1366
0
          result++;
1367
0
        }
1368
        /* class */
1369
0
        dpm.skipBytes(2);
1370
0
      } else {
1371
        /* type and class */
1372
0
        dpm.skipBytes(4);
1373
0
      }
1374
      /* ttl */
1375
0
      dpm.skipBytes(4);
1376
0
      dpm.skipRData();
1377
0
    }
1378
0
    const uint16_t arcount = ntohs(dh->arcount);
1379
0
    for(size_t n = 0; n < arcount; ++n) {
1380
0
      dpm.skipDomainName();
1381
0
      if (section == 3) {
1382
0
        uint16_t dnstype = dpm.get16BitInt();
1383
0
        if (dnstype == type) {
1384
0
          result++;
1385
0
        }
1386
        /* class */
1387
0
        dpm.skipBytes(2);
1388
0
      } else {
1389
        /* type and class */
1390
0
        dpm.skipBytes(4);
1391
0
      }
1392
      /* ttl */
1393
0
      dpm.skipBytes(4);
1394
0
      dpm.skipRData();
1395
0
    }
1396
0
  }
1397
0
  catch(...)
1398
0
  {
1399
0
  }
1400
0
  return result;
1401
0
}
1402
1403
bool getEDNSUDPPayloadSizeAndZ(const char* packet, size_t length, uint16_t* payloadSize, uint16_t* z)
1404
0
{
1405
0
  if (length < sizeof(dnsheader)) {
1406
0
    return false;
1407
0
  }
1408
1409
0
  *payloadSize = 0;
1410
0
  *z = 0;
1411
1412
0
  try
1413
0
  {
1414
0
    const dnsheader_aligned dh(packet);
1415
0
    if (dh->arcount == 0) {
1416
      // The OPT pseudo-RR, if present, has to be in the additional section (https://datatracker.ietf.org/doc/html/rfc6891#section-6.1.1)
1417
0
      return false;
1418
0
    }
1419
1420
0
    DNSPacketMangler dpm(const_cast<char*>(packet), length);
1421
1422
0
    const uint16_t qdcount = ntohs(dh->qdcount);
1423
0
    for(size_t n = 0; n < qdcount; ++n) {
1424
0
      dpm.skipDomainName();
1425
      /* type and class */
1426
0
      dpm.skipBytes(4);
1427
0
    }
1428
0
    const size_t numrecords = ntohs(dh->ancount) + ntohs(dh->nscount) + ntohs(dh->arcount);
1429
0
    for(size_t n = 0; n < numrecords; ++n) {
1430
0
      dpm.skipDomainName();
1431
0
      const auto dnstype = dpm.get16BitInt();
1432
1433
0
      if (dnstype == QType::OPT) {
1434
0
        const auto dnsclass = dpm.get16BitInt();
1435
        /* skip extended rcode and version */
1436
0
        dpm.skipBytes(2);
1437
0
        *z = dpm.get16BitInt();
1438
0
        *payloadSize = dnsclass;
1439
0
        return true;
1440
0
      }
1441
      /* skip class */
1442
0
      dpm.skipBytes(2);
1443
      /* TTL */
1444
0
      dpm.skipBytes(4);
1445
0
      dpm.skipRData();
1446
0
    }
1447
0
  }
1448
0
  catch(...)
1449
0
  {
1450
0
  }
1451
1452
0
  return false;
1453
0
}
1454
1455
bool visitDNSPacket(const std::string_view& packet, const std::function<bool(uint8_t, uint16_t, uint16_t, uint32_t, uint16_t, const char*)>& visitor)
1456
0
{
1457
0
  if (packet.size() < sizeof(dnsheader)) {
1458
0
    return false;
1459
0
  }
1460
1461
0
  try
1462
0
  {
1463
0
    const dnsheader_aligned dh(packet.data());
1464
0
    uint64_t numrecords = ntohs(dh->ancount) + ntohs(dh->nscount) + ntohs(dh->arcount);
1465
0
    PacketReader reader(packet);
1466
1467
0
    uint64_t n;
1468
0
    for (n = 0; n < ntohs(dh->qdcount) ; ++n) {
1469
0
      (void) reader.getName();
1470
      /* type and class */
1471
0
      reader.skip(4);
1472
0
    }
1473
1474
0
    for (n = 0; n < numrecords; ++n) {
1475
0
      (void) reader.getName();
1476
1477
0
      uint8_t section = n < ntohs(dh->ancount) ? 1 : (n < (ntohs(dh->ancount) + ntohs(dh->nscount)) ? 2 : 3);
1478
0
      uint16_t dnstype = reader.get16BitInt();
1479
0
      uint16_t dnsclass = reader.get16BitInt();
1480
1481
0
      if (dnstype == QType::OPT) {
1482
        // not getting near that one with a stick
1483
0
        break;
1484
0
      }
1485
1486
0
      uint32_t dnsttl = reader.get32BitInt();
1487
0
      uint16_t contentLength = reader.get16BitInt();
1488
0
      uint16_t pos = reader.getPosition();
1489
0
      reader.skip(contentLength);
1490
1491
0
      bool done = visitor(section, dnsclass, dnstype, dnsttl, contentLength, &packet.at(pos));
1492
0
      if (done) {
1493
0
        return true;
1494
0
      }
1495
0
    }
1496
0
  }
1497
0
  catch (...) {
1498
0
    return false;
1499
0
  }
1500
1501
0
  return true;
1502
0
}