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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
94 changes: 94 additions & 0 deletions python/tests/test_dataset_index.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
):
Expand Down
22 changes: 22 additions & 0 deletions python/tsfile/dataset/index.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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]:
Expand All @@ -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):
Expand All @@ -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 [
Expand All @@ -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)


Expand Down
Loading