/src/pdns/pdns/iputils.hh
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 | | #pragma once |
23 | | #include <string> |
24 | | #include <sys/socket.h> |
25 | | #include <netinet/in.h> |
26 | | #include <arpa/inet.h> |
27 | | #include <iostream> |
28 | | #include <cstdio> |
29 | | #include <functional> |
30 | | #include "pdnsexception.hh" |
31 | | #include "misc.hh" |
32 | | #include <netdb.h> |
33 | | #include <sstream> |
34 | | #include <sys/un.h> |
35 | | #include "expected.hh" |
36 | | |
37 | | #include "namespaces.hh" |
38 | | |
39 | | #ifdef __APPLE__ |
40 | | #include <libkern/OSByteOrder.h> |
41 | | |
42 | | #define htobe16(x) OSSwapHostToBigInt16(x) |
43 | | #define htole16(x) OSSwapHostToLittleInt16(x) |
44 | | #define be16toh(x) OSSwapBigToHostInt16(x) |
45 | | #define le16toh(x) OSSwapLittleToHostInt16(x) |
46 | | |
47 | | #define htobe32(x) OSSwapHostToBigInt32(x) |
48 | | #define htole32(x) OSSwapHostToLittleInt32(x) |
49 | | #define be32toh(x) OSSwapBigToHostInt32(x) |
50 | | #define le32toh(x) OSSwapLittleToHostInt32(x) |
51 | | |
52 | | #define htobe64(x) OSSwapHostToBigInt64(x) |
53 | | #define htole64(x) OSSwapHostToLittleInt64(x) |
54 | | #define be64toh(x) OSSwapBigToHostInt64(x) |
55 | | #define le64toh(x) OSSwapLittleToHostInt64(x) |
56 | | |
57 | | #if defined(CONNECT_DATA_IDEMPOTENT) && defined(CONNECT_RESUME_ON_READ_WRITE) |
58 | | #define CONNECTX_FASTOPEN 1 |
59 | | #endif |
60 | | |
61 | | #endif |
62 | | |
63 | | #ifdef __sun |
64 | | |
65 | | #define htobe16(x) BE_16(x) |
66 | | #define htole16(x) LE_16(x) |
67 | | #define be16toh(x) BE_IN16(&(x)) |
68 | | #define le16toh(x) LE_IN16(&(x)) |
69 | | |
70 | | #define htobe32(x) BE_32(x) |
71 | | #define htole32(x) LE_32(x) |
72 | | #define be32toh(x) BE_IN32(&(x)) |
73 | | #define le32toh(x) LE_IN32(&(x)) |
74 | | |
75 | | #define htobe64(x) BE_64(x) |
76 | | #define htole64(x) LE_64(x) |
77 | | #define be64toh(x) BE_IN64(&(x)) |
78 | | #define le64toh(x) LE_IN64(&(x)) |
79 | | |
80 | | #endif |
81 | | |
82 | | #ifdef __FreeBSD__ |
83 | | #include <sys/endian.h> |
84 | | #endif |
85 | | |
86 | | #if defined(__NetBSD__) && defined(IP_PKTINFO) && !defined(IP_SENDSRCADDR) |
87 | | // The IP_PKTINFO option in NetBSD was incompatible with Linux until a |
88 | | // change that also introduced IP_SENDSRCADDR for FreeBSD compatibility. |
89 | | #undef IP_PKTINFO |
90 | | #endif |
91 | | |
92 | | union ComboAddress |
93 | | { |
94 | | sockaddr_in sin4{}; |
95 | | sockaddr_in6 sin6; |
96 | | |
97 | | bool operator==(const ComboAddress& rhs) const |
98 | 0 | { |
99 | 0 | if (std::tie(sin4.sin_family, sin4.sin_port) != std::tie(rhs.sin4.sin_family, rhs.sin4.sin_port)) { |
100 | 0 | return false; |
101 | 0 | } |
102 | 0 | if (sin4.sin_family == AF_INET) { |
103 | 0 | return sin4.sin_addr.s_addr == rhs.sin4.sin_addr.s_addr; |
104 | 0 | } |
105 | 0 | return memcmp(&sin6.sin6_addr.s6_addr, &rhs.sin6.sin6_addr.s6_addr, sizeof(sin6.sin6_addr.s6_addr)) == 0; |
106 | 0 | } |
107 | | |
108 | | bool operator!=(const ComboAddress& rhs) const |
109 | 0 | { |
110 | 0 | return (!operator==(rhs)); |
111 | 0 | } |
112 | | |
113 | | bool operator<(const ComboAddress& rhs) const |
114 | 0 | { |
115 | 0 | if (sin4.sin_family == 0) { |
116 | 0 | return false; |
117 | 0 | } |
118 | 0 | if (std::tie(sin4.sin_family, sin4.sin_port) < std::tie(rhs.sin4.sin_family, rhs.sin4.sin_port)) { |
119 | 0 | return true; |
120 | 0 | } |
121 | 0 | if (std::tie(sin4.sin_family, sin4.sin_port) > std::tie(rhs.sin4.sin_family, rhs.sin4.sin_port)) { |
122 | 0 | return false; |
123 | 0 | } |
124 | 0 | if (sin4.sin_family == AF_INET) { |
125 | 0 | return sin4.sin_addr.s_addr < rhs.sin4.sin_addr.s_addr; |
126 | 0 | } |
127 | 0 | return memcmp(&sin6.sin6_addr.s6_addr, &rhs.sin6.sin6_addr.s6_addr, sizeof(sin6.sin6_addr.s6_addr)) < 0; |
128 | 0 | } |
129 | | |
130 | | bool operator>(const ComboAddress& rhs) const |
131 | 0 | { |
132 | 0 | return rhs.operator<(*this); |
133 | 0 | } |
134 | | |
135 | | struct addressPortOnlyHash |
136 | | { |
137 | | uint32_t operator()(const ComboAddress& address) const |
138 | 0 | { |
139 | 0 | // NOLINTBEGIN(cppcoreguidelines-pro-type-reinterpret-cast) |
140 | 0 | if (address.sin4.sin_family == AF_INET) { |
141 | 0 | const auto* start = reinterpret_cast<const unsigned char*>(&address.sin4.sin_addr.s_addr); |
142 | 0 | auto tmp = burtle(start, 4, 0); |
143 | 0 | return burtle(reinterpret_cast<const uint8_t*>(&address.sin4.sin_port), 2, tmp); |
144 | 0 | } |
145 | 0 | const auto* start = reinterpret_cast<const unsigned char*>(&address.sin6.sin6_addr.s6_addr); |
146 | 0 | auto tmp = burtle(start, 16, 0); |
147 | 0 | return burtle(reinterpret_cast<const unsigned char*>(&address.sin6.sin6_port), 2, tmp); |
148 | 0 | // NOLINTEND(cppcoreguidelines-pro-type-reinterpret-cast) |
149 | 0 | } |
150 | | }; |
151 | | |
152 | | struct addressOnlyHash |
153 | | { |
154 | | uint32_t operator()(const ComboAddress& address) const |
155 | 0 | { |
156 | 0 | const unsigned char* start = nullptr; |
157 | 0 | uint32_t len = 0; |
158 | 0 | // NOLINTBEGIN(cppcoreguidelines-pro-type-reinterpret-cast) |
159 | 0 | if (address.sin4.sin_family == AF_INET) { |
160 | 0 | start = reinterpret_cast<const unsigned char*>(&address.sin4.sin_addr.s_addr); |
161 | 0 | len = 4; |
162 | 0 | } |
163 | 0 | else { |
164 | 0 | start = reinterpret_cast<const unsigned char*>(&address.sin6.sin6_addr.s6_addr); |
165 | 0 | len = 16; |
166 | 0 | } |
167 | 0 | // NOLINTEND(cppcoreguidelines-pro-type-reinterpret-cast) |
168 | 0 | return burtle(start, len, 0); |
169 | 0 | } |
170 | | }; |
171 | | |
172 | | struct addressOnlyLessThan |
173 | | { |
174 | | bool operator()(const ComboAddress& lhs, const ComboAddress& rhs) const |
175 | 0 | { |
176 | 0 | if (lhs.sin4.sin_family < rhs.sin4.sin_family) { |
177 | 0 | return true; |
178 | 0 | } |
179 | 0 | if (lhs.sin4.sin_family > rhs.sin4.sin_family) { |
180 | 0 | return false; |
181 | 0 | } |
182 | 0 | if (lhs.sin4.sin_family == AF_INET) { |
183 | 0 | return lhs.sin4.sin_addr.s_addr < rhs.sin4.sin_addr.s_addr; |
184 | 0 | } |
185 | 0 | return memcmp(&lhs.sin6.sin6_addr.s6_addr, &rhs.sin6.sin6_addr.s6_addr, sizeof(lhs.sin6.sin6_addr.s6_addr)) < 0; |
186 | 0 | } |
187 | | }; |
188 | | |
189 | | struct addressOnlyEqual |
190 | | { |
191 | | bool operator()(const ComboAddress& lhs, const ComboAddress& rhs) const |
192 | 0 | { |
193 | 0 | if (lhs.sin4.sin_family != rhs.sin4.sin_family) { |
194 | 0 | return false; |
195 | 0 | } |
196 | 0 | if (lhs.sin4.sin_family == AF_INET) { |
197 | 0 | return lhs.sin4.sin_addr.s_addr == rhs.sin4.sin_addr.s_addr; |
198 | 0 | } |
199 | 0 | return memcmp(&lhs.sin6.sin6_addr.s6_addr, &rhs.sin6.sin6_addr.s6_addr, sizeof(lhs.sin6.sin6_addr.s6_addr)) == 0; |
200 | 0 | } |
201 | | }; |
202 | | |
203 | | [[nodiscard]] socklen_t getSocklen() const |
204 | 557k | { |
205 | 557k | if (sin4.sin_family == AF_INET) { |
206 | 359k | return sizeof(sin4); |
207 | 359k | } |
208 | 197k | return sizeof(sin6); |
209 | 557k | } |
210 | | |
211 | | ComboAddress() |
212 | 795k | { |
213 | 795k | sin4.sin_family = AF_INET; |
214 | 795k | sin4.sin_addr.s_addr = 0; |
215 | 795k | sin4.sin_port = 0; |
216 | 795k | sin6.sin6_scope_id = 0; |
217 | 795k | sin6.sin6_flowinfo = 0; |
218 | 795k | } |
219 | | |
220 | | ComboAddress(const struct sockaddr* socketAddress, socklen_t salen) |
221 | 0 | { |
222 | 0 | setSockaddr(socketAddress, salen); |
223 | 0 | }; |
224 | | |
225 | | ComboAddress(const struct sockaddr_in6* socketAddress) |
226 | 0 | { |
227 | 0 | // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast) |
228 | 0 | setSockaddr(reinterpret_cast<const struct sockaddr*>(socketAddress), sizeof(struct sockaddr_in6)); |
229 | 0 | }; |
230 | | |
231 | | ComboAddress(const struct sockaddr_in* socketAddress) |
232 | 0 | { |
233 | 0 | // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast) |
234 | 0 | setSockaddr(reinterpret_cast<const struct sockaddr*>(socketAddress), sizeof(struct sockaddr_in)); |
235 | 0 | }; |
236 | | |
237 | | void setSockaddr(const struct sockaddr* socketAddress, socklen_t salen) |
238 | 0 | { |
239 | 0 | if (salen > sizeof(struct sockaddr_in6)) { |
240 | 0 | throw PDNSException("ComboAddress can't handle other than sockaddr_in or sockaddr_in6"); |
241 | 0 | } |
242 | 0 | memcpy(this, socketAddress, salen); |
243 | 0 | } |
244 | | |
245 | | // 'port' sets a default value in case 'str' does not set a port |
246 | | explicit ComboAddress(const string& str, uint16_t port = 0) |
247 | 284k | { |
248 | 284k | memset(&sin6, 0, sizeof(sin6)); |
249 | 284k | sin4.sin_family = AF_INET; |
250 | 284k | sin4.sin_port = 0; |
251 | 284k | if (makeIPv4sockaddr(str, &sin4) != 0) { |
252 | 125k | sin6.sin6_family = AF_INET6; |
253 | 125k | if (makeIPv6sockaddr(str, &sin6) < 0) { |
254 | 36 | throw PDNSException("Unable to convert presentation address '" + str + "'"); |
255 | 36 | } |
256 | 125k | } |
257 | 284k | if (sin4.sin_port == 0) { // 'str' overrides port! |
258 | 277k | sin4.sin_port = htons(port); |
259 | 277k | } |
260 | 284k | } |
261 | | |
262 | | [[nodiscard]] bool isIPv6() const |
263 | 231k | { |
264 | 231k | return sin4.sin_family == AF_INET6; |
265 | 231k | } |
266 | | [[nodiscard]] bool isIPv4() const |
267 | 522k | { |
268 | 522k | return sin4.sin_family == AF_INET; |
269 | 522k | } |
270 | | |
271 | | [[nodiscard]] bool isMappedIPv4() const |
272 | 0 | { |
273 | 0 | if (sin4.sin_family != AF_INET6) { |
274 | 0 | return false; |
275 | 0 | } |
276 | 0 |
|
277 | 0 | int iter = 0; |
278 | 0 | // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast) |
279 | 0 | const auto* ptr = reinterpret_cast<const unsigned char*>(&sin6.sin6_addr.s6_addr); |
280 | 0 | for (iter = 0; iter < 10; ++iter) { |
281 | 0 | if (ptr[iter] != 0) { // NOLINT(cppcoreguidelines-pro-bounds-pointer-arithmetic) |
282 | 0 | return false; |
283 | 0 | } |
284 | 0 | } |
285 | 0 | for (; iter < 12; ++iter) { |
286 | 0 | if (ptr[iter] != 0xff) { // NOLINT(cppcoreguidelines-pro-bounds-pointer-arithmetic) |
287 | 0 | return false; |
288 | 0 | } |
289 | 0 | } |
290 | 0 | return true; |
291 | 0 | } |
292 | | |
293 | | [[nodiscard]] bool isUnspecified() const |
294 | 0 | { |
295 | 0 | static const ComboAddress unspecifiedV4("0.0.0.0:0"); |
296 | 0 | static const ComboAddress unspecifiedV6("[::]:0"); |
297 | 0 | const auto compare = ComboAddress::addressOnlyEqual(); |
298 | 0 | return compare(*this, unspecifiedV4) || compare(*this, unspecifiedV6); |
299 | 0 | } |
300 | | |
301 | | [[nodiscard]] ComboAddress mapToIPv4() const |
302 | 0 | { |
303 | 0 | if (!isMappedIPv4()) { |
304 | 0 | throw PDNSException("ComboAddress can't map non-mapped IPv6 address back to IPv4"); |
305 | 0 | } |
306 | 0 | ComboAddress ret; |
307 | 0 | ret.sin4.sin_family = AF_INET; |
308 | 0 | ret.sin4.sin_port = sin4.sin_port; |
309 | 0 |
|
310 | 0 | // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast) |
311 | 0 | const auto* ptr = reinterpret_cast<const unsigned char*>(&sin6.sin6_addr.s6_addr); |
312 | 0 | ptr += (sizeof(sin6.sin6_addr.s6_addr) - sizeof(ret.sin4.sin_addr.s_addr)); // NOLINT(cppcoreguidelines-pro-bounds-pointer-arithmetic) |
313 | 0 | memcpy(&ret.sin4.sin_addr.s_addr, ptr, sizeof(ret.sin4.sin_addr.s_addr)); |
314 | 0 | return ret; |
315 | 0 | } |
316 | | |
317 | | [[nodiscard]] string toString() const |
318 | 557k | { |
319 | 557k | std::array<char, 1024> host{}; |
320 | 557k | if (sin4.sin_family != 0) { |
321 | | // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast) |
322 | 557k | int retval = getnameinfo(reinterpret_cast<const struct sockaddr*>(this), getSocklen(), host.data(), host.size(), nullptr, 0, NI_NUMERICHOST); |
323 | 557k | if (retval == 0) { |
324 | 557k | return host.data(); |
325 | 557k | } |
326 | 0 | return "invalid " + string(gai_strerror(retval)); |
327 | 557k | } |
328 | 0 | return "invalid"; |
329 | 557k | } |
330 | | |
331 | | //! Ignores any interface specifiers possibly available in the sockaddr data. |
332 | | [[nodiscard]] string toStringNoInterface() const |
333 | 33.7k | { |
334 | 33.7k | std::array<char, 1024> host{}; |
335 | 33.7k | if (sin4.sin_family == AF_INET) { |
336 | 4.45k | const auto* ret = inet_ntop(sin4.sin_family, &sin4.sin_addr, host.data(), host.size()); |
337 | 4.45k | if (ret != nullptr) { |
338 | 4.45k | return host.data(); |
339 | 4.45k | } |
340 | 4.45k | } |
341 | 29.2k | else if (sin4.sin_family == AF_INET6) { |
342 | 29.2k | const auto* ret = inet_ntop(sin4.sin_family, &sin6.sin6_addr, host.data(), host.size()); |
343 | 29.2k | if (ret != nullptr) { |
344 | 29.2k | return host.data(); |
345 | 29.2k | } |
346 | 29.2k | } |
347 | 0 | else { |
348 | 0 | return "invalid"; |
349 | 0 | } |
350 | 0 | return "invalid " + stringerror(); |
351 | 33.7k | } |
352 | | |
353 | | [[nodiscard]] string toStringReversed() const |
354 | 0 | { |
355 | 0 | if (isIPv4()) { |
356 | 0 | const auto address = ntohl(sin4.sin_addr.s_addr); |
357 | 0 | auto aaa = (address >> 0) & 0xFF; |
358 | 0 | auto bbb = (address >> 8) & 0xFF; |
359 | 0 | auto ccc = (address >> 16) & 0xFF; |
360 | 0 | auto ddd = (address >> 24) & 0xFF; |
361 | 0 | return std::to_string(aaa) + "." + std::to_string(bbb) + "." + std::to_string(ccc) + "." + std::to_string(ddd); |
362 | 0 | } |
363 | 0 | const auto* addr = &sin6.sin6_addr; |
364 | 0 | std::stringstream res{}; |
365 | 0 | res << std::hex; |
366 | 0 | for (int i = 15; i >= 0; i--) { |
367 | 0 | auto byte = addr->s6_addr[i]; // NOLINT(cppcoreguidelines-pro-bounds-constant-array-index) |
368 | 0 | res << ((byte >> 0) & 0xF) << "."; |
369 | 0 | res << ((byte >> 4) & 0xF); |
370 | 0 | if (i != 0) { |
371 | 0 | res << "."; |
372 | 0 | } |
373 | 0 | } |
374 | 0 | return res.str(); |
375 | 0 | } |
376 | | |
377 | | [[nodiscard]] string toStringWithPort() const |
378 | 0 | { |
379 | 0 | if (sin4.sin_family == AF_INET) { |
380 | 0 | return toString() + ":" + std::to_string(ntohs(sin4.sin_port)); |
381 | 0 | } |
382 | 0 | return "[" + toString() + "]:" + std::to_string(ntohs(sin4.sin_port)); |
383 | 0 | } |
384 | | |
385 | | [[nodiscard]] string toStringWithPortExcept(int port) const |
386 | 0 | { |
387 | 0 | if (ntohs(sin4.sin_port) == port) { |
388 | 0 | return toString(); |
389 | 0 | } |
390 | 0 | if (sin4.sin_family == AF_INET) { |
391 | 0 | return toString() + ":" + std::to_string(ntohs(sin4.sin_port)); |
392 | 0 | } |
393 | 0 | return "[" + toString() + "]:" + std::to_string(ntohs(sin4.sin_port)); |
394 | 0 | } |
395 | | |
396 | | [[nodiscard]] string toLogString() const |
397 | 0 | { |
398 | 0 | return toStringWithPortExcept(53); |
399 | 0 | } |
400 | | |
401 | | [[nodiscard]] string toStructuredLogString() const |
402 | 0 | { |
403 | 0 | return toStringWithPort(); |
404 | 0 | } |
405 | | |
406 | | [[nodiscard]] string toByteString() const |
407 | 0 | { |
408 | 0 | // NOLINTBEGIN(cppcoreguidelines-pro-type-reinterpret-cast) |
409 | 0 | if (isIPv4()) { |
410 | 0 | return {reinterpret_cast<const char*>(&sin4.sin_addr.s_addr), sizeof(sin4.sin_addr.s_addr)}; |
411 | 0 | } |
412 | 0 | return {reinterpret_cast<const char*>(&sin6.sin6_addr.s6_addr), sizeof(sin6.sin6_addr.s6_addr)}; |
413 | 0 | // NOLINTEND(cppcoreguidelines-pro-type-reinterpret-cast) |
414 | 0 | } |
415 | | |
416 | | void truncate(unsigned int bits) noexcept; |
417 | | |
418 | | [[nodiscard]] uint16_t getNetworkOrderPort() const noexcept |
419 | 0 | { |
420 | 0 | return sin4.sin_port; |
421 | 0 | } |
422 | | [[nodiscard]] uint16_t getPort() const noexcept |
423 | 0 | { |
424 | 0 | return ntohs(getNetworkOrderPort()); |
425 | 0 | } |
426 | | void setPort(uint16_t port) |
427 | 20 | { |
428 | 20 | sin4.sin_port = htons(port); |
429 | 20 | } |
430 | | |
431 | | void reset() |
432 | 65 | { |
433 | 65 | memset(&sin6, 0, sizeof(sin6)); |
434 | 65 | } |
435 | | |
436 | | //! Get the total number of address bits (either 32 or 128 depending on IP version) |
437 | | [[nodiscard]] uint8_t getBits() const |
438 | 0 | { |
439 | 0 | if (isIPv4()) { |
440 | 0 | return 32; |
441 | 0 | } |
442 | 0 | if (isIPv6()) { |
443 | 0 | return 128; |
444 | 0 | } |
445 | 0 | return 0; |
446 | 0 | } |
447 | | /** Get the value of the bit at the provided bit index. When the index >= 0, |
448 | | the index is relative to the LSB starting at index zero. When the index < 0, |
449 | | the index is relative to the MSB starting at index -1 and counting down. |
450 | | */ |
451 | | [[nodiscard]] bool getBit(int index) const |
452 | 0 | { |
453 | 0 | if (isIPv4()) { |
454 | 0 | if (index >= 32) { |
455 | 0 | return false; |
456 | 0 | } |
457 | 0 | if (index < 0) { |
458 | 0 | if (index < -32) { |
459 | 0 | return false; |
460 | 0 | } |
461 | 0 | index = 32 + index; |
462 | 0 | } |
463 | 0 |
|
464 | 0 | uint32_t ls_addr = ntohl(sin4.sin_addr.s_addr); |
465 | 0 |
|
466 | 0 | return ((ls_addr & (1U << index)) != 0x00000000); |
467 | 0 | } |
468 | 0 | if (isIPv6()) { |
469 | 0 | if (index >= 128) { |
470 | 0 | return false; |
471 | 0 | } |
472 | 0 | if (index < 0) { |
473 | 0 | if (index < -128) { |
474 | 0 | return false; |
475 | 0 | } |
476 | 0 | index = 128 + index; |
477 | 0 | } |
478 | 0 |
|
479 | 0 | const auto* ls_addr = reinterpret_cast<const uint8_t*>(sin6.sin6_addr.s6_addr); // NOLINT(cppcoreguidelines-pro-type-reinterpret-cast) |
480 | 0 | uint8_t byte_idx = index / 8; |
481 | 0 | uint8_t bit_idx = index % 8; |
482 | 0 |
|
483 | 0 | return ((ls_addr[15 - byte_idx] & (1U << bit_idx)) != 0x00); // NOLINT(cppcoreguidelines-pro-bounds-pointer-arithmetic) |
484 | 0 | } |
485 | 0 | return false; |
486 | 0 | } |
487 | | |
488 | | /*! Returns a comma-separated string of IP addresses |
489 | | * |
490 | | * \param c An stl container with ComboAddresses |
491 | | * \param withPort Also print the port (default true) |
492 | | * \param portExcept Print the port, except when this is the port (default 53) |
493 | | */ |
494 | | template <template <class...> class Container, class... Args> |
495 | | static string caContainerToString(const Container<ComboAddress, Args...>& container, const bool withPort = true, const uint16_t portExcept = 53) |
496 | 849 | { |
497 | 849 | vector<string> strs; |
498 | 556k | for (const auto& address : container) { |
499 | 556k | if (withPort) { |
500 | 0 | strs.push_back(address.toStringWithPortExcept(portExcept)); |
501 | 0 | continue; |
502 | 0 | } |
503 | 556k | strs.push_back(address.toString()); |
504 | 556k | } |
505 | 849 | return boost::join(strs, ","); |
506 | 849 | }; |
507 | | }; |
508 | | |
509 | | union SockaddrWrapper |
510 | | { |
511 | | sockaddr_in sin4{}; |
512 | | sockaddr_in6 sin6; |
513 | | sockaddr_un sinun; |
514 | | |
515 | | [[nodiscard]] socklen_t getSocklen() const |
516 | 0 | { |
517 | 0 | if (sin4.sin_family == AF_INET) { |
518 | 0 | return sizeof(sin4); |
519 | 0 | } |
520 | 0 | if (sin6.sin6_family == AF_INET6) { |
521 | 0 | return sizeof(sin6); |
522 | 0 | } |
523 | 0 | if (sinun.sun_family == AF_UNIX) { |
524 | 0 | return sizeof(sinun); |
525 | 0 | } |
526 | 0 | return 0; |
527 | 0 | } |
528 | | |
529 | | SockaddrWrapper() |
530 | 0 | { |
531 | 0 | sin4.sin_family = AF_INET; |
532 | 0 | sin4.sin_addr.s_addr = 0; |
533 | 0 | sin4.sin_port = 0; |
534 | 0 | } |
535 | | |
536 | | SockaddrWrapper(const struct sockaddr* socketAddress, socklen_t salen) |
537 | 0 | { |
538 | 0 | setSockaddr(socketAddress, salen); |
539 | 0 | }; |
540 | | |
541 | | SockaddrWrapper(const struct sockaddr_in6* socketAddress) |
542 | 0 | { |
543 | 0 | // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast) |
544 | 0 | setSockaddr(reinterpret_cast<const struct sockaddr*>(socketAddress), sizeof(struct sockaddr_in6)); |
545 | 0 | }; |
546 | | |
547 | | SockaddrWrapper(const struct sockaddr_in* socketAddress) |
548 | 0 | { |
549 | 0 | // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast) |
550 | 0 | setSockaddr(reinterpret_cast<const struct sockaddr*>(socketAddress), sizeof(struct sockaddr_in)); |
551 | 0 | }; |
552 | | |
553 | | SockaddrWrapper(const struct sockaddr_un* socketAddress) |
554 | 0 | { |
555 | 0 | // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast) |
556 | 0 | setSockaddr(reinterpret_cast<const struct sockaddr*>(socketAddress), sizeof(struct sockaddr_un)); |
557 | 0 | }; |
558 | | |
559 | | void setSockaddr(const struct sockaddr* socketAddress, socklen_t salen) |
560 | 0 | { |
561 | 0 | if (salen > sizeof(struct sockaddr_un)) { |
562 | 0 | throw PDNSException("ComboAddress can't handle other than sockaddr_in, sockaddr_in6 or sockaddr_un"); |
563 | 0 | } |
564 | 0 | memcpy(this, socketAddress, salen); |
565 | 0 | } |
566 | | |
567 | | explicit SockaddrWrapper(const string& str, uint16_t port = 0) |
568 | 0 | { |
569 | 0 | memset(&sinun, 0, sizeof(sinun)); |
570 | 0 | sin4.sin_family = AF_INET; |
571 | 0 | sin4.sin_port = 0; |
572 | 0 | if (str == "\"\"" || str == "''") { |
573 | 0 | throw PDNSException("Stray quotation marks in address."); |
574 | 0 | } |
575 | 0 | if (makeIPv4sockaddr(str, &sin4) != 0) { |
576 | 0 | sin6.sin6_family = AF_INET6; |
577 | 0 | if (makeIPv6sockaddr(str, &sin6) < 0) { |
578 | 0 | sinun.sun_family = AF_UNIX; |
579 | 0 | // only attempt Unix socket address if address candidate does not contain a port |
580 | 0 | if (str.find(':') != string::npos || makeUNsockaddr(str, &sinun) < 0) { |
581 | 0 | throw PDNSException("Unable to convert presentation address '" + str + "'"); |
582 | 0 | } |
583 | 0 | } |
584 | 0 | } |
585 | 0 | if (sinun.sun_family != AF_UNIX && sin4.sin_port == 0) { // 'str' overrides port! |
586 | 0 | sin4.sin_port = htons(port); |
587 | 0 | } |
588 | 0 | } |
589 | | |
590 | | [[nodiscard]] bool isIPv6() const |
591 | 0 | { |
592 | 0 | return sin4.sin_family == AF_INET6; |
593 | 0 | } |
594 | | [[nodiscard]] bool isIPv4() const |
595 | 0 | { |
596 | 0 | return sin4.sin_family == AF_INET; |
597 | 0 | } |
598 | | [[nodiscard]] bool isUnixSocket() const |
599 | 0 | { |
600 | 0 | return sin4.sin_family == AF_UNIX; |
601 | 0 | } |
602 | | |
603 | | [[nodiscard]] string toString() const |
604 | 0 | { |
605 | 0 | if (sinun.sun_family == AF_UNIX) { |
606 | 0 | return static_cast<const char*>(sinun.sun_path); |
607 | 0 | } |
608 | 0 | std::array<char, 1024> host{}; |
609 | 0 | if (sin4.sin_family != 0) { |
610 | | // NOLINTNEXTLINE(cppcoreguidelines-pro-type-reinterpret-cast) |
611 | 0 | int retval = getnameinfo(reinterpret_cast<const struct sockaddr*>(this), getSocklen(), host.data(), host.size(), nullptr, 0, NI_NUMERICHOST); |
612 | 0 | if (retval == 0) { |
613 | 0 | return host.data(); |
614 | 0 | } |
615 | 0 | return "invalid " + string(gai_strerror(retval)); |
616 | 0 | } |
617 | 0 | return "invalid"; |
618 | 0 | } |
619 | | |
620 | | [[nodiscard]] string toStringWithPort() const |
621 | 0 | { |
622 | 0 | if (sinun.sun_family == AF_UNIX) { |
623 | 0 | return toString(); |
624 | 0 | } |
625 | 0 | if (sin4.sin_family == AF_INET) { |
626 | 0 | return toString() + ":" + std::to_string(ntohs(sin4.sin_port)); |
627 | 0 | } |
628 | 0 | return "[" + toString() + "]:" + std::to_string(ntohs(sin4.sin_port)); |
629 | 0 | } |
630 | | |
631 | | void reset() |
632 | 0 | { |
633 | 0 | memset(&sinun, 0, sizeof(sinun)); |
634 | 0 | } |
635 | | |
636 | | void tightenSocketPermissions(Logr::log_t slog, const std::string& gid) const; |
637 | | }; |
638 | | |
639 | | /** This exception is thrown by the Netmask class and by extension by the NetmaskGroup class */ |
640 | | class NetmaskException : public PDNSException |
641 | | { |
642 | | public: |
643 | | NetmaskException(const string& arg) : |
644 | 411 | PDNSException(arg) {} |
645 | | }; |
646 | | |
647 | | inline ComboAddress makeComboAddress(const string& str) |
648 | 21.8k | { |
649 | 21.8k | ComboAddress address; |
650 | 21.8k | address.sin4.sin_family = AF_INET; |
651 | 21.8k | if (makeIPv4sockaddr(str, &address.sin4) < 0) { |
652 | 18.9k | address.sin4.sin_family = AF_INET6; |
653 | 18.9k | if (makeIPv6sockaddr(str, &address.sin6) < 0) { |
654 | 411 | throw NetmaskException("Unable to convert '" + str + "' to a netmask"); |
655 | 411 | } |
656 | 18.9k | } |
657 | 21.4k | return address; |
658 | 21.8k | } |
659 | | |
660 | | inline ComboAddress makeComboAddressFromRaw(uint8_t version, const char* raw, size_t len) |
661 | 394k | { |
662 | 394k | ComboAddress address; |
663 | | |
664 | 394k | if (version == 4) { |
665 | 307k | address.sin4.sin_family = AF_INET; |
666 | 307k | if (len != sizeof(address.sin4.sin_addr)) { |
667 | 0 | throw NetmaskException("invalid raw address length"); |
668 | 0 | } |
669 | 307k | memcpy(&address.sin4.sin_addr, raw, sizeof(address.sin4.sin_addr)); |
670 | 307k | } |
671 | 87.8k | else if (version == 6) { |
672 | 87.8k | address.sin6.sin6_family = AF_INET6; |
673 | 87.8k | if (len != sizeof(address.sin6.sin6_addr)) { |
674 | 0 | throw NetmaskException("invalid raw address length"); |
675 | 0 | } |
676 | 87.8k | memcpy(&address.sin6.sin6_addr, raw, sizeof(address.sin6.sin6_addr)); |
677 | 87.8k | } |
678 | 0 | else { |
679 | 0 | throw NetmaskException("invalid address family"); |
680 | 0 | } |
681 | | |
682 | 394k | return address; |
683 | 394k | } |
684 | | |
685 | | inline ComboAddress makeComboAddressFromRaw(uint8_t version, const string& str) |
686 | 295k | { |
687 | 295k | return makeComboAddressFromRaw(version, str.c_str(), str.size()); |
688 | 295k | } |
689 | | |
690 | | /** This class represents a netmask and can be queried to see if a certain |
691 | | IP address is matched by this mask */ |
692 | | class Netmask |
693 | | { |
694 | | public: |
695 | | Netmask() |
696 | 25.7k | { |
697 | 25.7k | d_network.sin4.sin_family = 0; // disable this doing anything useful |
698 | 25.7k | d_network.sin4.sin_port = 0; // this guarantees d_network compares identical |
699 | 25.7k | } |
700 | | |
701 | | Netmask(const ComboAddress& network, uint8_t bits = 0xff) : |
702 | 36.5k | d_network(network) |
703 | 36.5k | { |
704 | 36.5k | d_network.sin4.sin_port = 0; |
705 | 36.5k | setBits(bits); |
706 | 36.5k | } |
707 | | |
708 | | Netmask(const sockaddr_in* network, uint8_t bits = 0xff) : |
709 | | d_network(network) |
710 | 0 | { |
711 | 0 | d_network.sin4.sin_port = 0; |
712 | 0 | setBits(bits); |
713 | 0 | } |
714 | | Netmask(const sockaddr_in6* network, uint8_t bits = 0xff) : |
715 | | d_network(network) |
716 | 0 | { |
717 | 0 | d_network.sin4.sin_port = 0; |
718 | 0 | setBits(bits); |
719 | 0 | } |
720 | | void setBits(uint8_t value) |
721 | 57.8k | { |
722 | 57.8k | d_bits = d_network.isIPv4() ? std::min(value, static_cast<uint8_t>(32U)) : std::min(value, static_cast<uint8_t>(128U)); |
723 | | |
724 | 57.8k | if (d_bits < 32) { |
725 | 14.2k | d_mask = ~(0xFFFFFFFF >> d_bits); |
726 | 14.2k | } |
727 | 43.5k | else { |
728 | | // note that d_mask is unused for IPv6 |
729 | 43.5k | d_mask = 0xFFFFFFFF; |
730 | 43.5k | } |
731 | | |
732 | 57.8k | if (isIPv4()) { |
733 | 10.0k | d_network.sin4.sin_addr.s_addr = htonl(ntohl(d_network.sin4.sin_addr.s_addr) & d_mask); |
734 | 10.0k | } |
735 | 47.8k | else if (isIPv6()) { |
736 | 47.8k | uint8_t bytes = d_bits / 8; |
737 | 47.8k | auto* address = reinterpret_cast<uint8_t*>(&d_network.sin6.sin6_addr.s6_addr); // NOLINT(cppcoreguidelines-pro-type-reinterpret-cast) |
738 | 47.8k | uint8_t bits = d_bits % 8; |
739 | 47.8k | auto mask = static_cast<uint8_t>(~(0xFF >> bits)); |
740 | | |
741 | 47.8k | if (bytes < sizeof(d_network.sin6.sin6_addr.s6_addr)) { |
742 | 9.35k | address[bytes] &= mask; // NOLINT(cppcoreguidelines-pro-bounds-pointer-arithmetic) |
743 | 9.35k | } |
744 | | |
745 | 180k | for (size_t idx = bytes + 1; idx < sizeof(d_network.sin6.sin6_addr.s6_addr); ++idx) { |
746 | 132k | address[idx] = 0; // NOLINT(cppcoreguidelines-pro-bounds-pointer-arithmetic) |
747 | 132k | } |
748 | 47.8k | } |
749 | 57.8k | } |
750 | | |
751 | | enum stringType : uint8_t |
752 | | { |
753 | | humanString, |
754 | | byteString, |
755 | | }; |
756 | | //! Constructor supplies the mask, which cannot be changed |
757 | | Netmask(const string& mask, stringType type = humanString) |
758 | 21.8k | { |
759 | 21.8k | if (type == byteString) { |
760 | 0 | uint8_t afi = mask.at(0); |
761 | 0 | size_t len = afi == 4 ? 4 : 16; |
762 | 0 | uint8_t bits = mask.at(len + 1); |
763 | |
|
764 | 0 | d_network = makeComboAddressFromRaw(afi, mask.substr(1, len)); |
765 | |
|
766 | 0 | setBits(bits); |
767 | 0 | } |
768 | 21.8k | else { |
769 | 21.8k | pair<string, string> split = splitField(mask, '/'); |
770 | 21.8k | d_network = makeComboAddress(split.first); |
771 | | |
772 | 21.8k | if (!split.second.empty()) { |
773 | 6.25k | setBits(pdns::checked_stoi<uint8_t>(split.second)); |
774 | 6.25k | } |
775 | 15.6k | else if (d_network.sin4.sin_family == AF_INET) { |
776 | 1.35k | setBits(32); |
777 | 1.35k | } |
778 | 14.2k | else { |
779 | 14.2k | setBits(128); |
780 | 14.2k | } |
781 | 21.8k | } |
782 | 21.8k | } |
783 | | |
784 | | [[nodiscard]] bool match(const ComboAddress& address) const |
785 | 0 | { |
786 | 0 | return match(&address); |
787 | 0 | } |
788 | | |
789 | | //! If this IP address in socket address matches |
790 | | bool match(const ComboAddress* address) const |
791 | 0 | { |
792 | 0 | if (d_network.sin4.sin_family != address->sin4.sin_family) { |
793 | 0 | return false; |
794 | 0 | } |
795 | 0 | if (d_network.sin4.sin_family == AF_INET) { |
796 | 0 | return match4(htonl((unsigned int)address->sin4.sin_addr.s_addr)); |
797 | 0 | } |
798 | 0 | if (d_network.sin6.sin6_family == AF_INET6) { |
799 | 0 | uint8_t bytes = d_bits / 8; |
800 | 0 | uint8_t index = 0; |
801 | 0 | // NOLINTBEGIN(cppcoreguidelines-pro-type-reinterpret-cast) |
802 | 0 | const auto* lhs = reinterpret_cast<const uint8_t*>(&d_network.sin6.sin6_addr.s6_addr); |
803 | 0 | const auto* rhs = reinterpret_cast<const uint8_t*>(&address->sin6.sin6_addr.s6_addr); |
804 | 0 | // NOLINTEND(cppcoreguidelines-pro-type-reinterpret-cast) |
805 | 0 |
|
806 | 0 | // NOLINTBEGIN(cppcoreguidelines-pro-bounds-pointer-arithmetic) |
807 | 0 | for (index = 0; index < bytes; ++index) { |
808 | 0 | if (lhs[index] != rhs[index]) { |
809 | 0 | return false; |
810 | 0 | } |
811 | 0 | } |
812 | 0 | // still here, now match remaining bits |
813 | 0 | uint8_t bits = d_bits % 8; |
814 | 0 | if (bits == 0) { |
815 | 0 | // no partial byte left to match, and lhs[index] would be one past the |
816 | 0 | // address for a /128 |
817 | 0 | return true; |
818 | 0 | } |
819 | 0 | auto mask = static_cast<uint8_t>(~(0xFF >> bits)); |
820 | 0 |
|
821 | 0 | return ((lhs[index]) == (rhs[index] & mask)); |
822 | 0 | // NOLINTEND(cppcoreguidelines-pro-bounds-pointer-arithmetic) |
823 | 0 | } |
824 | 0 | return false; |
825 | 0 | } |
826 | | |
827 | | //! If this ASCII IP address matches |
828 | | [[nodiscard]] bool match(const string& arg) const |
829 | 0 | { |
830 | 0 | ComboAddress address = makeComboAddress(arg); |
831 | 0 | return match(&address); |
832 | 0 | } |
833 | | |
834 | | //! If this IP address in native format matches |
835 | | [[nodiscard]] bool match4(uint32_t arg) const |
836 | 0 | { |
837 | 0 | return (arg & d_mask) == (ntohl(d_network.sin4.sin_addr.s_addr)); |
838 | 0 | } |
839 | | |
840 | | [[nodiscard]] string toString() const |
841 | 33.7k | { |
842 | 33.7k | return d_network.toStringNoInterface() + "/" + std::to_string((unsigned int)d_bits); |
843 | 33.7k | } |
844 | | |
845 | | [[nodiscard]] string toStringNoMask() const |
846 | 0 | { |
847 | 0 | return d_network.toStringNoInterface(); |
848 | 0 | } |
849 | | |
850 | | [[nodiscard]] string toByteString() const |
851 | 0 | { |
852 | 0 | ostringstream tmp; |
853 | 0 |
|
854 | 0 | tmp << (d_network.isIPv4() ? "\x04" : "\x06") |
855 | 0 | << d_network.toByteString() |
856 | 0 | << getBits(); |
857 | 0 |
|
858 | 0 | return tmp.str(); |
859 | 0 | } |
860 | | |
861 | | [[nodiscard]] const ComboAddress& getNetwork() const |
862 | 618k | { |
863 | 618k | return d_network; |
864 | 618k | } |
865 | | |
866 | | [[nodiscard]] const ComboAddress& getMaskedNetwork() const |
867 | 0 | { |
868 | 0 | return getNetwork(); |
869 | 0 | } |
870 | | |
871 | | [[nodiscard]] uint8_t getBits() const |
872 | 24.0k | { |
873 | 24.0k | return d_bits; |
874 | 24.0k | } |
875 | | |
876 | | [[nodiscard]] bool isIPv6() const |
877 | 47.8k | { |
878 | 47.8k | return d_network.sin6.sin6_family == AF_INET6; |
879 | 47.8k | } |
880 | | |
881 | | [[nodiscard]] bool isIPv4() const |
882 | 63.2k | { |
883 | 63.2k | return d_network.sin4.sin_family == AF_INET; |
884 | 63.2k | } |
885 | | |
886 | | bool operator<(const Netmask& rhs) const |
887 | 0 | { |
888 | 0 | if (empty() && !rhs.empty()) { |
889 | 0 | return false; |
890 | 0 | } |
891 | 0 | if (!empty() && rhs.empty()) { |
892 | 0 | return true; |
893 | 0 | } |
894 | 0 | if (d_bits > rhs.d_bits) { |
895 | 0 | return true; |
896 | 0 | } |
897 | 0 | if (d_bits < rhs.d_bits) { |
898 | 0 | return false; |
899 | 0 | } |
900 | 0 |
|
901 | 0 | return d_network < rhs.d_network; |
902 | 0 | } |
903 | | |
904 | | bool operator>(const Netmask& rhs) const |
905 | 0 | { |
906 | 0 | return rhs.operator<(*this); |
907 | 0 | } |
908 | | |
909 | | bool operator==(const Netmask& rhs) const |
910 | 0 | { |
911 | 0 | return std::tie(d_network, d_bits) == std::tie(rhs.d_network, rhs.d_bits); |
912 | 0 | } |
913 | | |
914 | | bool operator!=(const Netmask& rhs) const |
915 | 0 | { |
916 | 0 | return !operator==(rhs); |
917 | 0 | } |
918 | | |
919 | | [[nodiscard]] bool empty() const |
920 | 0 | { |
921 | 0 | return d_network.sin4.sin_family == 0; |
922 | 0 | } |
923 | | |
924 | | //! Get normalized version of the netmask. This means that all address bits below the network bits are zero. |
925 | | [[nodiscard]] Netmask getNormalized() const |
926 | 0 | { |
927 | 0 | return {getMaskedNetwork(), d_bits}; |
928 | 0 | } |
929 | | //! Get Netmask for super network of this one (i.e. with fewer network bits) |
930 | | [[nodiscard]] Netmask getSuper(uint8_t bits) const |
931 | 0 | { |
932 | 0 | return {d_network, std::min(d_bits, bits)}; |
933 | 0 | } |
934 | | |
935 | | //! Get the total number of address bits for this netmask (either 32 or 128 depending on IP version) |
936 | | [[nodiscard]] uint8_t getFullBits() const |
937 | 0 | { |
938 | 0 | return d_network.getBits(); |
939 | 0 | } |
940 | | |
941 | | /** Get the value of the bit at the provided bit index. When the index >= 0, |
942 | | the index is relative to the LSB starting at index zero. When the index < 0, |
943 | | the index is relative to the MSB starting at index -1 and counting down. |
944 | | When the index points outside the network bits, it always yields zero. |
945 | | */ |
946 | | [[nodiscard]] bool getBit(int bit) const |
947 | 0 | { |
948 | 0 | if (bit < -d_bits) { |
949 | 0 | return false; |
950 | 0 | } |
951 | 0 | if (bit >= 0) { |
952 | 0 | if (isIPv4()) { |
953 | 0 | if (bit >= 32 || bit < (32 - d_bits)) { |
954 | 0 | return false; |
955 | 0 | } |
956 | 0 | } |
957 | 0 | if (isIPv6()) { |
958 | 0 | if (bit >= 128 || bit < (128 - d_bits)) { |
959 | 0 | return false; |
960 | 0 | } |
961 | 0 | } |
962 | 0 | } |
963 | 0 | return d_network.getBit(bit); |
964 | 0 | } |
965 | | |
966 | | struct Hash |
967 | | { |
968 | | size_t operator()(const Netmask& netmask) const |
969 | 0 | { |
970 | 0 | return burtle(&netmask.d_bits, 1, ComboAddress::addressOnlyHash()(netmask.d_network)); |
971 | 0 | } |
972 | | }; |
973 | | |
974 | | private: |
975 | | ComboAddress d_network; |
976 | | uint32_t d_mask{0}; |
977 | | uint8_t d_bits{0}; |
978 | | }; |
979 | | |
980 | | namespace std |
981 | | { |
982 | | template <> |
983 | | struct hash<Netmask> |
984 | | { |
985 | | auto operator()(const Netmask& netmask) const |
986 | 0 | { |
987 | 0 | return Netmask::Hash{}(netmask); |
988 | 0 | } |
989 | | }; |
990 | | } |
991 | | |
992 | | /** Binary tree map implementation with <Netmask,T> pair. |
993 | | * |
994 | | * This is a binary tree implementation for storing attributes for IPv4 and IPv6 prefixes. |
995 | | * The most simple use case is simple NetmaskTree<bool> used by NetmaskGroup, which only |
996 | | * wants to know if given IP address is matched in the prefixes stored. |
997 | | * |
998 | | * This element is useful for anything that needs to *STORE* prefixes, and *MATCH* IP addresses |
999 | | * to a *LIST* of *PREFIXES*. Not the other way round. |
1000 | | * |
1001 | | * You can store IPv4 and IPv6 addresses to same tree, separate payload storage is kept per AFI. |
1002 | | * Network prefixes (Netmasks) are always recorded in normalized fashion, meaning that only |
1003 | | * the network bits are set. This is what is returned in the insert() and lookup() return |
1004 | | * values. |
1005 | | * |
1006 | | * Use swap if you need to move the tree to another NetmaskTree instance, it is WAY faster |
1007 | | * than using copy ctor or assignment operator, since it moves the nodes and tree root to |
1008 | | * new home instead of actually recreating the tree. |
1009 | | * |
1010 | | * Please see NetmaskGroup for example of simple use case. Other usecases can be found |
1011 | | * from GeoIPBackend and Sortlist, and from dnsdist. |
1012 | | */ |
1013 | | template <typename T, class K = Netmask> |
1014 | | class NetmaskTree |
1015 | | { |
1016 | | public: |
1017 | | class Iterator; |
1018 | | |
1019 | | using key_type = K; |
1020 | | using value_type = T; |
1021 | | using node_type = std::pair<const key_type, value_type>; |
1022 | | using size_type = size_t; |
1023 | | using iterator = class Iterator; |
1024 | | |
1025 | | private: |
1026 | | /** Single node in tree, internal use only. |
1027 | | */ |
1028 | | class TreeNode : boost::noncopyable |
1029 | | { |
1030 | | public: |
1031 | | explicit TreeNode() noexcept : |
1032 | | parent(nullptr), node(), assigned(false), d_bits(0) |
1033 | | { |
1034 | | } |
1035 | | explicit TreeNode(const key_type& key) : |
1036 | | parent(nullptr), node({key.getNormalized(), value_type()}), assigned(false), d_bits(key.getFullBits()) |
1037 | | { |
1038 | | } |
1039 | | |
1040 | | //<! Makes a left leaf node with specified key. |
1041 | | TreeNode* make_left(const key_type& key) |
1042 | 0 | { |
1043 | 0 | d_bits = node.first.getBits(); |
1044 | 0 | left = make_unique<TreeNode>(key); |
1045 | 0 | left->parent = this; |
1046 | 0 | return left.get(); |
1047 | 0 | } |
1048 | | |
1049 | | //<! Makes a right leaf node with specified key. |
1050 | | TreeNode* make_right(const key_type& key) |
1051 | 0 | { |
1052 | 0 | d_bits = node.first.getBits(); |
1053 | 0 | right = make_unique<TreeNode>(key); |
1054 | 0 | right->parent = this; |
1055 | 0 | return right.get(); |
1056 | 0 | } |
1057 | | |
1058 | | //<! Splits branch at indicated bit position by inserting key |
1059 | | TreeNode* split(const key_type& key, int bits) |
1060 | 0 | { |
1061 | 0 | if (parent == nullptr) { |
1062 | 0 | // not to be called on the root node |
1063 | 0 | throw std::logic_error( |
1064 | 0 | "NetmaskTree::TreeNode::split(): must not be called on root node"); |
1065 | 0 | } |
1066 | 0 |
|
1067 | 0 | // determine reference from parent |
1068 | 0 | unique_ptr<TreeNode>& parent_ref = (parent->left.get() == this ? parent->left : parent->right); |
1069 | 0 | if (parent_ref.get() != this) { |
1070 | 0 | throw std::logic_error( |
1071 | 0 | "NetmaskTree::TreeNode::split(): parent node reference is invalid"); |
1072 | 0 | } |
1073 | 0 |
|
1074 | 0 | // create new tree node for the new key and |
1075 | 0 | // attach the new node under our former parent |
1076 | 0 | auto new_intermediate_node = make_unique<TreeNode>(key); |
1077 | 0 | new_intermediate_node->d_bits = bits; |
1078 | 0 | new_intermediate_node->parent = parent; |
1079 | 0 | auto* new_intermediate_node_raw = new_intermediate_node.get(); |
1080 | 0 |
|
1081 | 0 | // hereafter new_intermediate points to "this" |
1082 | 0 | // ie the child of the new intermediate node |
1083 | 0 | std::swap(parent_ref, new_intermediate_node); |
1084 | 0 | // and we now assign this to current_node so |
1085 | 0 | // it's clear it no longer refers to the new |
1086 | 0 | // intermediate node |
1087 | 0 | std::unique_ptr<TreeNode> current_node = std::move(new_intermediate_node); |
1088 | 0 |
|
1089 | 0 | // attach "this" node below the new node |
1090 | 0 | // (left or right depending on bit) |
1091 | 0 | // technically the raw pointer escapes the duration of the |
1092 | 0 | // unique pointer, but just below we store the unique pointer |
1093 | 0 | // in the parent, so it lives as long as necessary |
1094 | 0 | // coverity[escape] |
1095 | 0 | current_node->parent = new_intermediate_node_raw; |
1096 | 0 | if (current_node->node.first.getBit(-1 - bits)) { |
1097 | 0 | new_intermediate_node_raw->right = std::move(current_node); |
1098 | 0 | } |
1099 | 0 | else { |
1100 | 0 | new_intermediate_node_raw->left = std::move(current_node); |
1101 | 0 | } |
1102 | 0 |
|
1103 | 0 | return new_intermediate_node_raw; |
1104 | 0 | } |
1105 | | |
1106 | | //<! Forks branch for new key at indicated bit position |
1107 | | TreeNode* fork(const key_type& key, int bits) |
1108 | 0 | { |
1109 | 0 | if (parent == nullptr) { |
1110 | 0 | // not to be called on the root node |
1111 | 0 | throw std::logic_error( |
1112 | 0 | "NetmaskTree::TreeNode::fork(): must not be called on root node"); |
1113 | 0 | } |
1114 | 0 |
|
1115 | 0 | // determine reference from parent |
1116 | 0 | unique_ptr<TreeNode>& parent_ref = (parent->left.get() == this ? parent->left : parent->right); |
1117 | 0 | if (parent_ref.get() != this) { |
1118 | 0 | throw std::logic_error( |
1119 | 0 | "NetmaskTree::TreeNode::fork(): parent node reference is invalid"); |
1120 | 0 | } |
1121 | 0 |
|
1122 | 0 | // create new tree node for the branch point |
1123 | 0 |
|
1124 | 0 | // the current node will now be a child of the new branch node |
1125 | 0 | // (hereafter new_child1 points to "this") |
1126 | 0 | unique_ptr<TreeNode> new_child1 = std::move(parent_ref); |
1127 | 0 | // attach the branch node under our former parent |
1128 | 0 | parent_ref = make_unique<TreeNode>(node.first.getSuper(bits)); |
1129 | 0 | auto* branch_node = parent_ref.get(); |
1130 | 0 | branch_node->d_bits = bits; |
1131 | 0 | branch_node->parent = parent; |
1132 | 0 |
|
1133 | 0 | // create second new leaf node for the new key |
1134 | 0 | unique_ptr<TreeNode> new_child2 = make_unique<TreeNode>(key); |
1135 | 0 | TreeNode* new_node = new_child2.get(); |
1136 | 0 |
|
1137 | 0 | // attach the new child nodes below the branch node |
1138 | 0 | // (left or right depending on bit) |
1139 | 0 | new_child1->parent = branch_node; |
1140 | 0 | new_child2->parent = branch_node; |
1141 | 0 | if (new_child1->node.first.getBit(-1 - bits)) { |
1142 | 0 | branch_node->right = std::move(new_child1); |
1143 | 0 | branch_node->left = std::move(new_child2); |
1144 | 0 | } |
1145 | 0 | else { |
1146 | 0 | branch_node->right = std::move(new_child2); |
1147 | 0 | branch_node->left = std::move(new_child1); |
1148 | 0 | } |
1149 | 0 | // now we have attached the new unique pointers to the tree: |
1150 | 0 | // - branch_node is below its parent |
1151 | 0 | // - new_child1 (ourselves) is below branch_node |
1152 | 0 | // - new_child2, the new leaf node, is below branch_node as well |
1153 | 0 |
|
1154 | 0 | return new_node; |
1155 | 0 | } |
1156 | | |
1157 | | //<! Traverse left branch depth-first |
1158 | | TreeNode* traverse_l() |
1159 | 0 | { |
1160 | 0 | TreeNode* tnode = this; |
1161 | 0 |
|
1162 | 0 | while (tnode->left) { |
1163 | 0 | tnode = tnode->left.get(); |
1164 | 0 | } |
1165 | 0 | return tnode; |
1166 | 0 | } |
1167 | | |
1168 | | //<! Traverse tree depth-first and in-order (L-N-R) |
1169 | | TreeNode* traverse_lnr() |
1170 | 0 | { |
1171 | 0 | TreeNode* tnode = this; |
1172 | 0 |
|
1173 | 0 | // precondition: descended left as deep as possible |
1174 | 0 | if (tnode->right) { |
1175 | 0 | // descend right |
1176 | 0 | tnode = tnode->right.get(); |
1177 | 0 | // descend left as deep as possible and return next node |
1178 | 0 | return tnode->traverse_l(); |
1179 | 0 | } |
1180 | 0 |
|
1181 | 0 | // ascend to parent |
1182 | 0 | while (tnode->parent != nullptr) { |
1183 | 0 | TreeNode* prev_child = tnode; |
1184 | 0 | tnode = tnode->parent; |
1185 | 0 |
|
1186 | 0 | // return this node, but only when we come from the left child branch |
1187 | 0 | if (tnode->left && tnode->left.get() == prev_child) { |
1188 | 0 | return tnode; |
1189 | 0 | } |
1190 | 0 | } |
1191 | 0 | return nullptr; |
1192 | 0 | } |
1193 | | |
1194 | | //<! Traverse only assigned nodes |
1195 | | TreeNode* traverse_lnr_assigned() |
1196 | 0 | { |
1197 | 0 | TreeNode* tnode = traverse_lnr(); |
1198 | 0 |
|
1199 | 0 | while (tnode != nullptr && !tnode->assigned) { |
1200 | 0 | tnode = tnode->traverse_lnr(); |
1201 | 0 | } |
1202 | 0 | return tnode; |
1203 | 0 | } |
1204 | | |
1205 | | unique_ptr<TreeNode> left; |
1206 | | unique_ptr<TreeNode> right; |
1207 | | TreeNode* parent; |
1208 | | |
1209 | | node_type node; |
1210 | | bool assigned; //<! Whether this node is assigned-to by the application |
1211 | | |
1212 | | int d_bits; //<! How many bits have been used so far |
1213 | | }; |
1214 | | |
1215 | | void cleanup_tree(TreeNode* node) |
1216 | 0 | { |
1217 | 0 | // only cleanup this node if it has no children and node not assigned |
1218 | 0 | if (!(node->left || node->right || node->assigned)) { |
1219 | 0 | // get parent node ptr |
1220 | 0 | TreeNode* pparent = node->parent; |
1221 | 0 | // delete this node |
1222 | 0 | if (pparent) { |
1223 | 0 | if (pparent->left.get() == node) { |
1224 | 0 | pparent->left.reset(); |
1225 | 0 | } |
1226 | 0 | else { |
1227 | 0 | pparent->right.reset(); |
1228 | 0 | } |
1229 | 0 | // now recurse up to the parent |
1230 | 0 | cleanup_tree(pparent); |
1231 | 0 | } |
1232 | 0 | } |
1233 | 0 | } |
1234 | | |
1235 | | void copyTree(const NetmaskTree& rhs) |
1236 | | { |
1237 | | try { |
1238 | | TreeNode* node = rhs.d_root.get(); |
1239 | | if (node != nullptr) { |
1240 | | node = node->traverse_l(); |
1241 | | } |
1242 | | while (node != nullptr) { |
1243 | | if (node->assigned) { |
1244 | | insert(node->node.first).second = node->node.second; |
1245 | | } |
1246 | | node = node->traverse_lnr(); |
1247 | | } |
1248 | | } |
1249 | | catch (const NetmaskException&) { |
1250 | | abort(); |
1251 | | } |
1252 | | catch (const std::logic_error&) { |
1253 | | abort(); |
1254 | | } |
1255 | | } |
1256 | | |
1257 | | public: |
1258 | | class Iterator |
1259 | | { |
1260 | | public: |
1261 | | using value_type = node_type; |
1262 | | using reference = node_type&; |
1263 | | using pointer = node_type*; |
1264 | | using iterator_category = std::forward_iterator_tag; |
1265 | | using difference_type = size_type; |
1266 | | |
1267 | | private: |
1268 | | friend class NetmaskTree; |
1269 | | |
1270 | | const NetmaskTree* d_tree; |
1271 | | TreeNode* d_node; |
1272 | | |
1273 | | Iterator(const NetmaskTree* tree, TreeNode* node) : |
1274 | | d_tree(tree), d_node(node) |
1275 | | { |
1276 | | } |
1277 | | |
1278 | | public: |
1279 | | Iterator() : |
1280 | | d_tree(nullptr), d_node(nullptr) {} |
1281 | | |
1282 | | Iterator& operator++() // prefix |
1283 | 0 | { |
1284 | 0 | if (d_node == nullptr) { |
1285 | 0 | throw std::logic_error( |
1286 | 0 | "NetmaskTree::Iterator::operator++: iterator is invalid"); |
1287 | 0 | } |
1288 | 0 | d_node = d_node->traverse_lnr_assigned(); |
1289 | 0 | return *this; |
1290 | 0 | } |
1291 | | Iterator operator++(int) // postfix |
1292 | | { |
1293 | | Iterator tmp(*this); |
1294 | | operator++(); |
1295 | | return tmp; |
1296 | | } |
1297 | | |
1298 | | reference operator*() |
1299 | 0 | { |
1300 | 0 | if (d_node == nullptr) { |
1301 | 0 | throw std::logic_error( |
1302 | 0 | "NetmaskTree::Iterator::operator*: iterator is invalid"); |
1303 | 0 | } |
1304 | 0 | return d_node->node; |
1305 | 0 | } |
1306 | | |
1307 | | pointer operator->() |
1308 | 0 | { |
1309 | 0 | if (d_node == nullptr) { |
1310 | 0 | throw std::logic_error( |
1311 | 0 | "NetmaskTree::Iterator::operator->: iterator is invalid"); |
1312 | 0 | } |
1313 | 0 | return &d_node->node; |
1314 | 0 | } |
1315 | | |
1316 | | bool operator==(const Iterator& rhs) const |
1317 | 0 | { |
1318 | 0 | return (d_tree == rhs.d_tree && d_node == rhs.d_node); |
1319 | 0 | } |
1320 | | bool operator!=(const Iterator& rhs) const |
1321 | 0 | { |
1322 | 0 | return !(*this == rhs); |
1323 | 0 | } |
1324 | | }; |
1325 | | |
1326 | | NetmaskTree() noexcept : |
1327 | | d_root(new TreeNode()), d_left(nullptr) |
1328 | | { |
1329 | | } |
1330 | | |
1331 | | NetmaskTree(const NetmaskTree& rhs) : |
1332 | | d_root(new TreeNode()), d_left(nullptr) |
1333 | | { |
1334 | | copyTree(rhs); |
1335 | | } |
1336 | | |
1337 | | ~NetmaskTree() = default; |
1338 | | |
1339 | | NetmaskTree& operator=(const NetmaskTree& rhs) |
1340 | | { |
1341 | | if (this != &rhs) { |
1342 | | clear(); |
1343 | | copyTree(rhs); |
1344 | | } |
1345 | | return *this; |
1346 | | } |
1347 | | |
1348 | | NetmaskTree(NetmaskTree&&) noexcept = default; |
1349 | | NetmaskTree& operator=(NetmaskTree&&) noexcept = default; |
1350 | | |
1351 | | [[nodiscard]] iterator begin() const |
1352 | 0 | { |
1353 | 0 | return Iterator(this, d_left); |
1354 | 0 | } |
1355 | | [[nodiscard]] iterator end() const |
1356 | 0 | { |
1357 | 0 | return Iterator(this, nullptr); |
1358 | 0 | } |
1359 | | iterator begin() |
1360 | | { |
1361 | | return Iterator(this, d_left); |
1362 | | } |
1363 | | iterator end() |
1364 | | { |
1365 | | return Iterator(this, nullptr); |
1366 | | } |
1367 | | |
1368 | | node_type& insert(const string& mask) |
1369 | | { |
1370 | | return insert(key_type(mask)); |
1371 | | } |
1372 | | |
1373 | | //<! Creates new value-pair in tree and returns it. |
1374 | | node_type& insert(const key_type& key) |
1375 | 0 | { |
1376 | 0 | TreeNode* node{}; |
1377 | 0 | bool is_left = true; |
1378 | 0 |
|
1379 | 0 | // we turn left on IPv4 and right on IPv6 |
1380 | 0 | if (key.isIPv4()) { |
1381 | 0 | node = d_root->left.get(); |
1382 | 0 | if (node == nullptr) { |
1383 | 0 |
|
1384 | 0 | d_root->left = make_unique<TreeNode>(key); |
1385 | 0 | node = d_root->left.get(); |
1386 | 0 | node->assigned = true; |
1387 | 0 | node->parent = d_root.get(); |
1388 | 0 | d_size++; |
1389 | 0 | d_left = node; |
1390 | 0 | return node->node; |
1391 | 0 | } |
1392 | 0 | } |
1393 | 0 | else if (key.isIPv6()) { |
1394 | 0 | node = d_root->right.get(); |
1395 | 0 | if (node == nullptr) { |
1396 | 0 |
|
1397 | 0 | d_root->right = make_unique<TreeNode>(key); |
1398 | 0 | node = d_root->right.get(); |
1399 | 0 | node->assigned = true; |
1400 | 0 | node->parent = d_root.get(); |
1401 | 0 | d_size++; |
1402 | 0 | if (!d_root->left) { |
1403 | 0 | d_left = node; |
1404 | 0 | } |
1405 | 0 | return node->node; |
1406 | 0 | } |
1407 | 0 | if (d_root->left) { |
1408 | 0 | is_left = false; |
1409 | 0 | } |
1410 | 0 | } |
1411 | 0 | else { |
1412 | 0 | throw NetmaskException("invalid address family"); |
1413 | 0 | } |
1414 | 0 |
|
1415 | 0 | // we turn left on 0 and right on 1 |
1416 | 0 | int bits = 0; |
1417 | 0 | for (; bits < key.getBits(); bits++) { |
1418 | 0 | bool vall = key.getBit(-1 - bits); |
1419 | 0 |
|
1420 | 0 | if (bits >= node->d_bits) { |
1421 | 0 | // the end of the current node is reached; continue with the next |
1422 | 0 | if (vall) { |
1423 | 0 | if (node->left || node->assigned) { |
1424 | 0 | is_left = false; |
1425 | 0 | } |
1426 | 0 | if (!node->right) { |
1427 | 0 | // the right branch doesn't exist yet; attach our key here |
1428 | 0 | node = node->make_right(key); |
1429 | 0 | break; |
1430 | 0 | } |
1431 | 0 | node = node->right.get(); |
1432 | 0 | } |
1433 | 0 | else { |
1434 | 0 | if (!node->left) { |
1435 | 0 | // the left branch doesn't exist yet; attach our key here |
1436 | 0 | node = node->make_left(key); |
1437 | 0 | break; |
1438 | 0 | } |
1439 | 0 | node = node->left.get(); |
1440 | 0 | } |
1441 | 0 | continue; |
1442 | 0 | } |
1443 | 0 | if (bits >= node->node.first.getBits()) { |
1444 | 0 | // the matching branch ends here, yet the key netmask has more bits; add a |
1445 | 0 | // child node below the existing branch leaf. |
1446 | 0 | if (vall) { |
1447 | 0 | if (node->assigned) { |
1448 | 0 | is_left = false; |
1449 | 0 | } |
1450 | 0 | node = node->make_right(key); |
1451 | 0 | } |
1452 | 0 | else { |
1453 | 0 | node = node->make_left(key); |
1454 | 0 | } |
1455 | 0 | break; |
1456 | 0 | } |
1457 | 0 | bool valr = node->node.first.getBit(-1 - bits); |
1458 | 0 | if (vall != valr) { |
1459 | 0 | if (vall) { |
1460 | 0 | is_left = false; |
1461 | 0 | } |
1462 | 0 | // the branch matches just upto this point, yet continues in a different |
1463 | 0 | // direction; fork the branch. |
1464 | 0 | node = node->fork(key, bits); |
1465 | 0 | break; |
1466 | 0 | } |
1467 | 0 | } |
1468 | 0 |
|
1469 | 0 | if (node->node.first.getBits() > key.getBits()) { |
1470 | 0 | // key is a super-network of the matching node; split the branch and |
1471 | 0 | // insert a node for the key above the matching node. |
1472 | 0 | node = node->split(key, key.getBits()); |
1473 | 0 | } |
1474 | 0 |
|
1475 | 0 | if (node->left) { |
1476 | 0 | is_left = false; |
1477 | 0 | } |
1478 | 0 |
|
1479 | 0 | node_type& value = node->node; |
1480 | 0 |
|
1481 | 0 | if (!node->assigned) { |
1482 | 0 | // only increment size if not assigned before |
1483 | 0 | d_size++; |
1484 | 0 | // update the pointer to the left-most tree node |
1485 | 0 | if (is_left) { |
1486 | 0 | d_left = node; |
1487 | 0 | } |
1488 | 0 | node->assigned = true; |
1489 | 0 | } |
1490 | 0 | else { |
1491 | 0 | // tree node exists for this value |
1492 | 0 | if (is_left && d_left != node) { |
1493 | 0 | throw std::logic_error( |
1494 | 0 | "NetmaskTree::insert(): lost track of left-most node in tree"); |
1495 | 0 | } |
1496 | 0 | } |
1497 | 0 |
|
1498 | 0 | return value; |
1499 | 0 | } |
1500 | | |
1501 | | //<! Creates or updates value |
1502 | | void insert_or_assign(const key_type& mask, const value_type& value) |
1503 | | { |
1504 | | insert(mask).second = value; |
1505 | | } |
1506 | | |
1507 | | void insert_or_assign(const string& mask, const value_type& value) |
1508 | | { |
1509 | | insert(key_type(mask)).second = value; |
1510 | | } |
1511 | | |
1512 | | //<! check if given key is present in TreeMap |
1513 | | [[nodiscard]] bool has_key(const key_type& key) const |
1514 | | { |
1515 | | const node_type* ptr = lookup(key); |
1516 | | return ptr && ptr->first == key; |
1517 | | } |
1518 | | |
1519 | | //<! Returns "best match" for key_type, which might not be value |
1520 | | [[nodiscard]] node_type* lookup(const key_type& value) const |
1521 | | { |
1522 | | if (empty()) { |
1523 | | return nullptr; |
1524 | | } |
1525 | | uint8_t max_bits = value.getBits(); |
1526 | | return lookupImpl(value, max_bits); |
1527 | | } |
1528 | | |
1529 | | //<! Perform best match lookup for value, using at most max_bits |
1530 | | [[nodiscard]] node_type* lookup(const ComboAddress& value, int max_bits = 128) const |
1531 | 0 | { |
1532 | 0 | if (empty()) { |
1533 | 0 | return nullptr; |
1534 | 0 | } |
1535 | 0 | uint8_t addr_bits = value.getBits(); |
1536 | 0 | if (max_bits < 0 || max_bits > addr_bits) { |
1537 | 0 | max_bits = addr_bits; |
1538 | 0 | } |
1539 | 0 |
|
1540 | 0 | return lookupImpl(key_type(value, max_bits), max_bits); |
1541 | 0 | } |
1542 | | |
1543 | | //<! Removes key from TreeMap. |
1544 | | void erase(const key_type& key) |
1545 | 0 | { |
1546 | 0 | TreeNode* node = nullptr; |
1547 | 0 |
|
1548 | 0 | if (key.isIPv4()) { |
1549 | 0 | node = d_root->left.get(); |
1550 | 0 | } |
1551 | 0 | else if (key.isIPv6()) { |
1552 | 0 | node = d_root->right.get(); |
1553 | 0 | } |
1554 | 0 | else { |
1555 | 0 | throw NetmaskException("invalid address family"); |
1556 | 0 | } |
1557 | 0 | // no tree, no value |
1558 | 0 | if (node == nullptr) { |
1559 | 0 | return; |
1560 | 0 | } |
1561 | 0 | int bits = 0; |
1562 | 0 | for (; node && bits < key.getBits(); bits++) { |
1563 | 0 | bool vall = key.getBit(-1 - bits); |
1564 | 0 | if (bits >= node->d_bits) { |
1565 | 0 | // the end of the current node is reached; continue with the next |
1566 | 0 | if (vall) { |
1567 | 0 | node = node->right.get(); |
1568 | 0 | } |
1569 | 0 | else { |
1570 | 0 | node = node->left.get(); |
1571 | 0 | } |
1572 | 0 | continue; |
1573 | 0 | } |
1574 | 0 | if (bits >= node->node.first.getBits()) { |
1575 | 0 | // the matching branch ends here |
1576 | 0 | if (key.getBits() != node->node.first.getBits()) { |
1577 | 0 | node = nullptr; |
1578 | 0 | } |
1579 | 0 | break; |
1580 | 0 | } |
1581 | 0 | bool valr = node->node.first.getBit(-1 - bits); |
1582 | 0 | if (vall != valr) { |
1583 | 0 | // the branch matches just upto this point, yet continues in a different |
1584 | 0 | // direction |
1585 | 0 | node = nullptr; |
1586 | 0 | break; |
1587 | 0 | } |
1588 | 0 | } |
1589 | 0 | if (node) { |
1590 | 0 | if (d_size == 0) { |
1591 | 0 | throw std::logic_error( |
1592 | 0 | "NetmaskTree::erase(): size of tree is zero before erase"); |
1593 | 0 | } |
1594 | 0 | d_size--; |
1595 | 0 | node->assigned = false; |
1596 | 0 | node->node.second = value_type(); |
1597 | 0 |
|
1598 | 0 | if (node == d_left) { |
1599 | 0 | d_left = d_left->traverse_lnr_assigned(); |
1600 | 0 | } |
1601 | 0 | cleanup_tree(node); |
1602 | 0 | } |
1603 | 0 | } |
1604 | | |
1605 | | void erase(const string& key) |
1606 | | { |
1607 | | erase(key_type(key)); |
1608 | | } |
1609 | | |
1610 | | //<! checks whether the container is empty. |
1611 | | [[nodiscard]] bool empty() const |
1612 | 0 | { |
1613 | 0 | return (d_size == 0); |
1614 | 0 | } |
1615 | | |
1616 | | //<! returns the number of elements |
1617 | | [[nodiscard]] size_type size() const |
1618 | 0 | { |
1619 | 0 | return d_size; |
1620 | 0 | } |
1621 | | |
1622 | | //<! See if given ComboAddress matches any prefix |
1623 | | [[nodiscard]] bool match(const ComboAddress& value) const |
1624 | | { |
1625 | | return (lookup(value) != nullptr); |
1626 | | } |
1627 | | |
1628 | | [[nodiscard]] bool match(const std::string& value) const |
1629 | | { |
1630 | | return match(ComboAddress(value)); |
1631 | | } |
1632 | | |
1633 | | //<! Clean out the tree |
1634 | | void clear() |
1635 | 0 | { |
1636 | 0 | d_root = make_unique<TreeNode>(); |
1637 | 0 | d_left = nullptr; |
1638 | 0 | d_size = 0; |
1639 | 0 | } |
1640 | | |
1641 | | //<! swaps the contents with another NetmaskTree |
1642 | | void swap(NetmaskTree& rhs) noexcept |
1643 | | { |
1644 | | std::swap(d_root, rhs.d_root); |
1645 | | std::swap(d_left, rhs.d_left); |
1646 | | std::swap(d_size, rhs.d_size); |
1647 | | } |
1648 | | |
1649 | | private: |
1650 | | [[nodiscard]] node_type* lookupImpl(const key_type& value, uint8_t max_bits) const |
1651 | 0 | { |
1652 | 0 | TreeNode* node = nullptr; |
1653 | 0 |
|
1654 | 0 | if (value.isIPv4()) { |
1655 | 0 | node = d_root->left.get(); |
1656 | 0 | } |
1657 | 0 | else if (value.isIPv6()) { |
1658 | 0 | node = d_root->right.get(); |
1659 | 0 | } |
1660 | 0 | else { |
1661 | 0 | throw NetmaskException("invalid address family"); |
1662 | 0 | } |
1663 | 0 | if (node == nullptr) { |
1664 | 0 | return nullptr; |
1665 | 0 | } |
1666 | 0 |
|
1667 | 0 | node_type* ret = nullptr; |
1668 | 0 |
|
1669 | 0 | int bits = 0; |
1670 | 0 | for (; bits < max_bits; bits++) { |
1671 | 0 | bool vall = value.getBit(-1 - bits); |
1672 | 0 | if (bits >= node->d_bits) { |
1673 | 0 | // the end of the current node is reached; continue with the next |
1674 | 0 | // (we keep track of last assigned node) |
1675 | 0 | if (node->assigned && bits == node->node.first.getBits()) { |
1676 | 0 | ret = &node->node; |
1677 | 0 | } |
1678 | 0 | if (vall) { |
1679 | 0 | if (!node->right) { |
1680 | 0 | break; |
1681 | 0 | } |
1682 | 0 | node = node->right.get(); |
1683 | 0 | } |
1684 | 0 | else { |
1685 | 0 | if (!node->left) { |
1686 | 0 | break; |
1687 | 0 | } |
1688 | 0 | node = node->left.get(); |
1689 | 0 | } |
1690 | 0 | continue; |
1691 | 0 | } |
1692 | 0 | if (bits >= node->node.first.getBits()) { |
1693 | 0 | // the matching branch ends here |
1694 | 0 | break; |
1695 | 0 | } |
1696 | 0 | bool valr = node->node.first.getBit(-1 - bits); |
1697 | 0 | if (vall != valr) { |
1698 | 0 | // the branch matches just upto this point, yet continues in a different |
1699 | 0 | // direction |
1700 | 0 | break; |
1701 | 0 | } |
1702 | 0 | } |
1703 | 0 | // needed if we did not find one in loop |
1704 | 0 | if (node->assigned && bits == node->node.first.getBits()) { |
1705 | 0 | ret = &node->node; |
1706 | 0 | } |
1707 | 0 | // this can be nullptr. |
1708 | 0 | return ret; |
1709 | 0 | } |
1710 | | |
1711 | | unique_ptr<TreeNode> d_root; //<! Root of our tree |
1712 | | TreeNode* d_left; |
1713 | | size_type d_size{0}; |
1714 | | }; |
1715 | | |
1716 | | /** This class represents a group of supplemental Netmask classes. An IP address matches |
1717 | | if it is matched by one or more of the Netmask objects within. |
1718 | | */ |
1719 | | class NetmaskGroup |
1720 | | { |
1721 | | public: |
1722 | | NetmaskGroup() noexcept = default; |
1723 | | |
1724 | | //! If this IP address is matched by any of the classes within |
1725 | | |
1726 | | bool match(const ComboAddress* address) const |
1727 | 0 | { |
1728 | 0 | const auto& ret = tree.lookup(*address); |
1729 | 0 | if (ret != nullptr) { |
1730 | 0 | return ret->second; |
1731 | 0 | } |
1732 | 0 | return false; |
1733 | 0 | } |
1734 | | |
1735 | | [[nodiscard]] bool match(const ComboAddress& address) const |
1736 | 0 | { |
1737 | 0 | return match(&address); |
1738 | 0 | } |
1739 | | |
1740 | | bool lookup(const ComboAddress* address, Netmask* nmp) const |
1741 | 0 | { |
1742 | 0 | const auto& ret = tree.lookup(*address); |
1743 | 0 | if (ret != nullptr) { |
1744 | 0 | if (nmp != nullptr) { |
1745 | 0 | *nmp = ret->first; |
1746 | 0 | } |
1747 | 0 | return ret->second; |
1748 | 0 | } |
1749 | 0 | return false; |
1750 | 0 | } |
1751 | | |
1752 | | bool lookup(const ComboAddress& address, Netmask* nmp) const |
1753 | 0 | { |
1754 | 0 | return lookup(&address, nmp); |
1755 | 0 | } |
1756 | | |
1757 | | //! Add this string to the list of possible matches |
1758 | | void addMask(const string& address, bool positive = true) |
1759 | 0 | { |
1760 | 0 | if (!address.empty() && address[0] == '!') { |
1761 | 0 | addMask(Netmask(address.substr(1)), false); |
1762 | 0 | } |
1763 | 0 | else { |
1764 | 0 | addMask(Netmask(address), positive); |
1765 | 0 | } |
1766 | 0 | } |
1767 | | |
1768 | | //! Add this Netmask to the list of possible matches |
1769 | | void addMask(const Netmask& netmask, bool positive = true) |
1770 | 0 | { |
1771 | 0 | tree.insert(netmask).second = positive; |
1772 | 0 | } |
1773 | | |
1774 | | void addMasks(const NetmaskGroup& group, std::optional<bool> positive) |
1775 | 0 | { |
1776 | 0 | for (const auto& entry : group.tree) { |
1777 | 0 | addMask(entry.first, positive ? *positive : entry.second); |
1778 | 0 | } |
1779 | 0 | } |
1780 | | |
1781 | | //! Delete this Netmask from the list of possible matches |
1782 | | void deleteMask(const Netmask& netmask) |
1783 | 0 | { |
1784 | 0 | tree.erase(netmask); |
1785 | 0 | } |
1786 | | |
1787 | | void deleteMasks(const NetmaskGroup& group) |
1788 | 0 | { |
1789 | 0 | for (const auto& entry : group.tree) { |
1790 | 0 | deleteMask(entry.first); |
1791 | 0 | } |
1792 | 0 | } |
1793 | | |
1794 | | void deleteMask(const std::string& address) |
1795 | 0 | { |
1796 | 0 | if (!address.empty()) { |
1797 | 0 | deleteMask(Netmask(address)); |
1798 | 0 | } |
1799 | 0 | } |
1800 | | |
1801 | | void clear() |
1802 | 0 | { |
1803 | 0 | tree.clear(); |
1804 | 0 | } |
1805 | | |
1806 | | [[nodiscard]] bool empty() const |
1807 | 0 | { |
1808 | 0 | return tree.empty(); |
1809 | 0 | } |
1810 | | |
1811 | | [[nodiscard]] size_t size() const |
1812 | 0 | { |
1813 | 0 | return tree.size(); |
1814 | 0 | } |
1815 | | |
1816 | | [[nodiscard]] string toString() const |
1817 | 0 | { |
1818 | 0 | ostringstream str; |
1819 | 0 | for (auto iter = tree.begin(); iter != tree.end(); ++iter) { |
1820 | 0 | if (iter != tree.begin()) { |
1821 | 0 | str << ", "; |
1822 | 0 | } |
1823 | 0 | if (!(iter->second)) { |
1824 | 0 | str << "!"; |
1825 | 0 | } |
1826 | 0 | str << iter->first.toString(); |
1827 | 0 | } |
1828 | 0 | return str.str(); |
1829 | 0 | } |
1830 | | |
1831 | | [[nodiscard]] std::vector<std::string> toStringVector() const |
1832 | 0 | { |
1833 | 0 | std::vector<std::string> out; |
1834 | 0 | out.reserve(tree.size()); |
1835 | 0 | for (const auto& entry : tree) { |
1836 | 0 | out.push_back((entry.second ? "" : "!") + entry.first.toString()); |
1837 | 0 | } |
1838 | 0 | return out; |
1839 | 0 | } |
1840 | | |
1841 | | void toMasks(const string& ips) |
1842 | 0 | { |
1843 | 0 | vector<string> parts; |
1844 | 0 | stringtok(parts, ips, ", \t"); |
1845 | 0 |
|
1846 | 0 | for (const auto& part : parts) { |
1847 | 0 | addMask(part); |
1848 | 0 | } |
1849 | 0 | } |
1850 | | |
1851 | | private: |
1852 | | NetmaskTree<bool> tree; |
1853 | | }; |
1854 | | |
1855 | | struct SComboAddress |
1856 | | { |
1857 | | SComboAddress(const ComboAddress& orig) : |
1858 | 0 | ca(orig) {} |
1859 | | ComboAddress ca; |
1860 | | bool operator<(const SComboAddress& rhs) const |
1861 | 0 | { |
1862 | 0 | return ComboAddress::addressOnlyLessThan()(ca, rhs.ca); |
1863 | 0 | } |
1864 | | operator const ComboAddress&() const |
1865 | 0 | { |
1866 | 0 | return ca; |
1867 | 0 | } |
1868 | | }; |
1869 | | |
1870 | | class NetworkError : public runtime_error |
1871 | | { |
1872 | | public: |
1873 | | NetworkError(const string& why = "Network Error") : |
1874 | | runtime_error(why.c_str()) |
1875 | 0 | {} |
1876 | | NetworkError(const char* why = "Network Error") : |
1877 | | runtime_error(why) |
1878 | 0 | {} |
1879 | | }; |
1880 | | |
1881 | | class AddressAndPortRange |
1882 | | { |
1883 | | public: |
1884 | | AddressAndPortRange() : |
1885 | | d_addrMask(0), d_portMask(0) |
1886 | | { |
1887 | | d_addr.sin4.sin_family = 0; // disable this doing anything useful |
1888 | | d_addr.sin4.sin_port = 0; // this guarantees d_network compares identical |
1889 | | } |
1890 | | |
1891 | | AddressAndPortRange(ComboAddress address, uint8_t addrMask, uint8_t portMask = 0) : |
1892 | | d_addr(address), d_addrMask(addrMask), d_portMask(portMask) |
1893 | 0 | { |
1894 | 0 | if (!d_addr.isIPv4()) { |
1895 | 0 | d_portMask = 0; |
1896 | 0 | } |
1897 | 0 |
|
1898 | 0 | uint16_t port = d_addr.getPort(); |
1899 | 0 | if (d_portMask < 16) { |
1900 | 0 | auto mask = static_cast<uint16_t>(~(0xFFFF >> d_portMask)); |
1901 | 0 | port = port & mask; |
1902 | 0 | } |
1903 | 0 |
|
1904 | 0 | if (d_addrMask < d_addr.getBits()) { |
1905 | 0 | if (d_portMask > 0) { |
1906 | 0 | throw std::runtime_error("Trying to create a AddressAndPortRange with a reduced address mask (" + std::to_string(d_addrMask) + ") and a port range (" + std::to_string(d_portMask) + ")"); |
1907 | 0 | } |
1908 | 0 | d_addr = Netmask(d_addr, d_addrMask).getMaskedNetwork(); |
1909 | 0 | } |
1910 | 0 | d_addr.setPort(port); |
1911 | 0 | } |
1912 | | |
1913 | | [[nodiscard]] uint8_t getFullBits() const |
1914 | 0 | { |
1915 | 0 | return d_addr.getBits() + 16; |
1916 | 0 | } |
1917 | | |
1918 | | [[nodiscard]] uint8_t getBits() const |
1919 | 0 | { |
1920 | 0 | if (d_addrMask < d_addr.getBits()) { |
1921 | 0 | return d_addrMask; |
1922 | 0 | } |
1923 | 0 |
|
1924 | 0 | return d_addr.getBits() + d_portMask; |
1925 | 0 | } |
1926 | | |
1927 | | /** Get the value of the bit at the provided bit index. When the index >= 0, |
1928 | | the index is relative to the LSB starting at index zero. When the index < 0, |
1929 | | the index is relative to the MSB starting at index -1 and counting down. |
1930 | | */ |
1931 | | [[nodiscard]] bool getBit(int index) const |
1932 | 0 | { |
1933 | 0 | if (index >= getFullBits()) { |
1934 | 0 | return false; |
1935 | 0 | } |
1936 | 0 | if (index < 0) { |
1937 | 0 | index = getFullBits() + index; |
1938 | 0 | } |
1939 | 0 |
|
1940 | 0 | if (index < 16) { |
1941 | 0 | /* we are into the port bits */ |
1942 | 0 | uint16_t port = d_addr.getPort(); |
1943 | 0 | return ((port & (1U << index)) != 0x0000); |
1944 | 0 | } |
1945 | 0 |
|
1946 | 0 | index -= 16; |
1947 | 0 |
|
1948 | 0 | return d_addr.getBit(index); |
1949 | 0 | } |
1950 | | |
1951 | | [[nodiscard]] bool isIPv4() const |
1952 | 0 | { |
1953 | 0 | return d_addr.isIPv4(); |
1954 | 0 | } |
1955 | | |
1956 | | [[nodiscard]] bool isIPv6() const |
1957 | 0 | { |
1958 | 0 | return d_addr.isIPv6(); |
1959 | 0 | } |
1960 | | |
1961 | | [[nodiscard]] AddressAndPortRange getNormalized() const |
1962 | 0 | { |
1963 | 0 | return {d_addr, d_addrMask, d_portMask}; |
1964 | 0 | } |
1965 | | |
1966 | | [[nodiscard]] AddressAndPortRange getSuper(uint8_t bits) const |
1967 | 0 | { |
1968 | 0 | if (bits <= d_addrMask) { |
1969 | 0 | return {d_addr, bits, 0}; |
1970 | 0 | } |
1971 | 0 | if (bits <= d_addrMask + d_portMask) { |
1972 | 0 | return {d_addr, d_addrMask, static_cast<uint8_t>(d_portMask - (bits - d_addrMask))}; |
1973 | 0 | } |
1974 | 0 |
|
1975 | 0 | return {d_addr, d_addrMask, d_portMask}; |
1976 | 0 | } |
1977 | | |
1978 | | [[nodiscard]] const ComboAddress& getNetwork() const |
1979 | 0 | { |
1980 | 0 | return d_addr; |
1981 | 0 | } |
1982 | | |
1983 | | [[nodiscard]] string toString() const |
1984 | 0 | { |
1985 | 0 | if (d_addrMask < d_addr.getBits() || d_portMask == 0) { |
1986 | 0 | return d_addr.toStringNoInterface() + "/" + std::to_string(d_addrMask); |
1987 | 0 | } |
1988 | 0 | return d_addr.toStringNoInterface() + ":" + std::to_string(d_addr.getPort()) + "/" + std::to_string(d_portMask); |
1989 | 0 | } |
1990 | | |
1991 | | [[nodiscard]] bool empty() const |
1992 | 0 | { |
1993 | 0 | return d_addr.sin4.sin_family == 0; |
1994 | 0 | } |
1995 | | |
1996 | | bool operator==(const AddressAndPortRange& rhs) const |
1997 | 0 | { |
1998 | 0 | return std::tie(d_addr, d_addrMask, d_portMask) == std::tie(rhs.d_addr, rhs.d_addrMask, rhs.d_portMask); |
1999 | 0 | } |
2000 | | |
2001 | | bool operator<(const AddressAndPortRange& rhs) const |
2002 | 0 | { |
2003 | 0 | if (empty() && !rhs.empty()) { |
2004 | 0 | return false; |
2005 | 0 | } |
2006 | 0 |
|
2007 | 0 | if (!empty() && rhs.empty()) { |
2008 | 0 | return true; |
2009 | 0 | } |
2010 | 0 |
|
2011 | 0 | if (d_addrMask > rhs.d_addrMask) { |
2012 | 0 | return true; |
2013 | 0 | } |
2014 | 0 |
|
2015 | 0 | if (d_addrMask < rhs.d_addrMask) { |
2016 | 0 | return false; |
2017 | 0 | } |
2018 | 0 |
|
2019 | 0 | if (d_addr < rhs.d_addr) { |
2020 | 0 | return true; |
2021 | 0 | } |
2022 | 0 |
|
2023 | 0 | if (d_addr > rhs.d_addr) { |
2024 | 0 | return false; |
2025 | 0 | } |
2026 | 0 |
|
2027 | 0 | if (d_portMask > rhs.d_portMask) { |
2028 | 0 | return true; |
2029 | 0 | } |
2030 | 0 |
|
2031 | 0 | if (d_portMask < rhs.d_portMask) { |
2032 | 0 | return false; |
2033 | 0 | } |
2034 | 0 |
|
2035 | 0 | return d_addr.getPort() < rhs.d_addr.getPort(); |
2036 | 0 | } |
2037 | | |
2038 | | bool operator>(const AddressAndPortRange& rhs) const |
2039 | 0 | { |
2040 | 0 | return rhs.operator<(*this); |
2041 | 0 | } |
2042 | | |
2043 | | struct hash |
2044 | | { |
2045 | | uint32_t operator()(const AddressAndPortRange& apr) const |
2046 | 0 | { |
2047 | 0 | ComboAddress::addressOnlyHash hashOp; |
2048 | 0 | uint16_t port = apr.d_addr.getPort(); |
2049 | 0 | /* it's fine to hash the whole address and port because the non-relevant parts have |
2050 | 0 | been masked to 0 */ |
2051 | 0 | return burtle(reinterpret_cast<const unsigned char*>(&port), sizeof(port), hashOp(apr.d_addr)); // NOLINT(cppcoreguidelines-pro-type-reinterpret-cast) |
2052 | 0 | } |
2053 | | }; |
2054 | | |
2055 | | private: |
2056 | | ComboAddress d_addr; |
2057 | | uint8_t d_addrMask; |
2058 | | /* only used for v4 addresses */ |
2059 | | uint8_t d_portMask; |
2060 | | }; |
2061 | | |
2062 | | int SSocket(int family, int type, int flags); |
2063 | | int SConnect(int sockfd, bool fastopen, const ComboAddress& remote); |
2064 | | /* tries to connect to remote for a maximum of timeout seconds. |
2065 | | sockfd should be set to non-blocking beforehand. |
2066 | | returns 0 on success (the socket is writable), throw a |
2067 | | runtime_error otherwise */ |
2068 | | int SConnectWithTimeout(int sockfd, bool fastopen, const ComboAddress& remote, const struct timeval& timeout); |
2069 | | int SBind(int sockfd, const ComboAddress& local); |
2070 | | int SAccept(int sockfd, ComboAddress& remote); |
2071 | | int SListen(int sockfd, int limit); |
2072 | | int SSetsockopt(int sockfd, int level, int opname, int value); |
2073 | | void setSocketIgnorePMTU(int sockfd, int family); |
2074 | | void setSocketForcePMTU(int sockfd, int family); |
2075 | | bool setReusePort(int sockfd); |
2076 | | |
2077 | | #if defined(IP_PKTINFO) |
2078 | | #define GEN_IP_PKTINFO IP_PKTINFO |
2079 | | #elif defined(IP_RECVDSTADDR) |
2080 | | #define GEN_IP_PKTINFO IP_RECVDSTADDR |
2081 | | #endif |
2082 | | |
2083 | | bool IsAnyAddress(const ComboAddress& addr); |
2084 | | bool HarvestDestinationAddress(const struct msghdr* msgh, ComboAddress* destination); |
2085 | | bool HarvestTimestamp(struct msghdr* msgh, struct timeval* timeval); |
2086 | | void fillMSGHdr(struct msghdr* msgh, struct iovec* iov, cmsgbuf_aligned* cbuf, size_t cbufsize, char* data, size_t datalen, ComboAddress* addr); |
2087 | | int sendOnNBSocket(int fileDesc, const struct msghdr* msgh); |
2088 | | [[nodiscard]] pdns::expected<size_t, int> sendMsgWithOptions(int socketDesc, const void* buffer, size_t len, const ComboAddress* dest, const ComboAddress* local, unsigned int localItf, int flags); |
2089 | | |
2090 | | /* requires a non-blocking, connected TCP socket */ |
2091 | | bool isTCPSocketUsable(int sock); |
2092 | | |
2093 | | extern template class NetmaskTree<bool>; |
2094 | | ComboAddress parseIPAndPort(const std::string& input, uint16_t port); |
2095 | | |
2096 | | std::set<std::string> getListOfNetworkInterfaces(); |
2097 | | std::vector<ComboAddress> getListOfAddressesOfNetworkInterface(const std::string& itf); |
2098 | | std::vector<Netmask> getListOfRangesOfNetworkInterface(const std::string& itf); |
2099 | | |
2100 | | /* These functions throw if the value was already set to a higher value, |
2101 | | or on error */ |
2102 | | void setSocketBuffer(int fileDesc, int optname, uint32_t size); |
2103 | | void setSocketReceiveBuffer(int fileDesc, uint32_t size); |
2104 | | void setSocketSendBuffer(int fileDesc, uint32_t size); |
2105 | | uint32_t raiseSocketReceiveBufferToMax(int socket); |
2106 | | uint32_t raiseSocketSendBufferToMax(int socket); |