diff --git a/python/tests/test_dataset_index.py b/python/tests/test_dataset_index.py index 4ac9f2ede..3d768540e 100644 --- a/python/tests/test_dataset_index.py +++ b/python/tests/test_dataset_index.py @@ -145,6 +145,100 @@ def test_build_publish_map_and_lookup(tmp_path): assert index.string(file_record[0]) == str(source) +@pytest.mark.parametrize("close_mode", ["explicit", "context", "runtime"]) +@pytest.mark.parametrize( + ("method", "args"), + [ + ("__enter__", ()), + ("count", (LOGICAL_SERIES,)), + ("record", (LOGICAL_SERIES, 0)), + ("records", (LOGICAL_SERIES,)), + ("records", (LOGICAL_SERIES, 0, 0)), + ("string_bytes", (0,)), + ("string", (0,)), + ("find_table_ids", ("root",)), + ("find_table_ids", ("missing",)), + ("find_device_id", (0, "root.")), + ("find_column_id", (0, "s1")), + ("find_series_id", (0, 0)), + ("describe_series", (0,)), + ("series_identity", (0,)), + ("find_series_span", (0, 0)), + ("locator_metadata", (0,)), + ("prepared_locator_metadata", (0, 0)), + ("device_route", (0,)), + ("table_name_id", (0,)), + ("column_name_id", (0,)), + ("device_time_bounds", (0,)), + ], +) +def test_index_reads_after_close_raise_runtime_error( + tmp_path, close_mode, method, args +): + source = tmp_path / "source.tsfile" + source.write_bytes(b"T" * 4096) + output = tmp_path / "dataset.tsidx" + write_index_atomic( + str(output), build_sections_from_dataframe(_synthetic_dataframe(str(source))) + ) + + if close_mode == "runtime": + runtime = DatasetRuntime(str(output)) + lease = runtime.lease() + index = runtime.index + lease.close() + assert runtime._torn_down + elif close_mode == "context": + with MappedDatasetIndex(str(output)) as index: + pass + else: + index = MappedDatasetIndex(str(output)) + index.close() + + index.close() + assert index._view is None + assert index._mmap is None + assert index._file is None + assert index._lookup is None + with pytest.raises(RuntimeError, match="Dataset Index is closed"): + result = getattr(index, method)(*args) + if method == "records": + list(result) + + +@pytest.mark.parametrize("started", [False, True]) +def test_records_iterator_rejects_reads_after_close(tmp_path, started): + source = tmp_path / "source.tsfile" + source.write_bytes(b"T" * 4096) + output = tmp_path / "dataset.tsidx" + write_index_atomic( + str(output), build_sections_from_dataframe(_synthetic_dataframe(str(source))) + ) + + with MappedDatasetIndex(str(output)) as index: + records = index.records(index_module.STRING_OFFSETS) + if started: + next(records) + + with pytest.raises(RuntimeError, match="Dataset Index is closed"): + next(records) + + +def test_closed_subset_rejects_named_read_after_runtime_teardown(tmp_path): + source = tmp_path / "part.tsfile" + _write_runtime_file(source, 0) + + with TsFileDataFrame(str(tmp_path), show_progress=False, use_index=True) as first: + with first[:1] as subset: + name = str(subset.list_timeseries()[0]) + runtime = first._runtime + first.close() + assert runtime._torn_down + assert runtime.index._view is None + with pytest.raises(RuntimeError, match="closed"): + subset[name][:] + + def test_index_lookup_is_required_and_does_not_unpack_python_record_tuples( tmp_path, monkeypatch ): diff --git a/python/tsfile/dataset/index.py b/python/tsfile/dataset/index.py index 9b4ba539c..ceee2f5f3 100644 --- a/python/tsfile/dataset/index.py +++ b/python/tsfile/dataset/index.py @@ -351,16 +351,23 @@ def close(self): self._file.close() self._file = None + def _assert_open(self): + if self._view is None: + raise RuntimeError("Dataset Index is closed") + def __enter__(self): + self._assert_open() return self def __exit__(self, *_): self.close() def count(self, section_type: int) -> int: + self._assert_open() return self._entries[section_type][4] def record(self, section_type: int, record_id: int) -> tuple: + self._assert_open() entry = self._entries[section_type] if record_id < 0 or record_id >= entry[4]: raise IndexError(record_id) @@ -377,6 +384,7 @@ def records(self, section_type: int, first: int = 0, count: Optional[int] = None yield self.record(section_type, record_id) def string_bytes(self, sid: int) -> bytes: + self._assert_open() offsets = self._entries[STRING_OFFSETS] strings = self._entries[STRING_BYTES] if sid < 0 or sid + 1 >= offsets[4]: @@ -388,6 +396,7 @@ def string_bytes(self, sid: int) -> bytes: return bytes(self._view[strings[2] + start : strings[2] + end]) def string(self, sid: int) -> str: + self._assert_open() return self._lookup.string(sid) def _equal_hash_range(self, section_type: int, hash_index: int, value_hash: int): @@ -408,6 +417,7 @@ def _equal_hash_range(self, section_type: int, hash_index: int, value_hash: int) return first, low def find_table_ids(self, name: str) -> List[int]: + self._assert_open() encoded = name.encode("utf-8") first, end = self._equal_hash_range(TABLE_NAME_INDEX, 0, name_hash(encoded)) return [ @@ -417,48 +427,60 @@ def find_table_ids(self, name: str) -> List[int]: ] def find_device_id(self, table_id: int, name: str) -> int: + self._assert_open() return self._lookup.find_device_id(table_id, name) def find_column_id(self, table_id: int, name: str) -> int: + self._assert_open() return self._lookup.find_column_id(table_id, name) def find_series_id(self, device_id: int, column_id: int) -> int: + self._assert_open() return self._lookup.find_series_id(device_id, column_id) def describe_series(self, series_id: int): """Return scalar route metadata for one logical series.""" + self._assert_open() return self._lookup.describe_series(series_id) def series_identity(self, series_id: int): """Return device and column ids for one logical series.""" + self._assert_open() return self._lookup.series_identity(series_id) def find_series_span(self, series_id: int, file_id: int): """Return one series span without exposing its full record tuple.""" + self._assert_open() return self._lookup.find_series_span(series_id, file_id) def locator_metadata(self, locator_id: int): """Return locator/device-span fields needed by the runtime reader.""" + self._assert_open() return self._lookup.locator_metadata(locator_id) def prepared_locator_metadata(self, file_id: int, locator_id: int): """Return generation and locator fields used by native prepare.""" + self._assert_open() return self._lookup.prepared_locator_metadata(file_id, locator_id) def device_route(self, device_id: int): """Return table id and logical-path string id for one device.""" + self._assert_open() return self._lookup.device_route(device_id) def table_name_id(self, table_id: int): """Return the string-pool id for one table name.""" + self._assert_open() return self._lookup.table_name_id(table_id) def column_name_id(self, column_id: int): """Return the string-pool id for one column name.""" + self._assert_open() return self._lookup.column_name_id(column_id) def device_time_bounds(self, device_id: int): """Return min/max timestamps for one device.""" + self._assert_open() return self._lookup.device_time_bounds(device_id)