diff --git a/HISTORY.rst b/HISTORY.rst index 8430252..2816968 100644 --- a/HISTORY.rst +++ b/HISTORY.rst @@ -6,6 +6,18 @@ History 3.3.0 ++++++++++++++++++ +* Fixed iteration over an IPv6 database with a network shorter than /96 + whose first bits are zero, such as ``::/1``. The readers raised + ``ValueError`` or skipped networks. The pure Python reader also raised + ``ValueError`` for a network that starts at ``::1:0:0``, such as + ``::1:0:0/96``. +* The pure Python reader now raises ``InvalidDatabaseError`` for a search tree + record that points before the data section. Previously, it returned an + empty map. +* Iterating over a database with a corrupt search tree, such as one with a + cycle, now raises ``InvalidDatabaseError``. Previously, the C extension + could corrupt memory, and the pure Python reader raised ``RecursionError`` + or returned part of the networks. * A second ``__init__`` on a pure Python ``Reader`` now closes the old database, and an iterator from before it raises ``ValueError``, as in the C extension. Before, the iterator walked the new database with node @@ -32,6 +44,11 @@ History during iteration, from another thread or from a signal handler. * Fixed a crash on free-threaded Python when two threads advanced the same iterator. + * An exhausted iterator now raises ``StopIteration`` after its ``Reader`` + closes, not ``ValueError``. + * The iterator now stops after any error, as the pure Python iterator + does. This includes an error in the data of one record. Previously, the + next call returned the remaining networks. * Added the ``node_byte_size`` and ``search_tree_size`` properties to ``Metadata``, as the pure Python ``Metadata`` has. diff --git a/extension/maxminddb.c b/extension/maxminddb.c index 570878d..ffa69d3 100644 --- a/extension/maxminddb.c +++ b/extension/maxminddb.c @@ -88,6 +88,8 @@ typedef struct { Reader_obj *reader; struct record *next; uint64_t generation; + // Set once next() has raised StopIteration or an error. + bool done; } ReaderIter_obj; typedef struct { @@ -139,6 +141,7 @@ static PyObject *metadata_value(PyObject *map, const char *key); static void set_error_from_cause(PyObject *type, const char *message); static PyObject *Metadata_node_byte_size(PyObject *self, void *closure); static PyObject *reader_iter_next(PyObject *self); +static void free_records(struct record *next); static bool format_sockaddr(struct sockaddr *addr, char *dst); static PyObject *from_entry_data_list(maxminddb_state *state, MMDB_entry_data_list_s **entry_data_list); @@ -915,6 +918,13 @@ static PyObject *ReaderIter_next(PyObject *self) { Py_BEGIN_CRITICAL_SECTION(self); #endif result = reader_iter_next(self); + // Stop after StopIteration or an error, as a generator does. + if (result == NULL) { + ReaderIter_obj *ri = (ReaderIter_obj *)self; + free_records(ri->next); + ri->next = NULL; + ri->done = true; + } #ifdef Py_GIL_DISABLED Py_END_CRITICAL_SECTION(); #endif @@ -929,6 +939,15 @@ static PyObject *reader_iter_next(PyObject *self) { ReaderIter_obj *ri = (ReaderIter_obj *)self; + // An iterator that is exhausted or that raised an error stays done, even + // after the reader closes. The list of pending records cannot show this: + // it is already empty when the last record is returned, and the next call + // must still report a closed or reopened reader, as the pure Python + // iterator does. + if (ri->done) { + return NULL; + } + if (reader_acquire_read_lock(ri->reader) != 0) { return NULL; } @@ -955,8 +974,11 @@ static PyObject *reader_iter_next(PyObject *self) { switch (cur->type) { case MMDB_RECORD_TYPE_INVALID: reader_release_read_lock(ri->reader); + // libmaxminddb before 1.14 returns this type for a record + // that points to the root or past the data section. Later + // versions fail in MMDB_read_node instead. PyErr_SetString(state->MaxMindDB_error, - "Invalid record when reading node"); + MMDB_strerror(MMDB_CORRUPT_SEARCH_TREE_ERROR)); free(cur); return NULL; case MMDB_RECORD_TYPE_SEARCH_NODE: { @@ -966,6 +988,18 @@ static PyObject *reader_iter_next(PyObject *self) { // These are aliased networks. Skip them. break; } + // Only a corrupt tree, such as one with a cycle, has a node at + // the full address depth. Without this check, an IPv4 tree + // gives a network longer than /32, and a cycle in either tree + // writes past the end of ip_packed at depth 128. + if (cur->depth >= ri->reader->mmdb->depth) { + reader_release_read_lock(ri->reader); + PyErr_SetString( + state->MaxMindDB_error, + MMDB_strerror(MMDB_CORRUPT_SEARCH_TREE_ERROR)); + free(cur); + return NULL; + } MMDB_search_node_s node; int status = MMDB_read_node( ri->reader->mmdb, (uint32_t)cur->record, &node); @@ -1050,7 +1084,9 @@ static PyObject *reader_iter_next(PyObject *self) { int ip_start = 0; Py_ssize_t ip_length = 4; if (depth == 128) { - if (is_ipv6(cur->ip_packed)) { + // A network shorter than /96 is IPv6, even if its first + // 96 bits are zero. + if (is_ipv6(cur->ip_packed) || cur->depth < 96) { // IPv6 address ip_length = 16; } else { @@ -1101,17 +1137,20 @@ static PyObject *reader_iter_next(PyObject *self) { return NULL; } -static void ReaderIter_dealloc(PyObject *self) { - ReaderIter_obj *ri = (ReaderIter_obj *)self; - - Py_DECREF(ri->reader); - - struct record *next = ri->next; +static void free_records(struct record *next) { while (next != NULL) { struct record *cur = next; next = cur->next; free(cur); } +} + +static void ReaderIter_dealloc(PyObject *self) { + ReaderIter_obj *ri = (ReaderIter_obj *)self; + + Py_DECREF(ri->reader); + + free_records(ri->next); PyTypeObject *type = Py_TYPE(self); PyObject_Del(self); Py_DECREF(type); diff --git a/maxminddb/reader.py b/maxminddb/reader.py index a310325..5ad7443 100644 --- a/maxminddb/reader.py +++ b/maxminddb/reader.py @@ -10,7 +10,7 @@ import contextlib import ipaddress from dataclasses import dataclass -from ipaddress import IPv4Address, IPv6Address +from ipaddress import IPv4Address, IPv4Network, IPv6Address, IPv6Network from typing import IO, TYPE_CHECKING, Any from maxminddb.const import MODE_AUTO, MODE_FD, MODE_FILE, MODE_MEMORY, MODE_MMAP @@ -29,6 +29,7 @@ _IPV4_MAX_NUM = 2**32 _REOPENED = "Attempt to iterate over a reopened MaxMind DB. Create a new iterator." _CLOSED = "Attempt to iterate over a closed MaxMind DB." +_CORRUPT_TREE = "The MaxMind DB file's search tree is corrupt" class Reader: @@ -48,6 +49,8 @@ class Reader: _metadata: Metadata _record_size: int _ipv4_start: int + _search_tree_size: int + _data_start: int # Incremented on each open, so an iterator can detect a reopen. _generation: int = 0 @@ -133,22 +136,22 @@ def _load( self._metadata = Metadata(**_metadata_fields(metadata, filename)) self._record_size = self._metadata.record_size + # _resolve_data_pointer uses these on every lookup. + self._search_tree_size = self._metadata.search_tree_size + self._data_start = ( + self._search_tree_size + self._DATA_SECTION_SEPARATOR_SIZE + ) + # Traversal reads nodes below node_count. Once the tree fits, those # reads need no length checks of their own. - tree_end = ( - self._metadata.search_tree_size + self._DATA_SECTION_SEPARATOR_SIZE - ) - if tree_end > self._buffer_size: + if self._data_start > self._buffer_size: msg = ( f"Error opening database file ({filename}). The search tree " "extends past the end of the file." ) raise InvalidDatabaseError(msg) # noqa: TRY301 - self._decoder = Decoder( - self._buffer, - self._metadata.search_tree_size + self._DATA_SECTION_SEPARATOR_SIZE, - ) + self._decoder = Decoder(self._buffer, self._data_start) self.closed = False ipv4_start = 0 @@ -236,23 +239,36 @@ def _iterate(self, generation: int) -> Iterator: yield record def _generate_children(self, node: int, depth: int, ip_acc: int) -> Iterator: - if ip_acc != 0 and node == self._ipv4_start: - # Skip nodes aliased to IPv4 - return - node_count = self._metadata.node_count + bits = 128 if self._metadata.ip_version == 6 else 32 if node > node_count: - bits = 128 if self._metadata.ip_version == 6 else 32 ip_acc <<= bits - depth - if ip_acc <= _IPV4_MAX_NUM and bits == 128: - depth -= 96 - yield ( - ipaddress.ip_network((ip_acc, depth)), - self._resolve_data_pointer( - node, - ), - ) + network: IPv4Network | IPv6Network + if bits == 32: + network = IPv4Network((ip_acc, depth)) + elif depth >= 96 and ip_acc < _IPV4_MAX_NUM: + # An IPv4 network in an IPv6 tree is at least /96, and its + # first 96 bits are zero. + network = IPv4Network((ip_acc, depth - 96)) + else: + network = IPv6Network((ip_acc, depth)) + yield (network, self._resolve_data_pointer(node)) elif node < node_count: + # Skip the IPv4 subtree when an address with a set bit in its first + # 96 bits leads to it, as the C extension does. Inside the IPv4 + # subtree, or in an IPv4 tree, a record that points back to it is a + # cycle. + if ( + node == self._ipv4_start + and bits == 128 + and ip_acc >> max(depth - 96, 0) != 0 + ): + return + # A node at the full address depth has no valid children, and no + # record can point to the root. Only a corrupt tree, such as one + # with a cycle, has either. + if depth >= bits or (node == 0 and depth > 0): + raise InvalidDatabaseError(_CORRUPT_TREE) left = self._read_node(node, 0) ip_acc <<= 1 depth += 1 @@ -307,11 +323,12 @@ def _read_node(self, node_number: int, index: int) -> int: raise InvalidDatabaseError(msg) def _resolve_data_pointer(self, pointer: int) -> Record: - resolved = pointer - self._metadata.node_count + self._metadata.search_tree_size + resolved = pointer - self._metadata.node_count + self._search_tree_size - if resolved >= self._buffer_size: - msg = "The MaxMind DB file's search tree is corrupt" - raise InvalidDatabaseError(msg) + # A pointer into the separator between the tree and the data section + # is as corrupt as one past the end, as libmaxminddb checks. + if resolved < self._data_start or resolved >= self._buffer_size: + raise InvalidDatabaseError(_CORRUPT_TREE) (data, _) = self._decoder.decode(resolved) return data diff --git a/tests/reader_test.py b/tests/reader_test.py index 80014be..45db487 100644 --- a/tests/reader_test.py +++ b/tests/reader_test.py @@ -146,6 +146,33 @@ def _database_with_metadata( return data[:start] + _encode_control(7, len(entries)) + items +def _database(records: tuple[int, ...], *, ip_version: int) -> bytes: + """Return a database with a 24-bit search tree and one data record. + + The tree has two records per node, and a record of node_count + 16 points + at the data record, the string "net". + """ + metadata = { + "binary_format_major_version": 2, + "binary_format_minor_version": 0, + "build_epoch": 1, + "database_type": "Test", + "description": {"en": "Test"}, + "ip_version": ip_version, + "languages": ["en"], + "node_count": len(records) // 2, + "record_size": 24, + } + tree = b"".join(record.to_bytes(3, "big") for record in records) + return ( + tree + + bytes(16) + + _encode_value("net") + + _METADATA_START_MARKER + + _encode_value(metadata) + ) + + def _encode_value(value: object, key: str = "") -> bytes: if isinstance(value, str): encoded = value.encode() @@ -783,6 +810,154 @@ def test_metadata_that_does_not_decode_is_rejected(self) -> None: UnicodeDecodeError, ) + def test_exhausted_iterator_stops_after_close(self) -> None: + with open_database( + f"{_TEST_DATA_DIR}/MaxMind-DB-test-ipv4-24.mmdb", + self.mode, + ) as reader: + iterator = iter(reader) + list(iterator) + self.assertEqual(next(iterator, "done"), "done") + + def test_search_node_at_full_depth_is_rejected(self) -> None: + # Nodes 0 to 32 form a left spine, so node 32 is a search node at + # depth 32, which an IPv4 tree cannot have. + data = 33 + 16 + records = [record for i in range(32) for record in (i + 1, data)] + records += [data, data] + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "too-deep.mmdb" + path.write_bytes(_database(tuple(records), ip_version=4)) + with ( + open_database(str(path), self.mode) as reader, + self.assertRaisesRegex(InvalidDatabaseError, "search tree is corrupt"), + ): + list(reader) + + def test_record_that_points_to_the_root_is_rejected(self) -> None: + # Node 1's right record points back to the root. The left records point + # at data, so a walk through the root again would yield networks that + # the tree does not have, such as 192.0.0.0/3. + records = (18, 1, 18, 0) + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "root-record.mmdb" + path.write_bytes(_database(records, ip_version=4)) + with open_database(str(path), self.mode) as reader: + seen: list[ipaddress.IPv4Network | ipaddress.IPv6Network] = [] + with self.assertRaisesRegex( + InvalidDatabaseError, + "search tree is corrupt", + ): + for network, _ in reader: + seen.append(network) + valid = { + ipaddress.ip_network("0.0.0.0/1"), + ipaddress.ip_network("128.0.0.0/2"), + } + self.assertLessEqual(set(seen), valid) + # Both readers yield node 0's left record before node 1. + self.assertIn(ipaddress.ip_network("0.0.0.0/1"), seen) + + def test_iterate_ipv6_networks_shorter_than_96_bits(self) -> None: + # One node whose two records point at the same data record, so the + # tree holds ::/1 and 8000::/1. + records = (17, 17) + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "short-ipv6.mmdb" + path.write_bytes(_database(records, ip_version=6)) + with open_database(str(path), self.mode) as reader: + self.assertEqual( + list(reader), + [ + (ipaddress.ip_network("::/1"), "net"), + (ipaddress.ip_network("8000::/1"), "net"), + ], + ) + + def test_iterate_ipv6_network_just_above_ipv4(self) -> None: + # A left spine of 96 nodes whose last right record is data holds only + # ::1:0:0/96. The first 96 bits are not all zero, so it is IPv6. + empty = 96 + data = empty + 16 + records = [record for i in range(95) for record in (i + 1, empty)] + records += [empty, data] + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "ipv6-above-ipv4.mmdb" + path.write_bytes(_database(tuple(records), ip_version=6)) + with open_database(str(path), self.mode) as reader: + self.assertEqual( + [network for network, _ in reader], + [ipaddress.ip_network("::1:0:0/96")], + ) + + def test_record_that_points_into_the_separator_is_rejected(self) -> None: + # The left record, node_count + 1, points into the 16-byte separator + # between the search tree and the data section. libmaxminddb before + # 1.14 reports it as bad data. + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "separator.mmdb" + path.write_bytes(_database((2, 17), ip_version=4)) + with ( + open_database(str(path), self.mode) as reader, + self.assertRaisesRegex( + InvalidDatabaseError, + "search tree is corrupt|contains bad data", + ), + ): + reader.get(self.ipf("1.1.1.1")) + + def test_cyclic_search_tree_is_rejected(self) -> None: + data = bytearray( + pathlib.Path(f"{_TEST_DATA_DIR}/MaxMind-DB-test-ipv4-24.mmdb").read_bytes(), + ) + # Point the left record of node 1 back at node 1. + data[6:9] = b"\x00\x00\x01" + with tempfile.TemporaryDirectory() as directory: + path = pathlib.Path(directory) / "cyclic.mmdb" + path.write_bytes(data) + with open_database(str(path), self.mode) as reader: + iterator = iter(reader) + with self.assertRaisesRegex( + InvalidDatabaseError, + "search tree is corrupt", + ): + list(iterator) + # The iterator stops after an error, as a generator does. + self.assertEqual(next(iterator, "done"), "done") + + # A record that points back to the root of an IPv4 tree. + broken = f"{_TEST_DATA_DIR}/MaxMind-DB-test-broken-search-tree-24.mmdb" + with ( + open_database(broken, self.mode) as reader, + self.assertRaisesRegex(InvalidDatabaseError, "search tree is corrupt"), + ): + list(reader) + + # A record in the IPv4 subtree of an IPv6 tree that points back to the + # IPv4 start node. Each record of that node, ::/97 and ::8000:0/97, + # points to it in turn. The right record makes the address nonzero. + source = f"{_TEST_DATA_DIR}/MaxMind-DB-test-mixed-24.mmdb" + with maxminddb.reader.Reader(source) as reader: + ipv4_start = reader._ipv4_start # noqa: SLF001 + original = pathlib.Path(source).read_bytes() + for index in (0, 1): + offset = ipv4_start * 6 + index * 3 + mixed = bytearray(original) + mixed[offset : offset + 3] = ipv4_start.to_bytes(3, "big") + with ( + self.subTest(index=index), + tempfile.TemporaryDirectory() as directory, + ): + path = pathlib.Path(directory) / "ipv4-cycle.mmdb" + path.write_bytes(mixed) + with ( + open_database(str(path), self.mode) as reader, + self.assertRaisesRegex( + InvalidDatabaseError, "search tree is corrupt" + ), + ): + list(reader) + def test_ip_validation(self) -> None: reader = open_database( "tests/data/test-data/MaxMind-DB-test-decoder.mmdb", @@ -1481,6 +1656,9 @@ def ip_network(*args, **kwargs): pass else: sys.exit("next() after close() did not raise ValueError") + # The error stops the iterator. + if next(iterator, None) is not None: + sys.exit("the iterator continued after the error") print("ok") """, )