From e5c29849cb1b52ebfc6f293011ba3be56e3fda23 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Tue, 29 Sep 2026 11:24:17 +0800 Subject: [PATCH 01/30] feat: save selective TransferQueue checkpoints by key Squashes the initial by-key checkpoint work with its follow-up refactors (helper extraction, restructured flow, fsync on publish). Signed-off-by: OutstanderWang --- .../e2e/test_data_dump_cross_topology_e2e.py | 136 ++++++ tests/e2e/test_data_dump_e2e.py | 427 ++++++++++++++++++ transfer_queue/__init__.py | 7 + transfer_queue/client.py | 102 +++++ transfer_queue/controller.py | 52 +++ transfer_queue/data_dump.py | 295 ++++++++++++ .../managers/simple_storage_manager.py | 123 +++++ transfer_queue/storage/simple_storage.py | 63 +++ transfer_queue/utils/zmq_utils.py | 6 + 9 files changed, 1211 insertions(+) create mode 100644 tests/e2e/test_data_dump_cross_topology_e2e.py create mode 100644 tests/e2e/test_data_dump_e2e.py create mode 100644 transfer_queue/data_dump.py diff --git a/tests/e2e/test_data_dump_cross_topology_e2e.py b/tests/e2e/test_data_dump_cross_topology_e2e.py new file mode 100644 index 00000000..cdf46960 --- /dev/null +++ b/tests/e2e/test_data_dump_cross_topology_e2e.py @@ -0,0 +1,136 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""A dump taken with N storage units must restore into a system with M storage units. + +This is the property that separates a data dump from a checkpoint. ``load_checkpoint`` +sends each storage unit file back to the unit at the same position, so it requires the +same unit count; ``load_data_by_key`` writes rows back by key and lets TransferQueue +route them for the current topology. + +Each test restarts TransferQueue with a different unit count, so this lives apart from +``test_data_dump_e2e.py``, whose fixtures hold one system for the whole module. + +Run with: + pytest tests/e2e/test_data_dump_cross_topology_e2e.py -v +""" + +import os +import shutil +import uuid +from pathlib import Path + +import pytest +import ray +import torch +from omegaconf import OmegaConf +from tensordict import TensorDict + +import transfer_queue as tq + +os.environ["RAY_DEDUP_LOGS"] = "0" + +_DEFAULT_DUMP_ROOT = "/apdcephfs_hldy/share_303541817/tq_dump_tests" + + +def _tq_config(num_storage_units: int) -> OmegaConf: + return OmegaConf.create( + { + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "total_storage_size": 400, + "num_data_storage_units": num_storage_units, + }, + }, + } + ) + + +@pytest.fixture(scope="module") +def ray_init(): + if not ray.is_initialized(): + ray.init(namespace="TestDataDumpCrossTopology") + yield + if ray.is_initialized(): + ray.shutdown() + + +@pytest.fixture +def dump_dir(request): + root = Path(os.environ.get("TQ_DUMP_TEST_ROOT", _DEFAULT_DUMP_ROOT)) / uuid.uuid4().hex + root.mkdir(parents=True) + yield root / "dump" + shutil.rmtree(root, ignore_errors=True) + + +def _row_input_ids(row: int) -> torch.Tensor: + return torch.tensor([row * 10, row * 10 + 1, row * 10 + 2]) + + +def _put_rows(partition_id: str, keys: list[str]) -> None: + tq.kv_batch_put( + keys=keys, + partition_id=partition_id, + fields=TensorDict( + {"input_ids": torch.stack([_row_input_ids(row) for row in range(len(keys))])}, + batch_size=len(keys), + ), + tags=[{"idx": row} for row in range(len(keys))], + ) + + +def _assert_rows_equal(actual: torch.Tensor, expected_rows: list[torch.Tensor]) -> None: + actual_rows = list(actual.unbind()) if actual.is_nested else list(actual) + assert len(actual_rows) == len(expected_rows) + for actual_row, expected_row in zip(actual_rows, expected_rows, strict=True): + assert torch.equal(actual_row, expected_row) + + +@pytest.mark.parametrize( + ("dump_units", "load_units"), + [(4, 2), (2, 4), (3, 3)], +) +def test_dump_restores_across_storage_unit_counts(ray_init, dump_dir, dump_units, load_units): + # Define test data + partition_id = "cross" + keys = [f"c{i}" for i in range(8)] + + # Dump with one topology + tq.init(_tq_config(dump_units)) + try: + _put_rows(partition_id, keys) + report = tq.dump_data_by_key(dump_dir, keys, partition_id) + assert report["rows_with_data"] == len(keys) + assert report["shards"] <= dump_units + finally: + tq.close() + + # Restore into a different topology + tq.init(_tq_config(load_units)) + try: + tq.load_data_by_key(dump_dir) + + # Check restored state: every row readable, payload and tag intact + retrieved = tq.kv_batch_get(keys=keys, partition_id=partition_id, select_fields=["input_ids"]) + _assert_rows_equal(retrieved["input_ids"], [_row_input_ids(row) for row in range(len(keys))]) + + controller = ray.get_actor("TransferQueueController", namespace="transfer_queue") + snapshot = ray.get(controller.get_partition_snapshot.remote(partition_id)) + for row, key in enumerate(keys): + assert snapshot.custom_meta[snapshot.keys_mapping[key]]["idx"] == row + finally: + tq.close() diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py new file mode 100644 index 00000000..ee2a5bbb --- /dev/null +++ b/tests/e2e/test_data_dump_e2e.py @@ -0,0 +1,427 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""End-to-end tests for dump_data_by_key / load_data_by_key. + +Each storage unit writes its own shard from the node it runs on, so +``TQ_DUMP_TEST_ROOT`` must point at a filesystem shared by the whole cluster. +A node-local path such as pytest's ``tmp_path`` fails on a multi-node cluster +with ``FileNotFoundError``. + +Run with: + pytest tests/e2e/test_data_dump_e2e.py -v +""" + +import json +import os +import shutil +import uuid +from pathlib import Path + +import pytest +import ray +import torch +from omegaconf import OmegaConf +from tensordict import NonTensorStack, TensorDict + +import transfer_queue as tq + +os.environ["RAY_DEDUP_LOGS"] = "0" + +_NUM_STORAGE_UNITS = 4 +_DEFAULT_DUMP_ROOT = "/apdcephfs_hldy/share_303541817/tq_dump_tests" + + +def _tq_config(num_storage_units: int) -> OmegaConf: + return OmegaConf.create( + { + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "total_storage_size": 400, + "num_data_storage_units": num_storage_units, + }, + }, + } + ) + + +# --------------------------------------------------------------------------- +# fixtures +# --------------------------------------------------------------------------- + + +@pytest.fixture(scope="module") +def ray_init(): + if not ray.is_initialized(): + ray.init(namespace="TestDataDumpE2E") + yield + if ray.is_initialized(): + ray.shutdown() + + +@pytest.fixture(scope="module") +def tq_system(ray_init): + tq.init(_tq_config(_NUM_STORAGE_UNITS)) + yield + tq.close() + + +@pytest.fixture +def controller(tq_system): + return ray.get_actor("TransferQueueController", namespace="transfer_queue") + + +@pytest.fixture(autouse=True) +def cleanup_partitions(controller): + yield + try: + for pid in ray.get(controller.list_partitions.remote()): + ray.get(controller.clear_partition.remote(pid)) + except Exception: + pass + + +@pytest.fixture(scope="module") +def shared_root(): + root = Path(os.environ.get("TQ_DUMP_TEST_ROOT", _DEFAULT_DUMP_ROOT)) / uuid.uuid4().hex + root.mkdir(parents=True) + yield root + shutil.rmtree(root, ignore_errors=True) + + +@pytest.fixture +def dump_dir(shared_root, request): + case = shared_root / request.node.name.replace("/", "_") + yield case / "dump" + shutil.rmtree(case, ignore_errors=True) + + +# --------------------------------------------------------------------------- +# helpers +# --------------------------------------------------------------------------- + + +def _row_input_ids(row: int) -> torch.Tensor: + """Deterministic per-row payload so a restored row can be traced to its key.""" + return torch.tensor([row * 10, row * 10 + 1, row * 10 + 2]) + + +def _put_rows(partition_id: str, keys: list[str]) -> None: + tq.kv_batch_put( + keys=keys, + partition_id=partition_id, + fields=TensorDict( + { + "input_ids": torch.stack([_row_input_ids(row) for row in range(len(keys))]), + "attention_mask": torch.ones(len(keys), 3), + }, + batch_size=len(keys), + ), + tags=[{"idx": row} for row in range(len(keys))], + ) + + +def _assert_rows_equal(actual: torch.Tensor, expected_rows: list[torch.Tensor]) -> None: + """Compare a retrieved field row by row. + + TransferQueue returns a batched field as a nested tensor, which ``torch.equal`` + cannot consume directly. + """ + actual_rows = list(actual.unbind()) if actual.is_nested else list(actual) + assert len(actual_rows) == len(expected_rows) + for actual_row, expected_row in zip(actual_rows, expected_rows, strict=True): + assert torch.equal(actual_row, expected_row) + + +def _keys_mapping(controller, partition_id: str) -> dict[str, int]: + snapshot = ray.get(controller.get_partition_snapshot.remote(partition_id)) + return dict(snapshot.keys_mapping) + + +# --------------------------------------------------------------------------- +# dump / load roundtrip +# --------------------------------------------------------------------------- + + +class TestDumpLoadRoundtrip: + def test_only_selected_keys_are_restored(self, tq_system, dump_dir, controller): + # Define test data + partition_id = "d_basic" + keys = [f"k{i}" for i in range(6)] + selected = ["k1", "k4"] + _put_rows(partition_id, keys) + + # Dump + report = tq.dump_data_by_key(dump_dir, selected, partition_id) + assert report == {"keys": 2, "rows_with_data": 2, "shards": report["shards"], "bytes": report["bytes"]} + + # Wipe, then restore + ray.get(controller.clear_partition.remote(partition_id)) + tq.load_data_by_key(dump_dir) + + # Check restored state + assert sorted(_keys_mapping(controller, partition_id)) == selected + retrieved = tq.kv_batch_get(keys=selected, partition_id=partition_id, select_fields=["input_ids"]) + _assert_rows_equal(retrieved["input_ids"], [_row_input_ids(keys.index(key)) for key in selected]) + + def test_tags_survive_the_roundtrip(self, tq_system, dump_dir, controller): + # Define test data + partition_id = "d_tags" + keys = [f"t{i}" for i in range(5)] + selected = ["t0", "t3"] + _put_rows(partition_id, keys) + + # Dump + wipe + load + tq.dump_data_by_key(dump_dir, selected, partition_id) + ray.get(controller.clear_partition.remote(partition_id)) + tq.load_data_by_key(dump_dir) + + # Check restored state + snapshot = ray.get(controller.get_partition_snapshot.remote(partition_id)) + for key in selected: + assert snapshot.custom_meta[snapshot.keys_mapping[key]]["idx"] == keys.index(key) + + def test_jagged_and_non_tensor_fields_survive(self, tq_system, dump_dir, controller): + """The shard stores per-row values; packing them back must reproduce the container.""" + # Define test data: variable-length rows plus a string field + partition_id = "d_jagged" + keys = ["j0", "j1", "j2"] + for row, key in enumerate(keys): + tq.kv_put( + key=key, + partition_id=partition_id, + fields=TensorDict( + { + "seq": torch.arange(row + 1, dtype=torch.float).unsqueeze(0), + "text": NonTensorStack(f"row-{row}"), + }, + batch_size=1, + ), + tag=None, + ) + + # Dump + wipe + load + tq.dump_data_by_key(dump_dir, keys, partition_id) + ray.get(controller.clear_partition.remote(partition_id)) + tq.load_data_by_key(dump_dir) + + # Check restored state + retrieved = tq.kv_batch_get(keys=keys, partition_id=partition_id, select_fields=["seq", "text"]) + for row, component in enumerate(retrieved["seq"].unbind()): + assert torch.equal(component, torch.arange(row + 1, dtype=torch.float)) + assert list(retrieved["text"]) == [f"row-{row}" for row in range(len(keys))] + + def test_heterogeneous_field_sets_are_grouped(self, tq_system, dump_dir, controller): + """A selective dump routinely mixes rows that finished different fields.""" + # Define test data: h1 has an extra field the others lack + partition_id = "d_hetero" + _put_rows(partition_id, ["h0", "h1", "h2"]) + tq.kv_put( + key="h1", + partition_id=partition_id, + fields=TensorDict({"routed_experts": torch.tensor([[7, 8]])}, batch_size=1), + tag=None, + ) + + # Dump + wipe + load + tq.dump_data_by_key(dump_dir, ["h0", "h1", "h2"], partition_id) + ray.get(controller.clear_partition.remote(partition_id)) + tq.load_data_by_key(dump_dir) + + # Check restored state: the extra field came back only on its own row + rows = tq.read_row_index(dump_dir)["rows"] + assert "routed_experts" in rows["h1"]["fields"] + assert "routed_experts" not in rows["h0"]["fields"] + retrieved = tq.kv_batch_get(keys=["h1"], partition_id=partition_id, select_fields=["routed_experts"]) + _assert_rows_equal(retrieved["routed_experts"], [torch.tensor([7, 8])]) + + def test_other_partitions_are_untouched(self, tq_system, dump_dir, controller): + """Restoring merges by key, unlike a checkpoint load which replaces everything.""" + # Define test data + _put_rows("d_target", ["a0", "a1"]) + _put_rows("d_bystander", ["b0", "b1"]) + + # Dump one partition, then restore it without wiping the other + tq.dump_data_by_key(dump_dir, ["a0"], "d_target") + ray.get(controller.clear_partition.remote("d_target")) + tq.load_data_by_key(dump_dir) + + # Check state: the bystander partition and its rows survived + assert sorted(ray.get(controller.list_partitions.remote())) == ["d_bystander", "d_target"] + retrieved = tq.kv_batch_get(keys=["b0", "b1"], partition_id="d_bystander", select_fields=["input_ids"]) + _assert_rows_equal(retrieved["input_ids"], [_row_input_ids(0), _row_input_ids(1)]) + + def test_row_without_produced_fields_keeps_its_key(self, tq_system, dump_dir, controller): + # Define test data: retrieve_meta with create=True registers a key with no fields + partition_id = "d_keyonly" + _put_rows(partition_id, ["p0"]) + client = tq.get_client() + client.kv_retrieve_meta(keys=["empty0"], partition_id=partition_id, create=True) + + # Dump both rows + report = tq.dump_data_by_key(dump_dir, ["p0", "empty0"], partition_id) + assert report["keys"] == 2 + assert report["rows_with_data"] == 1 + + # Wipe + load + ray.get(controller.clear_partition.remote(partition_id)) + tq.load_data_by_key(dump_dir) + + # Check restored state: the field-less row still exists + assert sorted(_keys_mapping(controller, partition_id)) == ["empty0", "p0"] + + def test_duplicate_keys_are_deduplicated(self, tq_system, dump_dir): + # Define test data + partition_id = "d_dupes" + _put_rows(partition_id, ["d0", "d1", "d2"]) + + # Dump with a repeated key + report = tq.dump_data_by_key(dump_dir, ["d1", "d1", "d2"], partition_id) + + # Check report and dump info + assert report["keys"] == 2 + with open(dump_dir / "dump_info.json", encoding="utf-8") as f: + assert json.load(f)["num_keys"] == 2 + + def test_empty_key_set_writes_a_readable_dump(self, tq_system, dump_dir, controller): + # Dump nothing + report = tq.dump_data_by_key(dump_dir, [], "d_empty") + + # Check saved state: still a complete, loadable dump + assert report["keys"] == 0 + assert report["rows_with_data"] == 0 + assert (dump_dir / "dump_info.json").exists() + + # Check that loading it is a no-op rather than an error + assert tq.load_data_by_key(dump_dir)["keys"] == 0 + + def test_live_partition_survives_a_dump(self, tq_system, dump_dir, controller): + # Define test data + partition_id = "d_nonmutating" + keys = [f"l{i}" for i in range(4)] + _put_rows(partition_id, keys) + before = _keys_mapping(controller, partition_id) + + # Dump a subset + tq.dump_data_by_key(dump_dir, ["l1"], partition_id) + + # Check live state: untouched rows keep their indexes and payloads + assert _keys_mapping(controller, partition_id) == before + retrieved = tq.kv_batch_get(keys=keys, partition_id=partition_id, select_fields=["input_ids"]) + _assert_rows_equal(retrieved["input_ids"], [_row_input_ids(row) for row in range(len(keys))]) + + def test_dump_replaces_preexisting_directory(self, tq_system, dump_dir): + # Define test data + partition_id = "d_replace" + _put_rows(partition_id, ["s0", "s1"]) + dump_dir.mkdir(parents=True) + (dump_dir / "stale.pkl").write_bytes(b"stale") + + # Dump + tq.dump_data_by_key(dump_dir, ["s0"], partition_id) + + # Check saved state + assert not (dump_dir / "stale.pkl").exists() + assert (dump_dir / "dump_info.json").exists() + + +# --------------------------------------------------------------------------- +# row index +# --------------------------------------------------------------------------- + + +class TestRowIndex: + def test_row_index_describes_keys_without_reading_payload(self, tq_system, dump_dir): + # Define test data + partition_id = "d_index" + keys = ["i0", "i1"] + _put_rows(partition_id, keys) + + # Dump + tq.dump_data_by_key(dump_dir, keys, partition_id) + + # Check the index + row_index = tq.read_row_index(dump_dir) + assert row_index["partition_id"] == partition_id + assert sorted(row_index["rows"]) == keys + for row, key in enumerate(keys): + assert row_index["rows"][key]["fields"] == ["attention_mask", "input_ids"] + assert row_index["rows"][key]["tag"] == {"idx": row} + + def test_read_row_index_rejects_a_missing_dump(self, tq_system, dump_dir): + with pytest.raises(FileNotFoundError, match="row_index.pt"): + tq.read_row_index(dump_dir) + + +# --------------------------------------------------------------------------- +# error handling +# --------------------------------------------------------------------------- + + +class TestDumpErrors: + def test_unknown_key_raises_and_leaves_no_directory(self, tq_system, dump_dir): + # Define test data + partition_id = "d_err_key" + _put_rows(partition_id, ["e0"]) + + # Dump a key that was never put + with pytest.raises(RuntimeError, match="keys not found"): + tq.dump_data_by_key(dump_dir, ["e0", "nope"], partition_id) + + # Check saved state: no partial directory left behind + assert not dump_dir.exists() + assert not dump_dir.with_name(dump_dir.name + ".tmp").exists() + + def test_unknown_partition_raises(self, tq_system, dump_dir): + _put_rows("d_err_part", ["e0"]) + with pytest.raises(RuntimeError, match="does not exist"): + tq.dump_data_by_key(dump_dir, ["e0"], "d_never_created") + + def test_load_rejects_a_dump_without_info(self, tq_system, dump_dir): + dump_dir.mkdir(parents=True) + with pytest.raises(FileNotFoundError, match="dump_info.json"): + tq.load_data_by_key(dump_dir) + + def test_load_rejects_a_missing_shard(self, tq_system, dump_dir): + # Define test data + dump + partition_id = "d_err_shard" + _put_rows(partition_id, ["m0", "m1"]) + tq.dump_data_by_key(dump_dir, ["m0", "m1"], partition_id) + + # Tamper: delete one shard file + shard = next((dump_dir / "shards").glob("shard_*.pkl")) + shard.unlink() + + with pytest.raises(FileNotFoundError, match="Missing dump shard"): + tq.load_data_by_key(dump_dir) + + def test_load_rejects_an_unknown_format_version(self, tq_system, dump_dir): + # Define test data + dump + partition_id = "d_err_version" + _put_rows(partition_id, ["v0"]) + tq.dump_data_by_key(dump_dir, ["v0"], partition_id) + + # Tamper: bump the format version beyond what this build reads + info_path = dump_dir / "dump_info.json" + with open(info_path, encoding="utf-8") as f: + info = json.load(f) + info["format_version"] = tq.data_dump.DUMP_FORMAT_VERSION + 1 + with open(info_path, "w", encoding="utf-8") as f: + json.dump(info, f) + + with pytest.raises(ValueError, match="Unsupported dump format version"): + tq.load_data_by_key(dump_dir) diff --git a/transfer_queue/__init__.py b/transfer_queue/__init__.py index 754bb4d8..97c180ce 100644 --- a/transfer_queue/__init__.py +++ b/transfer_queue/__init__.py @@ -16,6 +16,7 @@ import os from .client import TransferQueueClient +from .data_dump import dump_data_by_key, load_data_by_key, read_row_index from .dataloader import StreamingDataLoader, StreamingDataset from .interface import ( async_kv_batch_get, @@ -70,6 +71,12 @@ "save_checkpoint", "load_checkpoint", ] + + [ + # Selective Data Dump Interface + "dump_data_by_key", + "load_data_by_key", + "read_row_index", + ] + [ # High-Level StreamingDataLoader Interface "StreamingDataset", diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 0f7f13f1..13967260 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1124,6 +1124,73 @@ def _can_destroy_zmq_context(self) -> bool: return False return True + # ==================== Selective Data Dump API ==================== + @with_controller_socket + async def async_describe_rows_by_key( + self, + partition_id: str, + keys: list[str], + socket: zmq.asyncio.Socket | None = None, + ) -> dict[str, dict[str, Any]]: + """Asynchronously fetch the row metadata a selective dump needs, via ZMQ RPC. + + Args: + partition_id: Partition that owns ``keys``. + keys: Keys to describe, already deduplicated by the caller. + socket: ZMQ socket injected by @with_controller_socket. + + Returns: + ``{key: {"global_index": int, "fields": list[str], "tag": dict}}``. + + Raises: + RuntimeError: If the RPC fails, or the partition or a key is unknown. + """ + try: + assert socket is not None + request_msg = ZMQMessage.create( + request_type=ZMQRequestType.DESCRIBE_ROWS_BY_KEY, # type: ignore[arg-type] + sender_id=self.client_id, + receiver_id=self._controller.id, + body={"partition_id": partition_id, "keys": keys}, + ) + await socket.send_multipart(request_msg.serialize()) + response_serialized = await socket.recv_multipart(copy=False) + response_msg = ZMQMessage.deserialize(response_serialized) + if response_msg.request_type != ZMQRequestType.DESCRIBE_ROWS_BY_KEY_RESPONSE: + raise RuntimeError( + f"[{self.client_id}]: Unexpected response type {response_msg.request_type} " + f"from controller during row description" + ) + if not response_msg.body["success"]: + raise RuntimeError(response_msg.body["message"]) + return response_msg.body["rows"] + except Exception as e: + raise RuntimeError(f"[{self.client_id}]: Error in describe_rows_by_key: {str(e)}") from e + + async def async_dump_rows_by_index(self, shard_dir: str, global_indexes: list[int]) -> list[dict[str, Any]]: + """Asynchronously dump the given rows into per-storage-unit shards. + + Args: + shard_dir: Directory to write shard files into. + global_indexes: Global indexes to dump. + + Returns: + One entry per written shard. + + Raises: + RuntimeError: If the storage manager is not initialized, or a unit holds + no data for a row it was asked to dump. + NotImplementedError: If the storage backend does not support dumping. + """ + if not hasattr(self, "storage_manager") or self.storage_manager is None: + raise RuntimeError( + f"[{self.client_id}]: Storage manager not initialized. " + "Call initialize_storage_manager() before dump operations." + ) + if not hasattr(self.storage_manager, "dump_rows_by_index"): + raise NotImplementedError(f"{type(self.storage_manager).__name__} does not support selective data dump") + return await self.storage_manager.dump_rows_by_index(shard_dir, global_indexes) + # ==================== Checkpoint API ==================== @with_controller_socket async def async_save_controller_checkpoint( @@ -1308,6 +1375,8 @@ def wrapper(*args, **kwargs): self._kv_retrieve_meta = _make_sync(self.async_kv_retrieve_meta) self._kv_retrieve_keys = _make_sync(self.async_kv_retrieve_keys) self._kv_list = _make_sync(self.async_kv_list) + self._describe_rows_by_key = _make_sync(self.async_describe_rows_by_key) + self._dump_rows_by_index = _make_sync(self.async_dump_rows_by_index) self._save_controller_checkpoint = _make_sync(self.async_save_controller_checkpoint) self._load_controller_checkpoint = _make_sync(self.async_load_controller_checkpoint) self._save_storage_checkpoint = _make_sync(self.async_save_storage_checkpoint) @@ -1744,6 +1813,39 @@ def kv_list( return self._kv_list(partition_id=partition_id) + # ==================== Selective Data Dump API ==================== + def describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, dict[str, Any]]: + """Synchronously fetch the row metadata a selective dump needs, via ZMQ RPC. + + Args: + partition_id: Partition that owns ``keys``. + keys: Keys to describe, already deduplicated by the caller. + + Returns: + ``{key: {"global_index": int, "fields": list[str], "tag": dict}}``. + + Raises: + RuntimeError: If the RPC fails, or the partition or a key is unknown. + """ + return self._describe_rows_by_key(partition_id, keys) + + def dump_rows_by_index(self, shard_dir: str, global_indexes: list[int]) -> list[dict[str, Any]]: + """Synchronously dump the given rows into per-storage-unit shards. + + Args: + shard_dir: Directory to write shard files into. + global_indexes: Global indexes to dump. + + Returns: + One entry per written shard. + + Raises: + RuntimeError: If the storage manager is not initialized, or a unit holds + no data for a row it was asked to dump. + NotImplementedError: If the storage backend does not support dumping. + """ + return self._dump_rows_by_index(shard_dir, global_indexes) + # ==================== Checkpoint API ==================== def save_controller_checkpoint(self, path: str) -> None: """Synchronously save controller state to a file via ZMQ RPC. diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index b45a85a6..612942fd 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -1725,6 +1725,48 @@ def kv_retrieve_keys( return keys + # ==================== Selective Data Dump API ==================== + + def describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, dict[str, Any]]: + """Describe the key-addressed rows a selective dump needs, without copying them. + + Returns only metadata, so the caller can write a row index and route the payload + dump to the storage units that hold it. Deliberately avoids ``to_snapshot``: a + dump needs a handful of lookups per key, not a deep copy of the whole partition. + + Args: + partition_id: Partition that owns ``keys``. + keys: Keys to describe, already deduplicated by the caller. + + Returns: + ``{key: {"global_index": int, "fields": list[str], "tag": dict}}``. ``fields`` + is empty for a row that exists in metadata but has no produced field yet. + + Raises: + KeyError: The partition or any key does not exist. + """ + partition = self._get_partition(partition_id) + if partition is None: + raise KeyError(f"partition {partition_id!r} does not exist; existing partitions: {sorted(self.partitions)}") + + global_indexes = partition.kv_retrieve_indexes(keys) + missing_keys = [key for key, index in zip(keys, global_indexes, strict=True) if index is None] + if missing_keys: + raise KeyError(f"keys not found in partition {partition_id!r}: {missing_keys}") + + return { + key: { + "global_index": global_index, + "fields": sorted( + field_name + for field_name, field_meta in partition.field_metadata.items() + if global_index in field_meta.global_indexes + ), + "tag": partition.custom_meta.get(global_index, {}), + } + for key, global_index in zip(keys, global_indexes, strict=True) + } + def _init_zmq_socket(self): """Initialize ZMQ sockets for communication.""" self.zmq_context = zmq.Context() @@ -1948,6 +1990,7 @@ def _handle_request(self, request_msg: ZMQMessage, monitor: Any) -> ZMQMessage | ZMQRequestType.KV_RETRIEVE_META: self._handle_kv_retrieve_meta_request, ZMQRequestType.KV_RETRIEVE_KEYS: self._handle_kv_retrieve_keys_request, ZMQRequestType.KV_LIST: self._handle_kv_list_request, + ZMQRequestType.DESCRIBE_ROWS_BY_KEY: self._handle_describe_rows_by_key_request, ZMQRequestType.SAVE_CONTROLLER_CHECKPOINT: self._handle_save_controller_checkpoint_request, ZMQRequestType.LOAD_CONTROLLER_CHECKPOINT: self._handle_load_controller_checkpoint_request, } @@ -2165,6 +2208,15 @@ def _handle_kv_list_request(self, request_msg: ZMQMessage) -> ZMQMessage: {"partition_info": partition_info, "message": message}, ) + def _handle_describe_rows_by_key_request(self, request_msg: ZMQMessage) -> ZMQMessage: + params = request_msg.body + rows = self.describe_rows_by_key(params["partition_id"], params["keys"]) + return self._make_response( + request_msg, + ZMQRequestType.DESCRIBE_ROWS_BY_KEY_RESPONSE, + {"success": True, "rows": rows}, + ) + def _handle_save_controller_checkpoint_request(self, request_msg: ZMQMessage) -> ZMQMessage: self.save_checkpoint(request_msg.body["path"]) return self._make_response( diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py new file mode 100644 index 00000000..2336f802 --- /dev/null +++ b/transfer_queue/data_dump.py @@ -0,0 +1,295 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Selective data dump: persist the rows behind a set of keys, and read them back. + +This is deliberately not a checkpoint. It stores payload and the metadata needed to +address it by key, and nothing about the controller: no index manager, no sampler, no +partition snapshot. Restoring therefore goes through the ordinary ``kv_batch_put`` +path, which means a dump taken with N storage units restores into a system with M +storage units. Use ``save_checkpoint`` when you need a full system image instead. + +Layout:: + + / + dump_info.json # partition, counts, shard count + row_index.pt # key -> {global_index, fields, tag} + shards/ + shard_info.json # [{position, storage_unit_id, rows}] + shard__.pkl # {field_data: {field: {gidx: value}}, global_indexes} +""" + +import json +import os +import shutil +from pathlib import Path +from typing import Any + +import torch + +from transfer_queue.utils.logging_utils import get_logger + +logger = get_logger(__name__) + +DUMP_FORMAT_VERSION = 1 + +_DUMP_INFO_FILE = "dump_info.json" +_ROW_INDEX_FILE = "row_index.pt" +_SHARD_SUBDIR = "shards" +_SHARD_INFO_FILE = "shard_info.json" + + +def _fsync_file(file_object: Any) -> None: + """Push a just-written file out of page cache before the caller moves on.""" + file_object.flush() + os.fsync(file_object.fileno()) + + +def _fsync_directory(path: Path) -> None: + """Make a directory's own entries durable. + + fsync on a file says nothing about the directory entry naming it, so a crash can + lose a file that was itself fully synced. The rename that publishes the dump needs + the same treatment. + """ + fd = os.open(path, os.O_RDONLY) + try: + os.fsync(fd) + finally: + os.close(fd) + + +def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) -> dict[str, int]: + """Dump the rows addressed by ``keys`` into ``dump_dir``. + + Each storage unit pickles the rows it owns in its own process, so the payload never + passes through the caller. The caller writes only a small row index. + + The directory is replaced wholesale: the dump is staged in ``.tmp`` and + renamed over ``dump_dir``, so anything already there is destroyed. + + .. note:: + **Multi-node limitation**: dump_dir must reside on a shared network filesystem + (e.g. NFS, GPFS, Lustre) reachable from every storage unit, because each unit + writes its own shard from its own node. + + Callers must freeze writers for ``keys`` for the duration: TransferQueue has no + atomic snapshot, so a concurrent put can land between the row index and the shards. + + Args: + dump_dir: Directory to write the dump into. + keys: Keys to dump. Duplicates are dropped, first occurrence wins. + partition_id: Partition that owns ``keys``. + + Returns: + ``{"keys", "rows_with_data", "shards", "bytes"}``. + + Raises: + RuntimeError: TransferQueue is not initialized, the partition or a key does not + exist, or a storage unit holds no data for a row it was asked to dump. + """ + from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client + + if _TQ_CONTROLLER is None: + raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") + + unique_keys = list(dict.fromkeys(keys)) + dump_dir = Path(dump_dir) + tmp_dir = dump_dir.parent / (dump_dir.name + ".tmp") + if tmp_dir.exists(): + shutil.rmtree(tmp_dir) + tmp_dir.mkdir(parents=True) + + try: + client = _maybe_create_tq_client() + rows = client.describe_rows_by_key(partition_id, unique_keys) if unique_keys else {} + + # A row whose fields are all still unproduced has nothing for a storage unit to + # dump, but it keeps its key and tag so the restore can recreate the row. + indexes_with_data = sorted(row["global_index"] for row in rows.values() if row["fields"]) + + shard_records = ( + client.dump_rows_by_index(str(tmp_dir / _SHARD_SUBDIR), indexes_with_data) if indexes_with_data else [] + ) + shard_dir = tmp_dir / _SHARD_SUBDIR + shard_dir.mkdir(parents=True, exist_ok=True) + with open(shard_dir / _SHARD_INFO_FILE, "w", encoding="utf-8") as f: + json.dump(shard_records, f) + _fsync_file(f) + # The shards themselves were synced by the units that wrote them, but their + # directory entries were created here, on this node. + _fsync_directory(shard_dir) + + # torch.save rather than json: a tag is an arbitrary picklable dict, and this + # path must not fail on a tag that happens to hold a tensor. + with open(tmp_dir / _ROW_INDEX_FILE, "wb") as f: + torch.save({"partition_id": partition_id, "rows": rows}, f) + _fsync_file(f) + + with open(tmp_dir / _DUMP_INFO_FILE, "w", encoding="utf-8") as f: + json.dump( + { + "format_version": DUMP_FORMAT_VERSION, + "partition_id": partition_id, + "num_keys": len(unique_keys), + "num_rows_with_data": len(indexes_with_data), + "num_shards": len(shard_records), + }, + f, + indent=2, + ) + _fsync_file(f) + # Everything the dump claims is now durable, so the staging directory can be + # published. Syncing the parent makes the rename itself survive a crash. + _fsync_directory(tmp_dir) + + if dump_dir.exists(): + shutil.rmtree(dump_dir) + tmp_dir.rename(dump_dir) + _fsync_directory(dump_dir.parent) + except Exception: + if tmp_dir.exists(): + shutil.rmtree(tmp_dir) + raise + + total_bytes = sum(path.stat().st_size for path in dump_dir.rglob("*") if path.is_file()) + logger.info(f"Dumped {len(unique_keys)} keys of partition {partition_id} to {dump_dir}") + return { + "keys": len(unique_keys), + "rows_with_data": len(indexes_with_data), + "shards": len(shard_records), + "bytes": total_bytes, + } + + +def read_row_index(dump_dir: str | Path) -> dict[str, Any]: + """Read a dump's row index without touching its payload. + + Lets a caller answer questions about which keys and fields a dump holds, for + example whether a row satisfies the current data contract, before paying to + deserialize any shard. + + Args: + dump_dir: Directory previously written by ``dump_data_by_key``. + + Returns: + ``{"partition_id": str, "rows": {key: {"global_index", "fields", "tag"}}}``. + + Raises: + FileNotFoundError: The row index is missing. + """ + row_index_path = Path(dump_dir) / _ROW_INDEX_FILE + if not row_index_path.exists(): + raise FileNotFoundError(f"{_ROW_INDEX_FILE} not found in {dump_dir}") + return torch.load(row_index_path, weights_only=False) + + +def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: + """Restore a dump into the running TransferQueue. + + Rows are written back through ``kv_batch_put``, so restoring merges by key: rows + outside the dump are untouched, and the number of storage units may differ from + the one used to write the dump. Global indexes are reallocated, which is safe + because a dump is addressed by key. + + Shards are processed one at a time. Every field of a row lives in the shard of the + unit that owned it, so a shard is self-contained and peak memory stays at one shard. + + Args: + dump_dir: Directory previously written by ``dump_data_by_key``. + + Returns: + ``{"keys", "rows_with_data", "shards", "bytes"}``. + + Raises: + RuntimeError: TransferQueue is not initialized. + FileNotFoundError: The dump is incomplete. + ValueError: A shard disagrees with the row index. + """ + from transfer_queue.interface import _TQ_CONTROLLER, kv_batch_put + + if _TQ_CONTROLLER is None: + raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") + + dump_dir = Path(dump_dir) + info_path = dump_dir / _DUMP_INFO_FILE + if not info_path.exists(): + raise FileNotFoundError(f"{_DUMP_INFO_FILE} not found in {dump_dir}") + with open(info_path, encoding="utf-8") as f: + dump_info = json.load(f) + if dump_info["format_version"] != DUMP_FORMAT_VERSION: + raise ValueError( + f"Unsupported dump format version {dump_info['format_version']} in {dump_dir}; " + f"this build reads version {DUMP_FORMAT_VERSION}" + ) + + row_index = read_row_index(dump_dir) + partition_id = row_index["partition_id"] + rows = row_index["rows"] + + key_by_index = {row["global_index"]: key for key, row in rows.items()} + fields_by_index = {row["global_index"]: row["fields"] for key, row in rows.items()} + + shard_dir = dump_dir / _SHARD_SUBDIR + shard_info_path = shard_dir / _SHARD_INFO_FILE + if not shard_info_path.exists(): + raise FileNotFoundError(f"{_SHARD_INFO_FILE} not found in {shard_dir}") + with open(shard_info_path, encoding="utf-8") as f: + shard_records = json.load(f) + + from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager + + restored_indexes: set[int] = set() + for record in shard_records: + shard_path = shard_dir / f"shard_{record['position']}_{record['storage_unit_id']}.pkl" + if not shard_path.exists(): + raise FileNotFoundError(f"Missing dump shard: {shard_path}") + + for signature, batch in AsyncSimpleStorageManager.read_shard(str(shard_path), fields_by_index).items(): + global_indexes = batch["global_indexes"] + batch_keys = [key_by_index[global_index] for global_index in global_indexes] + kv_batch_put( + keys=batch_keys, + partition_id=partition_id, + fields=batch["fields"], + tags=[rows[key]["tag"] for key in batch_keys], + ) + restored_indexes.update(global_indexes) + + # Rows that had no produced field yet were never sent to a storage unit; recreate + # them from the row index so a caller holding their keys still finds them. + keys_without_data = [key for key, row in rows.items() if not row["fields"]] + if keys_without_data: + kv_batch_put( + keys=keys_without_data, + partition_id=partition_id, + fields=None, + tags=[rows[key]["tag"] for key in keys_without_data], + ) + + if len(restored_indexes) != dump_info["num_rows_with_data"]: + raise ValueError( + f"Dump restore row count mismatch in {dump_dir}: shards yielded {len(restored_indexes)} rows, " + f"{_DUMP_INFO_FILE} declares {dump_info['num_rows_with_data']}" + ) + + total_bytes = sum(path.stat().st_size for path in dump_dir.rglob("*") if path.is_file()) + logger.info(f"Restored {len(rows)} keys into partition {partition_id} from {dump_dir}") + return { + "keys": len(rows), + "rows_with_data": len(restored_indexes), + "shards": len(shard_records), + "bytes": total_bytes, + } diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index c89d1cfb..b0a9f9bd 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -15,6 +15,8 @@ import asyncio import os +import pickle +import socket import time import warnings from collections import defaultdict @@ -813,6 +815,127 @@ async def _load_single_storage_unit( f"[{self.storage_manager_id}]: Error restoring for storage unit {target_storage_unit}: {str(e)}" ) from e + @with_storage_unit_socket + async def _dump_single_shard( + self, + path: str, + target_storage_unit: str, + global_indexes: list[int], + socket: zmq.Socket = None, + ) -> int: + """Ask one storage unit to write the rows it owns into a shard file.""" + try: + request_msg = ZMQMessage.create( + request_type=ZMQRequestType.DUMP_ROWS, # type: ignore[arg-type] + sender_id=self.storage_manager_id, + receiver_id=target_storage_unit, + body={"path": path, "global_indexes": global_indexes}, + ) + await socket.send_multipart(request_msg.serialize(), copy=False) + messages = await socket.recv_multipart(copy=False) + response_msg = ZMQMessage.deserialize(messages) + if response_msg.request_type != ZMQRequestType.DUMP_ROWS_RESPONSE or not response_msg.body.get("success"): + raise RuntimeError( + f"Storage unit {target_storage_unit} failed to dump rows to {path}: " + f"{response_msg.body.get('message', 'unknown error')}" + ) + missing_rows = response_msg.body["missing_rows"] + if missing_rows: + # The controller reported these rows as produced, so a unit that has no + # data for them means the two disagree. Never write a half table. + raise RuntimeError( + f"Storage unit {target_storage_unit} holds no data for requested rows: {missing_rows[:20]}" + ) + return response_msg.body["dumped_rows"] + except Exception as e: + raise RuntimeError( + f"[{self.storage_manager_id}]: Error dumping shard from storage unit {target_storage_unit}: {str(e)}" + ) from e + + async def dump_rows_by_index(self, shard_dir: str, global_indexes: list[int]) -> list[dict[str, Any]]: + """Dump the given rows into one shard per storage unit, in parallel. + + Each unit pickles its own rows in its own process, so the payload never passes + through the caller. A unit that owns none of the rows is skipped rather than + writing an empty shard. + + Args: + shard_dir: Directory to write shard files into. + global_indexes: Global indexes to dump. + + Returns: + One entry per written shard: ``{"position", "storage_unit_id", "rows"}``. + + Raises: + RuntimeError: A unit holds no data for a row it was asked to dump. + """ + shard_dir_path = Path(shard_dir) + shard_dir_path.mkdir(parents=True, exist_ok=True) + + routing = self._group_by_hash(global_indexes) + targets = [(su_id, group.global_indexes) for su_id, group in routing.items()] + paths = [str(shard_dir_path / f"shard_{pos}_{su_id}.pkl") for pos, (su_id, _) in enumerate(targets)] + + dumped_rows = await asyncio.gather( + *( + self._dump_single_shard(path, target_storage_unit=su_id, global_indexes=indexes) + for path, (su_id, indexes) in zip(paths, targets, strict=True) + ) + ) + + logger.info( + f"[{self.storage_manager_id}]: dumped {sum(dumped_rows)} rows " + f"across {len(targets)} shards to {shard_dir_path}" + ) + return [ + {"position": pos, "storage_unit_id": su_id, "rows": rows} + for pos, ((su_id, _), rows) in enumerate(zip(targets, dumped_rows, strict=True)) + ] + + @staticmethod + def read_shard(path: str, fields_by_index: dict[int, list[str]]) -> dict[tuple[str, ...], dict[str, Any]]: + """Read one shard and regroup it into batches ready for a put. + + Rows are grouped by field signature because a single put accepts one + homogeneous TensorDict, and a selective dump routinely mixes rows that + finished different fields. + + Args: + path: Shard file written by ``dump_rows_by_index``. + fields_by_index: Field signature expected for each global index in the shard. + + Returns: + ``{field_signature: {"global_indexes": [...], "fields": TensorDict}}``. + + Raises: + ValueError: The shard is missing a field that the row index expects. + """ + with open(path, "rb") as f: + shard = pickle.load(f) + + field_data = shard["field_data"] + grouped: dict[tuple[str, ...], list[int]] = defaultdict(list) + for global_index in shard["global_indexes"]: + grouped[tuple(fields_by_index[global_index])].append(global_index) + + batches: dict[tuple[str, ...], dict[str, Any]] = {} + for signature, global_indexes in grouped.items(): + packed = {} + for field_name in signature: + values = field_data.get(field_name) + if values is None: + raise ValueError(f"shard {path} is missing field {field_name!r} required by the row index") + # Reuse the same packer the production get path uses, so a dumped row + # rebuilds into exactly the container it was read as while live. + packed[field_name] = AsyncSimpleStorageManager._pack_field_values( + [values[global_index] for global_index in global_indexes] + ) + batches[signature] = { + "global_indexes": global_indexes, + "fields": TensorDict(packed, batch_size=len(global_indexes)) if signature else None, + } + return batches + async def save_checkpoint(self, checkpoint_dir: str) -> None: """Dump all storage units to the storage_units/ subdirectory of checkpoint_dir. diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index df28521e..dd5af70b 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -1003,6 +1003,8 @@ def _process_one_worker_request(self, worker_socket: zmq.Socket, monitor: Any) - response_msg = self._handle_get_metrics() elif operation == ZMQRequestType.SAVE_STORAGE_CHECKPOINT: # type: ignore[arg-type] response_msg = self._handle_save_checkpoint(request_msg) + elif operation == ZMQRequestType.DUMP_ROWS: # type: ignore[arg-type] + response_msg = self._handle_dump_rows(request_msg) elif operation == ZMQRequestType.LOAD_STORAGE_CHECKPOINT: # type: ignore[arg-type] response_msg = self._handle_load_checkpoint(request_msg) else: @@ -1293,6 +1295,67 @@ def _handle_save_checkpoint(self, data_parts) -> ZMQMessage: body={"success": False, "message": str(e)}, ) + def _handle_dump_rows(self, data_parts) -> ZMQMessage: + """Serialize the requested rows of this unit into a self-contained shard. + + This runs inside the storage unit process, so the payload is pickled where it + already lives instead of being shipped to the caller first. The shard is keyed + by global index; the caller holds the row index that maps keys onto them. + + Args: + data_parts: ZMQMessage with ``path`` and ``global_indexes`` in body. + ``path`` must be reachable from the node running this actor, which + means a shared filesystem in a multi-node deployment. + ``global_indexes`` is the subset this unit owns, already routed by + the storage manager. + + Returns: + ZMQMessage with ``success=True``, ``dumped_rows`` and ``missing_rows`` on + success, or ``success=False`` and ``message`` on failure. ``missing_rows`` + lists requested rows this unit holds no data for. + """ + path = data_parts.body["path"] + requested_indexes = set(data_parts.body["global_indexes"]) + try: + field_data = {} + for field_name, values in self.storage_data.field_data.items(): + selected_values = { + global_index: values[global_index] for global_index in requested_indexes if global_index in values + } + if selected_values: + field_data[field_name] = selected_values + dumped_indexes = self.storage_data._active_keys & requested_indexes + shard = { + "storage_unit_id": self.storage_unit_id, + "field_data": field_data, + "global_indexes": sorted(dumped_indexes), + } + with open(path, "wb") as f: + pickle.dump(shard, f, protocol=pickle.HIGHEST_PROTOCOL) + # Report success only once the shard is on disk. Without this the call + # returns while the payload is still dirty page cache, and a node that + # dies before writeback leaves a dump whose manifest claims rows that + # cannot be read back. + f.flush() + os.fsync(f.fileno()) + logger.info(f"[{self.storage_unit_id}]: dumped {len(dumped_indexes)} rows to {path}") + return ZMQMessage.create( + request_type=ZMQRequestType.DUMP_ROWS_RESPONSE, # type: ignore[arg-type] + sender_id=self.storage_unit_id, + body={ + "success": True, + "dumped_rows": len(dumped_indexes), + "missing_rows": sorted(requested_indexes - self.storage_data._active_keys), + }, + ) + except Exception as e: + logger.error(f"[{self.storage_unit_id}]: dump rows failed: {e}") + return ZMQMessage.create( + request_type=ZMQRequestType.DUMP_ROWS_RESPONSE, # type: ignore[arg-type] + sender_id=self.storage_unit_id, + body={"success": False, "message": str(e)}, + ) + def _handle_load_checkpoint(self, data_parts) -> ZMQMessage: """Restore storage unit data directly from its checkpoint path. diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index a94cd3ab..41eebc4a 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -135,6 +135,12 @@ class ZMQRequestType(ExplicitEnum): LOAD_STORAGE_CHECKPOINT = "LOAD_STORAGE_CHECKPOINT" LOAD_STORAGE_CHECKPOINT_RESPONSE = "LOAD_STORAGE_CHECKPOINT_RESPONSE" + # SELECTIVE DATA DUMP + DESCRIBE_ROWS_BY_KEY = "DESCRIBE_ROWS_BY_KEY" + DESCRIBE_ROWS_BY_KEY_RESPONSE = "DESCRIBE_ROWS_BY_KEY_RESPONSE" + DUMP_ROWS = "DUMP_ROWS" + DUMP_ROWS_RESPONSE = "DUMP_ROWS_RESPONSE" + class ZMQServerInfo: """ From 6d51c3e3ddfdf927e34ee7dedd2d6ace9fe81e57 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 23 Sep 2026 15:56:43 +0800 Subject: [PATCH 02/30] fix: isolate tensor storage in selective dumps Signed-off-by: OutstanderWang --- tests/test_data_dump.py | 75 ++++++++++++++++++++++++ transfer_queue/data_dump.py | 3 +- transfer_queue/storage/simple_storage.py | 3 +- transfer_queue/utils/compact_pickle.py | 33 +++++++++++ 4 files changed, 112 insertions(+), 2 deletions(-) create mode 100644 tests/test_data_dump.py create mode 100644 transfer_queue/utils/compact_pickle.py diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py new file mode 100644 index 00000000..64597333 --- /dev/null +++ b/tests/test_data_dump.py @@ -0,0 +1,75 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Selective dump integrity and publication tests.""" + +import pickle +from types import SimpleNamespace + +import pytest +import torch + +from transfer_queue import data_dump, interface +from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData +from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType + + +@pytest.fixture +def unit(): + cls = SimpleStorageUnit.__ray_metadata__.modified_class + unit = cls.__new__(cls) + unit.storage_unit_id = "test_unit" + unit.storage_data = StorageUnitData() + return unit + + +@pytest.mark.parametrize("row_count", [1, 32]) +def test_dump_excludes_unselected_tensor_storage(unit, tmp_path, row_count): + batch = torch.arange(64 * 4096, dtype=torch.float32).reshape(64, 4096) + unit.storage_data.put_data({"x": batch}, list(range(64))) + path = tmp_path / "shard.pkl" + indexes = list(range(row_count)) + reply = unit._handle_dump_rows( + ZMQMessage.create( + request_type=ZMQRequestType.DUMP_ROWS, + sender_id="test", + body={"path": str(path), "global_indexes": indexes}, + ) + ) + assert reply.body["success"] + with path.open("rb") as f: + shard = pickle.load(f) + assert shard["global_indexes"] == indexes + for index, value in shard["field_data"]["x"].items(): + torch.testing.assert_close(value, batch[index]) + assert value.untyped_storage().nbytes() == value.numel() * value.element_size() + assert path.stat().st_size < row_count * batch[0].numel() * batch.element_size() * 2 + assert unit.storage_data.field_data["x"][0].untyped_storage().nbytes() == batch.numel() * batch.element_size() + + +def test_row_index_compacts_tensors_inside_tags(monkeypatch, tmp_path): + batch = torch.arange(64 * 4096).reshape(64, 4096) + tag = {"nested": [SimpleNamespace(value=batch[0])]} + client = SimpleNamespace( + describe_rows_by_key=lambda *_: { + "key": {"global_index": 0, "fields": [], "tag": tag}, + } + ) + monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) + monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: client) + data_dump.dump_data_by_key(tmp_path / "dump", ["key"], "p") + restored = data_dump.read_row_index(tmp_path / "dump")["rows"]["key"]["tag"]["nested"][0].value + torch.testing.assert_close(restored, batch[0]) + assert restored.untyped_storage().nbytes() == restored.numel() * restored.element_size() diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index 2336f802..e46a715a 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -39,6 +39,7 @@ import torch +from transfer_queue.utils import compact_pickle from transfer_queue.utils.logging_utils import get_logger logger = get_logger(__name__) @@ -135,7 +136,7 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - # torch.save rather than json: a tag is an arbitrary picklable dict, and this # path must not fail on a tag that happens to hold a tensor. with open(tmp_dir / _ROW_INDEX_FILE, "wb") as f: - torch.save({"partition_id": partition_id, "rows": rows}, f) + torch.save({"partition_id": partition_id, "rows": rows}, f, pickle_module=compact_pickle) _fsync_file(f) with open(tmp_dir / _DUMP_INFO_FILE, "w", encoding="utf-8") as f: diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index dd5af70b..79bcb7fc 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -47,6 +47,7 @@ import torch import zmq +from transfer_queue.utils import compact_pickle from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads, log_heavy_operation from transfer_queue.utils.enum_utils import Role from transfer_queue.utils.logging_utils import get_logger @@ -1331,7 +1332,7 @@ def _handle_dump_rows(self, data_parts) -> ZMQMessage: "global_indexes": sorted(dumped_indexes), } with open(path, "wb") as f: - pickle.dump(shard, f, protocol=pickle.HIGHEST_PROTOCOL) + compact_pickle.dump(shard, f) # Report success only once the shard is on disk. Without this the call # returns while the payload is still dirty page cache, and a node that # dies before writeback leaves a dump whose manifest claims rows that diff --git a/transfer_queue/utils/compact_pickle.py b/transfer_queue/utils/compact_pickle.py new file mode 100644 index 00000000..f73d3c0d --- /dev/null +++ b/transfer_queue/utils/compact_pickle.py @@ -0,0 +1,33 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Pickle selected values without retaining the storage of unselected tensor rows.""" + +import pickle + +import torch + + +class Pickler(pickle.Pickler): + def reducer_override(self, value): + # Tensor views can retain an entire batch, including unselected keys. This + # also handles tensors nested in picklable payload objects or metadata. + if isinstance(value, torch.Tensor): + return value.clone().__reduce_ex__(pickle.HIGHEST_PROTOCOL) + return NotImplemented + + +def dump(value, file): + Pickler(file, protocol=pickle.HIGHEST_PROTOCOL).dump(value) From 10902024e5ff7e6585e48a597e770cd11306a1e3 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 23 Sep 2026 15:57:36 +0800 Subject: [PATCH 03/30] fix: preserve previous dump until publication succeeds Signed-off-by: OutstanderWang --- tests/test_data_dump.py | 65 +++++++++++++++++++++++++++++++++++++ transfer_queue/data_dump.py | 35 +++++++++++++++++--- 2 files changed, 95 insertions(+), 5 deletions(-) diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index 64597333..93de772d 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -73,3 +73,68 @@ def test_row_index_compacts_tensors_inside_tags(monkeypatch, tmp_path): restored = data_dump.read_row_index(tmp_path / "dump")["rows"]["key"]["tag"]["nested"][0].value torch.testing.assert_close(restored, batch[0]) assert restored.untyped_storage().nbytes() == restored.numel() * restored.element_size() + + +@pytest.fixture +def empty_dump_client(monkeypatch): + monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) + monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: object()) + + +def test_failed_publication_preserves_previous_dump(empty_dump_client, monkeypatch, tmp_path): + dump = tmp_path / "dump" + data_dump.dump_data_by_key(dump, [], "old") + rename = type(dump).rename + + def fail_publish(path, target): + if path == tmp_path / "dump.tmp": + raise OSError("publication failed") + return rename(path, target) + + monkeypatch.setattr(type(dump), "rename", fail_publish) + with pytest.raises(OSError, match="publication failed"): + data_dump.dump_data_by_key(dump, [], "new") + assert data_dump.read_row_index(dump)["partition_id"] == "old" + assert not (tmp_path / "dump.tmp").exists() + + +@pytest.mark.parametrize("next_operation", ["read", "load", "dump"]) +def test_interrupted_publication_recovers_on_next_access(empty_dump_client, monkeypatch, tmp_path, next_operation): + dump = tmp_path / "dump" + data_dump.dump_data_by_key(dump, [], "old") + rename = type(dump).rename + + def interrupt_publish(path, target): + if path == tmp_path / "dump.tmp": + raise KeyboardInterrupt + return rename(path, target) + + with monkeypatch.context() as patch: + patch.setattr(type(dump), "rename", interrupt_publish) + with pytest.raises(KeyboardInterrupt): + data_dump.dump_data_by_key(dump, [], "new") + assert not dump.exists() + assert (tmp_path / "dump.old").exists() + if next_operation == "read": + assert data_dump.read_row_index(dump)["partition_id"] == "old" + elif next_operation == "load": + assert data_dump.load_data_by_key(dump)["keys"] == 0 + assert data_dump.read_row_index(dump)["partition_id"] == "old" + else: + data_dump.dump_data_by_key(dump, [], "replacement") + assert data_dump.read_row_index(dump)["partition_id"] == "replacement" + + +def test_backup_cleanup_failure_does_not_fail_published_dump(empty_dump_client, monkeypatch, tmp_path): + dump = tmp_path / "dump" + data_dump.dump_data_by_key(dump, [], "old") + rmtree = data_dump.shutil.rmtree + + def fail_cleanup(path): + if path == tmp_path / "dump.old": + raise OSError("cleanup failed") + return rmtree(path) + + monkeypatch.setattr(data_dump.shutil, "rmtree", fail_cleanup) + data_dump.dump_data_by_key(dump, [], "new") + assert data_dump.read_row_index(dump)["partition_id"] == "new" diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index e46a715a..0bb4f2ce 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -72,14 +72,22 @@ def _fsync_directory(path: Path) -> None: os.close(fd) +def _recover_dump(dump_dir: Path) -> None: + old_dir = dump_dir.with_name(dump_dir.name + ".old") + if not dump_dir.exists() and old_dir.exists(): + old_dir.rename(dump_dir) + _fsync_directory(dump_dir.parent) + + def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) -> dict[str, int]: """Dump the rows addressed by ``keys`` into ``dump_dir``. Each storage unit pickles the rows it owns in its own process, so the payload never passes through the caller. The caller writes only a small row index. - The directory is replaced wholesale: the dump is staged in ``.tmp`` and - renamed over ``dump_dir``, so anything already there is destroyed. + The directory is replaced wholesale. The previous dump is retained as ``.old`` + until publication is durable, and recovered on the next access after interruption. + Only one writer may publish to a given directory at a time. .. note:: **Multi-node limitation**: dump_dir must reside on a shared network filesystem @@ -107,8 +115,10 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") unique_keys = list(dict.fromkeys(keys)) - dump_dir = Path(dump_dir) + dump_dir = Path(dump_dir).absolute() + _recover_dump(dump_dir) tmp_dir = dump_dir.parent / (dump_dir.name + ".tmp") + old_dir = dump_dir.with_name(dump_dir.name + ".old") if tmp_dir.exists(): shutil.rmtree(tmp_dir) tmp_dir.mkdir(parents=True) @@ -157,14 +167,26 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - _fsync_directory(tmp_dir) if dump_dir.exists(): - shutil.rmtree(dump_dir) + if old_dir.exists(): + shutil.rmtree(old_dir) + dump_dir.rename(old_dir) + _fsync_directory(dump_dir.parent) tmp_dir.rename(dump_dir) _fsync_directory(dump_dir.parent) except Exception: + _recover_dump(dump_dir) if tmp_dir.exists(): shutil.rmtree(tmp_dir) raise + # Publication already succeeded; cleanup must not invalidate the new dump. + if old_dir.exists(): + try: + shutil.rmtree(old_dir) + _fsync_directory(dump_dir.parent) + except OSError: + logger.warning("Could not remove previous dump at %s", old_dir, exc_info=True) + total_bytes = sum(path.stat().st_size for path in dump_dir.rglob("*") if path.is_file()) logger.info(f"Dumped {len(unique_keys)} keys of partition {partition_id} to {dump_dir}") return { @@ -191,7 +213,9 @@ def read_row_index(dump_dir: str | Path) -> dict[str, Any]: Raises: FileNotFoundError: The row index is missing. """ - row_index_path = Path(dump_dir) / _ROW_INDEX_FILE + dump_dir = Path(dump_dir) + _recover_dump(dump_dir) + row_index_path = dump_dir / _ROW_INDEX_FILE if not row_index_path.exists(): raise FileNotFoundError(f"{_ROW_INDEX_FILE} not found in {dump_dir}") return torch.load(row_index_path, weights_only=False) @@ -225,6 +249,7 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") dump_dir = Path(dump_dir) + _recover_dump(dump_dir) info_path = dump_dir / _DUMP_INFO_FILE if not info_path.exists(): raise FileNotFoundError(f"{_DUMP_INFO_FILE} not found in {dump_dir}") From 16b341d7d840b584ffb5f7b4507477bdd86d6bf0 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 23 Sep 2026 16:00:21 +0800 Subject: [PATCH 04/30] test: use portable selective dump directories Signed-off-by: OutstanderWang --- tests/e2e/conftest.py | 32 +++++++++++++++++++ .../e2e/test_data_dump_cross_topology_e2e.py | 12 ++----- tests/e2e/test_data_dump_e2e.py | 18 ++--------- 3 files changed, 37 insertions(+), 25 deletions(-) create mode 100644 tests/e2e/conftest.py diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py new file mode 100644 index 00000000..b68de91a --- /dev/null +++ b/tests/e2e/conftest.py @@ -0,0 +1,32 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import os +import shutil +import tempfile +from pathlib import Path + +import pytest + + +@pytest.fixture(scope="module") +def dump_test_root(tmp_path_factory): + shared_root = os.environ.get("TQ_DUMP_TEST_ROOT") + if shared_root: + root = Path(tempfile.mkdtemp(prefix="tq-dump-", dir=shared_root)) + else: + root = tmp_path_factory.mktemp("tq-dump") + yield root + shutil.rmtree(root) diff --git a/tests/e2e/test_data_dump_cross_topology_e2e.py b/tests/e2e/test_data_dump_cross_topology_e2e.py index cdf46960..b5fdbd53 100644 --- a/tests/e2e/test_data_dump_cross_topology_e2e.py +++ b/tests/e2e/test_data_dump_cross_topology_e2e.py @@ -28,9 +28,6 @@ """ import os -import shutil -import uuid -from pathlib import Path import pytest import ray @@ -42,8 +39,6 @@ os.environ["RAY_DEDUP_LOGS"] = "0" -_DEFAULT_DUMP_ROOT = "/apdcephfs_hldy/share_303541817/tq_dump_tests" - def _tq_config(num_storage_units: int) -> OmegaConf: return OmegaConf.create( @@ -70,11 +65,8 @@ def ray_init(): @pytest.fixture -def dump_dir(request): - root = Path(os.environ.get("TQ_DUMP_TEST_ROOT", _DEFAULT_DUMP_ROOT)) / uuid.uuid4().hex - root.mkdir(parents=True) - yield root / "dump" - shutil.rmtree(root, ignore_errors=True) +def dump_dir(dump_test_root, request): + return dump_test_root / request.node.name / "dump" def _row_input_ids(row: int) -> torch.Tensor: diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index ee2a5bbb..1c97b08e 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -17,8 +17,7 @@ Each storage unit writes its own shard from the node it runs on, so ``TQ_DUMP_TEST_ROOT`` must point at a filesystem shared by the whole cluster. -A node-local path such as pytest's ``tmp_path`` fails on a multi-node cluster -with ``FileNotFoundError``. +Single-node runs default to pytest-managed temporary storage. Run with: pytest tests/e2e/test_data_dump_e2e.py -v @@ -27,8 +26,6 @@ import json import os import shutil -import uuid -from pathlib import Path import pytest import ray @@ -41,7 +38,6 @@ os.environ["RAY_DEDUP_LOGS"] = "0" _NUM_STORAGE_UNITS = 4 -_DEFAULT_DUMP_ROOT = "/apdcephfs_hldy/share_303541817/tq_dump_tests" def _tq_config(num_storage_units: int) -> OmegaConf: @@ -95,17 +91,9 @@ def cleanup_partitions(controller): pass -@pytest.fixture(scope="module") -def shared_root(): - root = Path(os.environ.get("TQ_DUMP_TEST_ROOT", _DEFAULT_DUMP_ROOT)) / uuid.uuid4().hex - root.mkdir(parents=True) - yield root - shutil.rmtree(root, ignore_errors=True) - - @pytest.fixture -def dump_dir(shared_root, request): - case = shared_root / request.node.name.replace("/", "_") +def dump_dir(dump_test_root, request): + case = dump_test_root / request.node.name.replace("/", "_") yield case / "dump" shutil.rmtree(case, ignore_errors=True) From 3a072403e5bc1cd419c9c028a4ada25d94961777 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 23 Sep 2026 16:20:36 +0800 Subject: [PATCH 05/30] feat: restore selective dumps directly on owner storage units Signed-off-by: OutstanderWang --- docs/checkpoint.md | 9 +- docs/data_dump.md | 103 ++++++++ .../e2e/test_data_dump_cross_topology_e2e.py | 7 +- tests/e2e/test_data_dump_e2e.py | 91 +++++++ tests/test_data_dump.py | 247 +++++++++++++++++- transfer_queue/client.py | 73 +++++- transfer_queue/data_dump.py | 194 +++++++++----- transfer_queue/storage/dump_io.py | 30 +++ .../managers/simple_storage_manager.py | 179 ++++++------- transfer_queue/storage/simple_storage.py | 137 ++++++---- transfer_queue/utils/compact_pickle.py | 4 + transfer_queue/utils/tensor_utils.py | 51 ++++ transfer_queue/utils/zmq_utils.py | 2 + 13 files changed, 894 insertions(+), 233 deletions(-) create mode 100644 docs/data_dump.md create mode 100644 transfer_queue/storage/dump_io.py diff --git a/docs/checkpoint.md b/docs/checkpoint.md index 8c0a2829..7da407ea 100644 --- a/docs/checkpoint.md +++ b/docs/checkpoint.md @@ -167,4 +167,11 @@ client.load_controller_checkpoint(...) # (2) controller restored second If step (1) partially succeeds and step (2) fails, the system is left in a mixed state: some storage units hold checkpoint data while the controller still reflects its pre-restore state. There is no rollback path. -**Workaround**: If `load_checkpoint` raises, call `tq.init()` again to reset the system to a clean state before retrying. \ No newline at end of file +**Workaround**: If `load_checkpoint` raises, call `tq.init()` again to reset the system to a clean state before retrying. + +## Exporting selected keys + +Use [selective data dumps](data_dump.md) when restoring selected keys into an +existing system or a different number of storage units. Version-2 SimpleStorage +restores use the same owner-side I/O pattern as checkpoint load, but read assigned +row ranges and merge values instead of replacing entire unit and controller state. diff --git a/docs/data_dump.md b/docs/data_dump.md new file mode 100644 index 00000000..0ae58b1f --- /dev/null +++ b/docs/data_dump.md @@ -0,0 +1,103 @@ +# Selective data dump + +`dump_data_by_key` persists selected keys, their produced fields and tags. +`load_data_by_key` merges them into a running TransferQueue. It preserves existing +key indexes and unrelated rows, fields and tag entries; new keys receive indexes +from the current controller. It does not restore sampler or consumption state. + +```python +import transfer_queue as tq + +tq.init() +tq.dump_data_by_key("/shared/dumps/selected", ["sample-1", "sample-2"], "train") +index = tq.read_row_index("/shared/dumps/selected") +tq.load_data_by_key("/shared/dumps/selected") +``` + +Pause writes and clears for these keys during both operations. A dump is not an +atomic snapshot of concurrent writers. Multiple publishers must not use the same +dump directory concurrently. + +## Distributed I/O + +On export, the storage manager groups source indexes by their current storage +owner and concurrently asks those units to write their records. Only units holding +selected rows participate. Tensor storage is compacted during serialization, so a +row view cannot include the rest of its original batch, including inside tags. + +On version-2 SimpleStorage restore: + +1. The caller reads the row index and shard manifest and validates all file ranges. +2. The controller resolves existing keys and allocates indexes for new keys. +3. The storage manager routes records by the **current** indexes, then sends one + load request to each participating unit concurrently. +4. Each target unit reads only its assigned byte ranges and merges those values + into local storage. Records are processed in batches of at most 128 rows per + shard; the caller never reads or forwards their payloads. +5. After all units succeed, the client publishes returned field schemas through + the controller and merges tags using the normal metadata update path. + +The number of source units can differ from the number of destination units. +Even a dump with one source shard can restore across several target units because +records are independently addressable. Empty rows are recreated from metadata. + +The dump directory must be on a filesystem accessible to every participating +storage unit. Local temporary storage suffices for single-node deployments. + +| Operation | State | Payload I/O | Unit count on restore | +| --- | --- | --- | --- | +| Checkpoint | Entire controller and storage state | Each unit reads/writes its whole file | Must match | +| Selective v2 dump | Selected fields and tags, merged by key | Each owner unit reads/writes its records | May differ | + +`DUMP_ROWS` and `LOAD_ROWS` are included in storage operation metrics. Unit logs +record loaded rows and bytes; the manager logs total bytes and participating units. +These count application reads, not filesystem read-ahead or physical disk traffic. + +## Format and compatibility + +New dumps use `format_version: 2`: + +```text +dump_info.json +row_index.pt +shards/ + shard_info.json + shard_0_.pkl + ... +``` + +Each shard is a sequence of independent pickle records containing a source global +index and a field/value mapping. `shard_info.json` records each source index's +`[offset, length]`. Source indexes only locate records; they are never reused as +current indexes without controller resolution. `row_index.pt` remains readable +with `read_row_index` without opening payload shards. + +Version-1 dumps remain readable using the prior caller-side KV put path. Restoring +to a backend without direct selective loading also uses KV puts. These compatibility +paths do not provide distributed file reads. Old builds that only understand +version 1 cannot read version-2 dumps. Export of nonempty dumps currently requires +SimpleStorage. + +## Failure behavior + +Publication writes and syncs `.tmp`, moves the old directory to `.old`, publishes +the new directory, and syncs its parent before deleting the backup. If publication +is interrupted while the main directory is absent, the next dump, load or row-index +read recovers `.old`. A backup-cleanup error does not invalidate a published dump. + +Restore is not transactional. All unit requests are awaited before returning an +error, and failed storage requests prevent publication of new ready metadata. +Earlier payload writes or earlier metadata updates can remain after a failure; +existing produced rows may already contain restored values. Correct the cause and +retry the same dump with writers paused. No unrelated partition is cleared. + +## Tests + +Run the selective E2E suite with its default pytest-managed temporary directory: + +```bash +python -m pytest -q tests/e2e/test_data_dump_e2e.py tests/e2e/test_data_dump_cross_topology_e2e.py +``` + +For a multi-node Ray cluster, set `TQ_DUMP_TEST_ROOT` to an existing shared directory. +Tests create and remove only their own child directories beneath that root. diff --git a/tests/e2e/test_data_dump_cross_topology_e2e.py b/tests/e2e/test_data_dump_cross_topology_e2e.py index b5fdbd53..ddfce9c5 100644 --- a/tests/e2e/test_data_dump_cross_topology_e2e.py +++ b/tests/e2e/test_data_dump_cross_topology_e2e.py @@ -94,7 +94,7 @@ def _assert_rows_equal(actual: torch.Tensor, expected_rows: list[torch.Tensor]) @pytest.mark.parametrize( ("dump_units", "load_units"), - [(4, 2), (2, 4), (3, 3)], + [(4, 2), (2, 4), (3, 3), (1, 4)], ) def test_dump_restores_across_storage_unit_counts(ray_init, dump_dir, dump_units, load_units): # Define test data @@ -114,7 +114,12 @@ def test_dump_restores_across_storage_unit_counts(ray_init, dump_dir, dump_units # Restore into a different topology tq.init(_tq_config(load_units)) try: + _put_rows("bystander", ["unrelated"]) + _put_rows(partition_id, [keys[3]]) + old_index = tq.get_client().kv_retrieve_meta([keys[3]], partition_id).global_indexes[0] tq.load_data_by_key(dump_dir) + assert tq.get_client().kv_retrieve_meta([keys[3]], partition_id).global_indexes == [old_index] + assert tq.kv_batch_get(["unrelated"], "bystander", ["input_ids"]).batch_size[0] == 1 # Check restored state: every row readable, payload and tag intact retrieved = tq.kv_batch_get(keys=keys, partition_id=partition_id, select_fields=["input_ids"]) diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index 1c97b08e..319dbf38 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -23,9 +23,13 @@ pytest tests/e2e/test_data_dump_e2e.py -v """ +import builtins import json import os +import pickle import shutil +from pathlib import Path +from unittest.mock import AsyncMock import pytest import ray @@ -413,3 +417,90 @@ def test_load_rejects_an_unknown_format_version(self, tq_system, dump_dir): with pytest.raises(ValueError, match="Unsupported dump format version"): tq.load_data_by_key(dump_dir) + + +def test_direct_load_bypasses_caller_payload_io(tq_system, dump_dir, controller, monkeypatch): + partition = "direct_load" + keys = [f"key-{i}" for i in range(16)] + _put_rows(partition, keys) + tq.dump_data_by_key(dump_dir, keys, partition) + before = _keys_mapping(controller, partition) + tq.kv_put(keys[0], partition, fields=TensorDict({"extra": torch.tensor([[42]])}, batch_size=1), tag={"keep": True}) + _put_rows(partition, ["bystander"]) + client = tq.get_client() + manager = client.storage_manager + original_load = manager._load_selected_rows + responses = [] + + async def load(*args, **kwargs): + response = await original_load(*args, **kwargs) + responses.append((kwargs["target_storage_unit"], response)) + return response + + real_open = builtins.open + + def no_payload_open(path, *args, **kwargs): + if isinstance(path, (str, Path)) and Path(path).name.startswith("shard_") and str(path).endswith(".pkl"): + raise AssertionError("Caller opened a payload shard") + return real_open(path, *args, **kwargs) + + monkeypatch.setattr(builtins, "open", no_payload_open) + monkeypatch.setattr(manager, "_load_selected_rows", load) + monkeypatch.setattr(manager, "put_data", AsyncMock(side_effect=AssertionError("Caller sent payload through put"))) + tq.load_data_by_key(dump_dir) + assert len({unit for unit, _ in responses}) == _NUM_STORAGE_UNITS + shard_bytes = sum(path.stat().st_size for path in (dump_dir / "shards").glob("shard_*.pkl")) + assert sum(response["bytes_read"] for _, response in responses) == shard_bytes + after = _keys_mapping(controller, partition) + assert all(after[key] == before[key] for key in keys) + assert "bystander" in after + actual = tq.kv_batch_get(keys, partition, select_fields=["input_ids"]) + _assert_rows_equal(actual["input_ids"], [_row_input_ids(i) for i in range(16)]) + extra = tq.kv_batch_get([keys[0]], partition, select_fields=["extra"]) + _assert_rows_equal(extra["extra"], [torch.tensor([42])]) + snapshot = ray.get(controller.get_partition_snapshot.remote(partition)) + assert snapshot.custom_meta[after[keys[0]]]["keep"] + assert snapshot.custom_meta[after[keys[0]]]["idx"] == 0 + + +def test_version_one_dump_remains_readable(tq_system, dump_dir, controller): + partition = "legacy" + dump_dir.mkdir(parents=True) + (dump_dir / "shards").mkdir() + torch.save( + { + "partition_id": partition, + "rows": { + "k": {"global_index": 100, "fields": ["x"], "tag": {"old": True}}, + "empty": {"global_index": 101, "fields": [], "tag": {}}, + }, + }, + dump_dir / "row_index.pt", + ) + (dump_dir / "dump_info.json").write_text( + json.dumps( + {"format_version": 1, "partition_id": partition, "num_keys": 2, "num_rows_with_data": 1, "num_shards": 1} + ) + ) + (dump_dir / "shards" / "shard_info.json").write_text( + json.dumps([{"position": 0, "storage_unit_id": "old", "rows": 1}]) + ) + with (dump_dir / "shards" / "shard_0_old.pkl").open("wb") as f: + pickle.dump({"global_indexes": [100], "field_data": {"x": {100: torch.tensor([7, 8])}}}, f) + tq.load_data_by_key(dump_dir) + assert sorted(_keys_mapping(controller, partition)) == ["empty", "k"] + _assert_rows_equal(tq.kv_batch_get(["k"], partition, select_fields=["x"])["x"], [torch.tensor([7, 8])]) + + +def test_corrupt_shard_keeps_new_rows_unproduced(tq_system, dump_dir, controller): + partition = "corrupt" + _put_rows(partition, ["key"]) + tq.dump_data_by_key(dump_dir, ["key"], partition) + path = next((dump_dir / "shards").glob("shard_*.pkl")) + path.write_bytes(b"!" * path.stat().st_size) + ray.get(controller.clear_partition.remote(partition)) + with pytest.raises(RuntimeError, match="failed to load rows"): + tq.load_data_by_key(dump_dir) + snapshot = ray.get(controller.get_partition_snapshot.remote(partition)) + assert "key" in snapshot.keys_mapping + assert not snapshot.field_metadata diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index 93de772d..82f5426f 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -15,13 +15,20 @@ """Selective dump integrity and publication tests.""" +import asyncio +import builtins +import io import pickle from types import SimpleNamespace +from unittest.mock import AsyncMock import pytest import torch from transfer_queue import data_dump, interface +from transfer_queue.client import AsyncTransferQueueClient +from transfer_queue.metadata import BatchMeta +from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType @@ -49,12 +56,15 @@ def test_dump_excludes_unselected_tensor_storage(unit, tmp_path, row_count): ) ) assert reply.body["success"] + assert set(reply.body["row_offsets"]) == set(indexes) with path.open("rb") as f: - shard = pickle.load(f) - assert shard["global_indexes"] == indexes - for index, value in shard["field_data"]["x"].items(): - torch.testing.assert_close(value, batch[index]) - assert value.untyped_storage().nbytes() == value.numel() * value.element_size() + for index, (offset, length) in reply.body["row_offsets"].items(): + f.seek(offset) + row = pickle.loads(f.read(length)) + assert row["global_index"] == index + value = row["fields"]["x"] + torch.testing.assert_close(value, batch[index]) + assert value.untyped_storage().nbytes() == value.numel() * value.element_size() assert path.stat().st_size < row_count * batch[0].numel() * batch.element_size() * 2 assert unit.storage_data.field_data["x"][0].untyped_storage().nbytes() == batch.numel() * batch.element_size() @@ -138,3 +148,230 @@ def fail_cleanup(path): monkeypatch.setattr(data_dump.shutil, "rmtree", fail_cleanup) data_dump.dump_data_by_key(dump, [], "new") assert data_dump.read_row_index(dump)["partition_id"] == "new" + + +def test_unit_reads_only_assigned_ranges_and_merges(unit, tmp_path, monkeypatch): + batch = torch.arange(8 * 4096, dtype=torch.float32).reshape(8, 4096) + unit.storage_data.put_data({"x": batch, "stale": batch}, list(range(8))) + path = tmp_path / "shard.pkl" + reply = unit._handle_dump_rows( + ZMQMessage.create( + request_type=ZMQRequestType.DUMP_ROWS, + sender_id="test", + body={ + "path": str(path), + "global_indexes": list(range(8)), + "fields_by_index": {index: ["x"] for index in range(8)}, + }, + ) + ) + assert reply.body["success"] + content = path.read_bytes() + reads = [] + + class TrackedFile(io.BytesIO): + name = str(path) + + def read(self, size=-1): + reads.append((self.tell(), size)) + return super().read(size) + + open_file = builtins.open + monkeypatch.setattr( + builtins, + "open", + lambda name, *a, **kw: TrackedFile(content) if str(name) == str(path) else open_file(name, *a, **kw), + ) + unit.storage_data = StorageUnitData() + unit.storage_data.put_data({"keep": ["value"]}, [101]) + records = [] + for source, target in [(1, 101), (6, 106)]: + offset, length = reply.body["row_offsets"][source] + records.append( + {"source_index": source, "target_index": target, "fields": ["x"], "offset": offset, "length": length} + ) + loaded = unit._handle_load_rows( + ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS, + sender_id="test", + body={"shards": [{"path": str(path), "records": records}]}, + ) + ) + assert loaded.body["success"], loaded.body + assert reads == [(row["offset"], row["length"]) for row in records] + assert loaded.body["bytes_read"] == sum(row["length"] for row in records) + assert set(unit.storage_data.field_data) == {"x", "keep"} + assert unit.storage_data.field_data["keep"][101] == "value" + for row in records: + torch.testing.assert_close(unit.storage_data.field_data["x"][row["target_index"]], batch[row["source_index"]]) + assert sorted(index for update in loaded.body["updates"] for index in update["global_indexes"]) == [101, 106] + + +@pytest.mark.asyncio +async def test_manager_loads_current_owners_concurrently(): + manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) + manager.storage_manager_id = "test" + manager.storage_unit_infos = dict.fromkeys(["u0", "u1", "u2", "u3"]) + manager.close = lambda: None + started = set() + ready = asyncio.Event() + seen = [] + + async def load(shards, target_storage_unit): + started.add(target_storage_unit) + if len(started) == 4: + ready.set() + await asyncio.wait_for(ready.wait(), timeout=2) + for shard in shards: + for row in shard["records"]: + assert target_storage_unit == f"u{row['target_index'] % 4}" + seen.append(row["source_index"]) + return {"updates": [], "bytes_read": 0} + + manager._load_selected_rows = load + records = [{"source_index": i, "target_index": 31 - i} for i in range(16)] + assert await manager.load_rows_by_index([{"path": "shard.pkl", "records": records}]) == [] + assert sorted(seen) == list(range(16)) + + +@pytest.mark.asyncio +async def test_failed_load_does_not_publish_ready_metadata(): + client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) + client.close = lambda: None + client.storage_manager = SimpleNamespace(load_rows_by_index=AsyncMock(side_effect=RuntimeError("read failed"))) + client.async_kv_retrieve_meta = AsyncMock(return_value=BatchMeta(global_indexes=[9], partition_ids=["p"])) + client._publish_loaded_rows = AsyncMock() + client.async_set_custom_meta = AsyncMock() + with pytest.raises(RuntimeError, match="read failed"): + await client.async_load_rows_by_key("p", {"k": {"tag": {}}}, []) + client._publish_loaded_rows.assert_not_called() + client.async_set_custom_meta.assert_not_called() + + +@pytest.mark.parametrize("problem", ["wrong_index", "truncated", "missing_field"]) +def test_unit_rejects_invalid_records(unit, tmp_path, problem): + path = tmp_path / "row.pkl" + path.write_bytes(pickle.dumps({"global_index": 1, "fields": {"x": "value"}})) + record = {"source_index": 1, "target_index": 2, "offset": 0, "length": path.stat().st_size, "fields": ["x"]} + if problem == "wrong_index": + record["source_index"] = 3 + elif problem == "truncated": + record["length"] += 1 + else: + record["fields"] = ["missing"] + reply = unit._handle_load_rows( + ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS, + sender_id="test", + body={"shards": [{"path": str(path), "records": [record]}]}, + ) + ) + assert not reply.body["success"] + assert not unit.storage_data._active_keys + + +def test_version_two_falls_back_to_kv_for_other_backends(unit, monkeypatch, tmp_path): + unit.storage_data.put_data({"x": [torch.tensor([7, 8])]}, [10]) + rows = { + "k": {"global_index": 10, "fields": ["x"], "tag": {"tag": 1}}, + "empty": {"global_index": 11, "fields": [], "tag": {}}, + } + + def dump(shard_dir, indexes, fields_by_index): + directory = type(tmp_path)(shard_dir) + directory.mkdir(parents=True) + response = unit._handle_dump_rows( + ZMQMessage.create( + request_type=ZMQRequestType.DUMP_ROWS, + sender_id="test", + body={ + "path": str(directory / "shard_0_unit.pkl"), + "global_indexes": indexes, + "fields_by_index": fields_by_index, + }, + ) + ) + assert response.body["success"] + return [ + { + "position": 0, + "storage_unit_id": "unit", + "rows": len(indexes), + "row_offsets": response.body["row_offsets"], + } + ] + + client = SimpleNamespace(describe_rows_by_key=lambda *_: rows, dump_rows_by_index=dump, storage_manager=object()) + monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) + monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: client) + data_dump.dump_data_by_key(tmp_path / "dump", list(rows), "p") + calls = [] + monkeypatch.setattr(interface, "kv_batch_put", lambda *args, **kwargs: calls.append((args, kwargs))) + data_dump.load_data_by_key(tmp_path / "dump") + assert calls[0][0][:2] == (["k"], "p") + torch.testing.assert_close(calls[0][0][2]["x"][0], torch.tensor([7, 8])) + assert calls[0][1]["tags"] == [{"tag": 1}] + assert calls[1] == ((["empty"], "p"), {"tags": [{}]}) + + +@pytest.mark.asyncio +async def test_load_waits_for_other_units_before_raising(): + manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) + manager.storage_manager_id = "test" + manager.storage_unit_infos = dict.fromkeys(["u0", "u1"]) + manager.close = lambda: None + failed = asyncio.Event() + finished = [] + + async def load(shards, target_storage_unit): + if target_storage_unit == "u0": + failed.set() + raise RuntimeError("unit failed") + await failed.wait() + await asyncio.sleep(0) + finished.append(target_storage_unit) + return {"bytes_read": 0, "updates": []} + + manager._load_selected_rows = load + with pytest.raises(RuntimeError, match="unit failed"): + await manager.load_rows_by_index([{"path": "shard", "records": [{"target_index": 0}, {"target_index": 1}]}]) + assert finished == ["u1"] + + +@pytest.mark.asyncio +async def test_load_rejects_failed_controller_metadata_ack(): + client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) + client.close = lambda: None + client._request_controller = AsyncMock(return_value=SimpleNamespace(body={"success": False})) + with pytest.raises(RuntimeError, match="Controller rejected"): + await AsyncTransferQueueClient._publish_loaded_rows.__wrapped__( + client, + "p", + [ + {"global_indexes": [3], "field_schema": {}}, + ], + ) + + +@pytest.mark.asyncio +async def test_dump_waits_for_writers_before_cleanup_can_start(tmp_path): + manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) + manager.storage_manager_id = "test" + manager.storage_unit_infos = dict.fromkeys(["u0", "u1"]) + manager.close = lambda: None + failed = asyncio.Event() + completed = [] + + async def dump(path, target_storage_unit, global_indexes, fields_by_index): + if target_storage_unit == "u0": + failed.set() + raise OSError("write failed") + await failed.wait() + await asyncio.sleep(0) + completed.append(target_storage_unit) + return {1: [0, 1]} + + manager._dump_single_shard = dump + with pytest.raises(OSError, match="write failed"): + await manager.dump_rows_by_index(str(tmp_path), [0, 1]) + assert completed == ["u1"] diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 13967260..11967b44 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1167,12 +1167,18 @@ async def async_describe_rows_by_key( except Exception as e: raise RuntimeError(f"[{self.client_id}]: Error in describe_rows_by_key: {str(e)}") from e - async def async_dump_rows_by_index(self, shard_dir: str, global_indexes: list[int]) -> list[dict[str, Any]]: + async def async_dump_rows_by_index( + self, + shard_dir: str, + global_indexes: list[int], + fields_by_index: dict[int, list[str]] | None = None, + ) -> list[dict[str, Any]]: """Asynchronously dump the given rows into per-storage-unit shards. Args: shard_dir: Directory to write shard files into. global_indexes: Global indexes to dump. + fields_by_index: Produced fields to persist; omitted for a raw storage dump. Returns: One entry per written shard. @@ -1189,7 +1195,50 @@ async def async_dump_rows_by_index(self, shard_dir: str, global_indexes: list[in ) if not hasattr(self.storage_manager, "dump_rows_by_index"): raise NotImplementedError(f"{type(self.storage_manager).__name__} does not support selective data dump") - return await self.storage_manager.dump_rows_by_index(shard_dir, global_indexes) + return await self.storage_manager.dump_rows_by_index(shard_dir, global_indexes, fields_by_index) + + async def async_load_rows_by_key( + self, + partition_id: str, + rows: dict[str, dict[str, Any]], + shards: list[dict[str, Any]], + ) -> None: + """Allocate current indexes, load payloads at owner units, then publish metadata.""" + manager = getattr(self, "storage_manager", None) + if manager is None or not hasattr(manager, "load_rows_by_index"): + raise NotImplementedError("Storage backend does not support direct selective load") + if not rows: + return + keys = list(rows) + metadata = await self.async_kv_retrieve_meta(keys, partition_id, create=True) + if metadata.size != len(keys): + raise RuntimeError("Selective load did not allocate every key") + target_indexes = dict(zip(keys, metadata.global_indexes, strict=True)) + for shard in shards: + for record in shard["records"]: + record["target_index"] = target_indexes[record["key"]] + updates = await manager.load_rows_by_index(shards) + await self._publish_loaded_rows(partition_id, updates) + metadata.update_custom_meta([rows[key]["tag"] for key in keys]) + await self.async_set_custom_meta(metadata) + + @with_controller_socket + async def _publish_loaded_rows( + self, + partition_id: str, + updates: list[dict[str, Any]], + socket: zmq.asyncio.Socket | None = None, + ) -> None: + # Loading must report a failed metadata update instead of silently succeeding. + for update in updates: + response = await self._request_controller( + socket=socket, + request_type=ZMQRequestType.NOTIFY_DATA_UPDATE, + response_type=ZMQRequestType.NOTIFY_DATA_UPDATE_ACK, + body={"partition_id": partition_id, **update}, + ) + if not response.body.get("success"): + raise RuntimeError(f"Controller rejected loaded row metadata for partition {partition_id!r}") # ==================== Checkpoint API ==================== @with_controller_socket @@ -1377,6 +1426,7 @@ def wrapper(*args, **kwargs): self._kv_list = _make_sync(self.async_kv_list) self._describe_rows_by_key = _make_sync(self.async_describe_rows_by_key) self._dump_rows_by_index = _make_sync(self.async_dump_rows_by_index) + self._load_rows_by_key = _make_sync(self.async_load_rows_by_key) self._save_controller_checkpoint = _make_sync(self.async_save_controller_checkpoint) self._load_controller_checkpoint = _make_sync(self.async_load_controller_checkpoint) self._save_storage_checkpoint = _make_sync(self.async_save_storage_checkpoint) @@ -1829,12 +1879,18 @@ def describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, """ return self._describe_rows_by_key(partition_id, keys) - def dump_rows_by_index(self, shard_dir: str, global_indexes: list[int]) -> list[dict[str, Any]]: + def dump_rows_by_index( + self, + shard_dir: str, + global_indexes: list[int], + fields_by_index: dict[int, list[str]] | None = None, + ) -> list[dict[str, Any]]: """Synchronously dump the given rows into per-storage-unit shards. Args: shard_dir: Directory to write shard files into. global_indexes: Global indexes to dump. + fields_by_index: Produced fields to persist; omitted for a raw storage dump. Returns: One entry per written shard. @@ -1844,7 +1900,16 @@ def dump_rows_by_index(self, shard_dir: str, global_indexes: list[int]) -> list[ no data for a row it was asked to dump. NotImplementedError: If the storage backend does not support dumping. """ - return self._dump_rows_by_index(shard_dir, global_indexes) + return self._dump_rows_by_index(shard_dir, global_indexes, fields_by_index) + + def load_rows_by_key( + self, + partition_id: str, + rows: dict[str, dict[str, Any]], + shards: list[dict[str, Any]], + ) -> None: + """Restore indexed dump records directly on the current storage owner units.""" + return self._load_rows_by_key(partition_id, rows, shards) # ==================== Checkpoint API ==================== def save_controller_checkpoint(self, path: str) -> None: diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index 0bb4f2ce..57458dfd 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -13,38 +13,41 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Selective data dump: persist the rows behind a set of keys, and read them back. +"""Persist selected rows and restore them by key without checkpointing controller state. -This is deliberately not a checkpoint. It stores payload and the metadata needed to -address it by key, and nothing about the controller: no index manager, no sampler, no -partition snapshot. Restoring therefore goes through the ordinary ``kv_batch_put`` -path, which means a dump taken with N storage units restores into a system with M -storage units. Use ``save_checkpoint`` when you need a full system image instead. +Version 2 stores independent row records in each storage-unit shard. The manifest +maps source indexes to byte offsets, so current owner units can read only their rows +when restoring into a different topology. Version 1 remains readable via KV puts. Layout:: / - dump_info.json # partition, counts, shard count + dump_info.json # version, partition, counts row_index.pt # key -> {global_index, fields, tag} shards/ - shard_info.json # [{position, storage_unit_id, rows}] - shard__.pkl # {field_data: {field: {gidx: value}}, global_indexes} + shard_info.json # unit, row count, source index -> [offset, length] + shard__.pkl # independent {global_index, fields} records """ import json import os +import pickle import shutil +from collections import defaultdict from pathlib import Path from typing import Any import torch +from tensordict import TensorDict +from transfer_queue.storage.dump_io import read_dump_row from transfer_queue.utils import compact_pickle from transfer_queue.utils.logging_utils import get_logger +from transfer_queue.utils.tensor_utils import pack_field_values logger = get_logger(__name__) -DUMP_FORMAT_VERSION = 1 +DUMP_FORMAT_VERSION = 2 _DUMP_INFO_FILE = "dump_info.json" _ROW_INDEX_FILE = "row_index.pt" @@ -132,7 +135,13 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - indexes_with_data = sorted(row["global_index"] for row in rows.values() if row["fields"]) shard_records = ( - client.dump_rows_by_index(str(tmp_dir / _SHARD_SUBDIR), indexes_with_data) if indexes_with_data else [] + client.dump_rows_by_index( + str(tmp_dir / _SHARD_SUBDIR), + indexes_with_data, + {row["global_index"]: row["fields"] for row in rows.values() if row["fields"]}, + ) + if indexes_with_data + else [] ) shard_dir = tmp_dir / _SHARD_SUBDIR shard_dir.mkdir(parents=True, exist_ok=True) @@ -222,100 +231,147 @@ def read_row_index(dump_dir: str | Path) -> dict[str, Any]: def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: - """Restore a dump into the running TransferQueue. - - Rows are written back through ``kv_batch_put``, so restoring merges by key: rows - outside the dump are untouched, and the number of storage units may differ from - the one used to write the dump. Global indexes are reallocated, which is safe - because a dump is addressed by key. + """Merge selected rows into the running system, preserving existing key indexes. - Shards are processed one at a time. Every field of a row lives in the shard of the - unit that owned it, so a shard is self-contained and peak memory stays at one shard. + SimpleStorage units read their assigned version-2 records directly and in parallel. + Version-1 dumps and other backends use the compatible KV put path. New keys receive + current indexes; unrelated rows and fields remain untouched. Writers and clears for + these keys must be paused during restore. Failure may leave partial payload writes; + retry the same dump after correcting the failure. Args: - dump_dir: Directory previously written by ``dump_data_by_key``. + dump_dir: Directory previously written by ``dump_data_by_key``. For direct + distributed loading it must be accessible from every storage unit. Returns: ``{"keys", "rows_with_data", "shards", "bytes"}``. Raises: - RuntimeError: TransferQueue is not initialized. + RuntimeError: TransferQueue is not initialized or a storage unit fails. FileNotFoundError: The dump is incomplete. - ValueError: A shard disagrees with the row index. + ValueError: The manifest or a row disagrees with the row index. """ - from transfer_queue.interface import _TQ_CONTROLLER, kv_batch_put + from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client if _TQ_CONTROLLER is None: raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") - dump_dir = Path(dump_dir) + dump_dir = Path(dump_dir).absolute() _recover_dump(dump_dir) info_path = dump_dir / _DUMP_INFO_FILE if not info_path.exists(): raise FileNotFoundError(f"{_DUMP_INFO_FILE} not found in {dump_dir}") with open(info_path, encoding="utf-8") as f: dump_info = json.load(f) - if dump_info["format_version"] != DUMP_FORMAT_VERSION: + if dump_info["format_version"] not in (1, DUMP_FORMAT_VERSION): raise ValueError( f"Unsupported dump format version {dump_info['format_version']} in {dump_dir}; " - f"this build reads version {DUMP_FORMAT_VERSION}" + f"this build reads versions 1 and {DUMP_FORMAT_VERSION}" ) row_index = read_row_index(dump_dir) partition_id = row_index["partition_id"] rows = row_index["rows"] - key_by_index = {row["global_index"]: key for key, row in rows.items()} - fields_by_index = {row["global_index"]: row["fields"] for key, row in rows.items()} - shard_dir = dump_dir / _SHARD_SUBDIR - shard_info_path = shard_dir / _SHARD_INFO_FILE - if not shard_info_path.exists(): - raise FileNotFoundError(f"{_SHARD_INFO_FILE} not found in {shard_dir}") - with open(shard_info_path, encoding="utf-8") as f: + with open(shard_dir / _SHARD_INFO_FILE, encoding="utf-8") as f: shard_records = json.load(f) - - from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager - - restored_indexes: set[int] = set() + if ( + len(rows) != dump_info["num_keys"] + or len(shard_records) != dump_info["num_shards"] + or partition_id != dump_info["partition_id"] + ): + raise ValueError("Dump manifest disagrees with the row index") + keys_by_index = {row["global_index"]: key for key, row in rows.items() if row["fields"]} + if len(keys_by_index) != dump_info["num_rows_with_data"]: + raise ValueError("Dump row count disagrees with the row index") + shards = [] + seen = set() for record in shard_records: - shard_path = shard_dir / f"shard_{record['position']}_{record['storage_unit_id']}.pkl" - if not shard_path.exists(): - raise FileNotFoundError(f"Missing dump shard: {shard_path}") - - for signature, batch in AsyncSimpleStorageManager.read_shard(str(shard_path), fields_by_index).items(): - global_indexes = batch["global_indexes"] - batch_keys = [key_by_index[global_index] for global_index in global_indexes] - kv_batch_put( - keys=batch_keys, - partition_id=partition_id, - fields=batch["fields"], - tags=[rows[key]["tag"] for key in batch_keys], - ) - restored_indexes.update(global_indexes) - - # Rows that had no produced field yet were never sent to a storage unit; recreate - # them from the row index so a caller holding their keys still finds them. - keys_without_data = [key for key, row in rows.items() if not row["fields"]] - if keys_without_data: - kv_batch_put( - keys=keys_without_data, - partition_id=partition_id, - fields=None, - tags=[rows[key]["tag"] for key in keys_without_data], - ) - - if len(restored_indexes) != dump_info["num_rows_with_data"]: - raise ValueError( - f"Dump restore row count mismatch in {dump_dir}: shards yielded {len(restored_indexes)} rows, " - f"{_DUMP_INFO_FILE} declares {dump_info['num_rows_with_data']}" - ) + path = shard_dir / f"shard_{record['position']}_{record['storage_unit_id']}.pkl" + if not path.is_file(): + raise FileNotFoundError(f"Missing dump shard: {path}") + records = [] + if dump_info["format_version"] == 2: + size = path.stat().st_size + offsets = record["row_offsets"] + if len(offsets) != record["rows"]: + raise ValueError(f"Dump shard row count mismatch: {path}") + for source_index, (offset, length) in offsets.items(): + source_index = int(source_index) + if source_index not in keys_by_index or source_index in seen: + raise ValueError(f"Unexpected or duplicated row {source_index} in {path}") + if offset < 0 or length <= 0 or offset + length > size: + raise ValueError(f"Invalid row range for {source_index} in {path}") + key = keys_by_index[source_index] + records.append( + { + "key": key, + "source_index": source_index, + "fields": rows[key]["fields"], + "offset": offset, + "length": length, + } + ) + seen.add(source_index) + shards.append({"path": str(path), "records": records}) + if dump_info["format_version"] == 2 and seen != set(keys_by_index): + raise ValueError("Dump shards do not contain every produced row") + + client = _maybe_create_tq_client() + if dump_info["format_version"] == 2 and hasattr(getattr(client, "storage_manager", None), "load_rows_by_index"): + client.load_rows_by_key(partition_id, rows, shards) + else: + _load_via_kv(partition_id, rows, shards, dump_info["format_version"]) total_bytes = sum(path.stat().st_size for path in dump_dir.rglob("*") if path.is_file()) logger.info(f"Restored {len(rows)} keys into partition {partition_id} from {dump_dir}") return { "keys": len(rows), - "rows_with_data": len(restored_indexes), + "rows_with_data": len(keys_by_index), "shards": len(shard_records), "bytes": total_bytes, } + + +def _load_via_kv(partition_id: str, rows: dict[str, Any], shards: list[dict], version: int) -> None: + """Retain v1 and non-SimpleStorage compatibility without changing their put contract.""" + from transfer_queue.interface import kv_batch_put + + keys_by_index = {row["global_index"]: key for key, row in rows.items()} + restored = set() + for shard in shards: + with open(shard["path"], "rb") as f: + if version == 1: + saved = pickle.load(f) + records = [ + {"key": keys_by_index[index], "source_index": index, "fields": rows[keys_by_index[index]]["fields"]} + for index in saved["global_indexes"] + ] + else: + records = shard["records"] + for start in range(0, len(records), 128): + groups = defaultdict(list) + for record in records[start : start + 128]: + index = record["source_index"] + fields = record["fields"] + if version == 1: + values = {name: saved["field_data"][name][index] for name in fields} + else: + values = read_dump_row(f, record["offset"], record["length"], index, fields) + groups[tuple(fields)].append((record["key"], values)) + restored.add(index) + for signature, batch in groups.items(): + keys = [key for key, _ in batch] + packed = {name: pack_field_values([values[name] for _, values in batch]) for name in signature} + kv_batch_put( + keys, + partition_id, + TensorDict(packed, batch_size=len(keys)), + tags=[rows[key]["tag"] for key in keys], + ) + if restored != {row["global_index"] for row in rows.values() if row["fields"]}: + raise ValueError("Dump restore row count mismatch") + keys = [key for key, row in rows.items() if not row["fields"]] + if keys: + kv_batch_put(keys, partition_id, tags=[rows[key]["tag"] for key in keys]) diff --git a/transfer_queue/storage/dump_io.py b/transfer_queue/storage/dump_io.py new file mode 100644 index 00000000..3bdcf7bf --- /dev/null +++ b/transfer_queue/storage/dump_io.py @@ -0,0 +1,30 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Read independently encoded rows without scanning unrelated shard data.""" + +import pickle + + +def read_dump_row(file, offset: int, length: int, global_index: int, fields: list[str]) -> dict: + """Read and validate exactly one record against its expected index and fields.""" + file.seek(offset) + payload = file.read(length) + if len(payload) != length: + raise ValueError(f"Truncated dump row {global_index} in {file.name}") + row = pickle.loads(payload) + if row["global_index"] != global_index or set(row["fields"]) != set(fields): + raise ValueError(f"Dump row {global_index} disagrees with the row index in {file.name}") + return row["fields"] diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index b0a9f9bd..f82f79b7 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -15,7 +15,6 @@ import asyncio import os -import pickle import socket import time import warnings @@ -39,6 +38,7 @@ ) from transfer_queue.utils.common import log_heavy_operation from transfer_queue.utils.logging_utils import get_logger +from transfer_queue.utils.tensor_utils import pack_field_values from transfer_queue.utils.zmq_utils import ( TQ_SOCKET_POOL_SIZE, ZMQMessage, @@ -540,55 +540,7 @@ async def _put_to_single_storage_unit( ) raise RuntimeError(f"Error in put to storage unit {target_storage_unit}: {type(e).__name__}: {e}") from e - @staticmethod - def _pack_field_values(values: list) -> torch.Tensor | NonTensorStack: - """ - Pack a list of per-sample values into a batched container. - - For pure tensor lists (no None), this tries nested tensor - (jagged layout first, then strided fallback), then falls back to - ``NonTensorStack``. Scalar tensors are stacked densely. - Mixed types, non-tensor values, or lists containing None placeholders - are grouped into a ``NonTensorStack``. - - Args: - values: List of per-sample values to pack. May contain None for - unfilled batch positions. - - Returns: - A ``torch.Tensor`` (nested or dense) when all values are tensors, - otherwise a ``NonTensorStack``. - - Raises: - ValueError: If *values* is empty. - """ - if not values: - raise ValueError("_pack_field_values received empty values list; caller should filter empty batches") - non_none = [v for v in values if v is not None] - if non_none and all(isinstance(v, torch.Tensor) for v in non_none): - if len(non_none) == len(values): - # Scalar tensors cannot be represented as jagged nested tensors; - # stack them densely to avoid noisy fallback warnings. - if all(v.dim() == 0 for v in non_none): - return torch.stack(non_none) - # Pure tensor list — try nested tensor - try: - return torch.nested.as_nested_tensor(values, layout=torch.jagged) - except (RuntimeError, TypeError) as e: - logger.warning( - f"Failed to pack nested tensor with jagged layout. " - f"Falling back to strided layout. Detailed error: {e}" - ) - try: - return torch.nested.as_nested_tensor(values, layout=torch.strided) - except (RuntimeError, TypeError) as e2: - logger.warning( - f"Failed to pack nested tensor with strided layout. " - f"Falling back to NonTensorStack. Detailed error: {e2}" - ) - return NonTensorStack(*values) - # Mixed tensor + None — cannot create nested tensor, fall through to NonTensorStack - return NonTensorStack(*values) + _pack_field_values = staticmethod(pack_field_values) async def get_data(self, metadata: BatchMeta) -> TensorDict: """ @@ -821,15 +773,16 @@ async def _dump_single_shard( path: str, target_storage_unit: str, global_indexes: list[int], + fields_by_index: dict[int, list[str]] | None = None, socket: zmq.Socket = None, - ) -> int: + ) -> dict[int, list[int]]: """Ask one storage unit to write the rows it owns into a shard file.""" try: request_msg = ZMQMessage.create( request_type=ZMQRequestType.DUMP_ROWS, # type: ignore[arg-type] sender_id=self.storage_manager_id, receiver_id=target_storage_unit, - body={"path": path, "global_indexes": global_indexes}, + body={"path": path, "global_indexes": global_indexes, "fields_by_index": fields_by_index}, ) await socket.send_multipart(request_msg.serialize(), copy=False) messages = await socket.recv_multipart(copy=False) @@ -846,13 +799,18 @@ async def _dump_single_shard( raise RuntimeError( f"Storage unit {target_storage_unit} holds no data for requested rows: {missing_rows[:20]}" ) - return response_msg.body["dumped_rows"] + return response_msg.body["row_offsets"] except Exception as e: raise RuntimeError( f"[{self.storage_manager_id}]: Error dumping shard from storage unit {target_storage_unit}: {str(e)}" ) from e - async def dump_rows_by_index(self, shard_dir: str, global_indexes: list[int]) -> list[dict[str, Any]]: + async def dump_rows_by_index( + self, + shard_dir: str, + global_indexes: list[int], + fields_by_index: dict[int, list[str]] | None = None, + ) -> list[dict[str, Any]]: """Dump the given rows into one shard per storage unit, in parallel. Each unit pickles its own rows in its own process, so the payload never passes @@ -862,9 +820,10 @@ async def dump_rows_by_index(self, shard_dir: str, global_indexes: list[int]) -> Args: shard_dir: Directory to write shard files into. global_indexes: Global indexes to dump. + fields_by_index: Produced fields to persist; omitted for a raw storage dump. Returns: - One entry per written shard: ``{"position", "storage_unit_id", "rows"}``. + One entry per written shard: ``{"position", "storage_unit_id", "rows", "row_offsets"}``. Raises: RuntimeError: A unit holds no data for a row it was asked to dump. @@ -876,65 +835,79 @@ async def dump_rows_by_index(self, shard_dir: str, global_indexes: list[int]) -> targets = [(su_id, group.global_indexes) for su_id, group in routing.items()] paths = [str(shard_dir_path / f"shard_{pos}_{su_id}.pkl") for pos, (su_id, _) in enumerate(targets)] - dumped_rows = await asyncio.gather( + row_offsets = await asyncio.gather( *( - self._dump_single_shard(path, target_storage_unit=su_id, global_indexes=indexes) + self._dump_single_shard( + path, + target_storage_unit=su_id, + global_indexes=indexes, + fields_by_index={index: fields_by_index[index] for index in indexes} if fields_by_index else None, + ) for path, (su_id, indexes) in zip(paths, targets, strict=True) - ) + ), + return_exceptions=True, ) + for result in row_offsets: + if isinstance(result, BaseException): + raise result logger.info( - f"[{self.storage_manager_id}]: dumped {sum(dumped_rows)} rows " + f"[{self.storage_manager_id}]: dumped {sum(len(offsets) for offsets in row_offsets)} rows " f"across {len(targets)} shards to {shard_dir_path}" ) return [ - {"position": pos, "storage_unit_id": su_id, "rows": rows} - for pos, ((su_id, _), rows) in enumerate(zip(targets, dumped_rows, strict=True)) + {"position": pos, "storage_unit_id": su_id, "rows": len(offsets), "row_offsets": offsets} + for pos, ((su_id, _), offsets) in enumerate(zip(targets, row_offsets, strict=True)) ] - @staticmethod - def read_shard(path: str, fields_by_index: dict[int, list[str]]) -> dict[tuple[str, ...], dict[str, Any]]: - """Read one shard and regroup it into batches ready for a put. - - Rows are grouped by field signature because a single put accepts one - homogeneous TensorDict, and a selective dump routinely mixes rows that - finished different fields. - - Args: - path: Shard file written by ``dump_rows_by_index``. - fields_by_index: Field signature expected for each global index in the shard. - - Returns: - ``{field_signature: {"global_indexes": [...], "fields": TensorDict}}``. - - Raises: - ValueError: The shard is missing a field that the row index expects. - """ - with open(path, "rb") as f: - shard = pickle.load(f) - - field_data = shard["field_data"] - grouped: dict[tuple[str, ...], list[int]] = defaultdict(list) - for global_index in shard["global_indexes"]: - grouped[tuple(fields_by_index[global_index])].append(global_index) - - batches: dict[tuple[str, ...], dict[str, Any]] = {} - for signature, global_indexes in grouped.items(): - packed = {} - for field_name in signature: - values = field_data.get(field_name) - if values is None: - raise ValueError(f"shard {path} is missing field {field_name!r} required by the row index") - # Reuse the same packer the production get path uses, so a dumped row - # rebuilds into exactly the container it was read as while live. - packed[field_name] = AsyncSimpleStorageManager._pack_field_values( - [values[global_index] for global_index in global_indexes] + async def load_rows_by_index(self, shards: list[dict[str, Any]]) -> list[dict[str, Any]]: + """Have current owner units read assigned byte ranges concurrently.""" + assignments = defaultdict(list) + for shard in shards: + rows = shard["records"] + for unit_id, group in self._group_by_hash([row["target_index"] for row in rows]).items(): + assignments[unit_id].append( + { + "path": shard["path"], + "records": [rows[pos] for pos in group.batch_positions], + } ) - batches[signature] = { - "global_indexes": global_indexes, - "fields": TensorDict(packed, batch_size=len(global_indexes)) if signature else None, - } - return batches + results = await asyncio.gather( + *(self._load_selected_rows(shards, target_storage_unit=unit_id) for unit_id, shards in assignments.items()), + return_exceptions=True, + ) + # Wait for every unit before returning an error; callers may retry or clean up. + for result in results: + if isinstance(result, BaseException): + raise result + logger.info( + "[%s]: loaded %s bytes across %s units", + self.storage_manager_id, + sum(result["bytes_read"] for result in results), + len(assignments), + ) + return [update for result in results for update in result["updates"]] + + @with_storage_unit_socket + async def _load_selected_rows( + self, + shards: list[dict[str, Any]], + target_storage_unit: str, + socket: zmq.Socket = None, + ) -> dict[str, Any]: + request = ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS, + sender_id=self.storage_manager_id, + receiver_id=target_storage_unit, + body={"shards": shards}, + ) + await socket.send_multipart(request.serialize(), copy=False) + response = ZMQMessage.deserialize(await socket.recv_multipart(copy=False)) + if response.request_type != ZMQRequestType.LOAD_ROWS_RESPONSE or not response.body.get("success"): + raise RuntimeError( + f"Storage unit {target_storage_unit} failed to load rows: {response.body.get('message')}" + ) + return response.body async def save_checkpoint(self, checkpoint_dir: str) -> None: """Dump all storage units to the storage_units/ subdirectory of checkpoint_dir. diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index 79bcb7fc..0ba9a27c 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -34,6 +34,7 @@ import shutil import time import weakref +from collections import defaultdict from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass from pathlib import Path @@ -46,12 +47,16 @@ import ray import torch import zmq +from tensordict import TensorDict +from transfer_queue.metadata import extract_field_schema +from transfer_queue.storage.dump_io import read_dump_row from transfer_queue.utils import compact_pickle from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads, log_heavy_operation from transfer_queue.utils.enum_utils import Role from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.perf_utils import IntervalPerfMonitor +from transfer_queue.utils.tensor_utils import pack_field_values from transfer_queue.utils.zmq_utils import ( STORAGE_CLIENT_IDENTITY_PREFIXES, ZMQMessage, @@ -1005,7 +1010,11 @@ def _process_one_worker_request(self, worker_socket: zmq.Socket, monitor: Any) - elif operation == ZMQRequestType.SAVE_STORAGE_CHECKPOINT: # type: ignore[arg-type] response_msg = self._handle_save_checkpoint(request_msg) elif operation == ZMQRequestType.DUMP_ROWS: # type: ignore[arg-type] - response_msg = self._handle_dump_rows(request_msg) + with monitor.measure(op_type="DUMP_ROWS"): + response_msg = self._handle_dump_rows(request_msg) + elif operation == ZMQRequestType.LOAD_ROWS: # type: ignore[arg-type] + with monitor.measure(op_type="LOAD_ROWS"): + response_msg = self._handle_load_rows(request_msg) elif operation == ZMQRequestType.LOAD_STORAGE_CHECKPOINT: # type: ignore[arg-type] response_msg = self._handle_load_checkpoint(request_msg) else: @@ -1241,7 +1250,7 @@ def _handle_get_metrics(self) -> ZMQMessage: # Include per-operation stats if Prometheus metrics are enabled if self._metrics is not None: op_stats = {} - for op_type in ("PUT_DATA", "GET_DATA", "CLEAR_DATA"): + for op_type in ("PUT_DATA", "GET_DATA", "CLEAR_DATA", "DUMP_ROWS", "LOAD_ROWS"): try: hist = self._metrics.request_duration.labels(op_type=op_type) counter = self._metrics.request_total.labels(op_type=op_type) @@ -1296,63 +1305,91 @@ def _handle_save_checkpoint(self, data_parts) -> ZMQMessage: body={"success": False, "message": str(e)}, ) - def _handle_dump_rows(self, data_parts) -> ZMQMessage: - """Serialize the requested rows of this unit into a self-contained shard. - - This runs inside the storage unit process, so the payload is pickled where it - already lives instead of being shipped to the caller first. The shard is keyed - by global index; the caller holds the row index that maps keys onto them. - - Args: - data_parts: ZMQMessage with ``path`` and ``global_indexes`` in body. - ``path`` must be reachable from the node running this actor, which - means a shared filesystem in a multi-node deployment. - ``global_indexes`` is the subset this unit owns, already routed by - the storage manager. - - Returns: - ZMQMessage with ``success=True``, ``dumped_rows`` and ``missing_rows`` on - success, or ``success=False`` and ``message`` on failure. ``missing_rows`` - lists requested rows this unit holds no data for. - """ - path = data_parts.body["path"] - requested_indexes = set(data_parts.body["global_indexes"]) + def _handle_dump_rows(self, request: ZMQMessage) -> ZMQMessage: + """Write independent row records and return their offsets, never their payloads.""" + path = request.body["path"] + indexes = set(request.body["global_indexes"]) try: - field_data = {} - for field_name, values in self.storage_data.field_data.items(): - selected_values = { - global_index: values[global_index] for global_index in requested_indexes if global_index in values - } - if selected_values: - field_data[field_name] = selected_values - dumped_indexes = self.storage_data._active_keys & requested_indexes - shard = { - "storage_unit_id": self.storage_unit_id, - "field_data": field_data, - "global_indexes": sorted(dumped_indexes), - } + missing = indexes - self.storage_data._active_keys + if missing: + raise ValueError(f"Storage holds no data for requested rows: {sorted(missing)[:20]}") + row_offsets = {} with open(path, "wb") as f: - compact_pickle.dump(shard, f) - # Report success only once the shard is on disk. Without this the call - # returns while the payload is still dirty page cache, and a node that - # dies before writeback leaves a dump whose manifest claims rows that - # cannot be read back. + for index in sorted(indexes): + described_fields = request.body.get("fields_by_index") + if described_fields is None: + fields = { + name: values[index] + for name, values in self.storage_data.field_data.items() + if index in values + } + else: + # Reused global indexes can retain fields no longer present in metadata. + fields = {name: self.storage_data.field_data[name][index] for name in described_fields[index]} + offset = f.tell() + compact_pickle.dump({"global_index": index, "fields": fields}, f) + row_offsets[index] = [offset, f.tell() - offset] f.flush() os.fsync(f.fileno()) - logger.info(f"[{self.storage_unit_id}]: dumped {len(dumped_indexes)} rows to {path}") + logger.info("[%s]: dumped %s rows to %s", self.storage_unit_id, len(indexes), path) return ZMQMessage.create( - request_type=ZMQRequestType.DUMP_ROWS_RESPONSE, # type: ignore[arg-type] + request_type=ZMQRequestType.DUMP_ROWS_RESPONSE, sender_id=self.storage_unit_id, - body={ - "success": True, - "dumped_rows": len(dumped_indexes), - "missing_rows": sorted(requested_indexes - self.storage_data._active_keys), - }, + body={"success": True, "dumped_rows": len(indexes), "missing_rows": [], "row_offsets": row_offsets}, + ) + except Exception as e: + logger.error("[%s]: dump rows failed: %s", self.storage_unit_id, e) + return ZMQMessage.create( + request_type=ZMQRequestType.DUMP_ROWS_RESPONSE, + sender_id=self.storage_unit_id, + body={"success": False, "message": str(e)}, + ) + + def _handle_load_rows(self, request: ZMQMessage) -> ZMQMessage: + """Seek to assigned records and merge them into local storage at current indexes.""" + updates = [] + bytes_read = 0 + try: + with limit_pytorch_auto_parallel_threads(TQ_NUM_THREADS): + for shard in request.body["shards"]: + records = sorted(shard["records"], key=lambda row: row["offset"]) + with open(shard["path"], "rb", buffering=0) as f: + # Bound temporary payload memory while retaining batched schema updates. + for start in range(0, len(records), 128): + groups = defaultdict(list) + for row in records[start : start + 128]: + fields = read_dump_row( + f, row["offset"], row["length"], row["source_index"], row["fields"] + ) + groups[tuple(row["fields"])].append((row["target_index"], fields)) + bytes_read += row["length"] + for signature, rows in groups.items(): + indexes = [index for index, _ in rows] + values = {name: [fields[name] for _, fields in rows] for name in signature} + packed = {name: pack_field_values(items) for name, items in values.items()} + schema = extract_field_schema(TensorDict(packed, batch_size=len(rows))) + for field in schema.values(): + if "per_sample_shapes" in field: + field["per_sample_shapes"] = dict( + zip(indexes, field["per_sample_shapes"], strict=True) + ) + self.storage_data.put_data(values, indexes) + updates.append({"global_indexes": indexes, "field_schema": schema}) + logger.info( + "[%s]: loaded %s rows (%s bytes)", + self.storage_unit_id, + sum(len(update["global_indexes"]) for update in updates), + bytes_read, + ) + return ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, + sender_id=self.storage_unit_id, + body={"success": True, "updates": updates, "bytes_read": bytes_read}, ) except Exception as e: - logger.error(f"[{self.storage_unit_id}]: dump rows failed: {e}") + logger.error("[%s]: load rows failed: %s", self.storage_unit_id, e) return ZMQMessage.create( - request_type=ZMQRequestType.DUMP_ROWS_RESPONSE, # type: ignore[arg-type] + request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, sender_id=self.storage_unit_id, body={"success": False, "message": str(e)}, ) diff --git a/transfer_queue/utils/compact_pickle.py b/transfer_queue/utils/compact_pickle.py index f73d3c0d..f83a5d60 100644 --- a/transfer_queue/utils/compact_pickle.py +++ b/transfer_queue/utils/compact_pickle.py @@ -21,7 +21,10 @@ class Pickler(pickle.Pickler): + """Clone tensor leaves during pickling, including those inside arbitrary containers.""" + def reducer_override(self, value): + """Serialize tensor values independently of their source storage.""" # Tensor views can retain an entire batch, including unselected keys. This # also handles tensors nested in picklable payload objects or metadata. if isinstance(value, torch.Tensor): @@ -30,4 +33,5 @@ def reducer_override(self, value): def dump(value, file): + """Write a standard pickle containing only the selected tensor values.""" Pickler(file, protocol=pickle.HIGHEST_PROTOCOL).dump(value) diff --git a/transfer_queue/utils/tensor_utils.py b/transfer_queue/utils/tensor_utils.py index b3b8fa06..a7a3ed7c 100644 --- a/transfer_queue/utils/tensor_utils.py +++ b/transfer_queue/utils/tensor_utils.py @@ -19,6 +19,7 @@ from functools import reduce import torch +from tensordict import NonTensorStack from torch import Tensor logger = logging.getLogger(__name__) @@ -180,3 +181,53 @@ def merge_contiguous_memory(ptrs: list[int], sizes: list[int]) -> tuple[list[int merged_sizes.append(current_size) return merged_ptrs, merged_sizes + + +def pack_field_values(values: list) -> torch.Tensor | NonTensorStack: + """ + Pack a list of per-sample values into a batched container. + + For pure tensor lists (no None), this tries nested tensor + (jagged layout first, then strided fallback), then falls back to + ``NonTensorStack``. Scalar tensors are stacked densely. + Mixed types, non-tensor values, or lists containing None placeholders + are grouped into a ``NonTensorStack``. + + Args: + values: List of per-sample values to pack. May contain None for + unfilled batch positions. + + Returns: + A ``torch.Tensor`` (nested or dense) when all values are tensors, + otherwise a ``NonTensorStack``. + + Raises: + ValueError: If *values* is empty. + """ + if not values: + raise ValueError("_pack_field_values received empty values list; caller should filter empty batches") + non_none = [v for v in values if v is not None] + if non_none and all(isinstance(v, torch.Tensor) for v in non_none): + if len(non_none) == len(values): + # Scalar tensors cannot be represented as jagged nested tensors; + # stack them densely to avoid noisy fallback warnings. + if all(v.dim() == 0 for v in non_none): + return torch.stack(non_none) + # Pure tensor list — try nested tensor + try: + return torch.nested.as_nested_tensor(values, layout=torch.jagged) + except (RuntimeError, TypeError) as e: + logger.warning( + f"Failed to pack nested tensor with jagged layout. " + f"Falling back to strided layout. Detailed error: {e}" + ) + try: + return torch.nested.as_nested_tensor(values, layout=torch.strided) + except (RuntimeError, TypeError) as e2: + logger.warning( + f"Failed to pack nested tensor with strided layout. " + f"Falling back to NonTensorStack. Detailed error: {e2}" + ) + return NonTensorStack(*values) + # Mixed tensor + None — cannot create nested tensor, fall through to NonTensorStack + return NonTensorStack(*values) diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index 41eebc4a..6bdb2cb9 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -140,6 +140,8 @@ class ZMQRequestType(ExplicitEnum): DESCRIBE_ROWS_BY_KEY_RESPONSE = "DESCRIBE_ROWS_BY_KEY_RESPONSE" DUMP_ROWS = "DUMP_ROWS" DUMP_ROWS_RESPONSE = "DUMP_ROWS_RESPONSE" + LOAD_ROWS = "LOAD_ROWS" + LOAD_ROWS_RESPONSE = "LOAD_ROWS_RESPONSE" class ZMQServerInfo: From 63b7cfd7f03c6087abbcc9db774f41c3e83ee3dd Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 24 Sep 2026 15:24:32 +0800 Subject: [PATCH 06/30] fix: preserve original schemas in selective dump restores Signed-off-by: OutstanderWang --- docs/data_dump.md | 15 ++- .../e2e/test_data_dump_cross_topology_e2e.py | 31 +++++- tests/e2e/test_data_dump_e2e.py | 10 ++ tests/test_data_dump.py | 35 ++++++- transfer_queue/client.py | 98 ++++++++++++------- transfer_queue/controller.py | 60 +++++++++++- transfer_queue/data_dump.py | 47 ++++++--- transfer_queue/storage/dump_io.py | 44 +++++++++ .../managers/simple_storage_manager.py | 2 +- transfer_queue/storage/simple_storage.py | 30 ++++-- transfer_queue/utils/zmq_utils.py | 2 + 11 files changed, 303 insertions(+), 71 deletions(-) diff --git a/docs/data_dump.md b/docs/data_dump.md index 0ae58b1f..03833fe3 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -25,7 +25,7 @@ owner and concurrently asks those units to write their records. Only units holdi selected rows participate. Tensor storage is compacted during serialization, so a row view cannot include the rest of its original batch, including inside tags. -On version-2 SimpleStorage restore: +On version-3 SimpleStorage restore: 1. The caller reads the row index and shard manifest and validates all file ranges. 2. The controller resolves existing keys and allocates indexes for new keys. @@ -55,7 +55,7 @@ These count application reads, not filesystem read-ahead or physical disk traffi ## Format and compatibility -New dumps use `format_version: 2`: +New dumps use `format_version: 3`: ```text dump_info.json @@ -72,10 +72,17 @@ index and a field/value mapping. `shard_info.json` records each source index's current indexes without controller resolution. `row_index.pt` remains readable with `read_row_index` without opening payload shards. -Version-1 dumps remain readable using the prior caller-side KV put path. Restoring +Version 3 also saves the original field schemas and only the selected nested row +shapes. Restore uses that schema regardless of target topology or batch boundaries; +non-tensor fields remain non-tensor even when a batch happens to contain only tensors. +Destination type conflicts are rejected before payload writes. + +Version-1 dumps remain readable using the prior caller-side KV put path. Version-2 +dumps retain direct reads, but lack original schemas and use the older inference +behavior; exact field-type preservation cannot be guaranteed for those files. Restoring to a backend without direct selective loading also uses KV puts. These compatibility paths do not provide distributed file reads. Old builds that only understand -version 1 cannot read version-2 dumps. Export of nonempty dumps currently requires +versions 1 or 2 cannot read version-3 dumps. Export of nonempty dumps currently requires SimpleStorage. ## Failure behavior diff --git a/tests/e2e/test_data_dump_cross_topology_e2e.py b/tests/e2e/test_data_dump_cross_topology_e2e.py index ddfce9c5..1fd0d0ee 100644 --- a/tests/e2e/test_data_dump_cross_topology_e2e.py +++ b/tests/e2e/test_data_dump_cross_topology_e2e.py @@ -33,7 +33,7 @@ import ray import torch from omegaconf import OmegaConf -from tensordict import TensorDict +from tensordict import NonTensorStack, TensorDict import transfer_queue as tq @@ -131,3 +131,32 @@ def test_dump_restores_across_storage_unit_counts(ray_init, dump_dir, dump_units assert snapshot.custom_meta[snapshot.keys_mapping[key]]["idx"] == row finally: tq.close() + + +@pytest.mark.parametrize("row_count", [2, 127, 128, 129, 130]) +def test_preserves_nontensor_schema_across_topology(ray_init, dump_dir, row_count): + keys = [f"k{i}" for i in range(row_count)] + values = [torch.tensor([i], dtype=torch.int64) for i in range(row_count - 1)] + [torch.tensor([1.5])] + if row_count > 2: + values[-2] = None + tq.init(_tq_config(1)) + try: + tq.kv_batch_put(keys, "mixed", TensorDict({"x": NonTensorStack(*values)}, batch_size=row_count)) + tq.dump_data_by_key(dump_dir, keys, "mixed") + finally: + tq.close() + tq.init(_tq_config(2)) + try: + tq.load_data_by_key(dump_dir) + schema = tq.get_client().kv_retrieve_meta(keys, "mixed").field_schema["x"] + assert schema["is_non_tensor"] + assert schema["dtype"] is None + assert not schema["is_nested"] + for key, expected in zip(keys, values, strict=True): + value = tq.kv_batch_get([key], "mixed", ["x"])["x"][0] + if expected is None: + assert value is None + else: + torch.testing.assert_close(value, expected) + finally: + tq.close() diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index 319dbf38..663b3396 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -504,3 +504,13 @@ def test_corrupt_shard_keeps_new_rows_unproduced(tq_system, dump_dir, controller snapshot = ray.get(controller.get_partition_snapshot.remote(partition)) assert "key" in snapshot.keys_mapping assert not snapshot.field_metadata + + +def test_incompatible_schema_rejected_before_writes(tq_system, dump_dir, controller): + tq.kv_batch_put(["k"], "schema", TensorDict({"x": torch.tensor([[3]], dtype=torch.int64)}, batch_size=1)) + tq.dump_data_by_key(dump_dir, ["k"], "schema") + tq.get_client().clear_partition("schema") + tq.kv_batch_put(["k"], "schema", TensorDict({"x": torch.tensor([[1.5]])}, batch_size=1)) + with pytest.raises(RuntimeError, match="dtype mismatch"): + tq.load_data_by_key(dump_dir) + _assert_rows_equal(tq.kv_batch_get(["k"], "schema", ["x"])["x"], [torch.tensor([1.5])]) diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index 82f5426f..9f88b8d3 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -73,8 +73,10 @@ def test_row_index_compacts_tensors_inside_tags(monkeypatch, tmp_path): batch = torch.arange(64 * 4096).reshape(64, 4096) tag = {"nested": [SimpleNamespace(value=batch[0])]} client = SimpleNamespace( - describe_rows_by_key=lambda *_: { - "key": {"global_index": 0, "fields": [], "tag": tag}, + describe_data_dump=lambda *_: { + "partition_id": "p", + "rows": {"key": {"global_index": 0, "fields": [], "tag": tag}}, + "field_schema": {}, } ) monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) @@ -88,7 +90,9 @@ def test_row_index_compacts_tensors_inside_tags(monkeypatch, tmp_path): @pytest.fixture def empty_dump_client(monkeypatch): monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) - monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: object()) + monkeypatch.setattr( + interface, "_maybe_create_tq_client", lambda: SimpleNamespace(validate_dump_schema=lambda *_: None) + ) def test_failed_publication_preserves_previous_dump(empty_dump_client, monkeypatch, tmp_path): @@ -301,7 +305,18 @@ def dump(shard_dir, indexes, fields_by_index): } ] - client = SimpleNamespace(describe_rows_by_key=lambda *_: rows, dump_rows_by_index=dump, storage_manager=object()) + client = SimpleNamespace( + describe_data_dump=lambda *_: { + "partition_id": "p", + "rows": rows, + "field_schema": { + "x": {"dtype": torch.int64, "shape": (2,), "is_nested": False, "is_non_tensor": False}, + }, + }, + validate_dump_schema=lambda *_: None, + dump_rows_by_index=dump, + storage_manager=object(), + ) monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: client) data_dump.dump_data_by_key(tmp_path / "dump", list(rows), "p") @@ -375,3 +390,15 @@ async def dump(path, target_storage_unit, global_indexes, fields_by_index): with pytest.raises(OSError, match="write failed"): await manager.dump_rows_by_index(str(tmp_path), [0, 1]) assert completed == ["u1"] + + +def test_rejected_schema_does_not_mark_new_row_ready(): + from transfer_queue.controller import DataPartitionStatus + + partition = DataPartitionStatus("p") + schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_non_tensor": False, "is_nested": False}} + assert partition.update_production_status([0], ["x"], schema) + conflict = {"x": {**schema["x"], "dtype": torch.float32}} + assert not partition.update_production_status([1], ["x"], conflict) + assert partition.production_status[1, partition.field_name_mapping["x"]] == 0 + assert partition.field_metadata["x"].global_indexes == {0} diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 11967b44..d6b96297 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1124,48 +1124,68 @@ def _can_destroy_zmq_context(self) -> bool: return False return True + async def _request_controller( + self, + socket: zmq.asyncio.Socket | None, + request_type: ZMQRequestType, + response_type: ZMQRequestType, + body: dict[str, Any], + ) -> ZMQMessage: + """Send one controller request and validate its response type.""" + assert socket is not None + request_msg = ZMQMessage.create( + request_type=request_type, # type: ignore[arg-type] + sender_id=self.client_id, + receiver_id=self._controller.id, + body=body, + ) + await socket.send_multipart(request_msg.serialize()) + response_serialized = await socket.recv_multipart(copy=False) + response_msg = ZMQMessage.deserialize(response_serialized) + logger.debug(f"[{self.client_id}]: Received {response_msg.request_type} from controller {self._controller.id}") + if response_msg.request_type != response_type: + message = response_msg.body.get("message", "Unknown error") + raise RuntimeError( + f"[{self.client_id}]: Expected {response_type}, got {response_msg.request_type} " + f"from controller {self._controller.id}: {message}" + ) + return response_msg + # ==================== Selective Data Dump API ==================== @with_controller_socket - async def async_describe_rows_by_key( + async def async_describe_data_dump( self, partition_id: str, keys: list[str], socket: zmq.asyncio.Socket | None = None, - ) -> dict[str, dict[str, Any]]: - """Asynchronously fetch the row metadata a selective dump needs, via ZMQ RPC. - - Args: - partition_id: Partition that owns ``keys``. - keys: Keys to describe, already deduplicated by the caller. - socket: ZMQ socket injected by @with_controller_socket. + ) -> dict[str, Any]: + """Fetch selected rows and their original field schemas without payloads.""" + response = await self._request_controller( + socket=socket, + request_type=ZMQRequestType.DESCRIBE_ROWS_BY_KEY, + response_type=ZMQRequestType.DESCRIBE_ROWS_BY_KEY_RESPONSE, + body={"partition_id": partition_id, "keys": keys}, + ) + return {name: response.body[name] for name in ("partition_id", "rows", "field_schema")} - Returns: - ``{key: {"global_index": int, "fields": list[str], "tag": dict}}``. + async def async_describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, dict[str, Any]]: + """Fetch key-addressed row metadata, preserving the existing row-only API.""" + return (await self.async_describe_data_dump(partition_id, keys))["rows"] - Raises: - RuntimeError: If the RPC fails, or the partition or a key is unknown. - """ - try: - assert socket is not None - request_msg = ZMQMessage.create( - request_type=ZMQRequestType.DESCRIBE_ROWS_BY_KEY, # type: ignore[arg-type] - sender_id=self.client_id, - receiver_id=self._controller.id, - body={"partition_id": partition_id, "keys": keys}, - ) - await socket.send_multipart(request_msg.serialize()) - response_serialized = await socket.recv_multipart(copy=False) - response_msg = ZMQMessage.deserialize(response_serialized) - if response_msg.request_type != ZMQRequestType.DESCRIBE_ROWS_BY_KEY_RESPONSE: - raise RuntimeError( - f"[{self.client_id}]: Unexpected response type {response_msg.request_type} " - f"from controller during row description" - ) - if not response_msg.body["success"]: - raise RuntimeError(response_msg.body["message"]) - return response_msg.body["rows"] - except Exception as e: - raise RuntimeError(f"[{self.client_id}]: Error in describe_rows_by_key: {str(e)}") from e + @with_controller_socket + async def async_validate_dump_schema( + self, + partition_id: str, + field_schema: dict, + socket: zmq.asyncio.Socket | None = None, + ) -> None: + """Reject incompatible destination fields before restoring payloads.""" + await self._request_controller( + socket=socket, + request_type=ZMQRequestType.VALIDATE_DUMP_SCHEMA, + response_type=ZMQRequestType.VALIDATE_DUMP_SCHEMA_RESPONSE, + body={"partition_id": partition_id, "field_schema": field_schema}, + ) async def async_dump_rows_by_index( self, @@ -1425,6 +1445,8 @@ def wrapper(*args, **kwargs): self._kv_retrieve_keys = _make_sync(self.async_kv_retrieve_keys) self._kv_list = _make_sync(self.async_kv_list) self._describe_rows_by_key = _make_sync(self.async_describe_rows_by_key) + self._describe_data_dump = _make_sync(self.async_describe_data_dump) + self._validate_dump_schema = _make_sync(self.async_validate_dump_schema) self._dump_rows_by_index = _make_sync(self.async_dump_rows_by_index) self._load_rows_by_key = _make_sync(self.async_load_rows_by_key) self._save_controller_checkpoint = _make_sync(self.async_save_controller_checkpoint) @@ -1879,6 +1901,14 @@ def describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, """ return self._describe_rows_by_key(partition_id, keys) + def describe_data_dump(self, partition_id: str, keys: list[str]) -> dict[str, Any]: + """Fetch the row index and selected field schemas for a selective dump.""" + return self._describe_data_dump(partition_id, keys) + + def validate_dump_schema(self, partition_id: str, field_schema: dict) -> None: + """Reject incompatible destination fields before restoring payloads.""" + return self._validate_dump_schema(partition_id, field_schema) + def dump_rows_by_index( self, shard_dir: str, diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index 612942fd..9adb43d8 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -540,14 +540,13 @@ def update_production_status( required_fields = len(self.field_name_mapping) self.ensure_fields_capacity(required_fields) - # Update production status + # Validate all field updates before changing readiness or field metadata. + self.validate_field_schema(field_schema) + self._update_field_metadata(global_indices, field_schema, custom_backend_meta) if self.production_status is not None and global_indices and field_names: field_indices = [self.field_name_mapping.get(f) for f in field_names] self.production_status[torch.tensor(global_indices)[:, None], torch.tensor(field_indices)] = 1 - # Update field metadata - self._update_field_metadata(global_indices, field_schema, custom_backend_meta) - # Save these global_indexes self.global_indexes.update(global_indices) @@ -557,6 +556,18 @@ def update_production_status( logger.error(f"Error updating production status for partition {self.partition_id}: {e}") return False + def validate_field_schema(self, field_schema: dict[str, dict[str, Any]]) -> None: + """Reject incompatible field types without changing metadata or readiness.""" + for name, incoming in field_schema.items(): + existing = self.field_metadata.get(name) + if existing is None: + continue + if bool(existing.is_non_tensor) != bool(incoming.get("is_non_tensor", False)): + raise ValueError(f"Field {name!r} tensor/non-tensor type mismatch") + dtype = incoming.get("dtype") + if dtype is not None and existing.dtype is not None and dtype != existing.dtype: + raise ValueError(f"Field {name!r} dtype mismatch: {existing.dtype} != {dtype}") + def _update_field_metadata( self, global_indexes: list[int], @@ -1991,6 +2002,7 @@ def _handle_request(self, request_msg: ZMQMessage, monitor: Any) -> ZMQMessage | ZMQRequestType.KV_RETRIEVE_KEYS: self._handle_kv_retrieve_keys_request, ZMQRequestType.KV_LIST: self._handle_kv_list_request, ZMQRequestType.DESCRIBE_ROWS_BY_KEY: self._handle_describe_rows_by_key_request, + ZMQRequestType.VALIDATE_DUMP_SCHEMA: self._handle_validate_dump_schema_request, ZMQRequestType.SAVE_CONTROLLER_CHECKPOINT: self._handle_save_controller_checkpoint_request, ZMQRequestType.LOAD_CONTROLLER_CHECKPOINT: self._handle_load_controller_checkpoint_request, } @@ -2233,6 +2245,46 @@ def _handle_load_controller_checkpoint_request(self, request_msg: ZMQMessage) -> {"success": True}, ) + def _make_response( + self, + request_msg: ZMQMessage, + response_type: ZMQRequestType, + body: dict[str, Any], + ) -> ZMQMessage: + """Build a controller response addressed to the request sender.""" + return ZMQMessage.create( + request_type=response_type, + sender_id=self.controller_id, + receiver_id=request_msg.sender_id, + body=body, + ) + + def _handle_describe_rows_by_key_request(self, request_msg: ZMQMessage) -> ZMQMessage: + params = request_msg.body + rows = self.describe_rows_by_key(params["partition_id"], params["keys"]) + partition = self.partitions[params["partition_id"]] + field_schema = {} + for name, meta in partition.field_metadata.items(): + indexes = [row["global_index"] for row in rows.values() if name in row["fields"]] + if not indexes: + continue + schema = meta.to_batch_schema(indexes) + if schema.get("is_nested"): + schema["per_sample_shapes"] = {index: meta.per_sample_shapes[index] for index in indexes} + field_schema[name] = schema + return self._make_response( + request_msg, + ZMQRequestType.DESCRIBE_ROWS_BY_KEY_RESPONSE, + {"success": True, "partition_id": params["partition_id"], "rows": rows, "field_schema": field_schema}, + ) + + def _handle_validate_dump_schema_request(self, request_msg: ZMQMessage) -> ZMQMessage: + params = request_msg.body + partition = self._get_partition(params["partition_id"]) + if partition is not None: + partition.validate_field_schema(params["field_schema"]) + return self._make_response(request_msg, ZMQRequestType.VALIDATE_DUMP_SCHEMA_RESPONSE, {"success": True}) + def get_zmq_server_info(self) -> ZMQServerInfo: """Get ZMQ server connection information.""" return self.zmq_server_info diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index 57458dfd..7b1a74ed 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -15,7 +15,7 @@ """Persist selected rows and restore them by key without checkpointing controller state. -Version 2 stores independent row records in each storage-unit shard. The manifest +Version 3 preserves field schemas alongside independent row records in each storage-unit shard. The manifest maps source indexes to byte offsets, so current owner units can read only their rows when restoring into a different topology. Version 1 remains readable via KV puts. @@ -40,14 +40,14 @@ import torch from tensordict import TensorDict -from transfer_queue.storage.dump_io import read_dump_row +from transfer_queue.storage.dump_io import pack_dump_field, read_dump_row, validate_dump_values from transfer_queue.utils import compact_pickle from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.tensor_utils import pack_field_values logger = get_logger(__name__) -DUMP_FORMAT_VERSION = 2 +DUMP_FORMAT_VERSION = 3 _DUMP_INFO_FILE = "dump_info.json" _ROW_INDEX_FILE = "row_index.pt" @@ -128,7 +128,16 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - try: client = _maybe_create_tq_client() - rows = client.describe_rows_by_key(partition_id, unique_keys) if unique_keys else {} + row_index = ( + client.describe_data_dump(partition_id, unique_keys) + if unique_keys + else { + "partition_id": partition_id, + "rows": {}, + "field_schema": {}, + } + ) + rows = row_index["rows"] # A row whose fields are all still unproduced has nothing for a storage unit to # dump, but it keeps its key and tag so the restore can recreate the row. @@ -155,7 +164,7 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - # torch.save rather than json: a tag is an arbitrary picklable dict, and this # path must not fail on a tag that happens to hold a tensor. with open(tmp_dir / _ROW_INDEX_FILE, "wb") as f: - torch.save({"partition_id": partition_id, "rows": rows}, f, pickle_module=compact_pickle) + torch.save(row_index, f, pickle_module=compact_pickle) _fsync_file(f) with open(tmp_dir / _DUMP_INFO_FILE, "w", encoding="utf-8") as f: @@ -233,7 +242,7 @@ def read_row_index(dump_dir: str | Path) -> dict[str, Any]: def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: """Merge selected rows into the running system, preserving existing key indexes. - SimpleStorage units read their assigned version-2 records directly and in parallel. + SimpleStorage units read their assigned indexed records directly and in parallel. Version-1 dumps and other backends use the compatible KV put path. New keys receive current indexes; unrelated rows and fields remain untouched. Writers and clears for these keys must be paused during restore. Failure may leave partial payload writes; @@ -263,10 +272,10 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: raise FileNotFoundError(f"{_DUMP_INFO_FILE} not found in {dump_dir}") with open(info_path, encoding="utf-8") as f: dump_info = json.load(f) - if dump_info["format_version"] not in (1, DUMP_FORMAT_VERSION): + if dump_info["format_version"] not in (1, 2, DUMP_FORMAT_VERSION): raise ValueError( f"Unsupported dump format version {dump_info['format_version']} in {dump_dir}; " - f"this build reads versions 1 and {DUMP_FORMAT_VERSION}" + f"this build reads versions 1 through {DUMP_FORMAT_VERSION}" ) row_index = read_row_index(dump_dir) @@ -292,7 +301,7 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: if not path.is_file(): raise FileNotFoundError(f"Missing dump shard: {path}") records = [] - if dump_info["format_version"] == 2: + if dump_info["format_version"] >= 2: size = path.stat().st_size offsets = record["row_offsets"] if len(offsets) != record["rows"]: @@ -314,12 +323,17 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: } ) seen.add(source_index) - shards.append({"path": str(path), "records": records}) - if dump_info["format_version"] == 2 and seen != set(keys_by_index): + shard = {"path": str(path), "records": records} + if dump_info["format_version"] >= 3: + shard["field_schema"] = row_index["field_schema"] + shards.append(shard) + if dump_info["format_version"] >= 2 and seen != set(keys_by_index): raise ValueError("Dump shards do not contain every produced row") client = _maybe_create_tq_client() - if dump_info["format_version"] == 2 and hasattr(getattr(client, "storage_manager", None), "load_rows_by_index"): + if dump_info["format_version"] >= 3: + client.validate_dump_schema(partition_id, row_index["field_schema"]) + if dump_info["format_version"] >= 2 and hasattr(getattr(client, "storage_manager", None), "load_rows_by_index"): client.load_rows_by_key(partition_id, rows, shards) else: _load_via_kv(partition_id, rows, shards, dump_info["format_version"]) @@ -359,11 +373,18 @@ def _load_via_kv(partition_id: str, rows: dict[str, Any], shards: list[dict], ve values = {name: saved["field_data"][name][index] for name in fields} else: values = read_dump_row(f, record["offset"], record["length"], index, fields) + if version >= 3: + validate_dump_values(values, shard["field_schema"], index) groups[tuple(fields)].append((record["key"], values)) restored.add(index) for signature, batch in groups.items(): keys = [key for key, _ in batch] - packed = {name: pack_field_values([values[name] for _, values in batch]) for name in signature} + packed = { + name: pack_dump_field([values[name] for _, values in batch], shard["field_schema"][name]) + if version >= 3 + else pack_field_values([values[name] for _, values in batch]) + for name in signature + } kv_batch_put( keys, partition_id, diff --git a/transfer_queue/storage/dump_io.py b/transfer_queue/storage/dump_io.py index 3bdcf7bf..0a202670 100644 --- a/transfer_queue/storage/dump_io.py +++ b/transfer_queue/storage/dump_io.py @@ -17,6 +17,9 @@ import pickle +import torch +from tensordict import NonTensorStack + def read_dump_row(file, offset: int, length: int, global_index: int, fields: list[str]) -> dict: """Read and validate exactly one record against its expected index and fields.""" @@ -28,3 +31,44 @@ def read_dump_row(file, offset: int, length: int, global_index: int, fields: lis if row["global_index"] != global_index or set(row["fields"]) != set(fields): raise ValueError(f"Dump row {global_index} disagrees with the row index in {file.name}") return row["fields"] + + +def validate_dump_values(values: dict, schema: dict, source_index: int) -> None: + """Validate persisted tensor values without changing their type or dtype.""" + for name, value in values.items(): + field = schema[name] + if field["is_non_tensor"]: + continue + shape = field["per_sample_shapes"][source_index] if field["is_nested"] else field["shape"] + actual_shape = tuple(value.shape) if isinstance(value, torch.Tensor) else None + # Existing dense scalar fields use a one-element metadata shape. + scalar = actual_shape == () and tuple(shape) == (1,) and not field["is_nested"] + if ( + not isinstance(value, torch.Tensor) + or value.dtype != field["dtype"] + or (actual_shape != tuple(shape) and not scalar) + ): + raise ValueError(f"Dump field {name!r} disagrees with its saved schema at row {source_index}") + + +def select_dump_schema(schema: dict, source_indexes: list[int], target_indexes: list[int], names: tuple) -> dict: + """Remap only the selected nested shapes to the current destination indexes.""" + selected = {} + for name in names: + field = dict(schema[name]) + if field["is_nested"]: + field["per_sample_shapes"] = { + target: field["per_sample_shapes"][source] + for source, target in zip(source_indexes, target_indexes, strict=True) + } + selected[name] = field + return selected + + +def pack_dump_field(values: list, schema: dict): + """Build fallback KV batches according to the original field contract.""" + if schema["is_non_tensor"]: + return NonTensorStack(*values) + if schema["is_nested"]: + return torch.nested.as_nested_tensor(values, layout=torch.jagged) + return torch.stack(values) diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index f82f79b7..06a88303 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -868,7 +868,7 @@ async def load_rows_by_index(self, shards: list[dict[str, Any]]) -> list[dict[st for unit_id, group in self._group_by_hash([row["target_index"] for row in rows]).items(): assignments[unit_id].append( { - "path": shard["path"], + **shard, "records": [rows[pos] for pos in group.batch_positions], } ) diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index 0ba9a27c..393d959e 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -50,7 +50,7 @@ from tensordict import TensorDict from transfer_queue.metadata import extract_field_schema -from transfer_queue.storage.dump_io import read_dump_row +from transfer_queue.storage.dump_io import read_dump_row, select_dump_schema, validate_dump_values from transfer_queue.utils import compact_pickle from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads, log_heavy_operation from transfer_queue.utils.enum_utils import Role @@ -1361,18 +1361,28 @@ def _handle_load_rows(self, request: ZMQMessage) -> ZMQMessage: fields = read_dump_row( f, row["offset"], row["length"], row["source_index"], row["fields"] ) - groups[tuple(row["fields"])].append((row["target_index"], fields)) + if "field_schema" in shard: + validate_dump_values(fields, shard["field_schema"], row["source_index"]) + groups[tuple(row["fields"])].append((row, fields)) bytes_read += row["length"] for signature, rows in groups.items(): - indexes = [index for index, _ in rows] + indexes = [row["target_index"] for row, _ in rows] values = {name: [fields[name] for _, fields in rows] for name in signature} - packed = {name: pack_field_values(items) for name, items in values.items()} - schema = extract_field_schema(TensorDict(packed, batch_size=len(rows))) - for field in schema.values(): - if "per_sample_shapes" in field: - field["per_sample_shapes"] = dict( - zip(indexes, field["per_sample_shapes"], strict=True) - ) + if "field_schema" in shard: + schema = select_dump_schema( + shard["field_schema"], + [row["source_index"] for row, _ in rows], + indexes, + signature, + ) + else: + packed = {name: pack_field_values(items) for name, items in values.items()} + schema = extract_field_schema(TensorDict(packed, batch_size=len(rows))) + for field in schema.values(): + if "per_sample_shapes" in field: + field["per_sample_shapes"] = dict( + zip(indexes, field["per_sample_shapes"], strict=True) + ) self.storage_data.put_data(values, indexes) updates.append({"global_indexes": indexes, "field_schema": schema}) logger.info( diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index 6bdb2cb9..b32273e4 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -142,6 +142,8 @@ class ZMQRequestType(ExplicitEnum): DUMP_ROWS_RESPONSE = "DUMP_ROWS_RESPONSE" LOAD_ROWS = "LOAD_ROWS" LOAD_ROWS_RESPONSE = "LOAD_ROWS_RESPONSE" + VALIDATE_DUMP_SCHEMA = "VALIDATE_DUMP_SCHEMA" + VALIDATE_DUMP_SCHEMA_RESPONSE = "VALIDATE_DUMP_SCHEMA_RESPONSE" class ZMQServerInfo: From 579ed35a6b89111a2e7e9478e850a22100fbe2f7 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 24 Sep 2026 15:34:46 +0800 Subject: [PATCH 07/30] fix: reserve restore indexes until remote operations settle Signed-off-by: OutstanderWang --- docs/data_dump.md | 43 +- tests/e2e/test_data_dump_e2e.py | 28 ++ tests/e2e/test_restore_timeout_e2e.py | 109 +++++ tests/test_data_dump.py | 36 +- tests/test_restore_lifecycle.py | 198 ++++++++ transfer_queue/__init__.py | 5 +- transfer_queue/client.py | 117 +++-- transfer_queue/controller.py | 446 ++++++++++++------ transfer_queue/data_dump.py | 63 ++- transfer_queue/interface.py | 1 + transfer_queue/storage/dump_io.py | 10 + .../managers/simple_storage_manager.py | 33 +- transfer_queue/storage/simple_storage.py | 80 ++++ transfer_queue/utils/zmq_utils.py | 11 + 14 files changed, 932 insertions(+), 248 deletions(-) create mode 100644 tests/e2e/test_restore_timeout_e2e.py create mode 100644 tests/test_restore_lifecycle.py diff --git a/docs/data_dump.md b/docs/data_dump.md index 03833fe3..0798e5b7 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -34,8 +34,10 @@ On version-3 SimpleStorage restore: 4. Each target unit reads only its assigned byte ranges and merges those values into local storage. Records are processed in batches of at most 128 rows per shard; the caller never reads or forwards their payloads. -5. After all units succeed, the client publishes returned field schemas through - the controller and merges tags using the normal metadata update path. +5. Each unit claims permission from the controller before writing, then reports + completion directly. The client commits the saved schemas and tags only after + every unit has completed. The controller reserves the destination partition + until commit or confirmed cancellation. The number of source units can differ from the number of destination units. Even a dump with one source shard can restore across several target units because @@ -47,7 +49,7 @@ storage unit. Local temporary storage suffices for single-node deployments. | Operation | State | Payload I/O | Unit count on restore | | --- | --- | --- | --- | | Checkpoint | Entire controller and storage state | Each unit reads/writes its whole file | Must match | -| Selective v2 dump | Selected fields and tags, merged by key | Each owner unit reads/writes its records | May differ | +| Selective v3 dump | Selected fields and tags, merged by key | Each owner unit reads/writes its records | May differ | `DUMP_ROWS` and `LOAD_ROWS` are included in storage operation metrics. Unit logs record loaded rows and bytes; the manager logs total bytes and participating units. @@ -79,9 +81,9 @@ Destination type conflicts are rejected before payload writes. Version-1 dumps remain readable using the prior caller-side KV put path. Version-2 dumps retain direct reads, but lack original schemas and use the older inference -behavior; exact field-type preservation cannot be guaranteed for those files. Restoring -to a backend without direct selective loading also uses KV puts. These compatibility -paths do not provide distributed file reads. Old builds that only understand +behavior; exact field-type preservation cannot be guaranteed for those files. +Restoring to a backend without direct selective loading uses KV puts. Version-1 +and KV fallback restores do not provide distributed file reads. Old builds that only understand versions 1 or 2 cannot read version-3 dumps. Export of nonempty dumps currently requires SimpleStorage. @@ -92,11 +94,30 @@ the new directory, and syncs its parent before deleting the backup. If publicati is interrupted while the main directory is absent, the next dump, load or row-index read recovers `.old`. A backup-cleanup error does not invalidate a published dump. -Restore is not transactional. All unit requests are awaited before returning an -error, and failed storage requests prevent publication of new ready metadata. -Earlier payload writes or earlier metadata updates can remain after a failure; -existing produced rows may already contain restored values. Correct the cause and -retry the same dump with writers paused. No unrelated partition is cleared. +Restore is not transactional: payload writes before a failure remain. Every load +has a unique ID. The controller blocks clearing/reusing its destination indexes +and conflicting KV puts while an operation is unresolved. Units must claim that ID +before writing; cancellation rejects requests that have not yet claimed permission. +A receive timeout never releases a writer that has already claimed permission. + +`RestorePendingError` means remote work is still running or its outcome is unknown. +The dump also retains a sibling `.restore` marker so an interrupted client cannot +silently allow its files to be replaced. After an interruption, call: + +```python +tq.recover_data_load("/shared/dumps/selected") +``` + +Recovery cancels unclaimed work, asks units to resend terminal results, and releases +indexes only after every claimed worker has finished. It does not publish partial +restores as ready or undo payload writes. Retry recovery while a unit is still busy; +a lost unit requires stopping the old TQ actors and restarting the whole TQ system. +Restarting only the controller while old storage actors run is unsupported. After +recovery succeeds, retry the dump or clear its keys. Unknown operations remain +reserved rather than guessing that a timeout stopped remote execution. + +Writers that already hold low-level metadata must remain paused throughout recovery. +The reservation covers the destination partition, so unrelated partitions can proceed. ## Tests diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index 663b3396..1b4335bf 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -514,3 +514,31 @@ def test_incompatible_schema_rejected_before_writes(tq_system, dump_dir, control with pytest.raises(RuntimeError, match="dtype mismatch"): tq.load_data_by_key(dump_dir) _assert_rows_equal(tq.kv_batch_get(["k"], "schema", ["x"])["x"], [torch.tensor([1.5])]) + + +def test_running_restore_blocks_clear_and_dump_until_recovery(tq_system, dump_dir, controller): + _put_rows("reserved", ["key"]) + tq.dump_data_by_key(dump_dir, ["key"], "reserved") + client = tq.get_client() + manager = client.storage_manager + rows = tq.read_row_index(dump_dir)["rows"] + tq.save_checkpoint(dump_dir.parent / "checkpoint") + units = list(manager.storage_unit_infos) + metadata = ray.get( + controller.begin_restore.remote("running-test", str(dump_dir.resolve()), "reserved", rows, units, {}) + ) + owner = units[metadata.global_indexes[0] % len(units)] + ray.get(controller.restore_unit.remote("running-test", owner, "claim")) + with pytest.raises(RuntimeError, match="unresolved"): + client.clear_partition("reserved") + with pytest.raises(tq.RestorePendingError): + tq.dump_data_by_key(dump_dir, ["key"], "reserved") + with pytest.raises(tq.RestorePendingError): + tq.load_checkpoint(dump_dir.parent / "checkpoint") + with pytest.raises(tq.RestorePendingError): + tq.recover_data_load(dump_dir) + ray.get(controller.restore_unit.remote("running-test", owner, "complete", {"success": False})) + tq.recover_data_load(dump_dir) + client.clear_partition("reserved") + tq.load_data_by_key(dump_dir) + assert tq.kv_batch_get(["key"], "reserved", ["input_ids"]).batch_size[0] == 1 diff --git a/tests/e2e/test_restore_timeout_e2e.py b/tests/e2e/test_restore_timeout_e2e.py new file mode 100644 index 00000000..a6a9facb --- /dev/null +++ b/tests/e2e/test_restore_timeout_e2e.py @@ -0,0 +1,109 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import threading +import time +from unittest.mock import patch + +import pytest +import ray +import torch +import zmq +import zmq.asyncio +from omegaconf import OmegaConf +from tensordict import TensorDict + +import transfer_queue as tq +from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager +from transfer_queue.utils.zmq_utils import ZMQMessage + + +def test_cancelled_delayed_load_cannot_overwrite_reused_index(tmp_path): + ray.init(namespace="review_timeout") + tq.init( + OmegaConf.create( + { + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": {"num_data_storage_units": 1, "total_storage_size": 20}, + } + } + ) + ) + try: + tq.kv_batch_put(["key"], "p", TensorDict({"x": torch.tensor([[1]])}, batch_size=1)) + tq.dump_data_by_key(tmp_path / "dump", ["key"], "p") + tq.kv_batch_put(["key"], "p", TensorDict({"x": torch.tensor([[2]])}, batch_size=1)) + manager = tq.get_client().storage_manager + real_info = next(iter(manager.storage_unit_infos.values())) + real_addr = real_info.to_addr("put_get_socket") + # Relay the request immediately, but delay delivery to the unit to model a queued load. + ready = threading.Event() + finished = threading.Event() + port = [] + + def relay(): + ctx = zmq.Context() + front = ctx.socket(zmq.ROUTER) + front.setsockopt(zmq.RCVTIMEO, 10000) + port.append(front.bind_to_random_port("tcp://127.0.0.1")) + back = ctx.socket(zmq.DEALER) + back.setsockopt(zmq.RCVTIMEO, 10000) + back.setsockopt(zmq.IDENTITY, b"TQ_STORAGE_review_relay") + back.connect(real_addr) + ready.set() + req = front.recv_multipart() + time.sleep(2) + back.send_multipart(req[1:]) + reply = back.recv_multipart() + assert not ZMQMessage.deserialize(reply).body["success"] + finished.set() + front.close(linger=0) + back.close(linger=0) + ctx.term() + + thread = threading.Thread(target=relay, daemon=True) + thread.start() + assert ready.wait(10) + + async def timeout_load(shards, target_storage_unit, restore): + ctx = zmq.asyncio.Context() + sock = ctx.socket(zmq.DEALER) + sock.setsockopt(zmq.RCVTIMEO, 100) + sock.connect(f"tcp://127.0.0.1:{port[0]}") + try: + return await AsyncSimpleStorageManager._load_selected_rows.__wrapped__( + manager, shards, target_storage_unit, restore=restore, socket=sock + ) + finally: + sock.close(linger=0) + ctx.term() + + with patch.object(manager, "_load_selected_rows", timeout_load): + with pytest.raises(zmq.error.Again): + tq.load_data_by_key(tmp_path / "dump") + value = tq.kv_batch_get(["key"], "p", ["x"])["x"][0].item() + + assert value == 2 + tq.get_client().clear_partition("p") + tq.kv_batch_put(["other"], "unrelated", TensorDict({"x": torch.tensor([[99]])}, batch_size=1)) + assert finished.wait(10) + thread.join(timeout=10) + later = tq.kv_batch_get(["other"], "unrelated", ["x"])["x"][0].item() + + assert later == 99 + finally: + tq.close() + ray.shutdown() diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index 9f88b8d3..421934c8 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -20,14 +20,11 @@ import io import pickle from types import SimpleNamespace -from unittest.mock import AsyncMock import pytest import torch from transfer_queue import data_dump, interface -from transfer_queue.client import AsyncTransferQueueClient -from transfer_queue.metadata import BatchMeta from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType @@ -194,7 +191,7 @@ def read(self, size=-1): records.append( {"source_index": source, "target_index": target, "fields": ["x"], "offset": offset, "length": length} ) - loaded = unit._handle_load_rows( + loaded = unit._load_rows( ZMQMessage.create( request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", @@ -238,20 +235,6 @@ async def load(shards, target_storage_unit): assert sorted(seen) == list(range(16)) -@pytest.mark.asyncio -async def test_failed_load_does_not_publish_ready_metadata(): - client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) - client.close = lambda: None - client.storage_manager = SimpleNamespace(load_rows_by_index=AsyncMock(side_effect=RuntimeError("read failed"))) - client.async_kv_retrieve_meta = AsyncMock(return_value=BatchMeta(global_indexes=[9], partition_ids=["p"])) - client._publish_loaded_rows = AsyncMock() - client.async_set_custom_meta = AsyncMock() - with pytest.raises(RuntimeError, match="read failed"): - await client.async_load_rows_by_key("p", {"k": {"tag": {}}}, []) - client._publish_loaded_rows.assert_not_called() - client.async_set_custom_meta.assert_not_called() - - @pytest.mark.parametrize("problem", ["wrong_index", "truncated", "missing_field"]) def test_unit_rejects_invalid_records(unit, tmp_path, problem): path = tmp_path / "row.pkl" @@ -263,7 +246,7 @@ def test_unit_rejects_invalid_records(unit, tmp_path, problem): record["length"] += 1 else: record["fields"] = ["missing"] - reply = unit._handle_load_rows( + reply = unit._load_rows( ZMQMessage.create( request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", @@ -353,21 +336,6 @@ async def load(shards, target_storage_unit): assert finished == ["u1"] -@pytest.mark.asyncio -async def test_load_rejects_failed_controller_metadata_ack(): - client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) - client.close = lambda: None - client._request_controller = AsyncMock(return_value=SimpleNamespace(body={"success": False})) - with pytest.raises(RuntimeError, match="Controller rejected"): - await AsyncTransferQueueClient._publish_loaded_rows.__wrapped__( - client, - "p", - [ - {"global_indexes": [3], "field_schema": {}}, - ], - ) - - @pytest.mark.asyncio async def test_dump_waits_for_writers_before_cleanup_can_start(tmp_path): manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) diff --git a/tests/test_restore_lifecycle.py b/tests/test_restore_lifecycle.py new file mode 100644 index 00000000..5b85acba --- /dev/null +++ b/tests/test_restore_lifecycle.py @@ -0,0 +1,198 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Restore reservations prevent late writes from corrupting reused indexes.""" + +from threading import RLock + +import pytest +import torch +import zmq + +from transfer_queue.controller import PartitionIndexManager, TransferQueueController +from transfer_queue.sampler import SequentialSampler +from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData +from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType + + +@pytest.fixture +def controller(): + cls = TransferQueueController.__ray_metadata__.modified_class + controller = cls.__new__(cls) + controller.controller_id = "test_controller" + controller.partitions = {} + controller.index_manager = PartitionIndexManager() + controller.sampler = SequentialSampler() + controller._restore_lock = RLock() + controller._restores = {} + controller._cancelled_restores = set() + controller._clearing_indexes = set() + return controller + + +def begin(controller, restore_id="r"): + return controller.begin_restore(restore_id, "/dump", "p", {"k": {"fields": ["x"], "tag": {}}}, ["u"], {}) + + +def test_cancel_before_claim_rejects_late_load_and_late_begin(controller): + metadata = begin(controller) + assert controller.finish_restore("r", commit=False)["finished"] + controller.clear_partition("p") + current = controller.kv_retrieve_meta(["other"], "other", create=True) + assert current.global_indexes == metadata.global_indexes + with pytest.raises(RuntimeError, match="no longer active"): + controller.restore_unit("r", "u", "claim") + with pytest.raises(RuntimeError, match="already cancelled"): + begin(controller) + controller.finish_restore("not-arrived", commit=False) + with pytest.raises(RuntimeError, match="already cancelled"): + begin(controller, "not-arrived") + + +def test_claimed_load_keeps_indexes_until_terminal_report(controller): + metadata = begin(controller) + controller.restore_unit("r", "u", "claim") + assert not controller.finish_restore("r", commit=False)["finished"] + for action in ( + lambda: controller.mark_clearing(metadata.global_indexes, ["p"]), + lambda: controller.clear_meta(metadata.global_indexes, ["p"]), + lambda: controller.clear_partition("p"), + lambda: controller.kv_retrieve_meta(["k"], "p", create=True), + lambda: begin(controller, "overlap"), + ): + with pytest.raises(RuntimeError, match="unresolved"): + action() + assert controller.list_restores("/dump") == ["r"] + assert controller.kv_retrieve_meta(["other"], "other", create=True).global_indexes != metadata.global_indexes + controller.restore_unit("r", "u", "complete", {"success": False}) + assert controller.finish_restore("r", commit=False)["finished"] + controller.clear_partition("p") + + +def test_success_publishes_schema_and_tags_only_at_finish(controller): + metadata = begin(controller) + index = metadata.global_indexes[0] + controller.restore_unit("r", "u", "claim") + schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_nested": False, "is_non_tensor": False}} + controller.restore_unit( + "r", + "u", + "complete", + { + "success": True, + "updates": [ + {"global_indexes": [index], "field_schema": schema}, + ], + }, + ) + assert not controller.partitions["p"].field_metadata + assert controller.finish_restore("r", commit=True)["finished"] + assert controller.partitions["p"].field_metadata["x"].global_indexes == {index} + assert not controller.list_restores("/dump") + + +def test_restore_cannot_start_in_the_middle_of_clear(controller): + metadata = controller.kv_retrieve_meta(["k"], "p", create=True) + controller.mark_clearing(metadata.global_indexes, ["p"]) + with pytest.raises(RuntimeError, match="unfinished clear"): + begin(controller) + controller.clear_partition("p") + begin(controller) + + +def test_unit_rejects_cancelled_permission_before_reading(controller): + begin(controller) + controller.finish_restore("r", commit=False) + cls = SimpleStorageUnit.__ray_metadata__.modified_class + unit = cls.__new__(cls) + unit.storage_unit_id = "u" + unit.storage_data = StorageUnitData() + unit._restore_results = {} + unit._restore_controller_request = lambda context, action, result=None: controller.restore_unit( + "r", "u", action, result + ) + unit._load_rows = lambda *_: pytest.fail("Cancelled request read payload") + response = unit._handle_load_rows( + ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS, + sender_id="test", + body={"restore": {"restore_id": "r"}, "shards": []}, + ) + ) + assert not response.body["success"] + + +def test_lost_completion_ack_is_recoverable_without_replaying_payload(controller): + begin(controller) + cls = SimpleStorageUnit.__ray_metadata__.modified_class + unit = cls.__new__(cls) + unit.storage_unit_id = "u" + unit._restore_results = {} + + def lose_completion(context, action, result=None): + if action == "claim": + controller.restore_unit("r", "u", action, result) + else: + raise zmq.error.Again() + + unit._restore_controller_request = lose_completion + unit._load_rows = lambda *_: ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, + sender_id="u", + body={"success": True, "updates": [], "bytes_read": 0}, + ) + unit._handle_load_rows( + ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", body={"restore": {"restore_id": "r"}, "shards": []} + ) + ) + assert not controller.finish_restore("r", commit=False)["finished"] + unit._restore_controller_request = lambda context, action, result=None: controller.restore_unit( + "r", "u", action, result + ) + response = unit._handle_report_restore( + ZMQMessage.create(request_type=ZMQRequestType.REPORT_RESTORE, sender_id="test", body={"restore_id": "r"}) + ) + assert response.body["success"] + assert controller.finish_restore("r", commit=False)["finished"] + + +def test_lost_claim_ack_can_be_settled_without_writing(controller): + begin(controller) + cls = SimpleStorageUnit.__ray_metadata__.modified_class + unit = cls.__new__(cls) + unit.storage_unit_id = "u" + unit._restore_results = {} + + def lose_ack(context, action, result=None): + controller.restore_unit("r", "u", action, result) + raise zmq.error.Again() + + unit._restore_controller_request = lose_ack + unit._load_rows = lambda *_: pytest.fail("Unacknowledged claim wrote data") + response = unit._handle_load_rows( + ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", body={"restore": {"restore_id": "r"}, "shards": []} + ) + ) + assert not response.body["success"] + assert not controller.finish_restore("r", commit=False)["finished"] + unit._restore_controller_request = lambda context, action, result=None: controller.restore_unit( + "r", "u", action, result + ) + assert unit._handle_report_restore( + ZMQMessage.create(request_type=ZMQRequestType.REPORT_RESTORE, sender_id="test", body={"restore_id": "r"}) + ).body["success"] + assert controller.finish_restore("r", commit=False)["finished"] diff --git a/transfer_queue/__init__.py b/transfer_queue/__init__.py index 97c180ce..a0fa4eee 100644 --- a/transfer_queue/__init__.py +++ b/transfer_queue/__init__.py @@ -16,7 +16,7 @@ import os from .client import TransferQueueClient -from .data_dump import dump_data_by_key, load_data_by_key, read_row_index +from .data_dump import dump_data_by_key, load_data_by_key, read_row_index, recover_data_load from .dataloader import StreamingDataLoader, StreamingDataset from .interface import ( async_kv_batch_get, @@ -45,6 +45,7 @@ from .sampler.seqlen_balanced_sampler import SeqlenBalancedSampler from .sampler.sequential_sampler import SequentialSampler from .storage import StorageKeyNotFoundError +from .storage.dump_io import RestorePendingError __all__ = ( [ @@ -76,6 +77,8 @@ "dump_data_by_key", "load_data_by_key", "read_row_index", + "recover_data_load", + "RestorePendingError", ] + [ # High-Level StreamingDataLoader Interface diff --git a/transfer_queue/client.py b/transfer_queue/client.py index d6b96297..ee226bc5 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -18,6 +18,7 @@ import threading import weakref from typing import Any, Callable +from uuid import uuid4 import torch import zmq @@ -26,6 +27,7 @@ from transfer_queue.metadata import BatchMeta from transfer_queue.storage import StorageManagerFactory +from transfer_queue.storage.dump_io import RestorePendingError from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.zmq_utils import ( @@ -1217,48 +1219,84 @@ async def async_dump_rows_by_index( raise NotImplementedError(f"{type(self.storage_manager).__name__} does not support selective data dump") return await self.storage_manager.dump_rows_by_index(shard_dir, global_indexes, fields_by_index) + def _restore_context(self, restore_id: str) -> dict: + return { + "restore_id": restore_id, + "controller_ip": self._controller.ip, + "controller_address": self._controller.to_addr("request_handle_socket"), + } + + @with_controller_socket + async def _restore_rpc(self, request_type: ZMQRequestType, body: dict, socket=None) -> dict: + response_type = ZMQRequestType(request_type.value + "_RESPONSE") + response = await self._request_controller(socket, request_type, response_type, body) + return response.body + async def async_load_rows_by_key( self, partition_id: str, rows: dict[str, dict[str, Any]], shards: list[dict[str, Any]], + dump_dir: str = "", + restore_id: str | None = None, ) -> None: - """Allocate current indexes, load payloads at owner units, then publish metadata.""" + """Reserve indexes until every unit has explicitly completed or been denied permission.""" manager = getattr(self, "storage_manager", None) if manager is None or not hasattr(manager, "load_rows_by_index"): raise NotImplementedError("Storage backend does not support direct selective load") if not rows: return - keys = list(rows) - metadata = await self.async_kv_retrieve_meta(keys, partition_id, create=True) - if metadata.size != len(keys): - raise RuntimeError("Selective load did not allocate every key") - target_indexes = dict(zip(keys, metadata.global_indexes, strict=True)) - for shard in shards: - for record in shard["records"]: - record["target_index"] = target_indexes[record["key"]] - updates = await manager.load_rows_by_index(shards) - await self._publish_loaded_rows(partition_id, updates) - metadata.update_custom_meta([rows[key]["tag"] for key in keys]) - await self.async_set_custom_meta(metadata) - - @with_controller_socket - async def _publish_loaded_rows( - self, - partition_id: str, - updates: list[dict[str, Any]], - socket: zmq.asyncio.Socket | None = None, - ) -> None: - # Loading must report a failed metadata update instead of silently succeeding. - for update in updates: - response = await self._request_controller( - socket=socket, - request_type=ZMQRequestType.NOTIFY_DATA_UPDATE, - response_type=ZMQRequestType.NOTIFY_DATA_UPDATE_ACK, - body={"partition_id": partition_id, **update}, + restore_id = restore_id or uuid4().hex + try: + response = await self._restore_rpc( + ZMQRequestType.BEGIN_RESTORE, + { + "restore_id": restore_id, + "dump_dir": dump_dir, + "partition_id": partition_id, + "rows": rows, + "units": list(manager.storage_unit_infos), + "schema": shards[0].get("field_schema", {}) if shards else {}, + }, ) - if not response.body.get("success"): - raise RuntimeError(f"Controller rejected loaded row metadata for partition {partition_id!r}") + metadata = response["metadata"] + target_indexes = dict(zip(rows, metadata.global_indexes, strict=True)) + for shard in shards: + for record in shard["records"]: + record["target_index"] = target_indexes[record["key"]] + await manager.load_rows_by_index(shards, self._restore_context(restore_id)) + result = await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": True}) + if not result["finished"]: + raise RestorePendingError(restore_id) + except BaseException as error: + try: + result = await asyncio.shield( + self._restore_rpc( + ZMQRequestType.FINISH_RESTORE, + {"restore_id": restore_id, "commit": False}, + ) + ) + except BaseException: + raise RestorePendingError(restore_id) from error + if not result["finished"]: + raise RestorePendingError(restore_id) from error + raise + + async def async_recover_data_load(self, dump_dir: str, restore_ids: list[str] | None = None) -> None: + """Cancel unclaimed loads and release only operations whose claimed workers have finished.""" + response = await self._restore_rpc(ZMQRequestType.LIST_RESTORES, {"dump_dir": dump_dir}) + for restore_id in set(response["restore_ids"]) | set(restore_ids or []): + await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": False}) + await self.storage_manager.report_restore(self._restore_context(restore_id)) + result = await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": False}) + if not result["finished"]: + raise RestorePendingError(restore_id) + + async def async_check_data_loads(self, dump_dir: str | None = None) -> None: + """Reject replacing a dump whose files may still be read by a storage unit.""" + response = await self._restore_rpc(ZMQRequestType.LIST_RESTORES, {"dump_dir": dump_dir}) + if response["restore_ids"]: + raise RestorePendingError(response["restore_ids"][0]) # ==================== Checkpoint API ==================== @with_controller_socket @@ -1449,6 +1487,8 @@ def wrapper(*args, **kwargs): self._validate_dump_schema = _make_sync(self.async_validate_dump_schema) self._dump_rows_by_index = _make_sync(self.async_dump_rows_by_index) self._load_rows_by_key = _make_sync(self.async_load_rows_by_key) + self._recover_data_load = _make_sync(self.async_recover_data_load) + self._check_data_loads = _make_sync(self.async_check_data_loads) self._save_controller_checkpoint = _make_sync(self.async_save_controller_checkpoint) self._load_controller_checkpoint = _make_sync(self.async_load_controller_checkpoint) self._save_storage_checkpoint = _make_sync(self.async_save_storage_checkpoint) @@ -1933,13 +1973,18 @@ def dump_rows_by_index( return self._dump_rows_by_index(shard_dir, global_indexes, fields_by_index) def load_rows_by_key( - self, - partition_id: str, - rows: dict[str, dict[str, Any]], - shards: list[dict[str, Any]], + self, partition_id: str, rows: dict, shards: list[dict], dump_dir: str = "", restore_id: str | None = None ) -> None: - """Restore indexed dump records directly on the current storage owner units.""" - return self._load_rows_by_key(partition_id, rows, shards) + """Restore selected payloads while the controller reserves destination indexes.""" + return self._load_rows_by_key(partition_id, rows, shards, dump_dir, restore_id) + + def recover_data_load(self, dump_dir: str, restore_ids: list[str] | None = None) -> None: + """Settle an interrupted load before retrying, clearing, or replacing its dump.""" + return self._recover_data_load(dump_dir, restore_ids) + + def check_data_loads(self, dump_dir: str | None = None) -> None: + """Reject replacing a dump while a remote restore may still read it.""" + return self._check_data_loads(dump_dir) # ==================== Checkpoint API ==================== def save_controller_checkpoint(self, path: str) -> None: diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index 9adb43d8..cb07518f 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -21,7 +21,7 @@ from dataclasses import dataclass, field from itertools import groupby from operator import itemgetter -from threading import Thread +from threading import RLock, Thread from typing import TYPE_CHECKING, Any, cast from uuid import uuid4 @@ -1008,6 +1008,10 @@ def __init__( # Partition-GlobalIndex management self.index_manager = PartitionIndexManager() # partition_id -> global_indexes + self._restore_lock = RLock() + self._restores: dict[str, dict[str, Any]] = {} + self._cancelled_restores: set[str] = set() + self._clearing_indexes: set[int] = set() # Connected storage managers tracking self._connected_storage_managers: set[str] = set() @@ -1490,30 +1494,34 @@ def mark_clearing(self, global_indexes: list[int], partition_ids: list[str]) -> global_indexes: global indexes to mark as pending deletion partition_ids: corresponding partition IDs for each global index """ - if global_indexes is None or partition_ids is None: - raise ValueError("global_indexes and partition_ids cannot be None") + with self._restore_lock: + if global_indexes is None or partition_ids is None: + raise ValueError("global_indexes and partition_ids cannot be None") + for pid in set(partition_ids): + self._assert_not_restoring(pid) - if len(global_indexes) != len(partition_ids): - raise ValueError( - f"global_indexes and partition_ids must have the same length, " - f"got {len(global_indexes)} and {len(partition_ids)}" - ) + if len(global_indexes) != len(partition_ids): + raise ValueError( + f"global_indexes and partition_ids must have the same length, " + f"got {len(global_indexes)} and {len(partition_ids)}" + ) - combined = list(zip(partition_ids, global_indexes, strict=True)) - combined.sort(key=itemgetter(0)) + combined = list(zip(partition_ids, global_indexes, strict=True)) + combined.sort(key=itemgetter(0)) - for partition_id, group in groupby(combined, key=itemgetter(0)): - partition = self._get_partition(partition_id) - if not partition: - logger.info( - f"[{self.controller_id}]: Trying to mark clearing in a non-existent partition {partition_id}. " - f"Skipping operation for this partition." - ) - continue - indexes = [idx for _, idx in group] - existing = list(set(indexes) & partition.global_indexes) - if existing and partition.production_status is not None: - partition.production_status[existing, :] = 0 + for partition_id, group in groupby(combined, key=itemgetter(0)): + partition = self._get_partition(partition_id) + if not partition: + logger.info( + f"[{self.controller_id}]: Trying to mark clearing in a non-existent partition {partition_id}. " + f"Skipping operation for this partition." + ) + continue + indexes = [idx for _, idx in group] + existing = list(set(indexes) & partition.global_indexes) + self._clearing_indexes.update(existing) + if existing and partition.production_status is not None: + partition.production_status[existing, :] = 0 def clear_partition(self, partition_id: str, clear_consumption: bool = True): """ @@ -1523,21 +1531,23 @@ def clear_partition(self, partition_id: str, clear_consumption: bool = True): partition_id: ID of the partition to clear clear_consumption: Whether to also clear consumption status """ + with self._restore_lock: + self._assert_not_restoring(partition_id) + logger.debug(f"[{self.controller_id}]: Clearing metadata in partition {partition_id}.") - logger.debug(f"[{self.controller_id}]: Clearing metadata in partition {partition_id}.") - - partition = self._get_partition(partition_id) - if not partition: - logger.warning( - f"[{self.controller_id}]: Trying to clear a non-existent partition {partition_id}. No action taken." - ) - return + partition = self._get_partition(partition_id) + if not partition: + logger.warning( + f"[{self.controller_id}]: Trying to clear a non-existent partition {partition_id}. No action taken." + ) + return - global_indexes_range = list(self.index_manager.get_indexes_for_partition(partition_id)) - partition.clear_data(global_indexes_range, clear_consumption) - self.index_manager.release_partition(partition_id) - self.partitions.pop(partition_id) - self.sampler.clear_cache(partition_id) + global_indexes_range = list(self.index_manager.get_indexes_for_partition(partition_id)) + partition.clear_data(global_indexes_range, clear_consumption) + self.index_manager.release_partition(partition_id) + self._clearing_indexes.difference_update(global_indexes_range) + self.partitions.pop(partition_id) + self.sampler.clear_cache(partition_id) def reset_consumption(self, partition_id: str, task_name: str | None = None): """ @@ -1574,54 +1584,57 @@ def clear_meta( partition_ids: IDs of the partitions to clear clear_consumption: Whether to also clear consumption status """ - - logger.debug( - f"[{self.controller_id}]: Clearing meta with global_indexes {global_indexes} in partition {partition_ids}" - ) - - if global_indexes is None or partition_ids is None: - raise ValueError("global_indexes and partition_ids cannot be None") - - if len(global_indexes) != len(partition_ids): - raise ValueError( - f"global_indexes and partition_ids must have the same length, " - f"got {len(global_indexes)} and {len(partition_ids)}" + with self._restore_lock: + logger.debug( + "[%s]: Clearing indexes %s in partitions %s", self.controller_id, global_indexes, partition_ids ) - combined = list(zip(partition_ids, global_indexes, strict=True)) - combined.sort(key=itemgetter(0)) + if global_indexes is None or partition_ids is None: + raise ValueError("global_indexes and partition_ids cannot be None") + for pid in set(partition_ids): + self._assert_not_restoring(pid) - for partition_id, group in groupby(combined, key=itemgetter(0)): - partition = self._get_partition(partition_id) - if not partition: - logger.info( - f"[{self.controller_id}]: Trying to clear data in a non-existent partition {partition_id}. " - f"Skipping operation for this partition." + if len(global_indexes) != len(partition_ids): + raise ValueError( + f"global_indexes and partition_ids must have the same length, " + f"got {len(global_indexes)} and {len(partition_ids)}" ) - continue - global_indexes_to_clear = [idx for _, idx in group] - existing_global_indexes = partition.global_indexes - non_existent_global_indexes = set(global_indexes_to_clear) - existing_global_indexes - if non_existent_global_indexes: - logger.info( - f"[{self.controller_id}]: Some global_indexes to be cleared do not exist in " - f"partition {partition_id}: {non_existent_global_indexes}. They will be ignored." - ) + combined = list(zip(partition_ids, global_indexes, strict=True)) + combined.sort(key=itemgetter(0)) - global_indexes_to_clear = list(set(global_indexes_to_clear) & existing_global_indexes) - if not global_indexes_to_clear: - logger.info( - f"[{self.controller_id}]: No existing global indexes to clear in partition {partition_id}. " - f"Skipping operation for this partition." - ) - continue + for partition_id, group in groupby(combined, key=itemgetter(0)): + partition = self._get_partition(partition_id) + if not partition: + logger.info( + f"[{self.controller_id}]: Trying to clear data in a non-existent partition {partition_id}. " + f"Skipping operation for this partition." + ) + continue + + global_indexes_to_clear = [idx for _, idx in group] + existing_global_indexes = partition.global_indexes + non_existent_global_indexes = set(global_indexes_to_clear) - existing_global_indexes + if non_existent_global_indexes: + logger.info( + f"[{self.controller_id}]: Some global_indexes to be cleared do not exist in " + f"partition {partition_id}: {non_existent_global_indexes}. They will be ignored." + ) - # Clear data from partition - partition.clear_data(global_indexes_to_clear, clear_consumption) + global_indexes_to_clear = list(set(global_indexes_to_clear) & existing_global_indexes) + if not global_indexes_to_clear: + logger.info( + f"[{self.controller_id}]: No existing global indexes to clear in partition {partition_id}. " + f"Skipping operation for this partition." + ) + continue + + # Clear data from partition + partition.clear_data(global_indexes_to_clear, clear_consumption) - # Release the specific indexes from index manager - self.index_manager.release_indexes(partition_id, global_indexes_to_clear) + # Release the specific indexes from index manager + self.index_manager.release_indexes(partition_id, global_indexes_to_clear) + self._clearing_indexes.difference_update(global_indexes_to_clear) def kv_retrieve_meta( self, @@ -1641,64 +1654,66 @@ def kv_retrieve_meta( Returns: metadata: BatchMeta of the requested keys """ + with self._restore_lock: + if create: + self._assert_not_restoring(partition_id) + logger.debug(f"[{self.controller_id}] Retrieve keys {keys} in partition {partition_id}") - logger.debug(f"[{self.controller_id}] Retrieve keys {keys} in partition {partition_id}") - - # Ensure partition exists - partition = self._get_partition(partition_id) - if partition is None: - if not create: - logger.warning( - f"[{self.controller_id}]: Partition {partition_id} not found. Returning empty BatchMeta." - ) - return BatchMeta.empty() - - self.create_partition(partition_id) + # Ensure partition exists partition = self._get_partition(partition_id) + if partition is None: + if not create: + logger.warning( + f"[{self.controller_id}]: Partition {partition_id} not found. Returning empty BatchMeta." + ) + return BatchMeta.empty() - assert partition is not None - global_indexes = partition.kv_retrieve_indexes(keys) - - none_indexes = [idx for idx, value in enumerate(global_indexes) if value is None] - if len(none_indexes) > 0: - if not create: - logger.warning( - f"Keys {[keys[i] for i in none_indexes]} were not found in partition {partition_id}. " - f"They will be excluded from the retrieved BatchMeta." - ) - else: - # create non-exist keys - batch_global_indexes = partition.activate_pre_allocated_indexes(len(none_indexes)) + self.create_partition(partition_id) + partition = self._get_partition(partition_id) - if len(batch_global_indexes) < len(none_indexes): - new_global_indexes = self.index_manager.allocate_indexes( - partition_id, count=(len(none_indexes) - len(batch_global_indexes)) + assert partition is not None + global_indexes = partition.kv_retrieve_indexes(keys) + + none_indexes = [idx for idx, value in enumerate(global_indexes) if value is None] + if len(none_indexes) > 0: + if not create: + logger.warning( + f"Keys {[keys[i] for i in none_indexes]} were not found in partition {partition_id}. " + f"They will be excluded from the retrieved BatchMeta." ) - batch_global_indexes.extend(new_global_indexes) + else: + # create non-exist keys + batch_global_indexes = partition.activate_pre_allocated_indexes(len(none_indexes)) - # register global_indexes in partition - partition.global_indexes.update(batch_global_indexes) + if len(batch_global_indexes) < len(none_indexes): + new_global_indexes = self.index_manager.allocate_indexes( + partition_id, count=(len(none_indexes) - len(batch_global_indexes)) + ) + batch_global_indexes.extend(new_global_indexes) - # register key-global_indexes mapping in partition - for i in range(len(none_indexes)): - global_indexes[none_indexes[i]] = batch_global_indexes[i] - partition.keys_mapping[keys[none_indexes[i]]] = batch_global_indexes[i] - partition.revert_keys_mapping[batch_global_indexes[i]] = keys[none_indexes[i]] + # register global_indexes in partition + partition.global_indexes.update(batch_global_indexes) - partition.ensure_samples_capacity(max(batch_global_indexes) + 1) + # register key-global_indexes mapping in partition + for i in range(len(none_indexes)): + global_indexes[none_indexes[i]] = batch_global_indexes[i] + partition.keys_mapping[keys[none_indexes[i]]] = batch_global_indexes[i] + partition.revert_keys_mapping[batch_global_indexes[i]] = keys[none_indexes[i]] - verified_global_indexes = [idx for idx in global_indexes if idx is not None] + partition.ensure_samples_capacity(max(batch_global_indexes) + 1) - # must fetch fields that the requested samples all have - col_mask = partition.production_status[verified_global_indexes, :].sum(dim=0).reshape(-1) == len( - verified_global_indexes - ) - data_fields = [] - for field_name, col_idx in partition.field_name_mapping.items(): - if col_idx < len(col_mask) and col_mask[col_idx]: - data_fields.append(field_name) + verified_global_indexes = [idx for idx in global_indexes if idx is not None] - return self.generate_batch_meta(partition_id, verified_global_indexes, data_fields, mode="force_fetch") + # must fetch fields that the requested samples all have + col_mask = partition.production_status[verified_global_indexes, :].sum(dim=0).reshape(-1) == len( + verified_global_indexes + ) + data_fields = [] + for field_name, col_idx in partition.field_name_mapping.items(): + if col_idx < len(col_mask) and col_mask[col_idx]: + data_fields.append(field_name) + + return self.generate_batch_meta(partition_id, verified_global_indexes, data_fields, mode="force_fetch") def kv_retrieve_keys( self, @@ -1778,6 +1793,106 @@ def describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, for key, global_index in zip(keys, global_indexes, strict=True) } + def _assert_not_restoring(self, partition_id: str | None = None) -> None: + for restore_id, restore in self._restores.items(): + if partition_id is None or restore["partition_id"] == partition_id: + raise RuntimeError( + f"Restore {restore_id} is unresolved; recover_data_load({restore['dump_dir']!r}) first" + ) + + def begin_restore( + self, restore_id: str, dump_dir: str, partition_id: str, rows: dict, units: list[str], schema: dict + ): + """Reserve a destination partition until every possible remote writer is settled.""" + with self._restore_lock: + if restore_id in self._cancelled_restores: + raise RuntimeError("Restore was already cancelled") + if restore_id in self._restores: + return self._restores[restore_id]["metadata"] + self._assert_not_restoring(partition_id) + partition = self._get_partition(partition_id) + if partition is not None: + if partition.global_indexes & self._clearing_indexes: + raise RuntimeError("Partition has an unfinished clear; complete it before restoring") + partition.validate_field_schema(schema) + keys = list(rows) + metadata = self.kv_retrieve_meta(keys, partition_id, create=True) + active_units = { + units[index % len(units)] + for key, index in zip(keys, metadata.global_indexes, strict=True) + if rows[key]["fields"] + } + metadata.update_custom_meta([rows[key]["tag"] for key in keys]) + self._restores[restore_id] = { + "partition_id": partition_id, + "dump_dir": dump_dir, + "metadata": metadata, + "units": {unit: "pending" for unit in active_units}, + "updates": {}, + "aborting": False, + } + return metadata + + def restore_unit(self, restore_id: str, unit_id: str, action: str, result: dict | None = None) -> None: + """Grant one execution and record its terminal result; unknown IDs never grant writes.""" + with self._restore_lock: + restore = self._restores.get(restore_id) + if restore is None or unit_id not in restore["units"]: + raise RuntimeError("Restore is no longer active") + state = restore["units"][unit_id] + if action == "claim": + if state != "pending" or restore["aborting"]: + raise RuntimeError(f"Restore unit cannot start in state {state}") + restore["units"][unit_id] = "running" + elif action == "complete" and state in ("running", "done", "failed") and result is not None: + restore["units"][unit_id] = "done" if result["success"] else "failed" + restore["updates"][unit_id] = result.get("updates", []) + else: + raise RuntimeError(f"Invalid restore transition {state} -> {action}") + + def finish_restore(self, restore_id: str, commit: bool) -> dict: + """Release indexes only after all claimed writers terminate; cancel unclaimed work.""" + with self._restore_lock: + restore = self._restores.get(restore_id) + if restore is None: + if commit: + raise RuntimeError("Restore is no longer active; metadata was not committed by this request") + self._cancelled_restores.add(restore_id) + return {"finished": True} + if not commit: + restore["aborting"] = True + for unit, state in restore["units"].items(): + if state == "pending": + restore["units"][unit] = "cancelled" + running = [unit for unit, state in restore["units"].items() if state in ("pending", "running")] + if running: + return {"finished": False, "units": running} + if commit: + if restore["aborting"] or any(state != "done" for state in restore["units"].values()): + raise RuntimeError("Restore has failed or was cancelled") + partition = self.partitions[restore["partition_id"]] + updates = [update for group in restore["updates"].values() for update in group] + for update in updates: + partition.validate_field_schema(update["field_schema"]) + for update in updates: + if not partition.update_production_status(update["global_indexes"], [], update["field_schema"]): + raise RuntimeError("Controller rejected restored metadata") + metadata = restore["metadata"] + partition.set_custom_meta(dict(zip(metadata.global_indexes, metadata.custom_meta, strict=True))) + else: + self._cancelled_restores.add(restore_id) + del self._restores[restore_id] + return {"finished": True} + + def list_restores(self, dump_dir: str | None) -> list[str]: + """Find unfinished operations for a dump, including after the initiating client exits.""" + with self._restore_lock: + return [ + restore_id + for restore_id, restore in self._restores.items() + if dump_dir is None or restore["dump_dir"] == dump_dir + ] + def _init_zmq_socket(self): """Initialize ZMQ sockets for communication.""" self.zmq_context = zmq.Context() @@ -2003,6 +2118,10 @@ def _handle_request(self, request_msg: ZMQMessage, monitor: Any) -> ZMQMessage | ZMQRequestType.KV_LIST: self._handle_kv_list_request, ZMQRequestType.DESCRIBE_ROWS_BY_KEY: self._handle_describe_rows_by_key_request, ZMQRequestType.VALIDATE_DUMP_SCHEMA: self._handle_validate_dump_schema_request, + ZMQRequestType.BEGIN_RESTORE: self._handle_begin_restore_request, + ZMQRequestType.RESTORE_UNIT: self._handle_restore_unit_request, + ZMQRequestType.FINISH_RESTORE: self._handle_finish_restore_request, + ZMQRequestType.LIST_RESTORES: self._handle_list_restores_request, ZMQRequestType.SAVE_CONTROLLER_CHECKPOINT: self._handle_save_controller_checkpoint_request, ZMQRequestType.LOAD_CONTROLLER_CHECKPOINT: self._handle_load_controller_checkpoint_request, } @@ -2278,6 +2397,22 @@ def _handle_describe_rows_by_key_request(self, request_msg: ZMQMessage) -> ZMQMe {"success": True, "partition_id": params["partition_id"], "rows": rows, "field_schema": field_schema}, ) + def _handle_begin_restore_request(self, request_msg: ZMQMessage) -> ZMQMessage: + metadata = self.begin_restore(**request_msg.body) + return self._make_response(request_msg, ZMQRequestType.BEGIN_RESTORE_RESPONSE, {"metadata": metadata}) + + def _handle_restore_unit_request(self, request_msg: ZMQMessage) -> ZMQMessage: + self.restore_unit(**request_msg.body) + return self._make_response(request_msg, ZMQRequestType.RESTORE_UNIT_RESPONSE, {"success": True}) + + def _handle_finish_restore_request(self, request_msg: ZMQMessage) -> ZMQMessage: + result = self.finish_restore(**request_msg.body) + return self._make_response(request_msg, ZMQRequestType.FINISH_RESTORE_RESPONSE, result) + + def _handle_list_restores_request(self, request_msg: ZMQMessage) -> ZMQMessage: + ids = self.list_restores(request_msg.body["dump_dir"]) + return self._make_response(request_msg, ZMQRequestType.LIST_RESTORES_RESPONSE, {"restore_ids": ids}) + def _handle_validate_dump_schema_request(self, request_msg: ZMQMessage) -> ZMQMessage: params = request_msg.body partition = self._get_partition(params["partition_id"]) @@ -2309,23 +2444,25 @@ def save_checkpoint(self, path: str) -> None: Raises: Exception: If serialization or file I/O fails. """ - try: - state = { - "controller_id": self.controller_id, - "partitions": {pid: p.to_snapshot() for pid, p in self.partitions.items()}, - "index_manager": { - "partition_to_indexes": dict(copy.deepcopy(self.index_manager.partition_to_indexes)), - "reusable_indexes": list(self.index_manager.reusable_indexes), - "global_index_counter": self.index_manager.global_index_counter, - "allocated_indexes": set(self.index_manager.allocated_indexes), - }, - "sampler": self.sampler.save_checkpoint(), - } - with open(path, "wb") as f: - pickle.dump(state, f, protocol=pickle.HIGHEST_PROTOCOL) - logger.info(f"[{self.controller_id}]: dumped to {path}") - except Exception as e: - raise RuntimeError(f"[{self.controller_id}]: save checkpoint failed: {e}") from e + with self._restore_lock: + self._assert_not_restoring() + try: + state = { + "controller_id": self.controller_id, + "partitions": {pid: p.to_snapshot() for pid, p in self.partitions.items()}, + "index_manager": { + "partition_to_indexes": dict(copy.deepcopy(self.index_manager.partition_to_indexes)), + "reusable_indexes": list(self.index_manager.reusable_indexes), + "global_index_counter": self.index_manager.global_index_counter, + "allocated_indexes": set(self.index_manager.allocated_indexes), + }, + "sampler": self.sampler.save_checkpoint(), + } + with open(path, "wb") as f: + pickle.dump(state, f, protocol=pickle.HIGHEST_PROTOCOL) + logger.info(f"[{self.controller_id}]: dumped to {path}") + except Exception as e: + raise RuntimeError(f"[{self.controller_id}]: save checkpoint failed: {e}") from e def load_checkpoint(self, path: str) -> None: """Restore controller state directly from a file. @@ -2336,24 +2473,27 @@ def load_checkpoint(self, path: str) -> None: Raises: Exception: If deserialization or file I/O fails. """ - try: - with open(path, "rb") as f: - state = pickle.load(f) + with self._restore_lock: + self._assert_not_restoring() + try: + with open(path, "rb") as f: + state = pickle.load(f) - self.controller_id = state["controller_id"] - self.partitions = state["partitions"] + self.controller_id = state["controller_id"] + self.partitions = state["partitions"] + self._clearing_indexes.clear() - im = state["index_manager"] - self.index_manager.partition_to_indexes = defaultdict(set, im["partition_to_indexes"]) - self.index_manager.reusable_indexes = im["reusable_indexes"] - self.index_manager.global_index_counter = im["global_index_counter"] - self.index_manager.allocated_indexes = im["allocated_indexes"] + im = state["index_manager"] + self.index_manager.partition_to_indexes = defaultdict(set, im["partition_to_indexes"]) + self.index_manager.reusable_indexes = im["reusable_indexes"] + self.index_manager.global_index_counter = im["global_index_counter"] + self.index_manager.allocated_indexes = im["allocated_indexes"] - self.sampler.load_checkpoint(state["sampler"]) + self.sampler.load_checkpoint(state["sampler"]) - logger.info(f"[{self.controller_id}]: restored from {path}") - except Exception as e: - raise RuntimeError(f"[{self.controller_id}]: load checkpoint failed: {e}") from e + logger.info(f"[{self.controller_id}]: restored from {path}") + except Exception as e: + raise RuntimeError(f"[{self.controller_id}]: load checkpoint failed: {e}") from e def register_sampler( self, diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index 7b1a74ed..22325c10 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -15,8 +15,8 @@ """Persist selected rows and restore them by key without checkpointing controller state. -Version 3 preserves field schemas alongside independent row records in each storage-unit shard. The manifest -maps source indexes to byte offsets, so current owner units can read only their rows +Version 3 preserves field schemas alongside independent row records in each shard. +The manifest maps source indexes to byte offsets, so current owner units read only their rows when restoring into a different topology. Version 1 remains readable via KV puts. Layout:: @@ -36,11 +36,12 @@ from collections import defaultdict from pathlib import Path from typing import Any +from uuid import uuid4 import torch from tensordict import TensorDict -from transfer_queue.storage.dump_io import pack_dump_field, read_dump_row, validate_dump_values +from transfer_queue.storage.dump_io import RestorePendingError, pack_dump_field, read_dump_row, validate_dump_values from transfer_queue.utils import compact_pickle from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.tensor_utils import pack_field_values @@ -118,7 +119,13 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") unique_keys = list(dict.fromkeys(keys)) - dump_dir = Path(dump_dir).absolute() + dump_dir = Path(dump_dir).resolve() + client = _maybe_create_tq_client() + if hasattr(client, "check_data_loads"): + client.check_data_loads(str(dump_dir)) + marker = dump_dir.with_name(dump_dir.name + ".restore") + if marker.exists(): + raise RestorePendingError(marker.read_text().strip()) _recover_dump(dump_dir) tmp_dir = dump_dir.parent / (dump_dir.name + ".tmp") old_dir = dump_dir.with_name(dump_dir.name + ".old") @@ -127,7 +134,6 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - tmp_dir.mkdir(parents=True) try: - client = _maybe_create_tq_client() row_index = ( client.describe_data_dump(partition_id, unique_keys) if unique_keys @@ -245,8 +251,8 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: SimpleStorage units read their assigned indexed records directly and in parallel. Version-1 dumps and other backends use the compatible KV put path. New keys receive current indexes; unrelated rows and fields remain untouched. Writers and clears for - these keys must be paused during restore. Failure may leave partial payload writes; - retry the same dump after correcting the failure. + these keys must be paused during restore. Failure may leave partial payload writes. + On RestorePendingError, call recover_data_load before retrying or clearing. Args: dump_dir: Directory previously written by ``dump_data_by_key``. For direct @@ -265,7 +271,13 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: if _TQ_CONTROLLER is None: raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") - dump_dir = Path(dump_dir).absolute() + dump_dir = Path(dump_dir).resolve() + client = _maybe_create_tq_client() + if hasattr(client, "check_data_loads"): + client.check_data_loads(str(dump_dir)) + marker = dump_dir.with_name(dump_dir.name + ".restore") + if marker.exists(): + raise RestorePendingError(marker.read_text().strip()) _recover_dump(dump_dir) info_path = dump_dir / _DUMP_INFO_FILE if not info_path.exists(): @@ -334,7 +346,21 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: if dump_info["format_version"] >= 3: client.validate_dump_schema(partition_id, row_index["field_schema"]) if dump_info["format_version"] >= 2 and hasattr(getattr(client, "storage_manager", None), "load_rows_by_index"): - client.load_rows_by_key(partition_id, rows, shards) + restore_id = uuid4().hex + if rows: + with marker.open("x") as f: + f.write(restore_id) + _fsync_file(f) + _fsync_directory(marker.parent) + try: + client.load_rows_by_key(partition_id, rows, shards, str(dump_dir), restore_id) + except RestorePendingError: + raise + except Exception: + marker.unlink(missing_ok=True) + raise + else: + marker.unlink(missing_ok=True) else: _load_via_kv(partition_id, rows, shards, dump_info["format_version"]) @@ -396,3 +422,22 @@ def _load_via_kv(partition_id: str, rows: dict[str, Any], shards: list[dict], ve keys = [key for key, row in rows.items() if not row["fields"]] if keys: kv_batch_put(keys, partition_id, tags=[rows[key]["tag"] for key in keys]) + + +def recover_data_load(dump_dir: str | Path) -> None: + """Settle an interrupted restore before retrying or releasing destination indexes. + + This cancels work that has not claimed permission and waits for known writers. + Running or unreachable units keep the reservation. Retry recovery once those + units can report completion; a lost unit requires restarting the whole TQ system. + Partial payload writes remain, but subsequent index reuse is safe after success. + """ + from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client + + if _TQ_CONTROLLER is None: + raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") + dump_dir = Path(dump_dir).resolve() + marker = dump_dir.with_name(dump_dir.name + ".restore") + ids = [marker.read_text().strip()] if marker.exists() else [] + _maybe_create_tq_client().recover_data_load(str(dump_dir), ids) + marker.unlink(missing_ok=True) diff --git a/transfer_queue/interface.py b/transfer_queue/interface.py index eca8762e..33abf997 100644 --- a/transfer_queue/interface.py +++ b/transfer_queue/interface.py @@ -1130,6 +1130,7 @@ def load_checkpoint( meta = json.load(f) client = _maybe_create_tq_client() + client.check_data_loads() controller_path = checkpoint_dir / _CONTROLLER_FILE if not controller_path.exists(): diff --git a/transfer_queue/storage/dump_io.py b/transfer_queue/storage/dump_io.py index 0a202670..0a2eb3a6 100644 --- a/transfer_queue/storage/dump_io.py +++ b/transfer_queue/storage/dump_io.py @@ -72,3 +72,13 @@ def pack_dump_field(values: list, schema: dict): if schema["is_nested"]: return torch.nested.as_nested_tensor(values, layout=torch.jagged) return torch.stack(values) + + +class RestorePendingError(RuntimeError): + """The controller still reserves indexes until remote restore activity is settled.""" + + def __init__(self, restore_id: str): + self.restore_id = restore_id + super().__init__( + f"Restore {restore_id} has an unknown outcome; run recover_data_load before retrying or clearing" + ) diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 06a88303..52e5b7b2 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -860,7 +860,9 @@ async def dump_rows_by_index( for pos, ((su_id, _), offsets) in enumerate(zip(targets, row_offsets, strict=True)) ] - async def load_rows_by_index(self, shards: list[dict[str, Any]]) -> list[dict[str, Any]]: + async def load_rows_by_index( + self, shards: list[dict[str, Any]], restore: dict | None = None + ) -> list[dict[str, Any]]: """Have current owner units read assigned byte ranges concurrently.""" assignments = defaultdict(list) for shard in shards: @@ -873,10 +875,15 @@ async def load_rows_by_index(self, shards: list[dict[str, Any]]) -> list[dict[st } ) results = await asyncio.gather( - *(self._load_selected_rows(shards, target_storage_unit=unit_id) for unit_id, shards in assignments.items()), + *( + self._load_selected_rows( + shards, target_storage_unit=unit_id, **({"restore": restore} if restore else {}) + ) + for unit_id, shards in assignments.items() + ), return_exceptions=True, ) - # Wait for every unit before returning an error; callers may retry or clean up. + # Local RPC completion is not remote completion; the controller retains reservations on timeout. for result in results: if isinstance(result, BaseException): raise result @@ -893,13 +900,14 @@ async def _load_selected_rows( self, shards: list[dict[str, Any]], target_storage_unit: str, + restore: dict | None = None, socket: zmq.Socket = None, ) -> dict[str, Any]: request = ZMQMessage.create( request_type=ZMQRequestType.LOAD_ROWS, sender_id=self.storage_manager_id, receiver_id=target_storage_unit, - body={"shards": shards}, + body={"shards": shards, "restore": restore}, ) await socket.send_multipart(request.serialize(), copy=False) response = ZMQMessage.deserialize(await socket.recv_multipart(copy=False)) @@ -909,6 +917,23 @@ async def _load_selected_rows( ) return response.body + async def report_restore(self, restore: dict) -> None: + """Ask units to resend cached terminal results; unknown units cannot grant release.""" + await asyncio.gather( + *(self._report_restore_unit(restore, target_storage_unit=unit) for unit in self.storage_unit_infos), + return_exceptions=True, + ) + + @with_storage_unit_socket + async def _report_restore_unit(self, restore: dict, target_storage_unit: str, socket: zmq.Socket = None) -> None: + request = ZMQMessage.create( + request_type=ZMQRequestType.REPORT_RESTORE, sender_id=self.storage_manager_id, body=restore + ) + await socket.send_multipart(request.serialize()) + response = ZMQMessage.deserialize(await socket.recv_multipart()) + if response.request_type != ZMQRequestType.REPORT_RESTORE_RESPONSE or not response.body.get("success"): + raise RuntimeError("Storage unit could not report its restore result") + async def save_checkpoint(self, checkpoint_dir: str) -> None: """Dump all storage units to the storage_units/ subdirectory of checkpoint_dir. diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index 393d959e..a2a45b4b 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -798,6 +798,7 @@ def __init__(self, config: dict[str, Any]): self.proxy_thread: Thread | None = None self.worker_thread: Thread | None = None + self._restore_results: dict[str, dict] = {} self._metrics: TQMetricsExporter | None = None self._init_zmq_socket() @@ -1015,6 +1016,8 @@ def _process_one_worker_request(self, worker_socket: zmq.Socket, monitor: Any) - elif operation == ZMQRequestType.LOAD_ROWS: # type: ignore[arg-type] with monitor.measure(op_type="LOAD_ROWS"): response_msg = self._handle_load_rows(request_msg) + elif operation == ZMQRequestType.REPORT_RESTORE: + response_msg = self._handle_report_restore(request_msg) elif operation == ZMQRequestType.LOAD_STORAGE_CHECKPOINT: # type: ignore[arg-type] response_msg = self._handle_load_checkpoint(request_msg) else: @@ -1345,7 +1348,84 @@ def _handle_dump_rows(self, request: ZMQMessage) -> ZMQMessage: body={"success": False, "message": str(e)}, ) + def _restore_controller_request(self, context: dict, action: str, result: dict | None = None) -> None: + socket = create_zmq_socket(self.zmq_context, zmq.DEALER, context["controller_ip"]) + try: + socket.setsockopt(zmq.RCVTIMEO, 10000) + socket.setsockopt(zmq.SNDTIMEO, 10000) + socket.connect(context["controller_address"]) + request = ZMQMessage.create( + request_type=ZMQRequestType.RESTORE_UNIT, + sender_id=self.storage_unit_id, + body={ + "restore_id": context["restore_id"], + "unit_id": self.storage_unit_id, + "action": action, + "result": result, + }, + ) + socket.send_multipart(request.serialize()) + response = ZMQMessage.deserialize(socket.recv_multipart()) + if response.request_type != ZMQRequestType.RESTORE_UNIT_RESPONSE: + raise RuntimeError(response.body.get("message", "Restore permission rejected")) + finally: + socket.close(linger=0) + + def _handle_report_restore(self, request: ZMQMessage) -> ZMQMessage: + context = request.body + result = self._restore_results.get(context["restore_id"]) + try: + if result is not None: + self._restore_controller_request(context, "complete", result) + self._restore_results.pop(context["restore_id"], None) + return ZMQMessage.create( + request_type=ZMQRequestType.REPORT_RESTORE_RESPONSE, + sender_id=self.storage_unit_id, + body={"success": True}, + ) + except Exception as e: + return ZMQMessage.create( + request_type=ZMQRequestType.REPORT_RESTORE_RESPONSE, + sender_id=self.storage_unit_id, + body={"success": False, "message": str(e)}, + ) + def _handle_load_rows(self, request: ZMQMessage) -> ZMQMessage: + """Claim permission before touching storage; report completion independently of the caller.""" + context = request.body.get("restore") + if context is None: + return ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, + sender_id=self.storage_unit_id, + body={"success": False, "message": "Missing restore reservation"}, + ) + try: + self._restore_controller_request(context, "claim") + except zmq.error.Again as e: + # A lost claim ACK may leave the controller in running state even though + # this worker will not write; retain a terminal result for recovery. + self._restore_results[context["restore_id"]] = {"success": False, "message": str(e)} + return ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, + sender_id=self.storage_unit_id, + body={"success": False, "message": "Restore claim outcome unknown"}, + ) + except Exception as e: + return ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, + sender_id=self.storage_unit_id, + body={"success": False, "message": str(e)}, + ) + response = self._load_rows(request) + self._restore_results[context["restore_id"]] = response.body + try: + self._restore_controller_request(context, "complete", response.body) + self._restore_results.pop(context["restore_id"], None) + except Exception as e: + logger.warning("[%s]: restore result retained for recovery: %s", self.storage_unit_id, e) + return response + + def _load_rows(self, request: ZMQMessage) -> ZMQMessage: """Seek to assigned records and merge them into local storage at current indexes.""" updates = [] bytes_read = 0 diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index b32273e4..1ef84c28 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -145,6 +145,17 @@ class ZMQRequestType(ExplicitEnum): VALIDATE_DUMP_SCHEMA = "VALIDATE_DUMP_SCHEMA" VALIDATE_DUMP_SCHEMA_RESPONSE = "VALIDATE_DUMP_SCHEMA_RESPONSE" + BEGIN_RESTORE = "BEGIN_RESTORE" + BEGIN_RESTORE_RESPONSE = "BEGIN_RESTORE_RESPONSE" + RESTORE_UNIT = "RESTORE_UNIT" + RESTORE_UNIT_RESPONSE = "RESTORE_UNIT_RESPONSE" + FINISH_RESTORE = "FINISH_RESTORE" + FINISH_RESTORE_RESPONSE = "FINISH_RESTORE_RESPONSE" + LIST_RESTORES = "LIST_RESTORES" + LIST_RESTORES_RESPONSE = "LIST_RESTORES_RESPONSE" + REPORT_RESTORE = "REPORT_RESTORE" + REPORT_RESTORE_RESPONSE = "REPORT_RESTORE_RESPONSE" + class ZMQServerInfo: """ From 06898dac54320fc0504635e82bf8672a9b97306d Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 24 Sep 2026 15:40:14 +0800 Subject: [PATCH 08/30] fix: serialize dump publication and recovery across processes Signed-off-by: OutstanderWang --- docs/data_dump.md | 13 +++- tests/test_dump_lock.py | 132 ++++++++++++++++++++++++++++++++++++ transfer_queue/data_dump.py | 40 ++++++++++- 3 files changed, 180 insertions(+), 5 deletions(-) create mode 100644 tests/test_dump_lock.py diff --git a/docs/data_dump.md b/docs/data_dump.md index 0798e5b7..80a83cde 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -15,8 +15,9 @@ tq.load_data_by_key("/shared/dumps/selected") ``` Pause writes and clears for these keys during both operations. A dump is not an -atomic snapshot of concurrent writers. Multiple publishers must not use the same -dump directory concurrently. +atomic snapshot of concurrent writers. Calls on the same dump path are serialized +by an exclusive lock in a stable sibling `.lock` file. Different dump paths remain +independent, and storage units within a load still read in parallel. ## Distributed I/O @@ -92,7 +93,13 @@ SimpleStorage. Publication writes and syncs `.tmp`, moves the old directory to `.old`, publishes the new directory, and syncs its parent before deleting the backup. If publication is interrupted while the main directory is absent, the next dump, load or row-index -read recovers `.old`. A backup-cleanup error does not invalidate a published dump. +read recovers `.old`. Readers perform recovery only while holding the same lock as +publishers, so a healthy rename window is never mistaken for a crashed writer. +The load keeps the lock through all remote reads; a pending load marker continues +to prevent replacement after a timeout or client exit. Do not delete the sibling +lock file: unlinking it can create two independent locks for the same dump. +The shared filesystem must provide cross-node advisory locking (not local-only +locks). A backup-cleanup error does not invalidate a published dump. Restore is not transactional: payload writes before a failure remain. Every load has a unique ID. The controller blocks clearing/reusing its destination indexes diff --git a/tests/test_dump_lock.py b/tests/test_dump_lock.py new file mode 100644 index 00000000..d79d5f9a --- /dev/null +++ b/tests/test_dump_lock.py @@ -0,0 +1,132 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Publication, reads, and recovery coordinate across processes.""" + +import multiprocessing +from pathlib import Path +from types import SimpleNamespace + +import pytest + +from transfer_queue import data_dump, interface + + +def _configure(): + interface._TQ_CONTROLLER = object() + interface._maybe_create_tq_client = lambda: SimpleNamespace(validate_dump_schema=lambda *_: None) + + +def _publisher(path, paused, release): + _configure() + rename = Path.rename + + def pause_at_publish(source, target): + if source == path.with_name(path.name + ".tmp"): + paused.set() + if not release.wait(20): + raise RuntimeError("Test publisher was not released") + return rename(source, target) + + Path.rename = pause_at_publish + data_dump.dump_data_by_key(path, [], "new") + + +def _reader(path, started, results): + started.set() + results.put(data_dump.read_row_index(path)["partition_id"]) + + +def _loader(path, paused, release): + _configure() + + def pause_load(*args): + paused.set() + if not release.wait(20): + raise RuntimeError("Test loader was not released") + + data_dump._load_via_kv = pause_load + data_dump.load_data_by_key(path) + + +@pytest.mark.parametrize("kill_writer", [False, True]) +def test_reader_waits_for_publisher_and_recovers_after_exit(tmp_path, monkeypatch, kill_writer): + monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) + monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: object()) + path = tmp_path / "dump" + data_dump.dump_data_by_key(path, [], "old") + context = multiprocessing.get_context("spawn") + paused, release, reader_started = context.Event(), context.Event(), context.Event() + results = context.Queue() + writer = context.Process(target=_publisher, args=(path, paused, release)) + reader = context.Process(target=_reader, args=(path, reader_started, results)) + writer.start() + try: + assert paused.wait(20) + reader.start() + assert reader_started.wait(20) + reader.join(0.2) + assert reader.is_alive(), "Reader rolled back a live publisher" + if kill_writer: + writer.terminate() + else: + release.set() + writer.join(20) + assert not writer.is_alive() + assert results.get(timeout=20) == ("old" if kill_writer else "new") + reader.join(20) + assert reader.exitcode == 0 + if not kill_writer: + assert writer.exitcode == 0 + assert path.with_name("dump.lock").exists() + finally: + # A terminated process may have held the Event's semaphore; do not reuse it. + for process in (reader, writer): + if process.pid and process.is_alive(): + process.terminate() + process.join(10) + results.close() + + +def test_publish_waits_until_load_finishes(tmp_path, monkeypatch): + monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) + monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: object()) + path = tmp_path / "dump" + data_dump.dump_data_by_key(path, [], "old") + context = multiprocessing.get_context("spawn") + load_paused, load_release, publish_paused, publish_release = [context.Event() for _ in range(4)] + loader = context.Process(target=_loader, args=(path, load_paused, load_release)) + writer = context.Process(target=_publisher, args=(path, publish_paused, publish_release)) + loader.start() + try: + assert load_paused.wait(20) + writer.start() + assert not publish_paused.wait(0.5) + assert path.exists() + load_release.set() + loader.join(20) + assert loader.exitcode == 0 + assert publish_paused.wait(20) + publish_release.set() + writer.join(20) + assert writer.exitcode == 0 + assert data_dump.read_row_index(path)["partition_id"] == "new" + finally: + load_release.set() + publish_release.set() + for process in (writer, loader): + if process.pid and process.is_alive(): + process.terminate() + process.join(10) diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index 22325c10..95cec709 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -29,11 +29,13 @@ shard__.pkl # independent {global_index, fields} records """ +import fcntl import json import os import pickle import shutil from collections import defaultdict +from contextlib import contextmanager from pathlib import Path from typing import Any from uuid import uuid4 @@ -76,6 +78,20 @@ def _fsync_directory(path: Path) -> None: os.close(fd) +@contextmanager +def _dump_lock(dump_dir: str | Path): + """Serialize publication and recovery using a stable sibling inode shared by all callers.""" + directory = Path(dump_dir).resolve() + directory.parent.mkdir(parents=True, exist_ok=True) + lock_path = directory.with_name(directory.name + ".lock") + with lock_path.open("a+b") as lock: + fcntl.flock(lock.fileno(), fcntl.LOCK_EX) + try: + yield directory + finally: + fcntl.flock(lock.fileno(), fcntl.LOCK_UN) + + def _recover_dump(dump_dir: Path) -> None: old_dir = dump_dir.with_name(dump_dir.name + ".old") if not dump_dir.exists() and old_dir.exists(): @@ -91,7 +107,7 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - The directory is replaced wholesale. The previous dump is retained as ``.old`` until publication is durable, and recovered on the next access after interruption. - Only one writer may publish to a given directory at a time. + Access to the same directory is serialized with a sibling file lock. .. note:: **Multi-node limitation**: dump_dir must reside on a shared network filesystem @@ -113,6 +129,11 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - RuntimeError: TransferQueue is not initialized, the partition or a key does not exist, or a storage unit holds no data for a row it was asked to dump. """ + with _dump_lock(dump_dir) as directory: + return _dump_data_by_key(directory, keys, partition_id) + + +def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dict[str, int]: from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client if _TQ_CONTROLLER is None: @@ -237,6 +258,11 @@ def read_row_index(dump_dir: str | Path) -> dict[str, Any]: Raises: FileNotFoundError: The row index is missing. """ + with _dump_lock(dump_dir) as directory: + return _read_row_index(directory) + + +def _read_row_index(dump_dir: Path) -> dict[str, Any]: dump_dir = Path(dump_dir) _recover_dump(dump_dir) row_index_path = dump_dir / _ROW_INDEX_FILE @@ -266,6 +292,11 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: FileNotFoundError: The dump is incomplete. ValueError: The manifest or a row disagrees with the row index. """ + with _dump_lock(dump_dir) as directory: + return _load_data_by_key(directory) + + +def _load_data_by_key(dump_dir: Path) -> dict[str, int]: from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client if _TQ_CONTROLLER is None: @@ -290,7 +321,7 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: f"this build reads versions 1 through {DUMP_FORMAT_VERSION}" ) - row_index = read_row_index(dump_dir) + row_index = _read_row_index(dump_dir) partition_id = row_index["partition_id"] rows = row_index["rows"] @@ -432,6 +463,11 @@ def recover_data_load(dump_dir: str | Path) -> None: units can report completion; a lost unit requires restarting the whole TQ system. Partial payload writes remain, but subsequent index reuse is safe after success. """ + with _dump_lock(dump_dir) as directory: + return _recover_data_load(directory) + + +def _recover_data_load(dump_dir: Path) -> None: from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client if _TQ_CONTROLLER is None: From 461109ef1d36402d79a13b90627183732dcf1715 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Mon, 28 Sep 2026 14:57:29 +0800 Subject: [PATCH 09/30] fix: use union syntax in selective dump test Signed-off-by: neowywang --- tests/e2e/test_data_dump_e2e.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index 1b4335bf..fc0b1bf5 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -440,7 +440,7 @@ async def load(*args, **kwargs): real_open = builtins.open def no_payload_open(path, *args, **kwargs): - if isinstance(path, (str, Path)) and Path(path).name.startswith("shard_") and str(path).endswith(".pkl"): + if isinstance(path, str | Path) and Path(path).name.startswith("shard_") and str(path).endswith(".pkl"): raise AssertionError("Caller opened a payload shard") return real_open(path, *args, **kwargs) From 99ff73a5d6ed4e60e312ec08836da779ac5771f9 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Mon, 28 Sep 2026 14:57:50 +0800 Subject: [PATCH 10/30] fix: narrow validated dump indexes to integers Signed-off-by: neowywang --- transfer_queue/controller.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index cb07518f..c9136c72 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -1790,7 +1790,7 @@ def describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, ), "tag": partition.custom_meta.get(global_index, {}), } - for key, global_index in zip(keys, global_indexes, strict=True) + for key, global_index in zip(keys, cast(list[int], global_indexes), strict=True) } def _assert_not_restoring(self, partition_id: str | None = None) -> None: From d51882da6e889a5066b4d764e8fb6c22cafec2fd Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Mon, 28 Sep 2026 14:57:50 +0800 Subject: [PATCH 11/30] fix: aggregate dump results after checking exceptions Signed-off-by: neowywang --- .../managers/simple_storage_manager.py | 26 +++++++++++-------- 1 file changed, 15 insertions(+), 11 deletions(-) diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 52e5b7b2..6b03d3fd 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -847,18 +847,18 @@ async def dump_rows_by_index( ), return_exceptions=True, ) - for result in row_offsets: - if isinstance(result, BaseException): - raise result + shards = [] + total_rows = 0 + for pos, ((su_id, _), offsets) in enumerate(zip(targets, row_offsets, strict=True)): + if isinstance(offsets, BaseException): + raise offsets + total_rows += len(offsets) + shards.append({"position": pos, "storage_unit_id": su_id, "rows": len(offsets), "row_offsets": offsets}) logger.info( - f"[{self.storage_manager_id}]: dumped {sum(len(offsets) for offsets in row_offsets)} rows " - f"across {len(targets)} shards to {shard_dir_path}" + f"[{self.storage_manager_id}]: dumped {total_rows} rows across {len(targets)} shards to {shard_dir_path}" ) - return [ - {"position": pos, "storage_unit_id": su_id, "rows": len(offsets), "row_offsets": offsets} - for pos, ((su_id, _), offsets) in enumerate(zip(targets, row_offsets, strict=True)) - ] + return shards async def load_rows_by_index( self, shards: list[dict[str, Any]], restore: dict | None = None @@ -884,16 +884,20 @@ async def load_rows_by_index( return_exceptions=True, ) # Local RPC completion is not remote completion; the controller retains reservations on timeout. + updates = [] + bytes_read = 0 for result in results: if isinstance(result, BaseException): raise result + bytes_read += result["bytes_read"] + updates.extend(result["updates"]) logger.info( "[%s]: loaded %s bytes across %s units", self.storage_manager_id, - sum(result["bytes_read"] for result in results), + bytes_read, len(assignments), ) - return [update for result in results for update in result["updates"]] + return updates @with_storage_unit_socket async def _load_selected_rows( From 2e08ab822bc703648e93256bd6aad7b26f5ce6d4 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Mon, 28 Sep 2026 14:57:50 +0800 Subject: [PATCH 12/30] fix: check backend recovery support before cancelling restores Signed-off-by: neowywang --- tests/test_restore_lifecycle.py | 18 ++++++++++++++++++ transfer_queue/client.py | 4 ++++ 2 files changed, 22 insertions(+) diff --git a/tests/test_restore_lifecycle.py b/tests/test_restore_lifecycle.py index 5b85acba..527c280c 100644 --- a/tests/test_restore_lifecycle.py +++ b/tests/test_restore_lifecycle.py @@ -16,11 +16,13 @@ """Restore reservations prevent late writes from corrupting reused indexes.""" from threading import RLock +from unittest.mock import AsyncMock import pytest import torch import zmq +from transfer_queue.client import AsyncTransferQueueClient from transfer_queue.controller import PartitionIndexManager, TransferQueueController from transfer_queue.sampler import SequentialSampler from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData @@ -46,6 +48,22 @@ def begin(controller, restore_id="r"): return controller.begin_restore(restore_id, "/dump", "p", {"k": {"fields": ["x"], "tag": {}}}, ["u"], {}) +@pytest.mark.asyncio +@pytest.mark.parametrize(("active_ids", "saved_ids"), [([], []), (["r"], []), ([], ["r"])]) +async def test_recovery_checks_backend_before_cancelling(active_ids, saved_ids): + client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) + client.storage_manager = object() + client._restore_rpc = AsyncMock(return_value={"restore_ids": active_ids}) + + if active_ids or saved_ids: + with pytest.raises(NotImplementedError, match="does not support selective load recovery"): + await client.async_recover_data_load("/dump", saved_ids) + else: + await client.async_recover_data_load("/dump", saved_ids) + + client._restore_rpc.assert_awaited_once_with(ZMQRequestType.LIST_RESTORES, {"dump_dir": "/dump"}) + + def test_cancel_before_claim_rejects_late_load_and_late_begin(controller): metadata = begin(controller) assert controller.finish_restore("r", commit=False)["finished"] diff --git a/transfer_queue/client.py b/transfer_queue/client.py index ee226bc5..14cc5778 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1286,6 +1286,10 @@ async def async_recover_data_load(self, dump_dir: str, restore_ids: list[str] | """Cancel unclaimed loads and release only operations whose claimed workers have finished.""" response = await self._restore_rpc(ZMQRequestType.LIST_RESTORES, {"dump_dir": dump_dir}) for restore_id in set(response["restore_ids"]) | set(restore_ids or []): + if not hasattr(self.storage_manager, "report_restore"): + raise NotImplementedError( + f"{type(self.storage_manager).__name__} does not support selective load recovery" + ) await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": False}) await self.storage_manager.report_restore(self._restore_context(restore_id)) result = await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": False}) From 704c3258307718a70b56a323a87b8cc1c109b12f Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Mon, 28 Sep 2026 16:28:00 +0800 Subject: [PATCH 13/30] fix: scope strict schema validation to selective restores Signed-off-by: neowywang --- tests/e2e/test_data_dump_e2e.py | 14 +++++++++++ tests/test_data_dump.py | 12 ---------- tests/test_restore_lifecycle.py | 41 +++++++++++++++++++++++++++++++++ transfer_queue/controller.py | 7 +++--- 4 files changed, 59 insertions(+), 15 deletions(-) diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index fc0b1bf5..e13aa8a9 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -516,6 +516,20 @@ def test_incompatible_schema_rejected_before_writes(tq_system, dump_dir, control _assert_rows_equal(tq.kv_batch_get(["k"], "schema", ["x"])["x"], [torch.tensor([1.5])]) +@pytest.mark.parametrize("tensor_first", [True, False]) +def test_regular_put_keeps_legacy_tensor_nontensor_acceptance(tq_system, controller, tensor_first): + values = [torch.tensor([[7]]), NonTensorStack("text")] + if not tensor_first: + values.reverse() + for key, value in zip(["first", "second"], values, strict=True): + tq.kv_batch_put([key], "legacy_put", TensorDict({"x": value}, batch_size=1)) + metadata = tq.get_client().kv_retrieve_meta(["first", "second"], "legacy_put") + assert metadata.is_ready + assert metadata.field_names == ["x"] + snapshot = ray.get(controller.get_partition_snapshot.remote("legacy_put")) + assert snapshot.field_metadata["x"].global_indexes == set(metadata.global_indexes) + + def test_running_restore_blocks_clear_and_dump_until_recovery(tq_system, dump_dir, controller): _put_rows("reserved", ["key"]) tq.dump_data_by_key(dump_dir, ["key"], "reserved") diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index 421934c8..cc9f5c62 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -358,15 +358,3 @@ async def dump(path, target_storage_unit, global_indexes, fields_by_index): with pytest.raises(OSError, match="write failed"): await manager.dump_rows_by_index(str(tmp_path), [0, 1]) assert completed == ["u1"] - - -def test_rejected_schema_does_not_mark_new_row_ready(): - from transfer_queue.controller import DataPartitionStatus - - partition = DataPartitionStatus("p") - schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_non_tensor": False, "is_nested": False}} - assert partition.update_production_status([0], ["x"], schema) - conflict = {"x": {**schema["x"], "dtype": torch.float32}} - assert not partition.update_production_status([1], ["x"], conflict) - assert partition.production_status[1, partition.field_name_mapping["x"]] == 0 - assert partition.field_metadata["x"].global_indexes == {0} diff --git a/tests/test_restore_lifecycle.py b/tests/test_restore_lifecycle.py index 527c280c..b0c2e4dd 100644 --- a/tests/test_restore_lifecycle.py +++ b/tests/test_restore_lifecycle.py @@ -130,6 +130,47 @@ def test_restore_cannot_start_in_the_middle_of_clear(controller): begin(controller) +def test_restore_rejects_schema_before_allocating_indexes(controller): + metadata = controller.kv_retrieve_meta(["existing"], "p", create=True) + partition = controller.partitions["p"] + schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_non_tensor": False, "is_nested": False}} + assert partition.update_production_status(metadata.global_indexes, [], schema) + conflict = {"x": {**schema["x"], "dtype": torch.float32}} + with pytest.raises(ValueError, match="dtype mismatch"): + controller.begin_restore("r", "/dump", "p", {"new": {"fields": ["x"], "tag": {}}}, ["u"], conflict) + assert set(partition.keys_mapping) == {"existing"} + assert not controller.list_restores("/dump") + + +def test_restore_validates_all_updates_before_publishing_readiness(controller): + metadata = controller.kv_retrieve_meta(["existing"], "p", create=True) + partition = controller.partitions["p"] + schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_non_tensor": False, "is_nested": False}} + assert partition.update_production_status(metadata.global_indexes, [], schema) + restored = begin(controller) + controller.restore_unit("r", "u", "claim") + controller.restore_unit( + "r", + "u", + "complete", + { + "success": True, + "updates": [ + {"global_indexes": restored.global_indexes, "field_schema": {"y": schema["x"]}}, + { + "global_indexes": restored.global_indexes, + "field_schema": {"x": {**schema["x"], "dtype": torch.float32}}, + }, + ], + }, + ) + with pytest.raises(ValueError, match="dtype mismatch"): + controller.finish_restore("r", commit=True) + assert "y" not in partition.field_metadata + assert not partition.production_status[restored.global_indexes].any() + assert controller.finish_restore("r", commit=False)["finished"] + + def test_unit_rejects_cancelled_permission_before_reading(controller): begin(controller) controller.finish_restore("r", commit=False) diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index c9136c72..84ed23e3 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -540,13 +540,14 @@ def update_production_status( required_fields = len(self.field_name_mapping) self.ensure_fields_capacity(required_fields) - # Validate all field updates before changing readiness or field metadata. - self.validate_field_schema(field_schema) - self._update_field_metadata(global_indices, field_schema, custom_backend_meta) + # Update production status if self.production_status is not None and global_indices and field_names: field_indices = [self.field_name_mapping.get(f) for f in field_names] self.production_status[torch.tensor(global_indices)[:, None], torch.tensor(field_indices)] = 1 + # Update field metadata + self._update_field_metadata(global_indices, field_schema, custom_backend_meta) + # Save these global_indexes self.global_indexes.update(global_indices) From 3a603fb982a3bf84d32ff235a6588548d20157bc Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Mon, 28 Sep 2026 16:34:23 +0800 Subject: [PATCH 14/30] fix: preserve pending restores after receive timeouts Signed-off-by: neowywang --- docs/data_dump.md | 29 ++++++--- tests/e2e/test_data_dump_e2e.py | 40 ++++++++++++ tests/e2e/test_restore_timeout_e2e.py | 76 ++++++++++++++++++++++- tests/test_restore_lifecycle.py | 88 ++++++++++++++++++++++++++- transfer_queue/client.py | 55 +++++++++++------ transfer_queue/controller.py | 31 +++++----- transfer_queue/data_dump.py | 20 +++--- 7 files changed, 284 insertions(+), 55 deletions(-) diff --git a/docs/data_dump.md b/docs/data_dump.md index 80a83cde..33e2b89b 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -106,25 +106,36 @@ has a unique ID. The controller blocks clearing/reusing its destination indexes and conflicting KV puts while an operation is unresolved. Units must claim that ID before writing; cancellation rejects requests that have not yet claimed permission. A receive timeout never releases a writer that has already claimed permission. +Timeouts leave the operation pending, even if every unit later succeeds. They do +not cancel it or require changing the timeout used by ordinary puts and gets. `RestorePendingError` means remote work is still running or its outcome is unknown. The dump also retains a sibling `.restore` marker so an interrupted client cannot silently allow its files to be replaced. After an interruption, call: ```python -tq.recover_data_load("/shared/dumps/selected") +committed = tq.recover_data_load("/shared/dumps/selected") ``` -Recovery cancels unclaimed work, asks units to resend terminal results, and releases -indexes only after every claimed worker has finished. It does not publish partial -restores as ready or undo payload writes. Retry recovery while a unit is still busy; -a lost unit requires stopping the old TQ actors and restarting the whole TQ system. -Restarting only the controller while old storage actors run is unsupported. After -recovery succeeds, retry the dump or clear its keys. Unknown operations remain -reserved rather than guessing that a timeout stopped remote execution. +Recovery asks units to resend terminal results and commits schemas and tags only +when all units succeeded. It returns `True` for committed loads (or no pending +load), and `False` for a failed or cancelled load after all claimed workers stopped. +Running or unknown work raises `RestorePendingError`; retry recovery later without +reloading payloads. The controller remembers terminal outcomes so a lost commit +reply can be confirmed safely by retrying recovery. + +To abandon the load explicitly, use `recover_data_load(dump_dir, cancel=True)`. +This denies unclaimed work and retains the reservation until claimed workers stop. +Cancellation cannot undo a load that has already committed. A failed unit also +cancels remaining unclaimed work; partial payload writes are never rolled back. +After cancellation settles, retry the load or clear its keys. A lost unit requires +stopping the old TQ actors and restarting the whole TQ system; restarting only the +controller while old storage actors run is unsupported. Writers that already hold low-level metadata must remain paused throughout recovery. -The reservation covers the destination partition, so unrelated partitions can proceed. +The controller reservation covers only the destination partition. Each storage unit +still serves requests on one worker thread: other partitions using that unit can +wait behind a load. The 128-row batches bound memory, not request latency. ## Tests diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index e13aa8a9..d7d616ba 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -530,6 +530,46 @@ def test_regular_put_keeps_legacy_tensor_nontensor_acceptance(tq_system, control assert snapshot.field_metadata["x"].global_indexes == set(metadata.global_indexes) +@pytest.mark.parametrize("lost_reply", ["load", "commit"]) +def test_recovery_commits_after_lost_reply(tq_system, dump_dir, monkeypatch, lost_reply): + import zmq + + from transfer_queue.utils.zmq_utils import ZMQRequestType + + _put_rows("lost_reply", ["key"]) + tq.dump_data_by_key(dump_dir, ["key"], "lost_reply") + client = tq.get_client() + client.clear_partition("lost_reply") + manager = client.storage_manager + original_load = manager._load_selected_rows + original_rpc = client._restore_rpc + loads = [] + + async def load(*args, **kwargs): + loads.append(kwargs["target_storage_unit"]) + result = await original_load(*args, **kwargs) + if lost_reply == "load": + raise zmq.error.Again() + return result + + async def rpc(action, body): + result = await original_rpc(action, body) + if lost_reply == "commit" and action == ZMQRequestType.FINISH_RESTORE and body["commit"]: + raise zmq.error.Again() + return result + + with monkeypatch.context() as patcher: + patcher.setattr(manager, "_load_selected_rows", load) + patcher.setattr(client, "_restore_rpc", rpc) + with pytest.raises(tq.RestorePendingError): + tq.load_data_by_key(dump_dir) + assert dump_dir.with_name(dump_dir.name + ".restore").exists() + assert tq.recover_data_load(dump_dir) is True + assert not dump_dir.with_name(dump_dir.name + ".restore").exists() + assert len(loads) == 1 + _assert_rows_equal(tq.kv_batch_get(["key"], "lost_reply", ["input_ids"])["input_ids"], [_row_input_ids(0)]) + + def test_running_restore_blocks_clear_and_dump_until_recovery(tq_system, dump_dir, controller): _put_rows("reserved", ["key"]) tq.dump_data_by_key(dump_dir, ["key"], "reserved") diff --git a/tests/e2e/test_restore_timeout_e2e.py b/tests/e2e/test_restore_timeout_e2e.py index a6a9facb..9113bc39 100644 --- a/tests/e2e/test_restore_timeout_e2e.py +++ b/tests/e2e/test_restore_timeout_e2e.py @@ -13,6 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import multiprocessing import threading import time from unittest.mock import patch @@ -30,6 +31,78 @@ from transfer_queue.utils.zmq_utils import ZMQMessage +def _check_claimed_load_survives_receive_timeout(tmp_path): + import transfer_queue.storage.managers.simple_storage_manager as manager_module + + ray.init(namespace="review_claimed_timeout") + try: + tq.init( + OmegaConf.create( + { + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "num_data_storage_units": 1, + "total_storage_size": 20, + }, + } + } + ) + ) + tq.kv_batch_put(["key"], "p", TensorDict({"x": torch.tensor([[7]])}, batch_size=1), tags=[{"saved": True}]) + dump_dir = tmp_path / "dump" + tq.dump_data_by_key(dump_dir, ["key"], "p") + tq.get_client().clear_partition("p") + actor = ray.get_actor("TransferQueueStorageUnit#0") + claimed = tmp_path / "claimed" + release = tmp_path / "release" + + def pause_after_claim(unit): + original = unit._load_rows + + def delayed(request): + claimed.touch() + deadline = time.monotonic() + 10 + while not release.exists() and time.monotonic() < deadline: + time.sleep(0.01) + return original(request) + + unit._load_rows = delayed + + ray.get(actor.__ray_call__.remote(pause_after_claim)) + with patch.object(manager_module, "TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT", 0.5): + with pytest.raises(tq.RestorePendingError): + tq.load_data_by_key(dump_dir) + assert claimed.exists() + assert dump_dir.with_name(dump_dir.name + ".restore").exists() + with pytest.raises(RuntimeError, match="unresolved"): + tq.get_client().clear_partition("p") + with pytest.raises(tq.RestorePendingError): + tq.dump_data_by_key(dump_dir, ["key"], "p") + release.touch() + assert tq.recover_data_load(dump_dir) is True + assert not dump_dir.with_name(dump_dir.name + ".restore").exists() + assert tq.kv_batch_get(["key"], "p", ["x"])["x"][0].item() == 7 + assert tq.get_client().kv_retrieve_meta(["key"], "p").custom_meta == [{"saved": True}] + finally: + tq.close() + ray.shutdown() + + +def test_claimed_load_survives_receive_timeout(tmp_path): + # Dynamic Ray actor instrumentation must not affect later in-process unit mocks. + context = multiprocessing.get_context("spawn") + process = context.Process(target=_check_claimed_load_survives_receive_timeout, args=(tmp_path,)) + process.start() + try: + process.join(60) + assert process.exitcode == 0 + finally: + if process.is_alive(): + process.terminate() + process.join(10) + + def test_cancelled_delayed_load_cannot_overwrite_reused_index(tmp_path): ray.init(namespace="review_timeout") tq.init( @@ -92,8 +165,9 @@ async def timeout_load(shards, target_storage_unit, restore): ctx.term() with patch.object(manager, "_load_selected_rows", timeout_load): - with pytest.raises(zmq.error.Again): + with pytest.raises(tq.RestorePendingError): tq.load_data_by_key(tmp_path / "dump") + assert tq.recover_data_load(tmp_path / "dump", cancel=True) is False value = tq.kv_batch_get(["key"], "p", ["x"])["x"][0].item() assert value == 2 diff --git a/tests/test_restore_lifecycle.py b/tests/test_restore_lifecycle.py index b0c2e4dd..24f48ae9 100644 --- a/tests/test_restore_lifecycle.py +++ b/tests/test_restore_lifecycle.py @@ -16,6 +16,7 @@ """Restore reservations prevent late writes from corrupting reused indexes.""" from threading import RLock +from types import SimpleNamespace from unittest.mock import AsyncMock import pytest @@ -25,6 +26,7 @@ from transfer_queue.client import AsyncTransferQueueClient from transfer_queue.controller import PartitionIndexManager, TransferQueueController from transfer_queue.sampler import SequentialSampler +from transfer_queue.storage.dump_io import RestorePendingError from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType @@ -39,7 +41,7 @@ def controller(): controller.sampler = SequentialSampler() controller._restore_lock = RLock() controller._restores = {} - controller._cancelled_restores = set() + controller._restore_outcomes = {} controller._clearing_indexes = set() return controller @@ -72,10 +74,10 @@ def test_cancel_before_claim_rejects_late_load_and_late_begin(controller): assert current.global_indexes == metadata.global_indexes with pytest.raises(RuntimeError, match="no longer active"): controller.restore_unit("r", "u", "claim") - with pytest.raises(RuntimeError, match="already cancelled"): + with pytest.raises(RuntimeError, match="already committed or cancelled"): begin(controller) controller.finish_restore("not-arrived", commit=False) - with pytest.raises(RuntimeError, match="already cancelled"): + with pytest.raises(RuntimeError, match="already committed or cancelled"): begin(controller, "not-arrived") @@ -255,3 +257,83 @@ def lose_ack(context, action, result=None): ZMQMessage.create(request_type=ZMQRequestType.REPORT_RESTORE, sender_id="test", body={"restore_id": "r"}) ).body["success"] assert controller.finish_restore("r", commit=False)["finished"] + + +def test_failed_unit_cancels_pending_units_but_waits_for_running_units(controller): + rows = {f"k{i}": {"fields": ["x"], "tag": {}} for i in range(3)} + controller.begin_restore("r", "/dump", "p", rows, ["u0", "u1", "u2"], {}) + controller.restore_unit("r", "u0", "claim") + controller.restore_unit("r", "u1", "claim") + controller.restore_unit("r", "u0", "complete", {"success": False}) + assert not controller.finish_restore("r", commit=True)["finished"] + with pytest.raises(RuntimeError, match="cannot start"): + controller.restore_unit("r", "u2", "claim") + with pytest.raises(RuntimeError, match="unresolved"): + controller.clear_partition("p") + controller.restore_unit("r", "u1", "complete", {"success": True, "updates": []}) + assert controller.finish_restore("r", commit=True) == {"finished": True, "committed": False} + assert not controller.partitions["p"].field_metadata + + +def test_unknown_restore_stays_pending_until_explicit_cancellation(controller): + assert not controller.finish_restore("unknown", commit=True)["finished"] + assert controller.finish_restore("unknown", commit=False) == {"finished": True, "committed": False} + with pytest.raises(RuntimeError, match="already committed or cancelled"): + begin(controller, "unknown") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("interruption", ["running_timeout", "lost_load_reply", "lost_commit_reply", "lost_complete"]) +async def test_pending_load_can_commit_without_replaying_payload(controller, interruption): + client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) + client._restore_context = lambda restore_id: {"restore_id": restore_id} + completed = {} + fail_commit_reply = interruption == "lost_commit_reply" + + async def rpc(action, body): + nonlocal fail_commit_reply + if action == ZMQRequestType.BEGIN_RESTORE: + return {"metadata": controller.begin_restore(**body)} + if action == ZMQRequestType.LIST_RESTORES: + return {"restore_ids": controller.list_restores(body["dump_dir"])} + result = controller.finish_restore(**body) + if body["commit"] and result["finished"] and fail_commit_reply: + fail_commit_reply = False + raise zmq.error.Again() + return result + + async def load(shards, context): + controller.restore_unit("r", "u", "claim") + index = controller._restores["r"]["metadata"].global_indexes[0] + schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_nested": False, "is_non_tensor": False}} + completed.update(success=True, updates=[{"global_indexes": [index], "field_schema": schema}]) + if interruption in ("lost_load_reply", "lost_commit_reply"): + controller.restore_unit("r", "u", "complete", completed) + if interruption in ("running_timeout", "lost_load_reply"): + raise zmq.error.Again() + + async def report(context): + controller.restore_unit("r", "u", "complete", completed) + + client._restore_rpc = rpc + client.storage_manager = SimpleNamespace( + storage_unit_infos={"u": None}, load_rows_by_index=AsyncMock(side_effect=load), report_restore=report + ) + rows = {"k": {"fields": ["x"], "tag": {"saved": True}}} + with pytest.raises(RestorePendingError): + await client.async_load_rows_by_key("p", rows, [], "/dump", "r") + if "r" in controller._restores: + assert not controller._restores["r"]["aborting"] + with pytest.raises(RuntimeError, match="unresolved"): + controller.clear_partition("p") + assert await client.async_recover_data_load("/dump", ["r"]) is True + assert await client.async_recover_data_load("/dump", ["r"]) is True + client.storage_manager.load_rows_by_index.assert_awaited_once() + partition = controller.partitions["p"] + index = partition.keys_mapping["k"] + assert partition.production_status[index, partition.field_name_mapping["x"]] == 1 + assert partition.custom_meta[index] == {"saved": True} + # Duplicate finalization must not publish metadata again after the key is cleared. + controller.clear_partition("p") + assert controller.finish_restore("r", commit=True) == {"finished": True, "committed": True} + assert "p" not in controller.partitions diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 14cc5778..a327acef 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1232,6 +1232,18 @@ async def _restore_rpc(self, request_type: ZMQRequestType, body: dict, socket=No response = await self._request_controller(socket, request_type, response_type, body) return response.body + async def _finish_data_load(self, restore_id: str, *, commit: bool) -> bool: + """Keep the reservation when completion is pending or its reply is lost.""" + try: + result = await self._restore_rpc( + ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": commit} + ) + except (zmq.error.Again, TimeoutError) as error: + raise RestorePendingError(restore_id) from error + if not result["finished"]: + raise RestorePendingError(restore_id) + return result["committed"] + async def async_load_rows_by_key( self, partition_id: str, @@ -1265,36 +1277,39 @@ async def async_load_rows_by_key( for record in shard["records"]: record["target_index"] = target_indexes[record["key"]] await manager.load_rows_by_index(shards, self._restore_context(restore_id)) - result = await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": True}) - if not result["finished"]: - raise RestorePendingError(restore_id) + if not await self._finish_data_load(restore_id, commit=True): + raise RuntimeError("Restore failed or was cancelled") + except RestorePendingError: + raise + except (zmq.error.Again, TimeoutError) as error: + raise RestorePendingError(restore_id) from error except BaseException as error: try: - result = await asyncio.shield( - self._restore_rpc( - ZMQRequestType.FINISH_RESTORE, - {"restore_id": restore_id, "commit": False}, - ) - ) + await asyncio.shield(self._finish_data_load(restore_id, commit=False)) except BaseException: raise RestorePendingError(restore_id) from error - if not result["finished"]: - raise RestorePendingError(restore_id) from error raise - async def async_recover_data_load(self, dump_dir: str, restore_ids: list[str] | None = None) -> None: - """Cancel unclaimed loads and release only operations whose claimed workers have finished.""" + async def async_recover_data_load( + self, dump_dir: str, restore_ids: list[str] | None = None, *, cancel: bool = False + ) -> bool: + """Finish successful loads, or explicitly cancel; return whether all loads committed.""" response = await self._restore_rpc(ZMQRequestType.LIST_RESTORES, {"dump_dir": dump_dir}) + committed = True for restore_id in set(response["restore_ids"]) | set(restore_ids or []): if not hasattr(self.storage_manager, "report_restore"): raise NotImplementedError( f"{type(self.storage_manager).__name__} does not support selective load recovery" ) - await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": False}) + if cancel: + try: + await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": False}) + except (zmq.error.Again, TimeoutError) as error: + raise RestorePendingError(restore_id) from error await self.storage_manager.report_restore(self._restore_context(restore_id)) - result = await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": False}) - if not result["finished"]: - raise RestorePendingError(restore_id) + outcome = await self._finish_data_load(restore_id, commit=not cancel) + committed = committed and outcome + return committed async def async_check_data_loads(self, dump_dir: str | None = None) -> None: """Reject replacing a dump whose files may still be read by a storage unit.""" @@ -1982,9 +1997,9 @@ def load_rows_by_key( """Restore selected payloads while the controller reserves destination indexes.""" return self._load_rows_by_key(partition_id, rows, shards, dump_dir, restore_id) - def recover_data_load(self, dump_dir: str, restore_ids: list[str] | None = None) -> None: - """Settle an interrupted load before retrying, clearing, or replacing its dump.""" - return self._recover_data_load(dump_dir, restore_ids) + def recover_data_load(self, dump_dir: str, restore_ids: list[str] | None = None, *, cancel: bool = False) -> bool: + """Finish or cancel interrupted loads; return whether every load committed.""" + return self._recover_data_load(dump_dir, restore_ids, cancel=cancel) def check_data_loads(self, dump_dir: str | None = None) -> None: """Reject replacing a dump while a remote restore may still read it.""" diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index 84ed23e3..1361ea36 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -1011,7 +1011,7 @@ def __init__( self.index_manager = PartitionIndexManager() # partition_id -> global_indexes self._restore_lock = RLock() self._restores: dict[str, dict[str, Any]] = {} - self._cancelled_restores: set[str] = set() + self._restore_outcomes: dict[str, bool] = {} self._clearing_indexes: set[int] = set() # Connected storage managers tracking @@ -1806,8 +1806,8 @@ def begin_restore( ): """Reserve a destination partition until every possible remote writer is settled.""" with self._restore_lock: - if restore_id in self._cancelled_restores: - raise RuntimeError("Restore was already cancelled") + if restore_id in self._restore_outcomes: + raise RuntimeError("Restore was already committed or cancelled") if restore_id in self._restores: return self._restores[restore_id]["metadata"] self._assert_not_restoring(partition_id) @@ -1838,6 +1838,8 @@ def restore_unit(self, restore_id: str, unit_id: str, action: str, result: dict """Grant one execution and record its terminal result; unknown IDs never grant writes.""" with self._restore_lock: restore = self._restores.get(restore_id) + if restore_id in self._restore_outcomes and action == "complete": + return if restore is None or unit_id not in restore["units"]: raise RuntimeError("Restore is no longer active") state = restore["units"][unit_id] @@ -1852,15 +1854,17 @@ def restore_unit(self, restore_id: str, unit_id: str, action: str, result: dict raise RuntimeError(f"Invalid restore transition {state} -> {action}") def finish_restore(self, restore_id: str, commit: bool) -> dict: - """Release indexes only after all claimed writers terminate; cancel unclaimed work.""" + """Settle a restore after its workers terminate, retaining the outcome for lost replies.""" with self._restore_lock: + if restore_id in self._restore_outcomes: + return {"finished": True, "committed": self._restore_outcomes[restore_id]} restore = self._restores.get(restore_id) if restore is None: if commit: - raise RuntimeError("Restore is no longer active; metadata was not committed by this request") - self._cancelled_restores.add(restore_id) - return {"finished": True} - if not commit: + return {"finished": False} + self._restore_outcomes[restore_id] = False + return {"finished": True, "committed": False} + if not commit or "failed" in restore["units"].values(): restore["aborting"] = True for unit, state in restore["units"].items(): if state == "pending": @@ -1868,9 +1872,8 @@ def finish_restore(self, restore_id: str, commit: bool) -> dict: running = [unit for unit, state in restore["units"].items() if state in ("pending", "running")] if running: return {"finished": False, "units": running} - if commit: - if restore["aborting"] or any(state != "done" for state in restore["units"].values()): - raise RuntimeError("Restore has failed or was cancelled") + committed = commit and not restore["aborting"] + if committed: partition = self.partitions[restore["partition_id"]] updates = [update for group in restore["updates"].values() for update in group] for update in updates: @@ -1880,10 +1883,10 @@ def finish_restore(self, restore_id: str, commit: bool) -> dict: raise RuntimeError("Controller rejected restored metadata") metadata = restore["metadata"] partition.set_custom_meta(dict(zip(metadata.global_indexes, metadata.custom_meta, strict=True))) - else: - self._cancelled_restores.add(restore_id) + # Retain only the terminal outcome, never payloads or per-row metadata. + self._restore_outcomes[restore_id] = committed del self._restores[restore_id] - return {"finished": True} + return {"finished": True, "committed": committed} def list_restores(self, dump_dir: str | None) -> list[str]: """Find unfinished operations for a dump, including after the initiating client exits.""" diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index 95cec709..a800ac8d 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -455,19 +455,22 @@ def _load_via_kv(partition_id: str, rows: dict[str, Any], shards: list[dict], ve kv_batch_put(keys, partition_id, tags=[rows[key]["tag"] for key in keys]) -def recover_data_load(dump_dir: str | Path) -> None: +def recover_data_load(dump_dir: str | Path, *, cancel: bool = False) -> bool: """Settle an interrupted restore before retrying or releasing destination indexes. - This cancels work that has not claimed permission and waits for known writers. - Running or unreachable units keep the reservation. Retry recovery once those - units can report completion; a lost unit requires restarting the whole TQ system. - Partial payload writes remain, but subsequent index reuse is safe after success. + By default, commit loads once every unit reports success. Use ``cancel=True`` + to deny unclaimed work and wait for claimed workers without publishing metadata. + Running or unknown work raises RestorePendingError and keeps the dump protected. + + Returns: + True if all loads committed or none needed recovery; False if any load was + cancelled or failed. Partial payload writes remain after cancellation. """ with _dump_lock(dump_dir) as directory: - return _recover_data_load(directory) + return _recover_data_load(directory, cancel=cancel) -def _recover_data_load(dump_dir: Path) -> None: +def _recover_data_load(dump_dir: Path, *, cancel: bool = False) -> bool: from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client if _TQ_CONTROLLER is None: @@ -475,5 +478,6 @@ def _recover_data_load(dump_dir: Path) -> None: dump_dir = Path(dump_dir).resolve() marker = dump_dir.with_name(dump_dir.name + ".restore") ids = [marker.read_text().strip()] if marker.exists() else [] - _maybe_create_tq_client().recover_data_load(str(dump_dir), ids) + committed = _maybe_create_tq_client().recover_data_load(str(dump_dir), ids, cancel=cancel) marker.unlink(missing_ok=True) + return committed From 87e5ac778f6e43b786460dec698af34bb7ebe0b1 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Mon, 28 Sep 2026 16:37:59 +0800 Subject: [PATCH 15/30] refactor: share storage routing with restore reservations Signed-off-by: neowywang --- tests/test_restore_lifecycle.py | 27 ++++++++++++++ transfer_queue/controller.py | 10 +++--- .../managers/simple_storage_manager.py | 20 ++--------- transfer_queue/utils/storage_routing.py | 35 +++++++++++++++++++ 4 files changed, 70 insertions(+), 22 deletions(-) create mode 100644 transfer_queue/utils/storage_routing.py diff --git a/tests/test_restore_lifecycle.py b/tests/test_restore_lifecycle.py index 24f48ae9..3e7a6516 100644 --- a/tests/test_restore_lifecycle.py +++ b/tests/test_restore_lifecycle.py @@ -27,6 +27,7 @@ from transfer_queue.controller import PartitionIndexManager, TransferQueueController from transfer_queue.sampler import SequentialSampler from transfer_queue.storage.dump_io import RestorePendingError +from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType @@ -282,6 +283,32 @@ def test_unknown_restore_stays_pending_until_explicit_cancellation(controller): begin(controller, "unknown") +@pytest.mark.parametrize("units", [["only"], ["u0", "u1"], ["u2", "u0", "u1"]]) +@pytest.mark.parametrize("has_payload", [True, False]) +def test_restore_reserves_only_current_storage_owners(controller, units, has_payload): + controller.kv_retrieve_meta(["other"], "unrelated", create=True) + existing = controller.kv_retrieve_meta(["keep", "gap", "empty", "last"], "p", create=True) + controller.clear_meta([existing.global_indexes[1]], ["p"]) + rows = { + key: {"fields": ["x"] if has_payload and key != "empty" else [], "tag": {}} + for key in ["last", "new", "empty", "keep"] + } + metadata = controller.begin_restore("r", "/dump", "p", rows, units, {}) + indexes = [index for key, index in zip(rows, metadata.global_indexes, strict=True) if rows[key]["fields"]] + manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) + manager.storage_unit_infos = dict.fromkeys(units) + manager.close = lambda: None + routed = manager._group_by_hash(indexes) + assert set(controller._restores["r"]["units"]) == set(routed) + assert metadata.global_indexes[0] == existing.global_indexes[-1] + assert metadata.global_indexes[-1] == existing.global_indexes[0] + for unit, group in routed.items(): + assert [indexes[pos] for pos in group.batch_positions] == group.global_indexes + controller.restore_unit("r", unit, "claim") + controller.restore_unit("r", unit, "complete", {"success": True, "updates": []}) + assert controller.finish_restore("r", commit=True) == {"finished": True, "committed": True} + + @pytest.mark.asyncio @pytest.mark.parametrize("interruption", ["running_timeout", "lost_load_reply", "lost_commit_reply", "lost_complete"]) async def test_pending_load_can_commit_without_replaying_payload(controller, interruption): diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index 1361ea36..7ead7a1c 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -39,6 +39,7 @@ from transfer_queue.utils.enum_utils import Role from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.perf_utils import IntervalPerfMonitor +from transfer_queue.utils.storage_routing import group_by_storage_unit from transfer_queue.utils.zmq_utils import ( ZMQMessage, ZMQRequestType, @@ -1818,11 +1819,10 @@ def begin_restore( partition.validate_field_schema(schema) keys = list(rows) metadata = self.kv_retrieve_meta(keys, partition_id, create=True) - active_units = { - units[index % len(units)] - for key, index in zip(keys, metadata.global_indexes, strict=True) - if rows[key]["fields"] - } + data_indexes = [ + index for key, index in zip(keys, metadata.global_indexes, strict=True) if rows[key]["fields"] + ] + active_units = group_by_storage_unit(data_indexes, units) metadata.update_custom_meta([rows[key]["tag"] for key in keys]) self._restores[restore_id] = { "partition_id": partition_id, diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 6b03d3fd..74aad6fe 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -23,7 +23,7 @@ from functools import partial from operator import itemgetter from pathlib import Path -from typing import Any, Callable, NamedTuple +from typing import Any, Callable import torch import zmq @@ -38,6 +38,7 @@ ) from transfer_queue.utils.common import log_heavy_operation from transfer_queue.utils.logging_utils import get_logger +from transfer_queue.utils.storage_routing import RoutingGroup, group_by_storage_unit from transfer_queue.utils.tensor_utils import pack_field_values from transfer_queue.utils.zmq_utils import ( TQ_SOCKET_POOL_SIZE, @@ -107,13 +108,6 @@ def _describe_unit_state(body: dict[str, Any]) -> str: ) -class RoutingGroup(NamedTuple): - """Routing result for a single storage unit.""" - - global_indexes: list[int] # global indexes routed to this SU - batch_positions: list[int] # corresponding positions in the original batch - - @StorageManagerFactory.register("SimpleStorage") class AsyncSimpleStorageManager(StorageManager): """Asynchronous storage manager that handles multiple storage units. @@ -219,15 +213,7 @@ def _group_by_hash(self, global_indexes: list[int]) -> dict[str, RoutingGroup]: NOTE: Dynamic SU scaling requires a data migration mechanism (not yet supported). """ - storage_unit_keys = list(self.storage_unit_infos.keys()) - num_units = len(storage_unit_keys) - gi_lists: dict[str, list[int]] = defaultdict(list) - pos_lists: dict[str, list[int]] = defaultdict(list) - for pos, global_idx in enumerate(global_indexes): - key = storage_unit_keys[global_idx % num_units] - gi_lists[key].append(global_idx) - pos_lists[key].append(pos) - return {key: RoutingGroup(gi_lists[key], pos_lists[key]) for key in gi_lists} + return group_by_storage_unit(global_indexes, list(self.storage_unit_infos)) def _describe_storage_unit(self, storage_unit_id: str) -> str: """Return ``ip:port`` for a storage unit, for use in diagnostics. diff --git a/transfer_queue/utils/storage_routing.py b/transfer_queue/utils/storage_routing.py new file mode 100644 index 00000000..f30e548a --- /dev/null +++ b/transfer_queue/utils/storage_routing.py @@ -0,0 +1,35 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from collections import defaultdict +from typing import NamedTuple + + +class RoutingGroup(NamedTuple): + """Indexes assigned to a storage unit and their positions in the input batch.""" + + global_indexes: list[int] + batch_positions: list[int] + + +def group_by_storage_unit(global_indexes: list[int], storage_unit_ids: list[str]) -> dict[str, RoutingGroup]: + """Route indexes using the ordered unit list shared by storage and restore reservations.""" + gi_lists: dict[str, list[int]] = defaultdict(list) + pos_lists: dict[str, list[int]] = defaultdict(list) + for pos, global_idx in enumerate(global_indexes): + key = storage_unit_ids[global_idx % len(storage_unit_ids)] + gi_lists[key].append(global_idx) + pos_lists[key].append(pos) + return {key: RoutingGroup(gi_lists[key], pos_lists[key]) for key in gi_lists} From eb6912a2c7e29cdbcd152e27406596f59005a9f7 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Mon, 28 Sep 2026 19:49:47 +0800 Subject: [PATCH 16/30] fix: recover missing nested dump schemas on storage owners Signed-off-by: neowywang --- docs/data_dump.md | 10 +++ tests/e2e/test_data_dump_e2e.py | 66 ++++++++++++++++++ tests/test_data_dump.py | 69 ++++++++++++++++++- transfer_queue/client.py | 10 ++- transfer_queue/controller.py | 2 +- transfer_queue/data_dump.py | 48 +++++++++++++ transfer_queue/storage/dump_io.py | 4 +- .../managers/simple_storage_manager.py | 36 +++++++--- transfer_queue/storage/simple_storage.py | 18 ++++- 9 files changed, 248 insertions(+), 15 deletions(-) diff --git a/docs/data_dump.md b/docs/data_dump.md index 33e2b89b..469b7e6b 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -80,6 +80,16 @@ shapes. Restore uses that schema regardless of target topology or batch boundari non-tensor fields remain non-tensor even when a batch happens to contain only tensors. Destination type conflicts are rejected before payload writes. +Legacy imports can leave nested fields with missing row shapes, for example when a +later checkpoint chunk wraps tensor values in `NonTensorStack`. Export asks the owner +units to inspect those values while writing their records. The caller merges only +shape/type metadata before publishing the dump; payloads still stay on the units. +Missing tensor shapes are filled from the actual values and their dtype is checked. +If a missing-shape row contains `None` or an object, the entire selected field is +saved as non-tensor, with a warning, so every unit restores the same field contract. +Fields already declared non-tensor remain non-tensor. This repairs the exported +schema without mutating live controller metadata or inventing shapes for objects. + Version-1 dumps remain readable using the prior caller-side KV put path. Version-2 dumps retain direct reads, but lack original schemas and use the older inference behavior; exact field-type preservation cannot be guaranteed for those files. diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index d7d616ba..a80743c7 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -492,6 +492,72 @@ def test_version_one_dump_remains_readable(tq_system, dump_dir, controller): _assert_rows_equal(tq.kv_batch_get(["k"], partition, select_fields=["x"])["x"], [torch.tensor([7, 8])]) +@pytest.mark.parametrize("row_count", [3, 127, 128, 129, 130]) +@pytest.mark.parametrize("last_kind", ["tensor", "none", "object"]) +def test_legacy_chunks_with_missing_nested_shapes_roundtrip(tq_system, dump_dir, controller, row_count, last_kind): + partition = "legacy_nested" + keys = [f"k{i}" for i in range(row_count)] + tensors = [torch.arange(i % 3 + 1, dtype=torch.int64) for i in range(row_count)] + last = tensors[-1] if last_kind == "tensor" else None if last_kind == "none" else {"pixels": tensors[-1]} + expected = [*tensors[:-1], last] + chunk = dump_dir.parent / "legacy-chunk.pt" + chunk.parent.mkdir(parents=True) + first = TensorDict( + { + "input_ids": torch.nested.as_nested_tensor(tensors[:-1], layout=torch.jagged), + "multi_modal_inputs#images": torch.nested.as_nested_tensor(tensors[:-1], layout=torch.jagged), + "wrapped": NonTensorStack(*tensors[:-1]), + }, + batch_size=row_count - 1, + ) + tail = TensorDict( + { + name: NonTensorStack(value) + for name, value in { + "input_ids": last, + "multi_modal_inputs#images": last, + "wrapped": tensors[-1], + }.items() + }, + batch_size=1, + ) + # Legacy wrappers reload each saved TensorDict without normalizing its container type. + for chunk_keys, fields in [(keys[:-1], first), (keys[-1:], tail)]: + torch.save({"keys": chunk_keys, "fields": fields}, chunk) + saved = torch.load(chunk, weights_only=False) + tq.kv_batch_put(saved["keys"], partition, saved["fields"], tags=[{"key": key} for key in chunk_keys]) + + snapshot = ray.get(controller.get_partition_snapshot.remote(partition)) + last_index = snapshot.keys_mapping[keys[-1]] + assert last_index in snapshot.field_metadata["input_ids"].global_indexes + assert last_index not in snapshot.field_metadata["input_ids"].per_sample_shapes + for target in [dump_dir, dump_dir.parent / "second-dump"]: + tq.dump_data_by_key(target, keys, partition) + index = tq.read_row_index(target) + for name in ["input_ids", "multi_modal_inputs#images"]: + schema = index["field_schema"][name] + assert schema["is_non_tensor"] == (last_kind != "tensor") + if last_kind == "tensor": + assert schema["is_nested"] + for key, value in zip(keys, expected, strict=True): + assert tuple(schema["per_sample_shapes"][index["rows"][key]["global_index"]]) == tuple(value.shape) + assert index["field_schema"]["wrapped"]["is_non_tensor"] + tq.kv_clear(keys, partition) + tq.load_data_by_key(target) + for start in range(0, row_count, 128): + end = min(start + 128, row_count) + restored = tq.kv_batch_get( + keys[start:end], partition, ["input_ids", "multi_modal_inputs#images", "wrapped"] + ) + for name in ["input_ids", "multi_modal_inputs#images"]: + for value, reference in zip(restored[name], expected[start:end], strict=True): + torch.testing.assert_close(value, reference) + for value, reference in zip(restored["wrapped"], tensors[start:end], strict=True): + torch.testing.assert_close(value, reference) + tags = tq.kv_list(partition)[partition] + assert tags == {key: {"key": key} for key in keys} + + def test_corrupt_shard_keeps_new_rows_unproduced(tq_system, dump_dir, controller): partition = "corrupt" _put_rows(partition, ["key"]) diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index cc9f5c62..09aaca0d 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -25,6 +25,7 @@ import torch from transfer_queue import data_dump, interface +from transfer_queue.storage.dump_io import validate_dump_values from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType @@ -345,16 +346,80 @@ async def test_dump_waits_for_writers_before_cleanup_can_start(tmp_path): failed = asyncio.Event() completed = [] - async def dump(path, target_storage_unit, global_indexes, fields_by_index): + async def dump(path, target_storage_unit, global_indexes, fields_by_index, missing_shapes): if target_storage_unit == "u0": failed.set() raise OSError("write failed") await failed.wait() await asyncio.sleep(0) completed.append(target_storage_unit) - return {1: [0, 1]} + return {"row_offsets": {1: [0, 1]}} manager._dump_single_shard = dump with pytest.raises(OSError, match="write failed"): await manager.dump_rows_by_index(str(tmp_path), [0, 1]) assert completed == ["u1"] + + +def test_dump_recovers_shapes_from_units_without_forwarding_payloads(unit, tmp_path): + unit.storage_data.put_data({"x": [torch.arange(2), torch.arange(3)], "y": [None, {"key": "value"}]}, [9, 10]) + response = unit._handle_dump_rows( + ZMQMessage.create( + request_type=ZMQRequestType.DUMP_ROWS, + sender_id="test", + body={ + "path": str(tmp_path / "shard.pkl"), + "global_indexes": [9, 10], + "missing_shapes": {9: ["x", "y"], 10: ["x", "y"]}, + }, + ) + ) + assert response.body["success"] + assert response.body["recovered_schema"] == { + 9: {"x": {"shape": (2,), "dtype": torch.int64}, "y": None}, + 10: {"x": {"shape": (3,), "dtype": torch.int64}, "y": None}, + } + + +@pytest.mark.parametrize("reverse_units", [False, True]) +def test_missing_shapes_merge_consistently_across_units(reverse_units): + schema = { + "x": {"dtype": torch.int64, "is_nested": True, "is_non_tensor": False, "per_sample_shapes": {1: None, 2: None}}, + } + shards = [ + {"recovered_schema": {1: {"x": {"shape": (2,), "dtype": torch.int64}}}}, + {"recovered_schema": {2: {"x": None}}}, + ] + if reverse_units: + shards.reverse() + data_dump._complete_dump_schema(schema, {1: ["x"], 2: ["x"]}, shards) + assert schema["x"] == {"dtype": None, "shape": None, "is_nested": False, "is_non_tensor": True} + assert all("recovered_schema" not in shard for shard in shards) + + +@pytest.mark.parametrize("problem", ["missing", "wrong_field", "wrong_dtype"]) +def test_dump_rejects_incomplete_schema_recovery(problem): + schema = {"x": {"dtype": torch.int64, "is_nested": True, "is_non_tensor": False, "per_sample_shapes": {9: None}}} + recovered = {9: {"x": {"shape": (2,), "dtype": torch.int64}}} + if problem == "missing": + recovered = {} + elif problem == "wrong_field": + recovered[9] = {"other": recovered[9]["x"]} + else: + recovered[9]["x"]["dtype"] = torch.float32 + with pytest.raises(ValueError): + data_dump._complete_dump_schema(schema, {9: ["x"]}, [{"recovered_schema": recovered}]) + + +def test_saved_missing_tensor_shape_is_reported_as_invalid_dump(): + schema = { + "x": { + "dtype": torch.int64, + "shape": None, + "is_nested": True, + "is_non_tensor": False, + "per_sample_shapes": {1: None}, + } + } + with pytest.raises(ValueError, match="has no saved shape at row 1"): + validate_dump_values({"x": torch.arange(2)}, schema, 1) diff --git a/transfer_queue/client.py b/transfer_queue/client.py index a327acef..8f7d325b 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1194,6 +1194,7 @@ async def async_dump_rows_by_index( shard_dir: str, global_indexes: list[int], fields_by_index: dict[int, list[str]] | None = None, + missing_shapes: dict[int, list[str]] | None = None, ) -> list[dict[str, Any]]: """Asynchronously dump the given rows into per-storage-unit shards. @@ -1201,6 +1202,7 @@ async def async_dump_rows_by_index( shard_dir: Directory to write shard files into. global_indexes: Global indexes to dump. fields_by_index: Produced fields to persist; omitted for a raw storage dump. + missing_shapes: Fields whose row shapes must be recovered from stored values. Returns: One entry per written shard. @@ -1217,7 +1219,9 @@ async def async_dump_rows_by_index( ) if not hasattr(self.storage_manager, "dump_rows_by_index"): raise NotImplementedError(f"{type(self.storage_manager).__name__} does not support selective data dump") - return await self.storage_manager.dump_rows_by_index(shard_dir, global_indexes, fields_by_index) + return await self.storage_manager.dump_rows_by_index( + shard_dir, global_indexes, fields_by_index, **({"missing_shapes": missing_shapes} if missing_shapes else {}) + ) def _restore_context(self, restore_id: str) -> dict: return { @@ -1973,6 +1977,7 @@ def dump_rows_by_index( shard_dir: str, global_indexes: list[int], fields_by_index: dict[int, list[str]] | None = None, + missing_shapes: dict[int, list[str]] | None = None, ) -> list[dict[str, Any]]: """Synchronously dump the given rows into per-storage-unit shards. @@ -1980,6 +1985,7 @@ def dump_rows_by_index( shard_dir: Directory to write shard files into. global_indexes: Global indexes to dump. fields_by_index: Produced fields to persist; omitted for a raw storage dump. + missing_shapes: Fields whose row shapes must be recovered from stored values. Returns: One entry per written shard. @@ -1989,7 +1995,7 @@ def dump_rows_by_index( no data for a row it was asked to dump. NotImplementedError: If the storage backend does not support dumping. """ - return self._dump_rows_by_index(shard_dir, global_indexes, fields_by_index) + return self._dump_rows_by_index(shard_dir, global_indexes, fields_by_index, missing_shapes=missing_shapes) def load_rows_by_key( self, partition_id: str, rows: dict, shards: list[dict], dump_dir: str = "", restore_id: str | None = None diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index 7ead7a1c..407596df 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -2393,7 +2393,7 @@ def _handle_describe_rows_by_key_request(self, request_msg: ZMQMessage) -> ZMQMe continue schema = meta.to_batch_schema(indexes) if schema.get("is_nested"): - schema["per_sample_shapes"] = {index: meta.per_sample_shapes[index] for index in indexes} + schema["per_sample_shapes"] = {index: meta.per_sample_shapes.get(index) for index in indexes} field_schema[name] = schema return self._make_response( request_msg, diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index a800ac8d..bf3daaff 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -165,6 +165,19 @@ def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dic } ) rows = row_index["rows"] + field_schema = row_index["field_schema"] + missing_shapes = {} + for row in rows.values(): + index = row["global_index"] + names = [ + name + for name in row["fields"] + if field_schema[name].get("is_nested") + and not field_schema[name].get("is_non_tensor") + and field_schema[name].get("per_sample_shapes", {}).get(index) is None + ] + if names: + missing_shapes[index] = names # A row whose fields are all still unproduced has nothing for a storage unit to # dump, but it keeps its key and tag so the restore can recreate the row. @@ -175,10 +188,12 @@ def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dic str(tmp_dir / _SHARD_SUBDIR), indexes_with_data, {row["global_index"]: row["fields"] for row in rows.values() if row["fields"]}, + **({"missing_shapes": missing_shapes} if missing_shapes else {}), ) if indexes_with_data else [] ) + _complete_dump_schema(field_schema, missing_shapes, shard_records) shard_dir = tmp_dir / _SHARD_SUBDIR shard_dir.mkdir(parents=True, exist_ok=True) with open(shard_dir / _SHARD_INFO_FILE, "w", encoding="utf-8") as f: @@ -242,6 +257,39 @@ def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dic } +def _complete_dump_schema(field_schema: dict, missing_shapes: dict[int, list[str]], shards: list[dict]) -> None: + """Complete legacy nested metadata before publishing any dump manifest.""" + recovered = {} + for shard in shards: + for index, fields in shard.pop("recovered_schema", {}).items(): + if index in recovered: + raise ValueError(f"Duplicated schema recovery for row {index}") + recovered[index] = fields + if set(recovered) != set(missing_shapes): + raise ValueError("Storage units did not return every requested missing row schema") + for index, names in missing_shapes.items(): + if set(recovered[index]) != set(names): + raise ValueError(f"Recovered schema fields disagree for row {index}") + + non_tensor_fields = {name for fields in recovered.values() for name, meta in fields.items() if meta is None} + for name in non_tensor_fields: + field_schema[name] = {"is_non_tensor": True, "is_nested": False, "shape": None, "dtype": None} + logger.warning( + "Dump field %r contains non-tensor values missing from legacy nested metadata; preserving as non-tensor", + name, + ) + for index, fields in recovered.items(): + for name, meta in fields.items(): + if name in non_tensor_fields: + continue + field = field_schema[name] + if meta["dtype"] != field["dtype"]: + raise ValueError(f"Dump field {name!r} dtype mismatch at row {index}") + field["per_sample_shapes"][index] = tuple(meta["shape"]) + if recovered: + logger.info("Recovered missing dump schemas for %s rows", len(recovered)) + + def read_row_index(dump_dir: str | Path) -> dict[str, Any]: """Read a dump's row index without touching its payload. diff --git a/transfer_queue/storage/dump_io.py b/transfer_queue/storage/dump_io.py index 0a2eb3a6..bc6cbead 100644 --- a/transfer_queue/storage/dump_io.py +++ b/transfer_queue/storage/dump_io.py @@ -39,7 +39,9 @@ def validate_dump_values(values: dict, schema: dict, source_index: int) -> None: field = schema[name] if field["is_non_tensor"]: continue - shape = field["per_sample_shapes"][source_index] if field["is_nested"] else field["shape"] + shape = field.get("per_sample_shapes", {}).get(source_index) if field["is_nested"] else field["shape"] + if shape is None: + raise ValueError(f"Dump field {name!r} has no saved shape at row {source_index}") actual_shape = tuple(value.shape) if isinstance(value, torch.Tensor) else None # Existing dense scalar fields use a one-element metadata shape. scalar = actual_shape == () and tuple(shape) == (1,) and not field["is_nested"] diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 74aad6fe..3d509b1f 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -760,15 +760,21 @@ async def _dump_single_shard( target_storage_unit: str, global_indexes: list[int], fields_by_index: dict[int, list[str]] | None = None, + missing_shapes: dict[int, list[str]] | None = None, socket: zmq.Socket = None, - ) -> dict[int, list[int]]: + ) -> dict[str, Any]: """Ask one storage unit to write the rows it owns into a shard file.""" try: request_msg = ZMQMessage.create( request_type=ZMQRequestType.DUMP_ROWS, # type: ignore[arg-type] sender_id=self.storage_manager_id, receiver_id=target_storage_unit, - body={"path": path, "global_indexes": global_indexes, "fields_by_index": fields_by_index}, + body={ + "path": path, + "global_indexes": global_indexes, + "fields_by_index": fields_by_index, + "missing_shapes": missing_shapes or {}, + }, ) await socket.send_multipart(request_msg.serialize(), copy=False) messages = await socket.recv_multipart(copy=False) @@ -785,7 +791,7 @@ async def _dump_single_shard( raise RuntimeError( f"Storage unit {target_storage_unit} holds no data for requested rows: {missing_rows[:20]}" ) - return response_msg.body["row_offsets"] + return response_msg.body except Exception as e: raise RuntimeError( f"[{self.storage_manager_id}]: Error dumping shard from storage unit {target_storage_unit}: {str(e)}" @@ -796,6 +802,7 @@ async def dump_rows_by_index( shard_dir: str, global_indexes: list[int], fields_by_index: dict[int, list[str]] | None = None, + missing_shapes: dict[int, list[str]] | None = None, ) -> list[dict[str, Any]]: """Dump the given rows into one shard per storage unit, in parallel. @@ -807,6 +814,7 @@ async def dump_rows_by_index( shard_dir: Directory to write shard files into. global_indexes: Global indexes to dump. fields_by_index: Produced fields to persist; omitted for a raw storage dump. + missing_shapes: Fields whose row shapes must be recovered from stored values. Returns: One entry per written shard: ``{"position", "storage_unit_id", "rows", "row_offsets"}``. @@ -821,13 +829,16 @@ async def dump_rows_by_index( targets = [(su_id, group.global_indexes) for su_id, group in routing.items()] paths = [str(shard_dir_path / f"shard_{pos}_{su_id}.pkl") for pos, (su_id, _) in enumerate(targets)] - row_offsets = await asyncio.gather( + results = await asyncio.gather( *( self._dump_single_shard( path, target_storage_unit=su_id, global_indexes=indexes, fields_by_index={index: fields_by_index[index] for index in indexes} if fields_by_index else None, + missing_shapes={index: missing_shapes[index] for index in indexes if index in missing_shapes} + if missing_shapes + else None, ) for path, (su_id, indexes) in zip(paths, targets, strict=True) ), @@ -835,11 +846,20 @@ async def dump_rows_by_index( ) shards = [] total_rows = 0 - for pos, ((su_id, _), offsets) in enumerate(zip(targets, row_offsets, strict=True)): - if isinstance(offsets, BaseException): - raise offsets + for pos, ((su_id, _), result) in enumerate(zip(targets, results, strict=True)): + if isinstance(result, BaseException): + raise result + offsets = result["row_offsets"] total_rows += len(offsets) - shards.append({"position": pos, "storage_unit_id": su_id, "rows": len(offsets), "row_offsets": offsets}) + shards.append( + { + "position": pos, + "storage_unit_id": su_id, + "rows": len(offsets), + "row_offsets": offsets, + "recovered_schema": result.get("recovered_schema", {}), + } + ) logger.info( f"[{self.storage_manager_id}]: dumped {total_rows} rows across {len(targets)} shards to {shard_dir_path}" diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index a2a45b4b..eff33529 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -1317,6 +1317,7 @@ def _handle_dump_rows(self, request: ZMQMessage) -> ZMQMessage: if missing: raise ValueError(f"Storage holds no data for requested rows: {sorted(missing)[:20]}") row_offsets = {} + recovered_schema = {} with open(path, "wb") as f: for index in sorted(indexes): described_fields = request.body.get("fields_by_index") @@ -1329,6 +1330,15 @@ def _handle_dump_rows(self, request: ZMQMessage) -> ZMQMessage: else: # Reused global indexes can retain fields no longer present in metadata. fields = {name: self.storage_data.field_data[name][index] for name in described_fields[index]} + missing_fields = request.body.get("missing_shapes", {}).get(index, []) + if missing_fields: + # Inspect values already being written; only shape/type metadata leaves the unit. + recovered_schema[index] = { + name: {"shape": tuple(fields[name].shape), "dtype": fields[name].dtype} + if isinstance(fields[name], torch.Tensor) + else None + for name in missing_fields + } offset = f.tell() compact_pickle.dump({"global_index": index, "fields": fields}, f) row_offsets[index] = [offset, f.tell() - offset] @@ -1338,7 +1348,13 @@ def _handle_dump_rows(self, request: ZMQMessage) -> ZMQMessage: return ZMQMessage.create( request_type=ZMQRequestType.DUMP_ROWS_RESPONSE, sender_id=self.storage_unit_id, - body={"success": True, "dumped_rows": len(indexes), "missing_rows": [], "row_offsets": row_offsets}, + body={ + "success": True, + "dumped_rows": len(indexes), + "missing_rows": [], + "row_offsets": row_offsets, + "recovered_schema": recovered_schema, + }, ) except Exception as e: logger.error("[%s]: dump rows failed: %s", self.storage_unit_id, e) From 52efe23b7e67521d5eba8a3d570cfdaa1ee58b68 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Mon, 28 Sep 2026 20:01:37 +0800 Subject: [PATCH 17/30] fix: maintain field schemas across wrapped tensor chunks Signed-off-by: neowywang --- docs/data_dump.md | 4 ++ tests/e2e/test_data_dump_e2e.py | 48 ++++++++++++++- tests/test_data_dump.py | 79 +++++++++++++++++++++++++ transfer_queue/controller.py | 14 +++++ transfer_queue/metadata.py | 19 +++++- transfer_queue/storage/managers/base.py | 23 +++---- 6 files changed, 173 insertions(+), 14 deletions(-) diff --git a/docs/data_dump.md b/docs/data_dump.md index 469b7e6b..edc207d3 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -89,6 +89,10 @@ If a missing-shape row contains `None` or an object, the entire selected field i saved as non-tensor, with a warning, so every unit restores the same field contract. Fields already declared non-tensor remain non-tensor. This repairs the exported schema without mutating live controller metadata or inventing shapes for objects. +New puts also carry tensor shape hints for homogeneous tensor values wrapped in +`NonTensorStack`. An existing tensor field uses those hints to keep its shape map +complete; real mixed values make the field non-tensor. A field originally declared +non-tensor stays non-tensor when later batches contain only tensors. Version-1 dumps remain readable using the prior caller-side KV put path. Version-2 dumps retain direct reads, but lack original schemas and use the older inference diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index a80743c7..13b7cf82 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -494,7 +494,10 @@ def test_version_one_dump_remains_readable(tq_system, dump_dir, controller): @pytest.mark.parametrize("row_count", [3, 127, 128, 129, 130]) @pytest.mark.parametrize("last_kind", ["tensor", "none", "object"]) -def test_legacy_chunks_with_missing_nested_shapes_roundtrip(tq_system, dump_dir, controller, row_count, last_kind): +@pytest.mark.parametrize("legacy_metadata", [False, True]) +def test_legacy_chunks_with_missing_nested_shapes_roundtrip( + tq_system, dump_dir, controller, row_count, last_kind, legacy_metadata, monkeypatch +): partition = "legacy_nested" keys = [f"k{i}" for i in range(row_count)] tensors = [torch.arange(i % 3 + 1, dtype=torch.int64) for i in range(row_count)] @@ -530,9 +533,44 @@ def test_legacy_chunks_with_missing_nested_shapes_roundtrip(tq_system, dump_dir, snapshot = ray.get(controller.get_partition_snapshot.remote(partition)) last_index = snapshot.keys_mapping[keys[-1]] assert last_index in snapshot.field_metadata["input_ids"].global_indexes - assert last_index not in snapshot.field_metadata["input_ids"].per_sample_shapes + if last_kind == "tensor": + assert tuple(snapshot.field_metadata["input_ids"].per_sample_shapes[last_index]) == tuple(tensors[-1].shape) + else: + assert snapshot.field_metadata["input_ids"].is_non_tensor + assert not snapshot.field_metadata["input_ids"].is_nested + if legacy_metadata: + from transfer_queue.controller import FieldMeta + + # Existing controller checkpoints can retain the old incomplete nested schema. + checkpoint = dump_dir.parent / "legacy-controller.pkl" + tq.get_client().save_controller_checkpoint(str(checkpoint)) + with checkpoint.open("rb") as file: + state = pickle.load(file) + for name in ["input_ids", "multi_modal_inputs#images"]: + state["partitions"][partition].field_metadata[name] = FieldMeta( + global_indexes=set(snapshot.keys_mapping.values()), + dtype=torch.int64, + is_nested=True, + is_non_tensor=False, + per_sample_shapes={ + snapshot.keys_mapping[key]: tuple(value.shape) + for key, value in zip(keys[:-1], tensors[:-1], strict=True) + }, + ) + with checkpoint.open("wb") as file: + pickle.dump(state, file) + tq.get_client().load_controller_checkpoint(str(checkpoint)) + open_file = builtins.open + + def no_shard_read(path, *args, **kwargs): + if isinstance(path, str | Path) and Path(path).name.startswith("shard_") and str(path).endswith(".pkl"): + pytest.fail("Dump schema repair read a payload shard in the caller") + return open_file(path, *args, **kwargs) + for target in [dump_dir, dump_dir.parent / "second-dump"]: - tq.dump_data_by_key(target, keys, partition) + with monkeypatch.context() as patcher: + patcher.setattr(builtins, "open", no_shard_read) + tq.dump_data_by_key(target, keys, partition) index = tq.read_row_index(target) for name in ["input_ids", "multi_modal_inputs#images"]: schema = index["field_schema"][name] @@ -542,6 +580,10 @@ def test_legacy_chunks_with_missing_nested_shapes_roundtrip(tq_system, dump_dir, for key, value in zip(keys, expected, strict=True): assert tuple(schema["per_sample_shapes"][index["rows"][key]["global_index"]]) == tuple(value.shape) assert index["field_schema"]["wrapped"]["is_non_tensor"] + if legacy_metadata and last_kind != "tensor" and target == dump_dir: + with pytest.raises(RuntimeError, match="tensor/non-tensor type mismatch"): + tq.load_data_by_key(target) + assert not target.with_name(target.name + ".restore").exists() tq.kv_clear(keys, partition) tq.load_data_by_key(target) for start in range(0, row_count, 128): diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index 09aaca0d..02ac5f29 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -423,3 +423,82 @@ def test_saved_missing_tensor_shape_is_reported_as_invalid_dump(): } with pytest.raises(ValueError, match="has no saved shape at row 1"): validate_dump_values({"x": torch.arange(2)}, schema, 1) + + +@pytest.mark.parametrize("value", [torch.arange(4), None, {"image": torch.arange(4)}]) +@pytest.mark.parametrize("wrapped_first", [False, True]) +def test_field_metadata_merges_wrapped_values_without_incomplete_nested_shapes(value, wrapped_first): + from tensordict import NonTensorStack, TensorDict + + from transfer_queue.controller import DataPartitionStatus + from transfer_queue.metadata import extract_field_schema + + batches = [ + TensorDict( + {"x": torch.nested.as_nested_tensor([torch.arange(2), torch.arange(3)], layout=torch.jagged)}, batch_size=2 + ), + TensorDict({"x": NonTensorStack(value)}, batch_size=1), + ] + if wrapped_first: + batches.reverse() + partition = DataPartitionStatus("p") + offset = 0 + for batch in batches: + indexes = list(range(offset, offset + batch.batch_size[0])) + schema = extract_field_schema(batch) + for field in schema.values(): + for part in [field, field.get("tensor_schema", {})]: + if "per_sample_shapes" in part: + part["per_sample_shapes"] = dict(zip(indexes, part["per_sample_shapes"], strict=True)) + assert partition.update_production_status(indexes, [], schema) + offset += batch.batch_size[0] + meta = partition.field_metadata["x"] + if not wrapped_first and isinstance(value, torch.Tensor): + assert meta.is_nested + assert not meta.is_non_tensor + assert meta.per_sample_shapes == {0: (2,), 1: (3,), 2: (4,)} + else: + assert meta.is_non_tensor + assert not meta.is_nested + assert meta.dtype is None + assert not meta.per_sample_shapes + assert meta.global_indexes == {0, 1, 2} + + +@pytest.mark.parametrize("shape", [(), (1,), (4,)]) +def test_wrapped_tensor_hints_preserve_dense_fields(shape): + from tensordict import NonTensorStack, TensorDict + + from transfer_queue.controller import DataPartitionStatus + from transfer_queue.metadata import extract_field_schema + + partition = DataPartitionStatus("p") + value = torch.ones(shape, dtype=torch.int64) + first = extract_field_schema(TensorDict({"x": value.unsqueeze(0)}, batch_size=1)) + second = extract_field_schema(TensorDict({"x": NonTensorStack(value)}, batch_size=1)) + assert second["x"]["is_non_tensor"] + assert partition.update_production_status([0], [], first) + assert partition.update_production_status([1], [], second) + meta = partition.field_metadata["x"] + assert not meta.is_nested + assert not meta.is_non_tensor + assert tuple(meta.shape) == (shape or (1,)) + assert meta.dtype == torch.int64 + + +def test_schema_hints_do_not_iterate_broadcast_nontensor_data(monkeypatch): + from tensordict import NonTensorData, TensorDict + + from transfer_queue.metadata import extract_field_schema + + data = TensorDict({"x": NonTensorData(data={"kind": "image"}, batch_size=(2,))}, batch_size=2) + original = NonTensorData.__getitem__ + + def bounded_getitem(self, index): + assert index == 0, "Schema extraction tried to iterate broadcast NonTensorData" + return original(self, index) + + monkeypatch.setattr(NonTensorData, "__getitem__", bounded_getitem) + field = extract_field_schema(data)["x"] + assert field["is_non_tensor"] + assert "tensor_schema" not in field diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index 407596df..eb5e121d 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -224,6 +224,20 @@ def update(self, incoming: dict[str, Any], incoming_global_indexes: list[int]) - Raises: ValueError: If incoming dtype conflicts with existing dtype. """ + if self.is_non_tensor: + self.global_indexes.update(incoming_global_indexes) + return + if incoming.get("is_non_tensor"): + tensor_schema = incoming.get("tensor_schema") + if tensor_schema is None or tensor_schema["dtype"] != self.dtype: + self.is_non_tensor = True + self.is_nested = False + self.dtype = self.shape = None + self.per_sample_shapes.clear() + self.global_indexes.update(incoming_global_indexes) + return + incoming = tensor_schema + # dtype consistency check new_dtype = incoming.get("dtype") if new_dtype is not None: diff --git a/transfer_queue/metadata.py b/transfer_queue/metadata.py index e89c1531..df25de54 100644 --- a/transfer_queue/metadata.py +++ b/transfer_queue/metadata.py @@ -23,7 +23,7 @@ import numpy as np import torch -from tensordict import TensorDict +from tensordict import NonTensorStack, TensorDict from transfer_queue.utils.logging_utils import get_logger @@ -185,6 +185,23 @@ def extract_field_schema(data: TensorDict) -> dict[str, dict[str, Any]]: # For nested tensors, record per-sample shapes if is_nested: field_meta["per_sample_shapes"] = [tuple(t.shape) for t in value.unbind()] + elif isinstance(value, NonTensorStack): + # A legacy chunk may wrap tensor rows without changing their values. Keep + # its non-tensor declaration; existing tensor fields can use this hint. + values = value.tolist() + if values and all(isinstance(item, torch.Tensor) and not item.is_nested for item in values): + dtype = values[0].dtype + if all(item.dtype == dtype for item in values): + shapes = [tuple(item.shape) for item in values] + nested = any(shape != shapes[0] for shape in shapes) + field_meta["tensor_schema"] = { + "dtype": dtype, + "shape": None if nested else shapes[0] or (1,), + "is_nested": nested, + "is_non_tensor": False, + } + if nested: + field_meta["tensor_schema"]["per_sample_shapes"] = shapes field_schema[field_name] = field_meta diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index 369f800e..785a0849 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -244,16 +244,19 @@ async def notify_data_update( normalized_field_schema = {} for field_name, field in field_schema.items(): field_copy = field.copy() - per_sample_shapes = field_copy.get("per_sample_shapes", None) - if isinstance(per_sample_shapes, list | tuple): - if len(per_sample_shapes) != len(global_indexes): - raise ValueError( - f"per_sample_shapes length ({len(per_sample_shapes)}) does not match " - f"number of global_indexes ({len(global_indexes)}) for field '{field_name}'. " - ) - field_copy["per_sample_shapes"] = { - global_indexes[i]: per_sample_shapes[i] for i in range(len(global_indexes)) - } + schemas = [field_copy] + if "tensor_schema" in field_copy: + field_copy["tensor_schema"] = dict(field_copy["tensor_schema"]) + schemas.append(field_copy["tensor_schema"]) + for schema in schemas: + per_sample_shapes = schema.get("per_sample_shapes") + if isinstance(per_sample_shapes, list | tuple): + if len(per_sample_shapes) != len(global_indexes): + raise ValueError( + f"per_sample_shapes length ({len(per_sample_shapes)}) does not match " + f"number of global_indexes ({len(global_indexes)}) for field '{field_name}'. " + ) + schema["per_sample_shapes"] = dict(zip(global_indexes, per_sample_shapes, strict=True)) normalized_field_schema[field_name] = field_copy request_msg = ZMQMessage.create( From ccbf583608302764b57acdb5ebedf17bbed35e3d Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Tue, 29 Sep 2026 11:55:11 +0800 Subject: [PATCH 18/30] [test] Shorten the live socket pool to trigger the restore receive timeout The socket pool applies its timeout when a socket connects, so patching TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT after tq.init() no longer reached the sockets the load reuses and the request waited out the paused unit. Signed-off-by: OutstanderWang --- tests/e2e/test_restore_timeout_e2e.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/tests/e2e/test_restore_timeout_e2e.py b/tests/e2e/test_restore_timeout_e2e.py index 9113bc39..2587b227 100644 --- a/tests/e2e/test_restore_timeout_e2e.py +++ b/tests/e2e/test_restore_timeout_e2e.py @@ -32,8 +32,6 @@ def _check_claimed_load_survives_receive_timeout(tmp_path): - import transfer_queue.storage.managers.simple_storage_manager as manager_module - ray.init(namespace="review_claimed_timeout") try: tq.init( @@ -70,7 +68,10 @@ def delayed(request): unit._load_rows = delayed ray.get(actor.__ray_call__.remote(pause_after_claim)) - with patch.object(manager_module, "TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT", 0.5): + # The pool fixes each socket's timeout at connect, so shorten the live pool itself; + # patching the module constant would only reach sockets opened after this point. + pool = tq.get_client().storage_manager.storage_rpc_pool + with patch.object(pool, "_timeout", 1), patch.object(pool, "_idle", {}): with pytest.raises(tq.RestorePendingError): tq.load_data_by_key(dump_dir) assert claimed.exists() From 41a562e73d1c3ea38cb07ae5e69adc91f8be31b8 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Tue, 29 Sep 2026 12:01:27 +0800 Subject: [PATCH 19/30] [fix] Drop methods duplicated when rebasing onto the handler-dispatch refactor The controller and client now route requests through a dispatch table and a shared _request_controller helper, so the copies carried over from the old elif chain shadowed those and left dead code behind. Signed-off-by: OutstanderWang --- transfer_queue/client.py | 27 ------------------- transfer_queue/controller.py | 23 ---------------- .../managers/simple_storage_manager.py | 1 - 3 files changed, 51 deletions(-) diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 8f7d325b..bcf0c21e 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1126,33 +1126,6 @@ def _can_destroy_zmq_context(self) -> bool: return False return True - async def _request_controller( - self, - socket: zmq.asyncio.Socket | None, - request_type: ZMQRequestType, - response_type: ZMQRequestType, - body: dict[str, Any], - ) -> ZMQMessage: - """Send one controller request and validate its response type.""" - assert socket is not None - request_msg = ZMQMessage.create( - request_type=request_type, # type: ignore[arg-type] - sender_id=self.client_id, - receiver_id=self._controller.id, - body=body, - ) - await socket.send_multipart(request_msg.serialize()) - response_serialized = await socket.recv_multipart(copy=False) - response_msg = ZMQMessage.deserialize(response_serialized) - logger.debug(f"[{self.client_id}]: Received {response_msg.request_type} from controller {self._controller.id}") - if response_msg.request_type != response_type: - message = response_msg.body.get("message", "Unknown error") - raise RuntimeError( - f"[{self.client_id}]: Expected {response_type}, got {response_msg.request_type} " - f"from controller {self._controller.id}: {message}" - ) - return response_msg - # ==================== Selective Data Dump API ==================== @with_controller_socket async def async_describe_data_dump( diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index eb5e121d..5752f359 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -2357,15 +2357,6 @@ def _handle_kv_list_request(self, request_msg: ZMQMessage) -> ZMQMessage: {"partition_info": partition_info, "message": message}, ) - def _handle_describe_rows_by_key_request(self, request_msg: ZMQMessage) -> ZMQMessage: - params = request_msg.body - rows = self.describe_rows_by_key(params["partition_id"], params["keys"]) - return self._make_response( - request_msg, - ZMQRequestType.DESCRIBE_ROWS_BY_KEY_RESPONSE, - {"success": True, "rows": rows}, - ) - def _handle_save_controller_checkpoint_request(self, request_msg: ZMQMessage) -> ZMQMessage: self.save_checkpoint(request_msg.body["path"]) return self._make_response( @@ -2382,20 +2373,6 @@ def _handle_load_controller_checkpoint_request(self, request_msg: ZMQMessage) -> {"success": True}, ) - def _make_response( - self, - request_msg: ZMQMessage, - response_type: ZMQRequestType, - body: dict[str, Any], - ) -> ZMQMessage: - """Build a controller response addressed to the request sender.""" - return ZMQMessage.create( - request_type=response_type, - sender_id=self.controller_id, - receiver_id=request_msg.sender_id, - body=body, - ) - def _handle_describe_rows_by_key_request(self, request_msg: ZMQMessage) -> ZMQMessage: params = request_msg.body rows = self.describe_rows_by_key(params["partition_id"], params["keys"]) diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 3d509b1f..e46137db 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -15,7 +15,6 @@ import asyncio import os -import socket import time import warnings from collections import defaultdict From b9824faaeed60242941e02fc83cda46d23f9b5a9 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Tue, 29 Sep 2026 14:09:47 +0800 Subject: [PATCH 20/30] fix: preserve per-unit restore report failures Signed-off-by: neowywang --- tests/test_restore_lifecycle.py | 67 +++++++++++++++++++ transfer_queue/client.py | 12 ++-- transfer_queue/storage/dump_io.py | 10 +-- .../managers/simple_storage_manager.py | 24 +++++-- 4 files changed, 99 insertions(+), 14 deletions(-) diff --git a/tests/test_restore_lifecycle.py b/tests/test_restore_lifecycle.py index 3e7a6516..ad480915 100644 --- a/tests/test_restore_lifecycle.py +++ b/tests/test_restore_lifecycle.py @@ -15,6 +15,7 @@ """Restore reservations prevent late writes from corrupting reused indexes.""" +import asyncio from threading import RLock from types import SimpleNamespace from unittest.mock import AsyncMock @@ -51,6 +52,72 @@ def begin(controller, restore_id="r"): return controller.begin_restore(restore_id, "/dump", "p", {"k": {"fields": ["x"], "tag": {}}}, ["u"], {}) +@pytest.mark.asyncio +async def test_report_failures_preserve_other_unit_reports(caplog): + manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) + manager.storage_unit_infos = dict.fromkeys(["rejected", "slow", "timeout"]) + manager.close = lambda: None + failed = asyncio.Event() + completed = [] + + async def report(restore, target_storage_unit): + if target_storage_unit == "rejected": + failed.set() + raise RuntimeError("Invalid restore transition pending -> complete") + await failed.wait() + if target_storage_unit == "timeout": + raise zmq.error.Again() + await asyncio.sleep(0) + completed.append(target_storage_unit) + + manager._report_restore_unit = report + errors = await manager.report_restore({"restore_id": "r"}) + assert completed == ["slow"] + assert set(errors) == {"rejected", "timeout"} + assert "pending -> complete" in errors["rejected"] + assert "Again" in errors["timeout"] + assert "Restore r: unit rejected" in caplog.text + assert "pending -> complete" in caplog.text + + +@pytest.mark.asyncio +async def test_report_preserves_controller_rejection_message(): + manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) + manager.storage_manager_id = "test" + manager.close = lambda: None + response = ZMQMessage.create( + request_type=ZMQRequestType.REPORT_RESTORE_RESPONSE, + sender_id="u", + body={"success": False, "message": "Invalid restore transition pending -> complete"}, + ) + socket = SimpleNamespace(send_multipart=AsyncMock(), recv_multipart=AsyncMock(return_value=response.serialize())) + with pytest.raises(RuntimeError, match="Restore r: storage unit u report failed: Invalid restore transition"): + await manager._report_restore_unit.__wrapped__(manager, {"restore_id": "r"}, "u", socket=socket) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("outcome", [None, False, True, "timeout"]) +async def test_recovery_surfaces_report_errors_only_while_outcome_is_unknown(outcome): + client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) + errors = {"u": "RuntimeError: completion rejected"} + client.storage_manager = SimpleNamespace(report_restore=AsyncMock(return_value=errors)) + client._restore_context = lambda restore_id: {"restore_id": restore_id} + reply = {"finished": outcome is not None, "committed": outcome} + client._restore_rpc = AsyncMock( + side_effect=[ + {"restore_ids": ["r"]}, + zmq.error.Again() if outcome == "timeout" else reply, + ] + ) + if outcome is None or outcome == "timeout": + with pytest.raises(RestorePendingError, match="completion rejected") as error: + await client.async_recover_data_load("/dump") + assert error.value.report_errors == errors + else: + assert await client.async_recover_data_load("/dump") is outcome + assert client._restore_rpc.await_args.args == (ZMQRequestType.FINISH_RESTORE, {"restore_id": "r", "commit": True}) + + @pytest.mark.asyncio @pytest.mark.parametrize(("active_ids", "saved_ids"), [([], []), (["r"], []), ([], ["r"])]) async def test_recovery_checks_backend_before_cancelling(active_ids, saved_ids): diff --git a/transfer_queue/client.py b/transfer_queue/client.py index bcf0c21e..08684314 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1209,16 +1209,18 @@ async def _restore_rpc(self, request_type: ZMQRequestType, body: dict, socket=No response = await self._request_controller(socket, request_type, response_type, body) return response.body - async def _finish_data_load(self, restore_id: str, *, commit: bool) -> bool: + async def _finish_data_load( + self, restore_id: str, *, commit: bool, report_errors: dict[str, str] | None = None + ) -> bool: """Keep the reservation when completion is pending or its reply is lost.""" try: result = await self._restore_rpc( ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": commit} ) except (zmq.error.Again, TimeoutError) as error: - raise RestorePendingError(restore_id) from error + raise RestorePendingError(restore_id, report_errors=report_errors) from error if not result["finished"]: - raise RestorePendingError(restore_id) + raise RestorePendingError(restore_id, report_errors=report_errors) return result["committed"] async def async_load_rows_by_key( @@ -1283,8 +1285,8 @@ async def async_recover_data_load( await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": False}) except (zmq.error.Again, TimeoutError) as error: raise RestorePendingError(restore_id) from error - await self.storage_manager.report_restore(self._restore_context(restore_id)) - outcome = await self._finish_data_load(restore_id, commit=not cancel) + report_errors = await self.storage_manager.report_restore(self._restore_context(restore_id)) + outcome = await self._finish_data_load(restore_id, commit=not cancel, report_errors=report_errors) committed = committed and outcome return committed diff --git a/transfer_queue/storage/dump_io.py b/transfer_queue/storage/dump_io.py index bc6cbead..1b17b0b8 100644 --- a/transfer_queue/storage/dump_io.py +++ b/transfer_queue/storage/dump_io.py @@ -79,8 +79,10 @@ def pack_dump_field(values: list, schema: dict): class RestorePendingError(RuntimeError): """The controller still reserves indexes until remote restore activity is settled.""" - def __init__(self, restore_id: str): + def __init__(self, restore_id: str, *, report_errors: dict[str, str] | None = None): self.restore_id = restore_id - super().__init__( - f"Restore {restore_id} has an unknown outcome; run recover_data_load before retrying or clearing" - ) + self.report_errors = report_errors or {} + message = f"Restore {restore_id} has an unknown outcome; run recover_data_load before retrying or clearing" + if self.report_errors: + message += f". Unit report failures: {self.report_errors}" + super().__init__(message) diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index e46137db..7ebaea3a 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -926,12 +926,23 @@ async def _load_selected_rows( ) return response.body - async def report_restore(self, restore: dict) -> None: - """Ask units to resend cached terminal results; unknown units cannot grant release.""" - await asyncio.gather( - *(self._report_restore_unit(restore, target_storage_unit=unit) for unit in self.storage_unit_infos), + async def report_restore(self, restore: dict) -> dict[str, str]: + """Wait for all unit reports and return failures without discarding successful reports.""" + units = list(self.storage_unit_infos) + results = await asyncio.gather( + *(self._report_restore_unit(restore, target_storage_unit=unit) for unit in units), return_exceptions=True, ) + errors = {} + for unit, result in zip(units, results, strict=True): + if isinstance(result, Exception): + errors[unit] = f"{type(result).__name__}: {result}" + logger.warning( + "Restore %s: unit %s could not report completion: %s", restore["restore_id"], unit, errors[unit] + ) + elif isinstance(result, BaseException): + raise result + return errors @with_storage_unit_socket async def _report_restore_unit(self, restore: dict, target_storage_unit: str, socket: zmq.Socket = None) -> None: @@ -941,7 +952,10 @@ async def _report_restore_unit(self, restore: dict, target_storage_unit: str, so await socket.send_multipart(request.serialize()) response = ZMQMessage.deserialize(await socket.recv_multipart()) if response.request_type != ZMQRequestType.REPORT_RESTORE_RESPONSE or not response.body.get("success"): - raise RuntimeError("Storage unit could not report its restore result") + message = response.body.get("message", f"Unexpected response: {response.request_type}") + raise RuntimeError( + f"Restore {restore['restore_id']}: storage unit {target_storage_unit} report failed: {message}" + ) async def save_checkpoint(self, checkpoint_dir: str) -> None: """Dump all storage units to the storage_units/ subdirectory of checkpoint_dir. From 888da5e377cbe78ce44494cd78308ec35e206dbb Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Tue, 29 Sep 2026 14:16:24 +0800 Subject: [PATCH 21/30] fix: explain stalled restores and settle failed claims Signed-off-by: neowywang --- docs/data_dump.md | 30 +++++- tests/e2e/test_restore_timeout_e2e.py | 42 ++++++++ tests/test_restore_lifecycle.py | 119 +++++++++++++++++++++++ transfer_queue/client.py | 7 +- transfer_queue/controller.py | 22 ++++- transfer_queue/storage/dump_io.py | 31 +++++- transfer_queue/storage/simple_storage.py | 6 +- 7 files changed, 245 insertions(+), 12 deletions(-) diff --git a/docs/data_dump.md b/docs/data_dump.md index edc207d3..462d1e31 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -134,9 +134,26 @@ committed = tq.recover_data_load("/shared/dumps/selected") Recovery asks units to resend terminal results and commits schemas and tags only when all units succeeded. It returns `True` for committed loads (or no pending load), and `False` for a failed or cancelled load after all claimed workers stopped. -Running or unknown work raises `RestorePendingError`; retry recovery later without -reloading payloads. The controller remembers terminal outcomes so a lost commit -reply can be confirmed safely by retrying recovery. +Unresolved work raises `RestorePendingError`, which includes `reason`, `unit_states` +and `report_errors`. Reports that fail retain the unit ID and original error message. +The controller remembers terminal outcomes so a lost commit reply can be confirmed +safely by retrying recovery; a redundant report failure cannot reverse that outcome. + +- `running` units have claimed permission but have not reported completion. Retry + recovery later; a reporting timeout does not prove that a worker has stopped. +- `pending` units have not claimed permission. A request may still be queued, so + recovery does not automatically cancel it. If the initiating client exited before + sending all requests, use `recover_data_load(dump_dir, cancel=True)` to abandon the + operation; simply repeating recovery cannot dispatch the missing requests. +- `reason="unknown_restore"` means the controller has no record of the ID in the + `.restore` marker. If the initiating client has stopped, use explicit cancellation. + After a whole-system restart, first ensure all old actors have stopped, then cancel + the stale operation. This marker blocks its dump path, not the new controller's + unrelated partitions or checkpoints. + +A unit whose claim timed out caches a failure confirming that it did not write. +Recovery accepts this failure even when the claim never reached the controller, +cancels unclaimed work and waits for any other claimed workers to finish. To abandon the load explicitly, use `recover_data_load(dump_dir, cancel=True)`. This denies unclaimed work and retains the reservation until claimed workers stop. @@ -145,6 +162,13 @@ cancels remaining unclaimed work; partial payload writes are never rolled back. After cancellation settles, retry the load or clear its keys. A lost unit requires stopping the old TQ actors and restarting the whole TQ system; restarting only the controller while old storage actors run is unsupported. +After restarting, cancel the old marker before loading the dump again: + +```python +# Run only after all old TQ actors have stopped and the new system is initialized. +tq.recover_data_load(dump_dir, cancel=True) +tq.load_data_by_key(dump_dir) +``` Writers that already hold low-level metadata must remain paused throughout recovery. The controller reservation covers only the destination partition. Each storage unit diff --git a/tests/e2e/test_restore_timeout_e2e.py b/tests/e2e/test_restore_timeout_e2e.py index 2587b227..51e95861 100644 --- a/tests/e2e/test_restore_timeout_e2e.py +++ b/tests/e2e/test_restore_timeout_e2e.py @@ -104,6 +104,48 @@ def test_claimed_load_survives_receive_timeout(tmp_path): process.join(10) +def test_whole_system_restart_explains_and_cancels_orphan_marker(tmp_path): + ray.init(namespace="restart_recovery") + config = OmegaConf.create( + { + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "num_data_storage_units": 1, + "total_storage_size": 20, + }, + } + } + ) + dump_dir = tmp_path / "dump" + marker = tmp_path / "dump.restore" + try: + tq.init(config) + tq.kv_put("key", "p", {"x": torch.tensor([7])}) + tq.dump_data_by_key(dump_dir, ["key"], "p") + controller = ray.get_actor("TransferQueueController", namespace="transfer_queue") + rows = tq.read_row_index(dump_dir)["rows"] + units = list(tq.get_client().storage_manager.storage_unit_infos) + ray.get(controller.begin_restore.remote("interrupted", str(dump_dir.resolve()), "p", rows, units, {})) + marker.write_text("interrupted") + tq.close() + tq.init(config) + for _ in range(2): + with pytest.raises(tq.RestorePendingError, match="cancel=True") as error: + tq.recover_data_load(dump_dir) + assert error.value.reason == "unknown_restore" + assert marker.exists() + tq.kv_put("unrelated", "other", {"x": torch.tensor([99])}) + tq.save_checkpoint(tmp_path / "checkpoint") + assert tq.recover_data_load(dump_dir, cancel=True) is False + assert not marker.exists() + tq.load_data_by_key(dump_dir) + assert tq.kv_batch_get(["key"], "p", ["x"])["x"][0].item() == 7 + finally: + tq.close() + ray.shutdown() + + def test_cancelled_delayed_load_cannot_overwrite_reused_index(tmp_path): ray.init(namespace="review_timeout") tq.init( diff --git a/tests/test_restore_lifecycle.py b/tests/test_restore_lifecycle.py index ad480915..78e4613d 100644 --- a/tests/test_restore_lifecycle.py +++ b/tests/test_restore_lifecycle.py @@ -52,6 +52,125 @@ def begin(controller, restore_id="r"): return controller.begin_restore(restore_id, "/dump", "p", {"k": {"fields": ["x"], "tag": {}}}, ["u"], {}) +@pytest.fixture +def recovery(controller): + client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) + client._restore_context = lambda restore_id: {"restore_id": restore_id} + cls = SimpleStorageUnit.__ray_metadata__.modified_class + unit = cls.__new__(cls) + unit.storage_unit_id = "u" + unit._restore_results = {} + unit._restore_controller_request = lambda context, action, result=None: controller.restore_unit( + context["restore_id"], "u", action, result + ) + unit._load_rows = lambda *_: pytest.fail("Unacknowledged claim must not write payload") + manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) + manager.storage_manager_id = "manager" + manager.storage_unit_infos = {"u": None} + manager.close = lambda: None + + async def report(context, target_storage_unit): + request = ZMQMessage.create(request_type=ZMQRequestType.REPORT_RESTORE, sender_id="test", body=context) + response = unit._handle_report_restore(request) + socket = SimpleNamespace( + send_multipart=AsyncMock(), recv_multipart=AsyncMock(return_value=response.serialize()) + ) + await manager._report_restore_unit.__wrapped__(manager, context, target_storage_unit, socket=socket) + + async def rpc(action, body): + if action == ZMQRequestType.LIST_RESTORES: + return {"restore_ids": controller.list_restores(body["dump_dir"])} + assert action == ZMQRequestType.FINISH_RESTORE + return controller.finish_restore(**body) + + manager._report_restore_unit = report + client._restore_rpc = rpc + client.storage_manager = manager + return client, unit + + +@pytest.mark.asyncio +@pytest.mark.parametrize("claim_arrived", [False, True]) +async def test_default_recovery_settles_failed_claim_without_payload(controller, recovery, claim_arrived): + client, unit = recovery + begin(controller) + original_request = unit._restore_controller_request + + def lose_claim_reply(context, action, result=None): + if claim_arrived: + original_request(context, action, result) + raise zmq.error.Again() + + unit._restore_controller_request = lose_claim_reply + response = unit._handle_load_rows( + ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS, + sender_id="test", + body={"restore": {"restore_id": "r"}, "shards": []}, + ) + ) + assert not response.body["success"] + unit._restore_controller_request = original_request + assert await client.async_recover_data_load("/dump", ["r"]) is False + assert not unit._restore_results + assert not controller.list_restores("/dump") + with pytest.raises(RuntimeError, match="no longer active"): + controller.restore_unit("r", "u", "claim") + controller.clear_partition("p") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("cancel", [False, True]) +async def test_pending_recovery_explains_options_without_cancelling_delayed_load(controller, recovery, cancel): + client, _ = recovery + begin(controller) + with pytest.raises(RestorePendingError, match="cancel=True") as error: + await client.async_recover_data_load("/dump") + assert error.value.reason == "unfinished_units" + assert error.value.unit_states == {"u": "pending"} + assert not controller._restores["r"]["aborting"] + if cancel: + assert await client.async_recover_data_load("/dump", cancel=True) is False + with pytest.raises(RuntimeError, match="no longer active"): + controller.restore_unit("r", "u", "claim") + else: + controller.restore_unit("r", "u", "claim") + controller.restore_unit("r", "u", "complete", {"success": True, "updates": []}) + assert await client.async_recover_data_load("/dump") is True + + +@pytest.mark.asyncio +async def test_failed_pending_claim_keeps_other_running_unit_reserved(controller, recovery): + client, unit = recovery + rows = {key: {"fields": ["x"], "tag": {}} for key in ["a", "b"]} + controller.begin_restore("r", "/dump", "p", rows, ["u", "other"], {}) + controller.restore_unit("r", "other", "claim") + unit._restore_results["r"] = {"success": False, "claim_failed": True, "message": "claim request lost"} + with pytest.raises(RestorePendingError) as error: + await client.async_recover_data_load("/dump") + assert error.value.unit_states == {"other": "running"} + with pytest.raises(RuntimeError, match="unresolved"): + controller.clear_partition("p") + controller.restore_unit("r", "other", "complete", {"success": True, "updates": []}) + assert await client.async_recover_data_load("/dump") is False + assert not controller.partitions["p"].field_metadata + + +@pytest.mark.parametrize( + "unit,result", + [ + ("u", {"success": True, "claim_failed": True}), + ("u", {"success": False}), + ("foreign", {"success": False, "claim_failed": True}), + ], +) +def test_pending_unit_cannot_report_success_or_unconfirmed_failure(controller, unit, result): + begin(controller) + with pytest.raises(RuntimeError): + controller.restore_unit("r", unit, "complete", result) + assert controller._restores["r"]["units"] == {"u": "pending"} + + @pytest.mark.asyncio async def test_report_failures_preserve_other_unit_reports(caplog): manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 08684314..f5d4c98b 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1220,7 +1220,12 @@ async def _finish_data_load( except (zmq.error.Again, TimeoutError) as error: raise RestorePendingError(restore_id, report_errors=report_errors) from error if not result["finished"]: - raise RestorePendingError(restore_id, report_errors=report_errors) + raise RestorePendingError( + restore_id, + reason=result.get("reason", "unknown_outcome"), + unit_states=result.get("unit_states"), + report_errors=report_errors, + ) return result["committed"] async def async_load_rows_by_key( diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index 5752f359..4f393baf 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -1861,6 +1861,15 @@ def restore_unit(self, restore_id: str, unit_id: str, action: str, result: dict if state != "pending" or restore["aborting"]: raise RuntimeError(f"Restore unit cannot start in state {state}") restore["units"][unit_id] = "running" + elif ( + action == "complete" + and state == "pending" + and result is not None + and result.get("claim_failed") is True + and result.get("success") is False + ): + # The claim may never have arrived; this worker confirms it did not write. + restore["units"][unit_id] = "failed" elif action == "complete" and state in ("running", "done", "failed") and result is not None: restore["units"][unit_id] = "done" if result["success"] else "failed" restore["updates"][unit_id] = result.get("updates", []) @@ -1875,7 +1884,7 @@ def finish_restore(self, restore_id: str, commit: bool) -> dict: restore = self._restores.get(restore_id) if restore is None: if commit: - return {"finished": False} + return {"finished": False, "reason": "unknown_restore"} self._restore_outcomes[restore_id] = False return {"finished": True, "committed": False} if not commit or "failed" in restore["units"].values(): @@ -1883,9 +1892,14 @@ def finish_restore(self, restore_id: str, commit: bool) -> dict: for unit, state in restore["units"].items(): if state == "pending": restore["units"][unit] = "cancelled" - running = [unit for unit, state in restore["units"].items() if state in ("pending", "running")] - if running: - return {"finished": False, "units": running} + unfinished = {unit: state for unit, state in restore["units"].items() if state in ("pending", "running")} + if unfinished: + return { + "finished": False, + "reason": "unfinished_units", + "units": list(unfinished), + "unit_states": unfinished, + } committed = commit and not restore["aborting"] if committed: partition = self.partitions[restore["partition_id"]] diff --git a/transfer_queue/storage/dump_io.py b/transfer_queue/storage/dump_io.py index 1b17b0b8..11772fef 100644 --- a/transfer_queue/storage/dump_io.py +++ b/transfer_queue/storage/dump_io.py @@ -77,12 +77,37 @@ def pack_dump_field(values: list, schema: dict): class RestorePendingError(RuntimeError): - """The controller still reserves indexes until remote restore activity is settled.""" + """Recovery is unresolved; expose unit states and reporting failures without releasing protection.""" - def __init__(self, restore_id: str, *, report_errors: dict[str, str] | None = None): + def __init__( + self, + restore_id: str, + *, + reason: str = "unknown_outcome", + unit_states: dict[str, str] | None = None, + report_errors: dict[str, str] | None = None, + ): self.restore_id = restore_id + self.reason = reason + self.unit_states = unit_states or {} self.report_errors = report_errors or {} - message = f"Restore {restore_id} has an unknown outcome; run recover_data_load before retrying or clearing" + message = f"Restore {restore_id} is unresolved ({reason})" + if reason == "unknown_restore": + message += ( + "; the controller has no record of this ID. If the initiating client has stopped, " + "or the whole TQ system was restarted after stopping all old actors, " + "use recover_data_load(dump_dir, cancel=True) to discard the stale operation" + ) + elif self.unit_states: + message += f"; unit states: {self.unit_states}. Pending units have not claimed permission; " + message += "running units have not reported completion. Retry recover_data_load to await completion" + if "pending" in self.unit_states.values(): + message += ( + "; if the initiating client exited before dispatching all requests, " + "use recover_data_load(dump_dir, cancel=True) instead of retrying indefinitely" + ) + else: + message += "; the remote outcome is unknown. Run recover_data_load before retrying or clearing" if self.report_errors: message += f". Unit report failures: {self.report_errors}" super().__init__(message) diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index eff33529..fbfc7a67 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -1420,7 +1420,11 @@ def _handle_load_rows(self, request: ZMQMessage) -> ZMQMessage: except zmq.error.Again as e: # A lost claim ACK may leave the controller in running state even though # this worker will not write; retain a terminal result for recovery. - self._restore_results[context["restore_id"]] = {"success": False, "message": str(e)} + self._restore_results[context["restore_id"]] = { + "success": False, + "claim_failed": True, + "message": f"Restore claim was not acknowledged; payload was not written: {e}", + } return ZMQMessage.create( request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, sender_id=self.storage_unit_id, From 2e66a787492a9167be0924a4467a80cc2b47af0c Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Tue, 29 Sep 2026 14:16:45 +0800 Subject: [PATCH 22/30] docs: clarify selective dump format compatibility Signed-off-by: neowywang --- docs/checkpoint.md | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/docs/checkpoint.md b/docs/checkpoint.md index 7da407ea..2d41c13b 100644 --- a/docs/checkpoint.md +++ b/docs/checkpoint.md @@ -172,6 +172,8 @@ If step (1) partially succeeds and step (2) fails, the system is left in a mixed ## Exporting selected keys Use [selective data dumps](data_dump.md) when restoring selected keys into an -existing system or a different number of storage units. Version-2 SimpleStorage -restores use the same owner-side I/O pattern as checkpoint load, but read assigned -row ranges and merge values instead of replacing entire unit and controller state. +existing system or a different number of storage units. SimpleStorage restores for +dump formats v2 and later read assigned row ranges directly on the current storage +owners and merge values instead of replacing entire unit and controller state. +New dumps use v3, which also preserves field schemas. Version-1 dumps use the +caller-side KV put fallback. From f8f62f1cab88b998affb38095472795f4ab1f3b3 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 8 Oct 2026 20:00:19 +0800 Subject: [PATCH 23/30] refactor: restore selected rows through the ordinary put contract Drop the restore reservation protocol (BEGIN/RESTORE_UNIT/FINISH/LIST/ REPORT_RESTORE, the .restore marker, RestorePendingError and recover_data_load). A load now resolves target indexes with kv_retrieve_meta(create=True), has the owner units read their records, publishes the saved schemas through notify_data_update once every unit has succeeded, and then writes the tags. Failure semantics match kv_batch_put: payload writes may be partial, and retrying is idempotent because existing keys keep their indexes. Storage units no longer call the controller, and the controller no longer depends on SimpleStorage routing, so the routing helper moves back into the manager. Fencing late writes is left to a follow-up designed together with the regular put path. Signed-off-by: OutstanderWang --- docs/data_dump.md | 75 +-- tests/e2e/test_data_dump_e2e.py | 74 +-- tests/e2e/test_restore_timeout_e2e.py | 226 ------- tests/test_data_dump.py | 18 +- tests/test_restore_lifecycle.py | 552 ------------------ transfer_queue/__init__.py | 5 +- transfer_queue/client.py | 131 +---- transfer_queue/controller.py | 463 +++++---------- transfer_queue/data_dump.py | 62 +- transfer_queue/interface.py | 1 - transfer_queue/storage/dump_io.py | 37 -- .../managers/simple_storage_manager.py | 85 ++- transfer_queue/storage/simple_storage.py | 84 --- transfer_queue/utils/storage_routing.py | 35 -- transfer_queue/utils/zmq_utils.py | 11 - 15 files changed, 250 insertions(+), 1609 deletions(-) delete mode 100644 tests/e2e/test_restore_timeout_e2e.py delete mode 100644 tests/test_restore_lifecycle.py delete mode 100644 transfer_queue/utils/storage_routing.py diff --git a/docs/data_dump.md b/docs/data_dump.md index 462d1e31..4625a183 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -35,10 +35,8 @@ On version-3 SimpleStorage restore: 4. Each target unit reads only its assigned byte ranges and merges those values into local storage. Records are processed in batches of at most 128 rows per shard; the caller never reads or forwards their payloads. -5. Each unit claims permission from the controller before writing, then reports - completion directly. The client commits the saved schemas and tags only after - every unit has completed. The controller reserves the destination partition - until commit or confirmed cancellation. +5. After every unit has succeeded, the storage manager publishes the saved schemas + to the controller, as an ordinary put does, and the client then writes the tags. The number of source units can differ from the number of destination units. Even a dump with one source shard can restore across several target units because @@ -109,71 +107,20 @@ the new directory, and syncs its parent before deleting the backup. If publicati is interrupted while the main directory is absent, the next dump, load or row-index read recovers `.old`. Readers perform recovery only while holding the same lock as publishers, so a healthy rename window is never mistaken for a crashed writer. -The load keeps the lock through all remote reads; a pending load marker continues -to prevent replacement after a timeout or client exit. Do not delete the sibling +The load keeps the lock through all remote reads. Do not delete the sibling lock file: unlinking it can create two independent locks for the same dump. The shared filesystem must provide cross-node advisory locking (not local-only locks). A backup-cleanup error does not invalidate a published dump. -Restore is not transactional: payload writes before a failure remain. Every load -has a unique ID. The controller blocks clearing/reusing its destination indexes -and conflicting KV puts while an operation is unresolved. Units must claim that ID -before writing; cancellation rejects requests that have not yet claimed permission. -A receive timeout never releases a writer that has already claimed permission. -Timeouts leave the operation pending, even if every unit later succeeds. They do -not cancel it or require changing the timeout used by ordinary puts and gets. +Restore has the failure semantics of `kv_batch_put`: it is not transactional, and +payload writes before a failure remain. Metadata is published only after every unit +has succeeded, so a failed load leaves those writes invisible. Existing keys keep +their indexes, so retrying the same load is idempotent; clearing the keys abandons it. +Like an ordinary put, a load does not fence late writes against indexes that are +cleared and reused while it runs, so keep writers and clears for these keys paused. -`RestorePendingError` means remote work is still running or its outcome is unknown. -The dump also retains a sibling `.restore` marker so an interrupted client cannot -silently allow its files to be replaced. After an interruption, call: - -```python -committed = tq.recover_data_load("/shared/dumps/selected") -``` - -Recovery asks units to resend terminal results and commits schemas and tags only -when all units succeeded. It returns `True` for committed loads (or no pending -load), and `False` for a failed or cancelled load after all claimed workers stopped. -Unresolved work raises `RestorePendingError`, which includes `reason`, `unit_states` -and `report_errors`. Reports that fail retain the unit ID and original error message. -The controller remembers terminal outcomes so a lost commit reply can be confirmed -safely by retrying recovery; a redundant report failure cannot reverse that outcome. - -- `running` units have claimed permission but have not reported completion. Retry - recovery later; a reporting timeout does not prove that a worker has stopped. -- `pending` units have not claimed permission. A request may still be queued, so - recovery does not automatically cancel it. If the initiating client exited before - sending all requests, use `recover_data_load(dump_dir, cancel=True)` to abandon the - operation; simply repeating recovery cannot dispatch the missing requests. -- `reason="unknown_restore"` means the controller has no record of the ID in the - `.restore` marker. If the initiating client has stopped, use explicit cancellation. - After a whole-system restart, first ensure all old actors have stopped, then cancel - the stale operation. This marker blocks its dump path, not the new controller's - unrelated partitions or checkpoints. - -A unit whose claim timed out caches a failure confirming that it did not write. -Recovery accepts this failure even when the claim never reached the controller, -cancels unclaimed work and waits for any other claimed workers to finish. - -To abandon the load explicitly, use `recover_data_load(dump_dir, cancel=True)`. -This denies unclaimed work and retains the reservation until claimed workers stop. -Cancellation cannot undo a load that has already committed. A failed unit also -cancels remaining unclaimed work; partial payload writes are never rolled back. -After cancellation settles, retry the load or clear its keys. A lost unit requires -stopping the old TQ actors and restarting the whole TQ system; restarting only the -controller while old storage actors run is unsupported. -After restarting, cancel the old marker before loading the dump again: - -```python -# Run only after all old TQ actors have stopped and the new system is initialized. -tq.recover_data_load(dump_dir, cancel=True) -tq.load_data_by_key(dump_dir) -``` - -Writers that already hold low-level metadata must remain paused throughout recovery. -The controller reservation covers only the destination partition. Each storage unit -still serves requests on one worker thread: other partitions using that unit can -wait behind a load. The 128-row batches bound memory, not request latency. +Each storage unit serves requests on one worker thread: other partitions using that +unit can wait behind a load. The 128-row batches bound memory, not request latency. ## Tests diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index 13b7cf82..7b6f14dd 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -583,7 +583,6 @@ def no_shard_read(path, *args, **kwargs): if legacy_metadata and last_kind != "tensor" and target == dump_dir: with pytest.raises(RuntimeError, match="tensor/non-tensor type mismatch"): tq.load_data_by_key(target) - assert not target.with_name(target.name + ".restore").exists() tq.kv_clear(keys, partition) tq.load_data_by_key(target) for start in range(0, row_count, 128): @@ -638,69 +637,28 @@ def test_regular_put_keeps_legacy_tensor_nontensor_acceptance(tq_system, control assert snapshot.field_metadata["x"].global_indexes == set(metadata.global_indexes) -@pytest.mark.parametrize("lost_reply", ["load", "commit"]) -def test_recovery_commits_after_lost_reply(tq_system, dump_dir, monkeypatch, lost_reply): - import zmq - - from transfer_queue.utils.zmq_utils import ZMQRequestType - - _put_rows("lost_reply", ["key"]) - tq.dump_data_by_key(dump_dir, ["key"], "lost_reply") +def test_failed_load_publishes_no_metadata_and_retry_is_idempotent(tq_system, dump_dir, controller, monkeypatch): + _put_rows("retry", ["key"]) + tq.dump_data_by_key(dump_dir, ["key"], "retry") client = tq.get_client() - client.clear_partition("lost_reply") + client.clear_partition("retry") manager = client.storage_manager original_load = manager._load_selected_rows - original_rpc = client._restore_rpc - loads = [] - async def load(*args, **kwargs): - loads.append(kwargs["target_storage_unit"]) - result = await original_load(*args, **kwargs) - if lost_reply == "load": - raise zmq.error.Again() - return result - - async def rpc(action, body): - result = await original_rpc(action, body) - if lost_reply == "commit" and action == ZMQRequestType.FINISH_RESTORE and body["commit"]: - raise zmq.error.Again() - return result + async def lose_reply(*args, **kwargs): + await original_load(*args, **kwargs) + raise RuntimeError("lost reply") with monkeypatch.context() as patcher: - patcher.setattr(manager, "_load_selected_rows", load) - patcher.setattr(client, "_restore_rpc", rpc) - with pytest.raises(tq.RestorePendingError): + patcher.setattr(manager, "_load_selected_rows", lose_reply) + with pytest.raises(RuntimeError, match="lost reply"): tq.load_data_by_key(dump_dir) - assert dump_dir.with_name(dump_dir.name + ".restore").exists() - assert tq.recover_data_load(dump_dir) is True - assert not dump_dir.with_name(dump_dir.name + ".restore").exists() - assert len(loads) == 1 - _assert_rows_equal(tq.kv_batch_get(["key"], "lost_reply", ["input_ids"])["input_ids"], [_row_input_ids(0)]) - + snapshot = ray.get(controller.get_partition_snapshot.remote("retry")) + index = snapshot.keys_mapping["key"] + assert not snapshot.field_metadata -def test_running_restore_blocks_clear_and_dump_until_recovery(tq_system, dump_dir, controller): - _put_rows("reserved", ["key"]) - tq.dump_data_by_key(dump_dir, ["key"], "reserved") - client = tq.get_client() - manager = client.storage_manager - rows = tq.read_row_index(dump_dir)["rows"] - tq.save_checkpoint(dump_dir.parent / "checkpoint") - units = list(manager.storage_unit_infos) - metadata = ray.get( - controller.begin_restore.remote("running-test", str(dump_dir.resolve()), "reserved", rows, units, {}) - ) - owner = units[metadata.global_indexes[0] % len(units)] - ray.get(controller.restore_unit.remote("running-test", owner, "claim")) - with pytest.raises(RuntimeError, match="unresolved"): - client.clear_partition("reserved") - with pytest.raises(tq.RestorePendingError): - tq.dump_data_by_key(dump_dir, ["key"], "reserved") - with pytest.raises(tq.RestorePendingError): - tq.load_checkpoint(dump_dir.parent / "checkpoint") - with pytest.raises(tq.RestorePendingError): - tq.recover_data_load(dump_dir) - ray.get(controller.restore_unit.remote("running-test", owner, "complete", {"success": False})) - tq.recover_data_load(dump_dir) - client.clear_partition("reserved") tq.load_data_by_key(dump_dir) - assert tq.kv_batch_get(["key"], "reserved", ["input_ids"]).batch_size[0] == 1 + snapshot = ray.get(controller.get_partition_snapshot.remote("retry")) + assert snapshot.keys_mapping["key"] == index + _assert_rows_equal(tq.kv_batch_get(["key"], "retry", ["input_ids"])["input_ids"], [_row_input_ids(0)]) + assert tq.kv_list("retry")["retry"] == {"key": {"idx": 0}} diff --git a/tests/e2e/test_restore_timeout_e2e.py b/tests/e2e/test_restore_timeout_e2e.py deleted file mode 100644 index 51e95861..00000000 --- a/tests/e2e/test_restore_timeout_e2e.py +++ /dev/null @@ -1,226 +0,0 @@ -# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. -# Copyright 2025 The TransferQueue Team -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -import multiprocessing -import threading -import time -from unittest.mock import patch - -import pytest -import ray -import torch -import zmq -import zmq.asyncio -from omegaconf import OmegaConf -from tensordict import TensorDict - -import transfer_queue as tq -from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager -from transfer_queue.utils.zmq_utils import ZMQMessage - - -def _check_claimed_load_survives_receive_timeout(tmp_path): - ray.init(namespace="review_claimed_timeout") - try: - tq.init( - OmegaConf.create( - { - "backend": { - "storage_backend": "SimpleStorage", - "SimpleStorage": { - "num_data_storage_units": 1, - "total_storage_size": 20, - }, - } - } - ) - ) - tq.kv_batch_put(["key"], "p", TensorDict({"x": torch.tensor([[7]])}, batch_size=1), tags=[{"saved": True}]) - dump_dir = tmp_path / "dump" - tq.dump_data_by_key(dump_dir, ["key"], "p") - tq.get_client().clear_partition("p") - actor = ray.get_actor("TransferQueueStorageUnit#0") - claimed = tmp_path / "claimed" - release = tmp_path / "release" - - def pause_after_claim(unit): - original = unit._load_rows - - def delayed(request): - claimed.touch() - deadline = time.monotonic() + 10 - while not release.exists() and time.monotonic() < deadline: - time.sleep(0.01) - return original(request) - - unit._load_rows = delayed - - ray.get(actor.__ray_call__.remote(pause_after_claim)) - # The pool fixes each socket's timeout at connect, so shorten the live pool itself; - # patching the module constant would only reach sockets opened after this point. - pool = tq.get_client().storage_manager.storage_rpc_pool - with patch.object(pool, "_timeout", 1), patch.object(pool, "_idle", {}): - with pytest.raises(tq.RestorePendingError): - tq.load_data_by_key(dump_dir) - assert claimed.exists() - assert dump_dir.with_name(dump_dir.name + ".restore").exists() - with pytest.raises(RuntimeError, match="unresolved"): - tq.get_client().clear_partition("p") - with pytest.raises(tq.RestorePendingError): - tq.dump_data_by_key(dump_dir, ["key"], "p") - release.touch() - assert tq.recover_data_load(dump_dir) is True - assert not dump_dir.with_name(dump_dir.name + ".restore").exists() - assert tq.kv_batch_get(["key"], "p", ["x"])["x"][0].item() == 7 - assert tq.get_client().kv_retrieve_meta(["key"], "p").custom_meta == [{"saved": True}] - finally: - tq.close() - ray.shutdown() - - -def test_claimed_load_survives_receive_timeout(tmp_path): - # Dynamic Ray actor instrumentation must not affect later in-process unit mocks. - context = multiprocessing.get_context("spawn") - process = context.Process(target=_check_claimed_load_survives_receive_timeout, args=(tmp_path,)) - process.start() - try: - process.join(60) - assert process.exitcode == 0 - finally: - if process.is_alive(): - process.terminate() - process.join(10) - - -def test_whole_system_restart_explains_and_cancels_orphan_marker(tmp_path): - ray.init(namespace="restart_recovery") - config = OmegaConf.create( - { - "backend": { - "storage_backend": "SimpleStorage", - "SimpleStorage": { - "num_data_storage_units": 1, - "total_storage_size": 20, - }, - } - } - ) - dump_dir = tmp_path / "dump" - marker = tmp_path / "dump.restore" - try: - tq.init(config) - tq.kv_put("key", "p", {"x": torch.tensor([7])}) - tq.dump_data_by_key(dump_dir, ["key"], "p") - controller = ray.get_actor("TransferQueueController", namespace="transfer_queue") - rows = tq.read_row_index(dump_dir)["rows"] - units = list(tq.get_client().storage_manager.storage_unit_infos) - ray.get(controller.begin_restore.remote("interrupted", str(dump_dir.resolve()), "p", rows, units, {})) - marker.write_text("interrupted") - tq.close() - tq.init(config) - for _ in range(2): - with pytest.raises(tq.RestorePendingError, match="cancel=True") as error: - tq.recover_data_load(dump_dir) - assert error.value.reason == "unknown_restore" - assert marker.exists() - tq.kv_put("unrelated", "other", {"x": torch.tensor([99])}) - tq.save_checkpoint(tmp_path / "checkpoint") - assert tq.recover_data_load(dump_dir, cancel=True) is False - assert not marker.exists() - tq.load_data_by_key(dump_dir) - assert tq.kv_batch_get(["key"], "p", ["x"])["x"][0].item() == 7 - finally: - tq.close() - ray.shutdown() - - -def test_cancelled_delayed_load_cannot_overwrite_reused_index(tmp_path): - ray.init(namespace="review_timeout") - tq.init( - OmegaConf.create( - { - "backend": { - "storage_backend": "SimpleStorage", - "SimpleStorage": {"num_data_storage_units": 1, "total_storage_size": 20}, - } - } - ) - ) - try: - tq.kv_batch_put(["key"], "p", TensorDict({"x": torch.tensor([[1]])}, batch_size=1)) - tq.dump_data_by_key(tmp_path / "dump", ["key"], "p") - tq.kv_batch_put(["key"], "p", TensorDict({"x": torch.tensor([[2]])}, batch_size=1)) - manager = tq.get_client().storage_manager - real_info = next(iter(manager.storage_unit_infos.values())) - real_addr = real_info.to_addr("put_get_socket") - # Relay the request immediately, but delay delivery to the unit to model a queued load. - ready = threading.Event() - finished = threading.Event() - port = [] - - def relay(): - ctx = zmq.Context() - front = ctx.socket(zmq.ROUTER) - front.setsockopt(zmq.RCVTIMEO, 10000) - port.append(front.bind_to_random_port("tcp://127.0.0.1")) - back = ctx.socket(zmq.DEALER) - back.setsockopt(zmq.RCVTIMEO, 10000) - back.setsockopt(zmq.IDENTITY, b"TQ_STORAGE_review_relay") - back.connect(real_addr) - ready.set() - req = front.recv_multipart() - time.sleep(2) - back.send_multipart(req[1:]) - reply = back.recv_multipart() - assert not ZMQMessage.deserialize(reply).body["success"] - finished.set() - front.close(linger=0) - back.close(linger=0) - ctx.term() - - thread = threading.Thread(target=relay, daemon=True) - thread.start() - assert ready.wait(10) - - async def timeout_load(shards, target_storage_unit, restore): - ctx = zmq.asyncio.Context() - sock = ctx.socket(zmq.DEALER) - sock.setsockopt(zmq.RCVTIMEO, 100) - sock.connect(f"tcp://127.0.0.1:{port[0]}") - try: - return await AsyncSimpleStorageManager._load_selected_rows.__wrapped__( - manager, shards, target_storage_unit, restore=restore, socket=sock - ) - finally: - sock.close(linger=0) - ctx.term() - - with patch.object(manager, "_load_selected_rows", timeout_load): - with pytest.raises(tq.RestorePendingError): - tq.load_data_by_key(tmp_path / "dump") - assert tq.recover_data_load(tmp_path / "dump", cancel=True) is False - value = tq.kv_batch_get(["key"], "p", ["x"])["x"][0].item() - - assert value == 2 - tq.get_client().clear_partition("p") - tq.kv_batch_put(["other"], "unrelated", TensorDict({"x": torch.tensor([[99]])}, batch_size=1)) - assert finished.wait(10) - thread.join(timeout=10) - later = tq.kv_batch_get(["other"], "unrelated", ["x"])["x"][0].item() - - assert later == 99 - finally: - tq.close() - ray.shutdown() diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index 02ac5f29..6f58e29f 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -192,7 +192,7 @@ def read(self, size=-1): records.append( {"source_index": source, "target_index": target, "fields": ["x"], "offset": offset, "length": length} ) - loaded = unit._load_rows( + loaded = unit._handle_load_rows( ZMQMessage.create( request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", @@ -232,7 +232,7 @@ async def load(shards, target_storage_unit): manager._load_selected_rows = load records = [{"source_index": i, "target_index": 31 - i} for i in range(16)] - assert await manager.load_rows_by_index([{"path": "shard.pkl", "records": records}]) == [] + assert await manager.load_rows_by_index("p", [{"path": "shard.pkl", "records": records}]) == 0 assert sorted(seen) == list(range(16)) @@ -247,7 +247,7 @@ def test_unit_rejects_invalid_records(unit, tmp_path, problem): record["length"] += 1 else: record["fields"] = ["missing"] - reply = unit._load_rows( + reply = unit._handle_load_rows( ZMQMessage.create( request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", @@ -329,12 +329,20 @@ async def load(shards, target_storage_unit): await failed.wait() await asyncio.sleep(0) finished.append(target_storage_unit) - return {"bytes_read": 0, "updates": []} + return {"bytes_read": 0, "updates": [{"global_indexes": [1], "field_schema": {}}]} + async def notify(*args): + notified.append(args) + + notified = [] manager._load_selected_rows = load + manager.notify_data_update = notify with pytest.raises(RuntimeError, match="unit failed"): - await manager.load_rows_by_index([{"path": "shard", "records": [{"target_index": 0}, {"target_index": 1}]}]) + await manager.load_rows_by_index( + "p", [{"path": "shard", "records": [{"target_index": 0}, {"target_index": 1}]}] + ) assert finished == ["u1"] + assert not notified @pytest.mark.asyncio diff --git a/tests/test_restore_lifecycle.py b/tests/test_restore_lifecycle.py deleted file mode 100644 index 78e4613d..00000000 --- a/tests/test_restore_lifecycle.py +++ /dev/null @@ -1,552 +0,0 @@ -# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. -# Copyright 2025 The TransferQueue Team -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Restore reservations prevent late writes from corrupting reused indexes.""" - -import asyncio -from threading import RLock -from types import SimpleNamespace -from unittest.mock import AsyncMock - -import pytest -import torch -import zmq - -from transfer_queue.client import AsyncTransferQueueClient -from transfer_queue.controller import PartitionIndexManager, TransferQueueController -from transfer_queue.sampler import SequentialSampler -from transfer_queue.storage.dump_io import RestorePendingError -from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager -from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData -from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType - - -@pytest.fixture -def controller(): - cls = TransferQueueController.__ray_metadata__.modified_class - controller = cls.__new__(cls) - controller.controller_id = "test_controller" - controller.partitions = {} - controller.index_manager = PartitionIndexManager() - controller.sampler = SequentialSampler() - controller._restore_lock = RLock() - controller._restores = {} - controller._restore_outcomes = {} - controller._clearing_indexes = set() - return controller - - -def begin(controller, restore_id="r"): - return controller.begin_restore(restore_id, "/dump", "p", {"k": {"fields": ["x"], "tag": {}}}, ["u"], {}) - - -@pytest.fixture -def recovery(controller): - client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) - client._restore_context = lambda restore_id: {"restore_id": restore_id} - cls = SimpleStorageUnit.__ray_metadata__.modified_class - unit = cls.__new__(cls) - unit.storage_unit_id = "u" - unit._restore_results = {} - unit._restore_controller_request = lambda context, action, result=None: controller.restore_unit( - context["restore_id"], "u", action, result - ) - unit._load_rows = lambda *_: pytest.fail("Unacknowledged claim must not write payload") - manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) - manager.storage_manager_id = "manager" - manager.storage_unit_infos = {"u": None} - manager.close = lambda: None - - async def report(context, target_storage_unit): - request = ZMQMessage.create(request_type=ZMQRequestType.REPORT_RESTORE, sender_id="test", body=context) - response = unit._handle_report_restore(request) - socket = SimpleNamespace( - send_multipart=AsyncMock(), recv_multipart=AsyncMock(return_value=response.serialize()) - ) - await manager._report_restore_unit.__wrapped__(manager, context, target_storage_unit, socket=socket) - - async def rpc(action, body): - if action == ZMQRequestType.LIST_RESTORES: - return {"restore_ids": controller.list_restores(body["dump_dir"])} - assert action == ZMQRequestType.FINISH_RESTORE - return controller.finish_restore(**body) - - manager._report_restore_unit = report - client._restore_rpc = rpc - client.storage_manager = manager - return client, unit - - -@pytest.mark.asyncio -@pytest.mark.parametrize("claim_arrived", [False, True]) -async def test_default_recovery_settles_failed_claim_without_payload(controller, recovery, claim_arrived): - client, unit = recovery - begin(controller) - original_request = unit._restore_controller_request - - def lose_claim_reply(context, action, result=None): - if claim_arrived: - original_request(context, action, result) - raise zmq.error.Again() - - unit._restore_controller_request = lose_claim_reply - response = unit._handle_load_rows( - ZMQMessage.create( - request_type=ZMQRequestType.LOAD_ROWS, - sender_id="test", - body={"restore": {"restore_id": "r"}, "shards": []}, - ) - ) - assert not response.body["success"] - unit._restore_controller_request = original_request - assert await client.async_recover_data_load("/dump", ["r"]) is False - assert not unit._restore_results - assert not controller.list_restores("/dump") - with pytest.raises(RuntimeError, match="no longer active"): - controller.restore_unit("r", "u", "claim") - controller.clear_partition("p") - - -@pytest.mark.asyncio -@pytest.mark.parametrize("cancel", [False, True]) -async def test_pending_recovery_explains_options_without_cancelling_delayed_load(controller, recovery, cancel): - client, _ = recovery - begin(controller) - with pytest.raises(RestorePendingError, match="cancel=True") as error: - await client.async_recover_data_load("/dump") - assert error.value.reason == "unfinished_units" - assert error.value.unit_states == {"u": "pending"} - assert not controller._restores["r"]["aborting"] - if cancel: - assert await client.async_recover_data_load("/dump", cancel=True) is False - with pytest.raises(RuntimeError, match="no longer active"): - controller.restore_unit("r", "u", "claim") - else: - controller.restore_unit("r", "u", "claim") - controller.restore_unit("r", "u", "complete", {"success": True, "updates": []}) - assert await client.async_recover_data_load("/dump") is True - - -@pytest.mark.asyncio -async def test_failed_pending_claim_keeps_other_running_unit_reserved(controller, recovery): - client, unit = recovery - rows = {key: {"fields": ["x"], "tag": {}} for key in ["a", "b"]} - controller.begin_restore("r", "/dump", "p", rows, ["u", "other"], {}) - controller.restore_unit("r", "other", "claim") - unit._restore_results["r"] = {"success": False, "claim_failed": True, "message": "claim request lost"} - with pytest.raises(RestorePendingError) as error: - await client.async_recover_data_load("/dump") - assert error.value.unit_states == {"other": "running"} - with pytest.raises(RuntimeError, match="unresolved"): - controller.clear_partition("p") - controller.restore_unit("r", "other", "complete", {"success": True, "updates": []}) - assert await client.async_recover_data_load("/dump") is False - assert not controller.partitions["p"].field_metadata - - -@pytest.mark.parametrize( - "unit,result", - [ - ("u", {"success": True, "claim_failed": True}), - ("u", {"success": False}), - ("foreign", {"success": False, "claim_failed": True}), - ], -) -def test_pending_unit_cannot_report_success_or_unconfirmed_failure(controller, unit, result): - begin(controller) - with pytest.raises(RuntimeError): - controller.restore_unit("r", unit, "complete", result) - assert controller._restores["r"]["units"] == {"u": "pending"} - - -@pytest.mark.asyncio -async def test_report_failures_preserve_other_unit_reports(caplog): - manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) - manager.storage_unit_infos = dict.fromkeys(["rejected", "slow", "timeout"]) - manager.close = lambda: None - failed = asyncio.Event() - completed = [] - - async def report(restore, target_storage_unit): - if target_storage_unit == "rejected": - failed.set() - raise RuntimeError("Invalid restore transition pending -> complete") - await failed.wait() - if target_storage_unit == "timeout": - raise zmq.error.Again() - await asyncio.sleep(0) - completed.append(target_storage_unit) - - manager._report_restore_unit = report - errors = await manager.report_restore({"restore_id": "r"}) - assert completed == ["slow"] - assert set(errors) == {"rejected", "timeout"} - assert "pending -> complete" in errors["rejected"] - assert "Again" in errors["timeout"] - assert "Restore r: unit rejected" in caplog.text - assert "pending -> complete" in caplog.text - - -@pytest.mark.asyncio -async def test_report_preserves_controller_rejection_message(): - manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) - manager.storage_manager_id = "test" - manager.close = lambda: None - response = ZMQMessage.create( - request_type=ZMQRequestType.REPORT_RESTORE_RESPONSE, - sender_id="u", - body={"success": False, "message": "Invalid restore transition pending -> complete"}, - ) - socket = SimpleNamespace(send_multipart=AsyncMock(), recv_multipart=AsyncMock(return_value=response.serialize())) - with pytest.raises(RuntimeError, match="Restore r: storage unit u report failed: Invalid restore transition"): - await manager._report_restore_unit.__wrapped__(manager, {"restore_id": "r"}, "u", socket=socket) - - -@pytest.mark.asyncio -@pytest.mark.parametrize("outcome", [None, False, True, "timeout"]) -async def test_recovery_surfaces_report_errors_only_while_outcome_is_unknown(outcome): - client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) - errors = {"u": "RuntimeError: completion rejected"} - client.storage_manager = SimpleNamespace(report_restore=AsyncMock(return_value=errors)) - client._restore_context = lambda restore_id: {"restore_id": restore_id} - reply = {"finished": outcome is not None, "committed": outcome} - client._restore_rpc = AsyncMock( - side_effect=[ - {"restore_ids": ["r"]}, - zmq.error.Again() if outcome == "timeout" else reply, - ] - ) - if outcome is None or outcome == "timeout": - with pytest.raises(RestorePendingError, match="completion rejected") as error: - await client.async_recover_data_load("/dump") - assert error.value.report_errors == errors - else: - assert await client.async_recover_data_load("/dump") is outcome - assert client._restore_rpc.await_args.args == (ZMQRequestType.FINISH_RESTORE, {"restore_id": "r", "commit": True}) - - -@pytest.mark.asyncio -@pytest.mark.parametrize(("active_ids", "saved_ids"), [([], []), (["r"], []), ([], ["r"])]) -async def test_recovery_checks_backend_before_cancelling(active_ids, saved_ids): - client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) - client.storage_manager = object() - client._restore_rpc = AsyncMock(return_value={"restore_ids": active_ids}) - - if active_ids or saved_ids: - with pytest.raises(NotImplementedError, match="does not support selective load recovery"): - await client.async_recover_data_load("/dump", saved_ids) - else: - await client.async_recover_data_load("/dump", saved_ids) - - client._restore_rpc.assert_awaited_once_with(ZMQRequestType.LIST_RESTORES, {"dump_dir": "/dump"}) - - -def test_cancel_before_claim_rejects_late_load_and_late_begin(controller): - metadata = begin(controller) - assert controller.finish_restore("r", commit=False)["finished"] - controller.clear_partition("p") - current = controller.kv_retrieve_meta(["other"], "other", create=True) - assert current.global_indexes == metadata.global_indexes - with pytest.raises(RuntimeError, match="no longer active"): - controller.restore_unit("r", "u", "claim") - with pytest.raises(RuntimeError, match="already committed or cancelled"): - begin(controller) - controller.finish_restore("not-arrived", commit=False) - with pytest.raises(RuntimeError, match="already committed or cancelled"): - begin(controller, "not-arrived") - - -def test_claimed_load_keeps_indexes_until_terminal_report(controller): - metadata = begin(controller) - controller.restore_unit("r", "u", "claim") - assert not controller.finish_restore("r", commit=False)["finished"] - for action in ( - lambda: controller.mark_clearing(metadata.global_indexes, ["p"]), - lambda: controller.clear_meta(metadata.global_indexes, ["p"]), - lambda: controller.clear_partition("p"), - lambda: controller.kv_retrieve_meta(["k"], "p", create=True), - lambda: begin(controller, "overlap"), - ): - with pytest.raises(RuntimeError, match="unresolved"): - action() - assert controller.list_restores("/dump") == ["r"] - assert controller.kv_retrieve_meta(["other"], "other", create=True).global_indexes != metadata.global_indexes - controller.restore_unit("r", "u", "complete", {"success": False}) - assert controller.finish_restore("r", commit=False)["finished"] - controller.clear_partition("p") - - -def test_success_publishes_schema_and_tags_only_at_finish(controller): - metadata = begin(controller) - index = metadata.global_indexes[0] - controller.restore_unit("r", "u", "claim") - schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_nested": False, "is_non_tensor": False}} - controller.restore_unit( - "r", - "u", - "complete", - { - "success": True, - "updates": [ - {"global_indexes": [index], "field_schema": schema}, - ], - }, - ) - assert not controller.partitions["p"].field_metadata - assert controller.finish_restore("r", commit=True)["finished"] - assert controller.partitions["p"].field_metadata["x"].global_indexes == {index} - assert not controller.list_restores("/dump") - - -def test_restore_cannot_start_in_the_middle_of_clear(controller): - metadata = controller.kv_retrieve_meta(["k"], "p", create=True) - controller.mark_clearing(metadata.global_indexes, ["p"]) - with pytest.raises(RuntimeError, match="unfinished clear"): - begin(controller) - controller.clear_partition("p") - begin(controller) - - -def test_restore_rejects_schema_before_allocating_indexes(controller): - metadata = controller.kv_retrieve_meta(["existing"], "p", create=True) - partition = controller.partitions["p"] - schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_non_tensor": False, "is_nested": False}} - assert partition.update_production_status(metadata.global_indexes, [], schema) - conflict = {"x": {**schema["x"], "dtype": torch.float32}} - with pytest.raises(ValueError, match="dtype mismatch"): - controller.begin_restore("r", "/dump", "p", {"new": {"fields": ["x"], "tag": {}}}, ["u"], conflict) - assert set(partition.keys_mapping) == {"existing"} - assert not controller.list_restores("/dump") - - -def test_restore_validates_all_updates_before_publishing_readiness(controller): - metadata = controller.kv_retrieve_meta(["existing"], "p", create=True) - partition = controller.partitions["p"] - schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_non_tensor": False, "is_nested": False}} - assert partition.update_production_status(metadata.global_indexes, [], schema) - restored = begin(controller) - controller.restore_unit("r", "u", "claim") - controller.restore_unit( - "r", - "u", - "complete", - { - "success": True, - "updates": [ - {"global_indexes": restored.global_indexes, "field_schema": {"y": schema["x"]}}, - { - "global_indexes": restored.global_indexes, - "field_schema": {"x": {**schema["x"], "dtype": torch.float32}}, - }, - ], - }, - ) - with pytest.raises(ValueError, match="dtype mismatch"): - controller.finish_restore("r", commit=True) - assert "y" not in partition.field_metadata - assert not partition.production_status[restored.global_indexes].any() - assert controller.finish_restore("r", commit=False)["finished"] - - -def test_unit_rejects_cancelled_permission_before_reading(controller): - begin(controller) - controller.finish_restore("r", commit=False) - cls = SimpleStorageUnit.__ray_metadata__.modified_class - unit = cls.__new__(cls) - unit.storage_unit_id = "u" - unit.storage_data = StorageUnitData() - unit._restore_results = {} - unit._restore_controller_request = lambda context, action, result=None: controller.restore_unit( - "r", "u", action, result - ) - unit._load_rows = lambda *_: pytest.fail("Cancelled request read payload") - response = unit._handle_load_rows( - ZMQMessage.create( - request_type=ZMQRequestType.LOAD_ROWS, - sender_id="test", - body={"restore": {"restore_id": "r"}, "shards": []}, - ) - ) - assert not response.body["success"] - - -def test_lost_completion_ack_is_recoverable_without_replaying_payload(controller): - begin(controller) - cls = SimpleStorageUnit.__ray_metadata__.modified_class - unit = cls.__new__(cls) - unit.storage_unit_id = "u" - unit._restore_results = {} - - def lose_completion(context, action, result=None): - if action == "claim": - controller.restore_unit("r", "u", action, result) - else: - raise zmq.error.Again() - - unit._restore_controller_request = lose_completion - unit._load_rows = lambda *_: ZMQMessage.create( - request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, - sender_id="u", - body={"success": True, "updates": [], "bytes_read": 0}, - ) - unit._handle_load_rows( - ZMQMessage.create( - request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", body={"restore": {"restore_id": "r"}, "shards": []} - ) - ) - assert not controller.finish_restore("r", commit=False)["finished"] - unit._restore_controller_request = lambda context, action, result=None: controller.restore_unit( - "r", "u", action, result - ) - response = unit._handle_report_restore( - ZMQMessage.create(request_type=ZMQRequestType.REPORT_RESTORE, sender_id="test", body={"restore_id": "r"}) - ) - assert response.body["success"] - assert controller.finish_restore("r", commit=False)["finished"] - - -def test_lost_claim_ack_can_be_settled_without_writing(controller): - begin(controller) - cls = SimpleStorageUnit.__ray_metadata__.modified_class - unit = cls.__new__(cls) - unit.storage_unit_id = "u" - unit._restore_results = {} - - def lose_ack(context, action, result=None): - controller.restore_unit("r", "u", action, result) - raise zmq.error.Again() - - unit._restore_controller_request = lose_ack - unit._load_rows = lambda *_: pytest.fail("Unacknowledged claim wrote data") - response = unit._handle_load_rows( - ZMQMessage.create( - request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", body={"restore": {"restore_id": "r"}, "shards": []} - ) - ) - assert not response.body["success"] - assert not controller.finish_restore("r", commit=False)["finished"] - unit._restore_controller_request = lambda context, action, result=None: controller.restore_unit( - "r", "u", action, result - ) - assert unit._handle_report_restore( - ZMQMessage.create(request_type=ZMQRequestType.REPORT_RESTORE, sender_id="test", body={"restore_id": "r"}) - ).body["success"] - assert controller.finish_restore("r", commit=False)["finished"] - - -def test_failed_unit_cancels_pending_units_but_waits_for_running_units(controller): - rows = {f"k{i}": {"fields": ["x"], "tag": {}} for i in range(3)} - controller.begin_restore("r", "/dump", "p", rows, ["u0", "u1", "u2"], {}) - controller.restore_unit("r", "u0", "claim") - controller.restore_unit("r", "u1", "claim") - controller.restore_unit("r", "u0", "complete", {"success": False}) - assert not controller.finish_restore("r", commit=True)["finished"] - with pytest.raises(RuntimeError, match="cannot start"): - controller.restore_unit("r", "u2", "claim") - with pytest.raises(RuntimeError, match="unresolved"): - controller.clear_partition("p") - controller.restore_unit("r", "u1", "complete", {"success": True, "updates": []}) - assert controller.finish_restore("r", commit=True) == {"finished": True, "committed": False} - assert not controller.partitions["p"].field_metadata - - -def test_unknown_restore_stays_pending_until_explicit_cancellation(controller): - assert not controller.finish_restore("unknown", commit=True)["finished"] - assert controller.finish_restore("unknown", commit=False) == {"finished": True, "committed": False} - with pytest.raises(RuntimeError, match="already committed or cancelled"): - begin(controller, "unknown") - - -@pytest.mark.parametrize("units", [["only"], ["u0", "u1"], ["u2", "u0", "u1"]]) -@pytest.mark.parametrize("has_payload", [True, False]) -def test_restore_reserves_only_current_storage_owners(controller, units, has_payload): - controller.kv_retrieve_meta(["other"], "unrelated", create=True) - existing = controller.kv_retrieve_meta(["keep", "gap", "empty", "last"], "p", create=True) - controller.clear_meta([existing.global_indexes[1]], ["p"]) - rows = { - key: {"fields": ["x"] if has_payload and key != "empty" else [], "tag": {}} - for key in ["last", "new", "empty", "keep"] - } - metadata = controller.begin_restore("r", "/dump", "p", rows, units, {}) - indexes = [index for key, index in zip(rows, metadata.global_indexes, strict=True) if rows[key]["fields"]] - manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) - manager.storage_unit_infos = dict.fromkeys(units) - manager.close = lambda: None - routed = manager._group_by_hash(indexes) - assert set(controller._restores["r"]["units"]) == set(routed) - assert metadata.global_indexes[0] == existing.global_indexes[-1] - assert metadata.global_indexes[-1] == existing.global_indexes[0] - for unit, group in routed.items(): - assert [indexes[pos] for pos in group.batch_positions] == group.global_indexes - controller.restore_unit("r", unit, "claim") - controller.restore_unit("r", unit, "complete", {"success": True, "updates": []}) - assert controller.finish_restore("r", commit=True) == {"finished": True, "committed": True} - - -@pytest.mark.asyncio -@pytest.mark.parametrize("interruption", ["running_timeout", "lost_load_reply", "lost_commit_reply", "lost_complete"]) -async def test_pending_load_can_commit_without_replaying_payload(controller, interruption): - client = AsyncTransferQueueClient.__new__(AsyncTransferQueueClient) - client._restore_context = lambda restore_id: {"restore_id": restore_id} - completed = {} - fail_commit_reply = interruption == "lost_commit_reply" - - async def rpc(action, body): - nonlocal fail_commit_reply - if action == ZMQRequestType.BEGIN_RESTORE: - return {"metadata": controller.begin_restore(**body)} - if action == ZMQRequestType.LIST_RESTORES: - return {"restore_ids": controller.list_restores(body["dump_dir"])} - result = controller.finish_restore(**body) - if body["commit"] and result["finished"] and fail_commit_reply: - fail_commit_reply = False - raise zmq.error.Again() - return result - - async def load(shards, context): - controller.restore_unit("r", "u", "claim") - index = controller._restores["r"]["metadata"].global_indexes[0] - schema = {"x": {"dtype": torch.int64, "shape": (1,), "is_nested": False, "is_non_tensor": False}} - completed.update(success=True, updates=[{"global_indexes": [index], "field_schema": schema}]) - if interruption in ("lost_load_reply", "lost_commit_reply"): - controller.restore_unit("r", "u", "complete", completed) - if interruption in ("running_timeout", "lost_load_reply"): - raise zmq.error.Again() - - async def report(context): - controller.restore_unit("r", "u", "complete", completed) - - client._restore_rpc = rpc - client.storage_manager = SimpleNamespace( - storage_unit_infos={"u": None}, load_rows_by_index=AsyncMock(side_effect=load), report_restore=report - ) - rows = {"k": {"fields": ["x"], "tag": {"saved": True}}} - with pytest.raises(RestorePendingError): - await client.async_load_rows_by_key("p", rows, [], "/dump", "r") - if "r" in controller._restores: - assert not controller._restores["r"]["aborting"] - with pytest.raises(RuntimeError, match="unresolved"): - controller.clear_partition("p") - assert await client.async_recover_data_load("/dump", ["r"]) is True - assert await client.async_recover_data_load("/dump", ["r"]) is True - client.storage_manager.load_rows_by_index.assert_awaited_once() - partition = controller.partitions["p"] - index = partition.keys_mapping["k"] - assert partition.production_status[index, partition.field_name_mapping["x"]] == 1 - assert partition.custom_meta[index] == {"saved": True} - # Duplicate finalization must not publish metadata again after the key is cleared. - controller.clear_partition("p") - assert controller.finish_restore("r", commit=True) == {"finished": True, "committed": True} - assert "p" not in controller.partitions diff --git a/transfer_queue/__init__.py b/transfer_queue/__init__.py index a0fa4eee..97c180ce 100644 --- a/transfer_queue/__init__.py +++ b/transfer_queue/__init__.py @@ -16,7 +16,7 @@ import os from .client import TransferQueueClient -from .data_dump import dump_data_by_key, load_data_by_key, read_row_index, recover_data_load +from .data_dump import dump_data_by_key, load_data_by_key, read_row_index from .dataloader import StreamingDataLoader, StreamingDataset from .interface import ( async_kv_batch_get, @@ -45,7 +45,6 @@ from .sampler.seqlen_balanced_sampler import SeqlenBalancedSampler from .sampler.sequential_sampler import SequentialSampler from .storage import StorageKeyNotFoundError -from .storage.dump_io import RestorePendingError __all__ = ( [ @@ -77,8 +76,6 @@ "dump_data_by_key", "load_data_by_key", "read_row_index", - "recover_data_load", - "RestorePendingError", ] + [ # High-Level StreamingDataLoader Interface diff --git a/transfer_queue/client.py b/transfer_queue/client.py index f5d4c98b..66a1c726 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -18,7 +18,6 @@ import threading import weakref from typing import Any, Callable -from uuid import uuid4 import torch import zmq @@ -27,7 +26,6 @@ from transfer_queue.metadata import BatchMeta from transfer_queue.storage import StorageManagerFactory -from transfer_queue.storage.dump_io import RestorePendingError from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.zmq_utils import ( @@ -1196,110 +1194,31 @@ async def async_dump_rows_by_index( shard_dir, global_indexes, fields_by_index, **({"missing_shapes": missing_shapes} if missing_shapes else {}) ) - def _restore_context(self, restore_id: str) -> dict: - return { - "restore_id": restore_id, - "controller_ip": self._controller.ip, - "controller_address": self._controller.to_addr("request_handle_socket"), - } - - @with_controller_socket - async def _restore_rpc(self, request_type: ZMQRequestType, body: dict, socket=None) -> dict: - response_type = ZMQRequestType(request_type.value + "_RESPONSE") - response = await self._request_controller(socket, request_type, response_type, body) - return response.body - - async def _finish_data_load( - self, restore_id: str, *, commit: bool, report_errors: dict[str, str] | None = None - ) -> bool: - """Keep the reservation when completion is pending or its reply is lost.""" - try: - result = await self._restore_rpc( - ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": commit} - ) - except (zmq.error.Again, TimeoutError) as error: - raise RestorePendingError(restore_id, report_errors=report_errors) from error - if not result["finished"]: - raise RestorePendingError( - restore_id, - reason=result.get("reason", "unknown_outcome"), - unit_states=result.get("unit_states"), - report_errors=report_errors, - ) - return result["committed"] - async def async_load_rows_by_key( self, partition_id: str, rows: dict[str, dict[str, Any]], shards: list[dict[str, Any]], - dump_dir: str = "", - restore_id: str | None = None, - ) -> None: - """Reserve indexes until every unit has explicitly completed or been denied permission.""" + ) -> int: + """Restore selected payloads at the indexes their keys resolve to now; return bytes read. + + Failure semantics match ``kv_batch_put``: new keys stay registered and payload + writes may be partial, and retrying is idempotent because keys keep their indexes. + """ manager = getattr(self, "storage_manager", None) if manager is None or not hasattr(manager, "load_rows_by_index"): raise NotImplementedError("Storage backend does not support direct selective load") if not rows: - return - restore_id = restore_id or uuid4().hex - try: - response = await self._restore_rpc( - ZMQRequestType.BEGIN_RESTORE, - { - "restore_id": restore_id, - "dump_dir": dump_dir, - "partition_id": partition_id, - "rows": rows, - "units": list(manager.storage_unit_infos), - "schema": shards[0].get("field_schema", {}) if shards else {}, - }, - ) - metadata = response["metadata"] - target_indexes = dict(zip(rows, metadata.global_indexes, strict=True)) - for shard in shards: - for record in shard["records"]: - record["target_index"] = target_indexes[record["key"]] - await manager.load_rows_by_index(shards, self._restore_context(restore_id)) - if not await self._finish_data_load(restore_id, commit=True): - raise RuntimeError("Restore failed or was cancelled") - except RestorePendingError: - raise - except (zmq.error.Again, TimeoutError) as error: - raise RestorePendingError(restore_id) from error - except BaseException as error: - try: - await asyncio.shield(self._finish_data_load(restore_id, commit=False)) - except BaseException: - raise RestorePendingError(restore_id) from error - raise - - async def async_recover_data_load( - self, dump_dir: str, restore_ids: list[str] | None = None, *, cancel: bool = False - ) -> bool: - """Finish successful loads, or explicitly cancel; return whether all loads committed.""" - response = await self._restore_rpc(ZMQRequestType.LIST_RESTORES, {"dump_dir": dump_dir}) - committed = True - for restore_id in set(response["restore_ids"]) | set(restore_ids or []): - if not hasattr(self.storage_manager, "report_restore"): - raise NotImplementedError( - f"{type(self.storage_manager).__name__} does not support selective load recovery" - ) - if cancel: - try: - await self._restore_rpc(ZMQRequestType.FINISH_RESTORE, {"restore_id": restore_id, "commit": False}) - except (zmq.error.Again, TimeoutError) as error: - raise RestorePendingError(restore_id) from error - report_errors = await self.storage_manager.report_restore(self._restore_context(restore_id)) - outcome = await self._finish_data_load(restore_id, commit=not cancel, report_errors=report_errors) - committed = committed and outcome - return committed - - async def async_check_data_loads(self, dump_dir: str | None = None) -> None: - """Reject replacing a dump whose files may still be read by a storage unit.""" - response = await self._restore_rpc(ZMQRequestType.LIST_RESTORES, {"dump_dir": dump_dir}) - if response["restore_ids"]: - raise RestorePendingError(response["restore_ids"][0]) + return 0 + metadata = await self.async_kv_retrieve_meta(list(rows), partition_id, create=True) + target_indexes = dict(zip(rows, metadata.global_indexes, strict=True)) + for shard in shards: + for record in shard["records"]: + record["target_index"] = target_indexes[record["key"]] + bytes_read = await manager.load_rows_by_index(partition_id, shards) + metadata.update_custom_meta([row["tag"] for row in rows.values()]) + await self.async_set_custom_meta(metadata) + return bytes_read # ==================== Checkpoint API ==================== @with_controller_socket @@ -1490,8 +1409,6 @@ def wrapper(*args, **kwargs): self._validate_dump_schema = _make_sync(self.async_validate_dump_schema) self._dump_rows_by_index = _make_sync(self.async_dump_rows_by_index) self._load_rows_by_key = _make_sync(self.async_load_rows_by_key) - self._recover_data_load = _make_sync(self.async_recover_data_load) - self._check_data_loads = _make_sync(self.async_check_data_loads) self._save_controller_checkpoint = _make_sync(self.async_save_controller_checkpoint) self._load_controller_checkpoint = _make_sync(self.async_load_controller_checkpoint) self._save_storage_checkpoint = _make_sync(self.async_save_storage_checkpoint) @@ -1977,19 +1894,9 @@ def dump_rows_by_index( """ return self._dump_rows_by_index(shard_dir, global_indexes, fields_by_index, missing_shapes=missing_shapes) - def load_rows_by_key( - self, partition_id: str, rows: dict, shards: list[dict], dump_dir: str = "", restore_id: str | None = None - ) -> None: - """Restore selected payloads while the controller reserves destination indexes.""" - return self._load_rows_by_key(partition_id, rows, shards, dump_dir, restore_id) - - def recover_data_load(self, dump_dir: str, restore_ids: list[str] | None = None, *, cancel: bool = False) -> bool: - """Finish or cancel interrupted loads; return whether every load committed.""" - return self._recover_data_load(dump_dir, restore_ids, cancel=cancel) - - def check_data_loads(self, dump_dir: str | None = None) -> None: - """Reject replacing a dump while a remote restore may still read it.""" - return self._check_data_loads(dump_dir) + def load_rows_by_key(self, partition_id: str, rows: dict, shards: list[dict]) -> int: + """Restore selected payloads at the indexes their keys resolve to now; return bytes read.""" + return self._load_rows_by_key(partition_id, rows, shards) # ==================== Checkpoint API ==================== def save_controller_checkpoint(self, path: str) -> None: diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index 4f393baf..1fe06baf 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -21,7 +21,7 @@ from dataclasses import dataclass, field from itertools import groupby from operator import itemgetter -from threading import RLock, Thread +from threading import Thread from typing import TYPE_CHECKING, Any, cast from uuid import uuid4 @@ -39,7 +39,6 @@ from transfer_queue.utils.enum_utils import Role from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.perf_utils import IntervalPerfMonitor -from transfer_queue.utils.storage_routing import group_by_storage_unit from transfer_queue.utils.zmq_utils import ( ZMQMessage, ZMQRequestType, @@ -1024,10 +1023,6 @@ def __init__( # Partition-GlobalIndex management self.index_manager = PartitionIndexManager() # partition_id -> global_indexes - self._restore_lock = RLock() - self._restores: dict[str, dict[str, Any]] = {} - self._restore_outcomes: dict[str, bool] = {} - self._clearing_indexes: set[int] = set() # Connected storage managers tracking self._connected_storage_managers: set[str] = set() @@ -1510,34 +1505,30 @@ def mark_clearing(self, global_indexes: list[int], partition_ids: list[str]) -> global_indexes: global indexes to mark as pending deletion partition_ids: corresponding partition IDs for each global index """ - with self._restore_lock: - if global_indexes is None or partition_ids is None: - raise ValueError("global_indexes and partition_ids cannot be None") - for pid in set(partition_ids): - self._assert_not_restoring(pid) + if global_indexes is None or partition_ids is None: + raise ValueError("global_indexes and partition_ids cannot be None") - if len(global_indexes) != len(partition_ids): - raise ValueError( - f"global_indexes and partition_ids must have the same length, " - f"got {len(global_indexes)} and {len(partition_ids)}" - ) + if len(global_indexes) != len(partition_ids): + raise ValueError( + f"global_indexes and partition_ids must have the same length, " + f"got {len(global_indexes)} and {len(partition_ids)}" + ) - combined = list(zip(partition_ids, global_indexes, strict=True)) - combined.sort(key=itemgetter(0)) + combined = list(zip(partition_ids, global_indexes, strict=True)) + combined.sort(key=itemgetter(0)) - for partition_id, group in groupby(combined, key=itemgetter(0)): - partition = self._get_partition(partition_id) - if not partition: - logger.info( - f"[{self.controller_id}]: Trying to mark clearing in a non-existent partition {partition_id}. " - f"Skipping operation for this partition." - ) - continue - indexes = [idx for _, idx in group] - existing = list(set(indexes) & partition.global_indexes) - self._clearing_indexes.update(existing) - if existing and partition.production_status is not None: - partition.production_status[existing, :] = 0 + for partition_id, group in groupby(combined, key=itemgetter(0)): + partition = self._get_partition(partition_id) + if not partition: + logger.info( + f"[{self.controller_id}]: Trying to mark clearing in a non-existent partition {partition_id}. " + f"Skipping operation for this partition." + ) + continue + indexes = [idx for _, idx in group] + existing = list(set(indexes) & partition.global_indexes) + if existing and partition.production_status is not None: + partition.production_status[existing, :] = 0 def clear_partition(self, partition_id: str, clear_consumption: bool = True): """ @@ -1547,23 +1538,21 @@ def clear_partition(self, partition_id: str, clear_consumption: bool = True): partition_id: ID of the partition to clear clear_consumption: Whether to also clear consumption status """ - with self._restore_lock: - self._assert_not_restoring(partition_id) - logger.debug(f"[{self.controller_id}]: Clearing metadata in partition {partition_id}.") - partition = self._get_partition(partition_id) - if not partition: - logger.warning( - f"[{self.controller_id}]: Trying to clear a non-existent partition {partition_id}. No action taken." - ) - return + logger.debug(f"[{self.controller_id}]: Clearing metadata in partition {partition_id}.") - global_indexes_range = list(self.index_manager.get_indexes_for_partition(partition_id)) - partition.clear_data(global_indexes_range, clear_consumption) - self.index_manager.release_partition(partition_id) - self._clearing_indexes.difference_update(global_indexes_range) - self.partitions.pop(partition_id) - self.sampler.clear_cache(partition_id) + partition = self._get_partition(partition_id) + if not partition: + logger.warning( + f"[{self.controller_id}]: Trying to clear a non-existent partition {partition_id}. No action taken." + ) + return + + global_indexes_range = list(self.index_manager.get_indexes_for_partition(partition_id)) + partition.clear_data(global_indexes_range, clear_consumption) + self.index_manager.release_partition(partition_id) + self.partitions.pop(partition_id) + self.sampler.clear_cache(partition_id) def reset_consumption(self, partition_id: str, task_name: str | None = None): """ @@ -1600,57 +1589,54 @@ def clear_meta( partition_ids: IDs of the partitions to clear clear_consumption: Whether to also clear consumption status """ - with self._restore_lock: - logger.debug( - "[%s]: Clearing indexes %s in partitions %s", self.controller_id, global_indexes, partition_ids + + logger.debug( + f"[{self.controller_id}]: Clearing meta with global_indexes {global_indexes} in partition {partition_ids}" + ) + + if global_indexes is None or partition_ids is None: + raise ValueError("global_indexes and partition_ids cannot be None") + + if len(global_indexes) != len(partition_ids): + raise ValueError( + f"global_indexes and partition_ids must have the same length, " + f"got {len(global_indexes)} and {len(partition_ids)}" ) - if global_indexes is None or partition_ids is None: - raise ValueError("global_indexes and partition_ids cannot be None") - for pid in set(partition_ids): - self._assert_not_restoring(pid) + combined = list(zip(partition_ids, global_indexes, strict=True)) + combined.sort(key=itemgetter(0)) - if len(global_indexes) != len(partition_ids): - raise ValueError( - f"global_indexes and partition_ids must have the same length, " - f"got {len(global_indexes)} and {len(partition_ids)}" + for partition_id, group in groupby(combined, key=itemgetter(0)): + partition = self._get_partition(partition_id) + if not partition: + logger.info( + f"[{self.controller_id}]: Trying to clear data in a non-existent partition {partition_id}. " + f"Skipping operation for this partition." ) + continue - combined = list(zip(partition_ids, global_indexes, strict=True)) - combined.sort(key=itemgetter(0)) - - for partition_id, group in groupby(combined, key=itemgetter(0)): - partition = self._get_partition(partition_id) - if not partition: - logger.info( - f"[{self.controller_id}]: Trying to clear data in a non-existent partition {partition_id}. " - f"Skipping operation for this partition." - ) - continue - - global_indexes_to_clear = [idx for _, idx in group] - existing_global_indexes = partition.global_indexes - non_existent_global_indexes = set(global_indexes_to_clear) - existing_global_indexes - if non_existent_global_indexes: - logger.info( - f"[{self.controller_id}]: Some global_indexes to be cleared do not exist in " - f"partition {partition_id}: {non_existent_global_indexes}. They will be ignored." - ) + global_indexes_to_clear = [idx for _, idx in group] + existing_global_indexes = partition.global_indexes + non_existent_global_indexes = set(global_indexes_to_clear) - existing_global_indexes + if non_existent_global_indexes: + logger.info( + f"[{self.controller_id}]: Some global_indexes to be cleared do not exist in " + f"partition {partition_id}: {non_existent_global_indexes}. They will be ignored." + ) - global_indexes_to_clear = list(set(global_indexes_to_clear) & existing_global_indexes) - if not global_indexes_to_clear: - logger.info( - f"[{self.controller_id}]: No existing global indexes to clear in partition {partition_id}. " - f"Skipping operation for this partition." - ) - continue + global_indexes_to_clear = list(set(global_indexes_to_clear) & existing_global_indexes) + if not global_indexes_to_clear: + logger.info( + f"[{self.controller_id}]: No existing global indexes to clear in partition {partition_id}. " + f"Skipping operation for this partition." + ) + continue - # Clear data from partition - partition.clear_data(global_indexes_to_clear, clear_consumption) + # Clear data from partition + partition.clear_data(global_indexes_to_clear, clear_consumption) - # Release the specific indexes from index manager - self.index_manager.release_indexes(partition_id, global_indexes_to_clear) - self._clearing_indexes.difference_update(global_indexes_to_clear) + # Release the specific indexes from index manager + self.index_manager.release_indexes(partition_id, global_indexes_to_clear) def kv_retrieve_meta( self, @@ -1670,66 +1656,64 @@ def kv_retrieve_meta( Returns: metadata: BatchMeta of the requested keys """ - with self._restore_lock: - if create: - self._assert_not_restoring(partition_id) - logger.debug(f"[{self.controller_id}] Retrieve keys {keys} in partition {partition_id}") - # Ensure partition exists + logger.debug(f"[{self.controller_id}] Retrieve keys {keys} in partition {partition_id}") + + # Ensure partition exists + partition = self._get_partition(partition_id) + if partition is None: + if not create: + logger.warning( + f"[{self.controller_id}]: Partition {partition_id} not found. Returning empty BatchMeta." + ) + return BatchMeta.empty() + + self.create_partition(partition_id) partition = self._get_partition(partition_id) - if partition is None: - if not create: - logger.warning( - f"[{self.controller_id}]: Partition {partition_id} not found. Returning empty BatchMeta." - ) - return BatchMeta.empty() - self.create_partition(partition_id) - partition = self._get_partition(partition_id) + assert partition is not None + global_indexes = partition.kv_retrieve_indexes(keys) - assert partition is not None - global_indexes = partition.kv_retrieve_indexes(keys) - - none_indexes = [idx for idx, value in enumerate(global_indexes) if value is None] - if len(none_indexes) > 0: - if not create: - logger.warning( - f"Keys {[keys[i] for i in none_indexes]} were not found in partition {partition_id}. " - f"They will be excluded from the retrieved BatchMeta." - ) - else: - # create non-exist keys - batch_global_indexes = partition.activate_pre_allocated_indexes(len(none_indexes)) + none_indexes = [idx for idx, value in enumerate(global_indexes) if value is None] + if len(none_indexes) > 0: + if not create: + logger.warning( + f"Keys {[keys[i] for i in none_indexes]} were not found in partition {partition_id}. " + f"They will be excluded from the retrieved BatchMeta." + ) + else: + # create non-exist keys + batch_global_indexes = partition.activate_pre_allocated_indexes(len(none_indexes)) - if len(batch_global_indexes) < len(none_indexes): - new_global_indexes = self.index_manager.allocate_indexes( - partition_id, count=(len(none_indexes) - len(batch_global_indexes)) - ) - batch_global_indexes.extend(new_global_indexes) + if len(batch_global_indexes) < len(none_indexes): + new_global_indexes = self.index_manager.allocate_indexes( + partition_id, count=(len(none_indexes) - len(batch_global_indexes)) + ) + batch_global_indexes.extend(new_global_indexes) - # register global_indexes in partition - partition.global_indexes.update(batch_global_indexes) + # register global_indexes in partition + partition.global_indexes.update(batch_global_indexes) - # register key-global_indexes mapping in partition - for i in range(len(none_indexes)): - global_indexes[none_indexes[i]] = batch_global_indexes[i] - partition.keys_mapping[keys[none_indexes[i]]] = batch_global_indexes[i] - partition.revert_keys_mapping[batch_global_indexes[i]] = keys[none_indexes[i]] + # register key-global_indexes mapping in partition + for i in range(len(none_indexes)): + global_indexes[none_indexes[i]] = batch_global_indexes[i] + partition.keys_mapping[keys[none_indexes[i]]] = batch_global_indexes[i] + partition.revert_keys_mapping[batch_global_indexes[i]] = keys[none_indexes[i]] - partition.ensure_samples_capacity(max(batch_global_indexes) + 1) + partition.ensure_samples_capacity(max(batch_global_indexes) + 1) - verified_global_indexes = [idx for idx in global_indexes if idx is not None] + verified_global_indexes = [idx for idx in global_indexes if idx is not None] - # must fetch fields that the requested samples all have - col_mask = partition.production_status[verified_global_indexes, :].sum(dim=0).reshape(-1) == len( - verified_global_indexes - ) - data_fields = [] - for field_name, col_idx in partition.field_name_mapping.items(): - if col_idx < len(col_mask) and col_mask[col_idx]: - data_fields.append(field_name) + # must fetch fields that the requested samples all have + col_mask = partition.production_status[verified_global_indexes, :].sum(dim=0).reshape(-1) == len( + verified_global_indexes + ) + data_fields = [] + for field_name, col_idx in partition.field_name_mapping.items(): + if col_idx < len(col_mask) and col_mask[col_idx]: + data_fields.append(field_name) - return self.generate_batch_meta(partition_id, verified_global_indexes, data_fields, mode="force_fetch") + return self.generate_batch_meta(partition_id, verified_global_indexes, data_fields, mode="force_fetch") def kv_retrieve_keys( self, @@ -1809,122 +1793,6 @@ def describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, for key, global_index in zip(keys, cast(list[int], global_indexes), strict=True) } - def _assert_not_restoring(self, partition_id: str | None = None) -> None: - for restore_id, restore in self._restores.items(): - if partition_id is None or restore["partition_id"] == partition_id: - raise RuntimeError( - f"Restore {restore_id} is unresolved; recover_data_load({restore['dump_dir']!r}) first" - ) - - def begin_restore( - self, restore_id: str, dump_dir: str, partition_id: str, rows: dict, units: list[str], schema: dict - ): - """Reserve a destination partition until every possible remote writer is settled.""" - with self._restore_lock: - if restore_id in self._restore_outcomes: - raise RuntimeError("Restore was already committed or cancelled") - if restore_id in self._restores: - return self._restores[restore_id]["metadata"] - self._assert_not_restoring(partition_id) - partition = self._get_partition(partition_id) - if partition is not None: - if partition.global_indexes & self._clearing_indexes: - raise RuntimeError("Partition has an unfinished clear; complete it before restoring") - partition.validate_field_schema(schema) - keys = list(rows) - metadata = self.kv_retrieve_meta(keys, partition_id, create=True) - data_indexes = [ - index for key, index in zip(keys, metadata.global_indexes, strict=True) if rows[key]["fields"] - ] - active_units = group_by_storage_unit(data_indexes, units) - metadata.update_custom_meta([rows[key]["tag"] for key in keys]) - self._restores[restore_id] = { - "partition_id": partition_id, - "dump_dir": dump_dir, - "metadata": metadata, - "units": {unit: "pending" for unit in active_units}, - "updates": {}, - "aborting": False, - } - return metadata - - def restore_unit(self, restore_id: str, unit_id: str, action: str, result: dict | None = None) -> None: - """Grant one execution and record its terminal result; unknown IDs never grant writes.""" - with self._restore_lock: - restore = self._restores.get(restore_id) - if restore_id in self._restore_outcomes and action == "complete": - return - if restore is None or unit_id not in restore["units"]: - raise RuntimeError("Restore is no longer active") - state = restore["units"][unit_id] - if action == "claim": - if state != "pending" or restore["aborting"]: - raise RuntimeError(f"Restore unit cannot start in state {state}") - restore["units"][unit_id] = "running" - elif ( - action == "complete" - and state == "pending" - and result is not None - and result.get("claim_failed") is True - and result.get("success") is False - ): - # The claim may never have arrived; this worker confirms it did not write. - restore["units"][unit_id] = "failed" - elif action == "complete" and state in ("running", "done", "failed") and result is not None: - restore["units"][unit_id] = "done" if result["success"] else "failed" - restore["updates"][unit_id] = result.get("updates", []) - else: - raise RuntimeError(f"Invalid restore transition {state} -> {action}") - - def finish_restore(self, restore_id: str, commit: bool) -> dict: - """Settle a restore after its workers terminate, retaining the outcome for lost replies.""" - with self._restore_lock: - if restore_id in self._restore_outcomes: - return {"finished": True, "committed": self._restore_outcomes[restore_id]} - restore = self._restores.get(restore_id) - if restore is None: - if commit: - return {"finished": False, "reason": "unknown_restore"} - self._restore_outcomes[restore_id] = False - return {"finished": True, "committed": False} - if not commit or "failed" in restore["units"].values(): - restore["aborting"] = True - for unit, state in restore["units"].items(): - if state == "pending": - restore["units"][unit] = "cancelled" - unfinished = {unit: state for unit, state in restore["units"].items() if state in ("pending", "running")} - if unfinished: - return { - "finished": False, - "reason": "unfinished_units", - "units": list(unfinished), - "unit_states": unfinished, - } - committed = commit and not restore["aborting"] - if committed: - partition = self.partitions[restore["partition_id"]] - updates = [update for group in restore["updates"].values() for update in group] - for update in updates: - partition.validate_field_schema(update["field_schema"]) - for update in updates: - if not partition.update_production_status(update["global_indexes"], [], update["field_schema"]): - raise RuntimeError("Controller rejected restored metadata") - metadata = restore["metadata"] - partition.set_custom_meta(dict(zip(metadata.global_indexes, metadata.custom_meta, strict=True))) - # Retain only the terminal outcome, never payloads or per-row metadata. - self._restore_outcomes[restore_id] = committed - del self._restores[restore_id] - return {"finished": True, "committed": committed} - - def list_restores(self, dump_dir: str | None) -> list[str]: - """Find unfinished operations for a dump, including after the initiating client exits.""" - with self._restore_lock: - return [ - restore_id - for restore_id, restore in self._restores.items() - if dump_dir is None or restore["dump_dir"] == dump_dir - ] - def _init_zmq_socket(self): """Initialize ZMQ sockets for communication.""" self.zmq_context = zmq.Context() @@ -2150,10 +2018,6 @@ def _handle_request(self, request_msg: ZMQMessage, monitor: Any) -> ZMQMessage | ZMQRequestType.KV_LIST: self._handle_kv_list_request, ZMQRequestType.DESCRIBE_ROWS_BY_KEY: self._handle_describe_rows_by_key_request, ZMQRequestType.VALIDATE_DUMP_SCHEMA: self._handle_validate_dump_schema_request, - ZMQRequestType.BEGIN_RESTORE: self._handle_begin_restore_request, - ZMQRequestType.RESTORE_UNIT: self._handle_restore_unit_request, - ZMQRequestType.FINISH_RESTORE: self._handle_finish_restore_request, - ZMQRequestType.LIST_RESTORES: self._handle_list_restores_request, ZMQRequestType.SAVE_CONTROLLER_CHECKPOINT: self._handle_save_controller_checkpoint_request, ZMQRequestType.LOAD_CONTROLLER_CHECKPOINT: self._handle_load_controller_checkpoint_request, } @@ -2406,22 +2270,6 @@ def _handle_describe_rows_by_key_request(self, request_msg: ZMQMessage) -> ZMQMe {"success": True, "partition_id": params["partition_id"], "rows": rows, "field_schema": field_schema}, ) - def _handle_begin_restore_request(self, request_msg: ZMQMessage) -> ZMQMessage: - metadata = self.begin_restore(**request_msg.body) - return self._make_response(request_msg, ZMQRequestType.BEGIN_RESTORE_RESPONSE, {"metadata": metadata}) - - def _handle_restore_unit_request(self, request_msg: ZMQMessage) -> ZMQMessage: - self.restore_unit(**request_msg.body) - return self._make_response(request_msg, ZMQRequestType.RESTORE_UNIT_RESPONSE, {"success": True}) - - def _handle_finish_restore_request(self, request_msg: ZMQMessage) -> ZMQMessage: - result = self.finish_restore(**request_msg.body) - return self._make_response(request_msg, ZMQRequestType.FINISH_RESTORE_RESPONSE, result) - - def _handle_list_restores_request(self, request_msg: ZMQMessage) -> ZMQMessage: - ids = self.list_restores(request_msg.body["dump_dir"]) - return self._make_response(request_msg, ZMQRequestType.LIST_RESTORES_RESPONSE, {"restore_ids": ids}) - def _handle_validate_dump_schema_request(self, request_msg: ZMQMessage) -> ZMQMessage: params = request_msg.body partition = self._get_partition(params["partition_id"]) @@ -2453,25 +2301,23 @@ def save_checkpoint(self, path: str) -> None: Raises: Exception: If serialization or file I/O fails. """ - with self._restore_lock: - self._assert_not_restoring() - try: - state = { - "controller_id": self.controller_id, - "partitions": {pid: p.to_snapshot() for pid, p in self.partitions.items()}, - "index_manager": { - "partition_to_indexes": dict(copy.deepcopy(self.index_manager.partition_to_indexes)), - "reusable_indexes": list(self.index_manager.reusable_indexes), - "global_index_counter": self.index_manager.global_index_counter, - "allocated_indexes": set(self.index_manager.allocated_indexes), - }, - "sampler": self.sampler.save_checkpoint(), - } - with open(path, "wb") as f: - pickle.dump(state, f, protocol=pickle.HIGHEST_PROTOCOL) - logger.info(f"[{self.controller_id}]: dumped to {path}") - except Exception as e: - raise RuntimeError(f"[{self.controller_id}]: save checkpoint failed: {e}") from e + try: + state = { + "controller_id": self.controller_id, + "partitions": {pid: p.to_snapshot() for pid, p in self.partitions.items()}, + "index_manager": { + "partition_to_indexes": dict(copy.deepcopy(self.index_manager.partition_to_indexes)), + "reusable_indexes": list(self.index_manager.reusable_indexes), + "global_index_counter": self.index_manager.global_index_counter, + "allocated_indexes": set(self.index_manager.allocated_indexes), + }, + "sampler": self.sampler.save_checkpoint(), + } + with open(path, "wb") as f: + pickle.dump(state, f, protocol=pickle.HIGHEST_PROTOCOL) + logger.info(f"[{self.controller_id}]: dumped to {path}") + except Exception as e: + raise RuntimeError(f"[{self.controller_id}]: save checkpoint failed: {e}") from e def load_checkpoint(self, path: str) -> None: """Restore controller state directly from a file. @@ -2482,27 +2328,24 @@ def load_checkpoint(self, path: str) -> None: Raises: Exception: If deserialization or file I/O fails. """ - with self._restore_lock: - self._assert_not_restoring() - try: - with open(path, "rb") as f: - state = pickle.load(f) + try: + with open(path, "rb") as f: + state = pickle.load(f) - self.controller_id = state["controller_id"] - self.partitions = state["partitions"] - self._clearing_indexes.clear() + self.controller_id = state["controller_id"] + self.partitions = state["partitions"] - im = state["index_manager"] - self.index_manager.partition_to_indexes = defaultdict(set, im["partition_to_indexes"]) - self.index_manager.reusable_indexes = im["reusable_indexes"] - self.index_manager.global_index_counter = im["global_index_counter"] - self.index_manager.allocated_indexes = im["allocated_indexes"] + im = state["index_manager"] + self.index_manager.partition_to_indexes = defaultdict(set, im["partition_to_indexes"]) + self.index_manager.reusable_indexes = im["reusable_indexes"] + self.index_manager.global_index_counter = im["global_index_counter"] + self.index_manager.allocated_indexes = im["allocated_indexes"] - self.sampler.load_checkpoint(state["sampler"]) + self.sampler.load_checkpoint(state["sampler"]) - logger.info(f"[{self.controller_id}]: restored from {path}") - except Exception as e: - raise RuntimeError(f"[{self.controller_id}]: load checkpoint failed: {e}") from e + logger.info(f"[{self.controller_id}]: restored from {path}") + except Exception as e: + raise RuntimeError(f"[{self.controller_id}]: load checkpoint failed: {e}") from e def register_sampler( self, diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index bf3daaff..709c2bf7 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -38,12 +38,11 @@ from contextlib import contextmanager from pathlib import Path from typing import Any -from uuid import uuid4 import torch from tensordict import TensorDict -from transfer_queue.storage.dump_io import RestorePendingError, pack_dump_field, read_dump_row, validate_dump_values +from transfer_queue.storage.dump_io import pack_dump_field, read_dump_row, validate_dump_values from transfer_queue.utils import compact_pickle from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.tensor_utils import pack_field_values @@ -142,11 +141,6 @@ def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dic unique_keys = list(dict.fromkeys(keys)) dump_dir = Path(dump_dir).resolve() client = _maybe_create_tq_client() - if hasattr(client, "check_data_loads"): - client.check_data_loads(str(dump_dir)) - marker = dump_dir.with_name(dump_dir.name + ".restore") - if marker.exists(): - raise RestorePendingError(marker.read_text().strip()) _recover_dump(dump_dir) tmp_dir = dump_dir.parent / (dump_dir.name + ".tmp") old_dir = dump_dir.with_name(dump_dir.name + ".old") @@ -325,8 +319,8 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: SimpleStorage units read their assigned indexed records directly and in parallel. Version-1 dumps and other backends use the compatible KV put path. New keys receive current indexes; unrelated rows and fields remain untouched. Writers and clears for - these keys must be paused during restore. Failure may leave partial payload writes. - On RestorePendingError, call recover_data_load before retrying or clearing. + these keys must be paused during restore. As with ``kv_batch_put``, a failure may + leave new keys registered and payload partially written; retrying is idempotent. Args: dump_dir: Directory previously written by ``dump_data_by_key``. For direct @@ -351,12 +345,6 @@ def _load_data_by_key(dump_dir: Path) -> dict[str, int]: raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") dump_dir = Path(dump_dir).resolve() - client = _maybe_create_tq_client() - if hasattr(client, "check_data_loads"): - client.check_data_loads(str(dump_dir)) - marker = dump_dir.with_name(dump_dir.name + ".restore") - if marker.exists(): - raise RestorePendingError(marker.read_text().strip()) _recover_dump(dump_dir) info_path = dump_dir / _DUMP_INFO_FILE if not info_path.exists(): @@ -425,21 +413,7 @@ def _load_data_by_key(dump_dir: Path) -> dict[str, int]: if dump_info["format_version"] >= 3: client.validate_dump_schema(partition_id, row_index["field_schema"]) if dump_info["format_version"] >= 2 and hasattr(getattr(client, "storage_manager", None), "load_rows_by_index"): - restore_id = uuid4().hex - if rows: - with marker.open("x") as f: - f.write(restore_id) - _fsync_file(f) - _fsync_directory(marker.parent) - try: - client.load_rows_by_key(partition_id, rows, shards, str(dump_dir), restore_id) - except RestorePendingError: - raise - except Exception: - marker.unlink(missing_ok=True) - raise - else: - marker.unlink(missing_ok=True) + client.load_rows_by_key(partition_id, rows, shards) else: _load_via_kv(partition_id, rows, shards, dump_info["format_version"]) @@ -501,31 +475,3 @@ def _load_via_kv(partition_id: str, rows: dict[str, Any], shards: list[dict], ve keys = [key for key, row in rows.items() if not row["fields"]] if keys: kv_batch_put(keys, partition_id, tags=[rows[key]["tag"] for key in keys]) - - -def recover_data_load(dump_dir: str | Path, *, cancel: bool = False) -> bool: - """Settle an interrupted restore before retrying or releasing destination indexes. - - By default, commit loads once every unit reports success. Use ``cancel=True`` - to deny unclaimed work and wait for claimed workers without publishing metadata. - Running or unknown work raises RestorePendingError and keeps the dump protected. - - Returns: - True if all loads committed or none needed recovery; False if any load was - cancelled or failed. Partial payload writes remain after cancellation. - """ - with _dump_lock(dump_dir) as directory: - return _recover_data_load(directory, cancel=cancel) - - -def _recover_data_load(dump_dir: Path, *, cancel: bool = False) -> bool: - from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client - - if _TQ_CONTROLLER is None: - raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") - dump_dir = Path(dump_dir).resolve() - marker = dump_dir.with_name(dump_dir.name + ".restore") - ids = [marker.read_text().strip()] if marker.exists() else [] - committed = _maybe_create_tq_client().recover_data_load(str(dump_dir), ids, cancel=cancel) - marker.unlink(missing_ok=True) - return committed diff --git a/transfer_queue/interface.py b/transfer_queue/interface.py index 33abf997..eca8762e 100644 --- a/transfer_queue/interface.py +++ b/transfer_queue/interface.py @@ -1130,7 +1130,6 @@ def load_checkpoint( meta = json.load(f) client = _maybe_create_tq_client() - client.check_data_loads() controller_path = checkpoint_dir / _CONTROLLER_FILE if not controller_path.exists(): diff --git a/transfer_queue/storage/dump_io.py b/transfer_queue/storage/dump_io.py index 11772fef..8840eaca 100644 --- a/transfer_queue/storage/dump_io.py +++ b/transfer_queue/storage/dump_io.py @@ -74,40 +74,3 @@ def pack_dump_field(values: list, schema: dict): if schema["is_nested"]: return torch.nested.as_nested_tensor(values, layout=torch.jagged) return torch.stack(values) - - -class RestorePendingError(RuntimeError): - """Recovery is unresolved; expose unit states and reporting failures without releasing protection.""" - - def __init__( - self, - restore_id: str, - *, - reason: str = "unknown_outcome", - unit_states: dict[str, str] | None = None, - report_errors: dict[str, str] | None = None, - ): - self.restore_id = restore_id - self.reason = reason - self.unit_states = unit_states or {} - self.report_errors = report_errors or {} - message = f"Restore {restore_id} is unresolved ({reason})" - if reason == "unknown_restore": - message += ( - "; the controller has no record of this ID. If the initiating client has stopped, " - "or the whole TQ system was restarted after stopping all old actors, " - "use recover_data_load(dump_dir, cancel=True) to discard the stale operation" - ) - elif self.unit_states: - message += f"; unit states: {self.unit_states}. Pending units have not claimed permission; " - message += "running units have not reported completion. Retry recover_data_load to await completion" - if "pending" in self.unit_states.values(): - message += ( - "; if the initiating client exited before dispatching all requests, " - "use recover_data_load(dump_dir, cancel=True) instead of retrying indefinitely" - ) - else: - message += "; the remote outcome is unknown. Run recover_data_load before retrying or clearing" - if self.report_errors: - message += f". Unit report failures: {self.report_errors}" - super().__init__(message) diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 7ebaea3a..5a46c1ce 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -22,7 +22,7 @@ from functools import partial from operator import itemgetter from pathlib import Path -from typing import Any, Callable +from typing import Any, Callable, NamedTuple import torch import zmq @@ -37,7 +37,6 @@ ) from transfer_queue.utils.common import log_heavy_operation from transfer_queue.utils.logging_utils import get_logger -from transfer_queue.utils.storage_routing import RoutingGroup, group_by_storage_unit from transfer_queue.utils.tensor_utils import pack_field_values from transfer_queue.utils.zmq_utils import ( TQ_SOCKET_POOL_SIZE, @@ -107,6 +106,13 @@ def _describe_unit_state(body: dict[str, Any]) -> str: ) +class RoutingGroup(NamedTuple): + """Routing result for a single storage unit.""" + + global_indexes: list[int] # global indexes routed to this SU + batch_positions: list[int] # corresponding positions in the original batch + + @StorageManagerFactory.register("SimpleStorage") class AsyncSimpleStorageManager(StorageManager): """Asynchronous storage manager that handles multiple storage units. @@ -212,7 +218,15 @@ def _group_by_hash(self, global_indexes: list[int]) -> dict[str, RoutingGroup]: NOTE: Dynamic SU scaling requires a data migration mechanism (not yet supported). """ - return group_by_storage_unit(global_indexes, list(self.storage_unit_infos)) + storage_unit_keys = list(self.storage_unit_infos.keys()) + num_units = len(storage_unit_keys) + gi_lists: dict[str, list[int]] = defaultdict(list) + pos_lists: dict[str, list[int]] = defaultdict(list) + for pos, global_idx in enumerate(global_indexes): + key = storage_unit_keys[global_idx % num_units] + gi_lists[key].append(global_idx) + pos_lists[key].append(pos) + return {key: RoutingGroup(gi_lists[key], pos_lists[key]) for key in gi_lists} def _describe_storage_unit(self, storage_unit_id: str) -> str: """Return ``ip:port`` for a storage unit, for use in diagnostics. @@ -865,10 +879,15 @@ async def dump_rows_by_index( ) return shards - async def load_rows_by_index( - self, shards: list[dict[str, Any]], restore: dict | None = None - ) -> list[dict[str, Any]]: - """Have current owner units read assigned byte ranges concurrently.""" + async def load_rows_by_index(self, partition_id: str, shards: list[dict[str, Any]]) -> int: + """Have current owner units read assigned byte ranges concurrently, then publish metadata. + + Metadata is published only after every unit succeeded, so a failed load leaves its + payload writes invisible; retrying rewrites the same target indexes. + + Returns: + Payload bytes read by the units. + """ assignments = defaultdict(list) for shard in shards: rows = shard["records"] @@ -880,43 +899,36 @@ async def load_rows_by_index( } ) results = await asyncio.gather( - *( - self._load_selected_rows( - shards, target_storage_unit=unit_id, **({"restore": restore} if restore else {}) - ) - for unit_id, shards in assignments.items() - ), + *(self._load_selected_rows(shards, target_storage_unit=unit_id) for unit_id, shards in assignments.items()), return_exceptions=True, ) - # Local RPC completion is not remote completion; the controller retains reservations on timeout. - updates = [] - bytes_read = 0 for result in results: if isinstance(result, BaseException): raise result - bytes_read += result["bytes_read"] - updates.extend(result["updates"]) + for result in results: + for update in result["updates"]: + await self.notify_data_update(partition_id, update["global_indexes"], update["field_schema"]) + bytes_read = sum(result["bytes_read"] for result in results) logger.info( "[%s]: loaded %s bytes across %s units", self.storage_manager_id, bytes_read, len(assignments), ) - return updates + return bytes_read @with_storage_unit_socket async def _load_selected_rows( self, shards: list[dict[str, Any]], target_storage_unit: str, - restore: dict | None = None, socket: zmq.Socket = None, ) -> dict[str, Any]: request = ZMQMessage.create( request_type=ZMQRequestType.LOAD_ROWS, sender_id=self.storage_manager_id, receiver_id=target_storage_unit, - body={"shards": shards, "restore": restore}, + body={"shards": shards}, ) await socket.send_multipart(request.serialize(), copy=False) response = ZMQMessage.deserialize(await socket.recv_multipart(copy=False)) @@ -926,37 +938,6 @@ async def _load_selected_rows( ) return response.body - async def report_restore(self, restore: dict) -> dict[str, str]: - """Wait for all unit reports and return failures without discarding successful reports.""" - units = list(self.storage_unit_infos) - results = await asyncio.gather( - *(self._report_restore_unit(restore, target_storage_unit=unit) for unit in units), - return_exceptions=True, - ) - errors = {} - for unit, result in zip(units, results, strict=True): - if isinstance(result, Exception): - errors[unit] = f"{type(result).__name__}: {result}" - logger.warning( - "Restore %s: unit %s could not report completion: %s", restore["restore_id"], unit, errors[unit] - ) - elif isinstance(result, BaseException): - raise result - return errors - - @with_storage_unit_socket - async def _report_restore_unit(self, restore: dict, target_storage_unit: str, socket: zmq.Socket = None) -> None: - request = ZMQMessage.create( - request_type=ZMQRequestType.REPORT_RESTORE, sender_id=self.storage_manager_id, body=restore - ) - await socket.send_multipart(request.serialize()) - response = ZMQMessage.deserialize(await socket.recv_multipart()) - if response.request_type != ZMQRequestType.REPORT_RESTORE_RESPONSE or not response.body.get("success"): - message = response.body.get("message", f"Unexpected response: {response.request_type}") - raise RuntimeError( - f"Restore {restore['restore_id']}: storage unit {target_storage_unit} report failed: {message}" - ) - async def save_checkpoint(self, checkpoint_dir: str) -> None: """Dump all storage units to the storage_units/ subdirectory of checkpoint_dir. diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index fbfc7a67..8bb42843 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -798,7 +798,6 @@ def __init__(self, config: dict[str, Any]): self.proxy_thread: Thread | None = None self.worker_thread: Thread | None = None - self._restore_results: dict[str, dict] = {} self._metrics: TQMetricsExporter | None = None self._init_zmq_socket() @@ -1016,8 +1015,6 @@ def _process_one_worker_request(self, worker_socket: zmq.Socket, monitor: Any) - elif operation == ZMQRequestType.LOAD_ROWS: # type: ignore[arg-type] with monitor.measure(op_type="LOAD_ROWS"): response_msg = self._handle_load_rows(request_msg) - elif operation == ZMQRequestType.REPORT_RESTORE: - response_msg = self._handle_report_restore(request_msg) elif operation == ZMQRequestType.LOAD_STORAGE_CHECKPOINT: # type: ignore[arg-type] response_msg = self._handle_load_checkpoint(request_msg) else: @@ -1364,88 +1361,7 @@ def _handle_dump_rows(self, request: ZMQMessage) -> ZMQMessage: body={"success": False, "message": str(e)}, ) - def _restore_controller_request(self, context: dict, action: str, result: dict | None = None) -> None: - socket = create_zmq_socket(self.zmq_context, zmq.DEALER, context["controller_ip"]) - try: - socket.setsockopt(zmq.RCVTIMEO, 10000) - socket.setsockopt(zmq.SNDTIMEO, 10000) - socket.connect(context["controller_address"]) - request = ZMQMessage.create( - request_type=ZMQRequestType.RESTORE_UNIT, - sender_id=self.storage_unit_id, - body={ - "restore_id": context["restore_id"], - "unit_id": self.storage_unit_id, - "action": action, - "result": result, - }, - ) - socket.send_multipart(request.serialize()) - response = ZMQMessage.deserialize(socket.recv_multipart()) - if response.request_type != ZMQRequestType.RESTORE_UNIT_RESPONSE: - raise RuntimeError(response.body.get("message", "Restore permission rejected")) - finally: - socket.close(linger=0) - - def _handle_report_restore(self, request: ZMQMessage) -> ZMQMessage: - context = request.body - result = self._restore_results.get(context["restore_id"]) - try: - if result is not None: - self._restore_controller_request(context, "complete", result) - self._restore_results.pop(context["restore_id"], None) - return ZMQMessage.create( - request_type=ZMQRequestType.REPORT_RESTORE_RESPONSE, - sender_id=self.storage_unit_id, - body={"success": True}, - ) - except Exception as e: - return ZMQMessage.create( - request_type=ZMQRequestType.REPORT_RESTORE_RESPONSE, - sender_id=self.storage_unit_id, - body={"success": False, "message": str(e)}, - ) - def _handle_load_rows(self, request: ZMQMessage) -> ZMQMessage: - """Claim permission before touching storage; report completion independently of the caller.""" - context = request.body.get("restore") - if context is None: - return ZMQMessage.create( - request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, - sender_id=self.storage_unit_id, - body={"success": False, "message": "Missing restore reservation"}, - ) - try: - self._restore_controller_request(context, "claim") - except zmq.error.Again as e: - # A lost claim ACK may leave the controller in running state even though - # this worker will not write; retain a terminal result for recovery. - self._restore_results[context["restore_id"]] = { - "success": False, - "claim_failed": True, - "message": f"Restore claim was not acknowledged; payload was not written: {e}", - } - return ZMQMessage.create( - request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, - sender_id=self.storage_unit_id, - body={"success": False, "message": "Restore claim outcome unknown"}, - ) - except Exception as e: - return ZMQMessage.create( - request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, - sender_id=self.storage_unit_id, - body={"success": False, "message": str(e)}, - ) - response = self._load_rows(request) - self._restore_results[context["restore_id"]] = response.body - try: - self._restore_controller_request(context, "complete", response.body) - self._restore_results.pop(context["restore_id"], None) - except Exception as e: - logger.warning("[%s]: restore result retained for recovery: %s", self.storage_unit_id, e) - return response - - def _load_rows(self, request: ZMQMessage) -> ZMQMessage: """Seek to assigned records and merge them into local storage at current indexes.""" updates = [] bytes_read = 0 diff --git a/transfer_queue/utils/storage_routing.py b/transfer_queue/utils/storage_routing.py deleted file mode 100644 index f30e548a..00000000 --- a/transfer_queue/utils/storage_routing.py +++ /dev/null @@ -1,35 +0,0 @@ -# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. -# Copyright 2025 The TransferQueue Team -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -from collections import defaultdict -from typing import NamedTuple - - -class RoutingGroup(NamedTuple): - """Indexes assigned to a storage unit and their positions in the input batch.""" - - global_indexes: list[int] - batch_positions: list[int] - - -def group_by_storage_unit(global_indexes: list[int], storage_unit_ids: list[str]) -> dict[str, RoutingGroup]: - """Route indexes using the ordered unit list shared by storage and restore reservations.""" - gi_lists: dict[str, list[int]] = defaultdict(list) - pos_lists: dict[str, list[int]] = defaultdict(list) - for pos, global_idx in enumerate(global_indexes): - key = storage_unit_ids[global_idx % len(storage_unit_ids)] - gi_lists[key].append(global_idx) - pos_lists[key].append(pos) - return {key: RoutingGroup(gi_lists[key], pos_lists[key]) for key in gi_lists} diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index 1ef84c28..b32273e4 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -145,17 +145,6 @@ class ZMQRequestType(ExplicitEnum): VALIDATE_DUMP_SCHEMA = "VALIDATE_DUMP_SCHEMA" VALIDATE_DUMP_SCHEMA_RESPONSE = "VALIDATE_DUMP_SCHEMA_RESPONSE" - BEGIN_RESTORE = "BEGIN_RESTORE" - BEGIN_RESTORE_RESPONSE = "BEGIN_RESTORE_RESPONSE" - RESTORE_UNIT = "RESTORE_UNIT" - RESTORE_UNIT_RESPONSE = "RESTORE_UNIT_RESPONSE" - FINISH_RESTORE = "FINISH_RESTORE" - FINISH_RESTORE_RESPONSE = "FINISH_RESTORE_RESPONSE" - LIST_RESTORES = "LIST_RESTORES" - LIST_RESTORES_RESPONSE = "LIST_RESTORES_RESPONSE" - REPORT_RESTORE = "REPORT_RESTORE" - REPORT_RESTORE_RESPONSE = "REPORT_RESTORE_RESPONSE" - class ZMQServerInfo: """ From 247824bb0ddfc8ae4116c8269971953d8406806f Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 8 Oct 2026 20:02:27 +0800 Subject: [PATCH 24/30] fix: give selective dump and load their own long-timeout socket pool DUMP_ROWS and LOAD_ROWS send one request per storage unit covering all of its selected rows, so their duration grows with the selection size and the shared filesystem rather than with a batch. They used the put/get pool and its 200 s timeout, so a healthy unit reading a large selection from a slow filesystem failed the load. Serve them from a separate pool whose timeout is TQ_SIMPLE_STORAGE_DUMP_TIMEOUT (3600 s by default), leaving the put/get timeout unchanged. Signed-off-by: OutstanderWang --- docs/data_dump.md | 3 +++ tests/e2e/test_data_dump_e2e.py | 11 +++++++++ .../managers/simple_storage_manager.py | 23 +++++++++++++++++-- 3 files changed, 35 insertions(+), 2 deletions(-) diff --git a/docs/data_dump.md b/docs/data_dump.md index 4625a183..8e4095a4 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -121,6 +121,9 @@ cleared and reused while it runs, so keep writers and clears for these keys paus Each storage unit serves requests on one worker thread: other partitions using that unit can wait behind a load. The 128-row batches bound memory, not request latency. +Dump and load requests use their own connection pool, whose timeout is +`TQ_SIMPLE_STORAGE_DUMP_TIMEOUT` (3600 seconds by default) rather than the put/get +timeout, because each unit receives a single request covering all of its rows. ## Tests diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index 7b6f14dd..d05dddfd 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -662,3 +662,14 @@ async def lose_reply(*args, **kwargs): assert snapshot.keys_mapping["key"] == index _assert_rows_equal(tq.kv_batch_get(["key"], "retry", ["input_ids"])["input_ids"], [_row_input_ids(0)]) assert tq.kv_list("retry")["retry"] == {"key": {"idx": 0}} + + +def test_dump_and_load_do_not_use_the_put_get_timeout_pool(tq_system, dump_dir, monkeypatch): + _put_rows("own_pool", ["key"]) + manager = tq.get_client().storage_manager + # Any lease from the ordinary put/get pool now fails, so only the dump pool can serve these. + monkeypatch.setattr(manager, "storage_rpc_pool", None) + tq.dump_data_by_key(dump_dir, ["key"], "own_pool") + tq.load_data_by_key(dump_dir) + monkeypatch.undo() + _assert_rows_equal(tq.kv_batch_get(["key"], "own_pool", ["input_ids"])["input_ids"], [_row_input_ids(0)]) diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 5a46c1ce..94586226 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -59,6 +59,10 @@ # Timeout for the post-failure probe, which only has to answer whether the unit still serves. TQ_SIMPLE_STORAGE_PROBE_TIMEOUT = int(os.environ.get("TQ_SIMPLE_STORAGE_PROBE_TIMEOUT", 10)) +# A selective dump or load sends one request per unit covering all of its rows, so its duration +# grows with the selection and shared-filesystem speed rather than with a batch. +TQ_SIMPLE_STORAGE_DUMP_TIMEOUT = int(os.environ.get("TQ_SIMPLE_STORAGE_DUMP_TIMEOUT", 3600)) + class StorageUnitTimeout(RuntimeError): """A storage unit did not answer within the send/recv timeout. @@ -105,6 +109,12 @@ def _describe_unit_state(body: dict[str, Any]) -> str: resolve_target=lambda args, kwargs: kwargs.get("target_storage_unit"), ) +with_storage_unit_dump_socket = with_zmq_socket( + get_peer=lambda self, target: self.storage_unit_infos[target], + get_pool=lambda self: self.storage_dump_pool, + resolve_target=lambda args, kwargs: kwargs.get("target_storage_unit"), +) + class RoutingGroup(NamedTuple): """Routing result for a single storage unit.""" @@ -142,6 +152,12 @@ def __init__( timeout=TQ_SIMPLE_STORAGE_PROBE_TIMEOUT, maxsize=1, ) + self.storage_dump_pool = ZMQSocketPool( + self.zmq_context, + f"{self.storage_manager_id}_dump", + "put_get_socket", + timeout=TQ_SIMPLE_STORAGE_DUMP_TIMEOUT, + ) self.config = config server_infos: ZMQServerInfo | dict[str, ZMQServerInfo] | None = config.get("zmq_info", None) @@ -766,7 +782,7 @@ async def _load_single_storage_unit( f"[{self.storage_manager_id}]: Error restoring for storage unit {target_storage_unit}: {str(e)}" ) from e - @with_storage_unit_socket + @with_storage_unit_dump_socket async def _dump_single_shard( self, path: str, @@ -917,7 +933,7 @@ async def load_rows_by_index(self, partition_id: str, shards: list[dict[str, Any ) return bytes_read - @with_storage_unit_socket + @with_storage_unit_dump_socket async def _load_selected_rows( self, shards: list[dict[str, Any]], @@ -1003,4 +1019,7 @@ def close(self) -> None: probe_pool = getattr(self, "storage_probe_pool", None) if probe_pool is not None: probe_pool.close() + dump_pool = getattr(self, "storage_dump_pool", None) + if dump_pool is not None: + dump_pool.close() super().close() From 79b1413de18c148fd9041a4ff93d166d0f4575f7 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 8 Oct 2026 20:10:00 +0800 Subject: [PATCH 25/30] fix: publish dumps once and read them without locks Replacing a dump in place needed a sibling .old backup, recovery on every access, and a cross-node flock on a sibling .lock file. Readers therefore needed write access next to the dump, and a read could rename .old back into place. Stage each dump in a uniquely named sibling directory, write dump_info.json last, and rename it into place once synced. An existing target is refused, so a published dump never changes: loads need no lock, readers need no write access, and of two dumps racing to one path only the first rename succeeds. Signed-off-by: OutstanderWang --- docs/data_dump.md | 20 ++--- tests/e2e/test_data_dump_e2e.py | 36 +++++---- tests/test_data_dump.py | 75 ++++++------------ tests/test_dump_lock.py | 132 -------------------------------- transfer_queue/data_dump.py | 75 ++++-------------- 5 files changed, 68 insertions(+), 270 deletions(-) delete mode 100644 tests/test_dump_lock.py diff --git a/docs/data_dump.md b/docs/data_dump.md index 8e4095a4..03f5dbdb 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -15,9 +15,9 @@ tq.load_data_by_key("/shared/dumps/selected") ``` Pause writes and clears for these keys during both operations. A dump is not an -atomic snapshot of concurrent writers. Calls on the same dump path are serialized -by an exclusive lock in a stable sibling `.lock` file. Different dump paths remain -independent, and storage units within a load still read in parallel. +atomic snapshot of concurrent writers. Each dump goes to a new directory: an existing +target is refused, so a published dump never changes and can be loaded from a +read-only location. ## Distributed I/O @@ -102,15 +102,11 @@ SimpleStorage. ## Failure behavior -Publication writes and syncs `.tmp`, moves the old directory to `.old`, publishes -the new directory, and syncs its parent before deleting the backup. If publication -is interrupted while the main directory is absent, the next dump, load or row-index -read recovers `.old`. Readers perform recovery only while holding the same lock as -publishers, so a healthy rename window is never mistaken for a crashed writer. -The load keeps the lock through all remote reads. Do not delete the sibling -lock file: unlinking it can create two independent locks for the same dump. -The shared filesystem must provide cross-node advisory locking (not local-only -locks). A backup-cleanup error does not invalidate a published dump. +A dump is staged in a uniquely named sibling `.tmp-` directory, with +`dump_info.json` written last, and renamed into place once everything is synced. +A failed dump removes its staging directory; a crash can leave one behind, which is +never read and can be deleted. If two dumps race to one path, only the first rename +succeeds. Readers need no lock and no write access. Restore has the failure semantics of `kv_batch_put`: it is not transactional, and payload writes before a failure remain. Metadata is published only after every unit diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index d05dddfd..b4bb674d 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -316,19 +316,28 @@ def test_live_partition_survives_a_dump(self, tq_system, dump_dir, controller): retrieved = tq.kv_batch_get(keys=keys, partition_id=partition_id, select_fields=["input_ids"]) _assert_rows_equal(retrieved["input_ids"], [_row_input_ids(row) for row in range(len(keys))]) - def test_dump_replaces_preexisting_directory(self, tq_system, dump_dir): - # Define test data - partition_id = "d_replace" - _put_rows(partition_id, ["s0", "s1"]) - dump_dir.mkdir(parents=True) - (dump_dir / "stale.pkl").write_bytes(b"stale") - - # Dump + def test_dump_refuses_an_existing_directory(self, tq_system, dump_dir): + partition_id = "d_existing" + _put_rows(partition_id, ["s0"]) tq.dump_data_by_key(dump_dir, ["s0"], partition_id) - - # Check saved state - assert not (dump_dir / "stale.pkl").exists() - assert (dump_dir / "dump_info.json").exists() + with pytest.raises(FileExistsError): + tq.dump_data_by_key(dump_dir, ["s0"], partition_id) + assert sorted(path.name for path in dump_dir.parent.iterdir()) == [dump_dir.name] + assert tq.read_row_index(dump_dir)["partition_id"] == partition_id + + def test_published_dump_loads_from_a_read_only_directory(self, tq_system, dump_dir): + partition_id = "d_read_only" + _put_rows(partition_id, ["s0"]) + tq.dump_data_by_key(dump_dir, ["s0"], partition_id) + tq.kv_clear(["s0"], partition_id) + dump_dir.parent.chmod(0o555) + try: + assert sorted(tq.read_row_index(dump_dir)["rows"]) == ["s0"] + tq.load_data_by_key(dump_dir) + assert sorted(path.name for path in dump_dir.parent.iterdir()) == [dump_dir.name] + finally: + dump_dir.parent.chmod(0o755) + _assert_rows_equal(tq.kv_batch_get(["s0"], partition_id, ["input_ids"])["input_ids"], [_row_input_ids(0)]) # --------------------------------------------------------------------------- @@ -375,8 +384,7 @@ def test_unknown_key_raises_and_leaves_no_directory(self, tq_system, dump_dir): tq.dump_data_by_key(dump_dir, ["e0", "nope"], partition_id) # Check saved state: no partial directory left behind - assert not dump_dir.exists() - assert not dump_dir.with_name(dump_dir.name + ".tmp").exists() + assert not list(dump_dir.parent.iterdir()) def test_unknown_partition_raises(self, tq_system, dump_dir): _put_rows("d_err_part", ["e0"]) diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index 6f58e29f..1819adb1 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -93,63 +93,32 @@ def empty_dump_client(monkeypatch): ) -def test_failed_publication_preserves_previous_dump(empty_dump_client, monkeypatch, tmp_path): - dump = tmp_path / "dump" - data_dump.dump_data_by_key(dump, [], "old") - rename = type(dump).rename - - def fail_publish(path, target): - if path == tmp_path / "dump.tmp": - raise OSError("publication failed") - return rename(path, target) - - monkeypatch.setattr(type(dump), "rename", fail_publish) - with pytest.raises(OSError, match="publication failed"): - data_dump.dump_data_by_key(dump, [], "new") - assert data_dump.read_row_index(dump)["partition_id"] == "old" - assert not (tmp_path / "dump.tmp").exists() - +def test_failed_dump_leaves_no_staging_directory(empty_dump_client, monkeypatch, tmp_path): + def fail(path): + raise OSError("disk full") -@pytest.mark.parametrize("next_operation", ["read", "load", "dump"]) -def test_interrupted_publication_recovers_on_next_access(empty_dump_client, monkeypatch, tmp_path, next_operation): - dump = tmp_path / "dump" - data_dump.dump_data_by_key(dump, [], "old") - rename = type(dump).rename - - def interrupt_publish(path, target): - if path == tmp_path / "dump.tmp": - raise KeyboardInterrupt - return rename(path, target) - - with monkeypatch.context() as patch: - patch.setattr(type(dump), "rename", interrupt_publish) - with pytest.raises(KeyboardInterrupt): - data_dump.dump_data_by_key(dump, [], "new") - assert not dump.exists() - assert (tmp_path / "dump.old").exists() - if next_operation == "read": - assert data_dump.read_row_index(dump)["partition_id"] == "old" - elif next_operation == "load": - assert data_dump.load_data_by_key(dump)["keys"] == 0 - assert data_dump.read_row_index(dump)["partition_id"] == "old" - else: - data_dump.dump_data_by_key(dump, [], "replacement") - assert data_dump.read_row_index(dump)["partition_id"] == "replacement" + monkeypatch.setattr(data_dump, "_fsync_directory", fail) + with pytest.raises(OSError, match="disk full"): + data_dump.dump_data_by_key(tmp_path / "dump", [], "p") + assert not list(tmp_path.iterdir()) -def test_backup_cleanup_failure_does_not_fail_published_dump(empty_dump_client, monkeypatch, tmp_path): +def test_racing_dumps_to_one_path_publish_only_the_first(empty_dump_client, monkeypatch, tmp_path): dump = tmp_path / "dump" - data_dump.dump_data_by_key(dump, [], "old") - rmtree = data_dump.shutil.rmtree - - def fail_cleanup(path): - if path == tmp_path / "dump.old": - raise OSError("cleanup failed") - return rmtree(path) - - monkeypatch.setattr(data_dump.shutil, "rmtree", fail_cleanup) - data_dump.dump_data_by_key(dump, [], "new") - assert data_dump.read_row_index(dump)["partition_id"] == "new" + fsync = data_dump._fsync_directory + raced = [] + + def publish_another_dump_first(path): + if path.name.startswith("dump.tmp-") and not raced: + raced.append(path) + data_dump.dump_data_by_key(dump, [], "first") + fsync(path) + + monkeypatch.setattr(data_dump, "_fsync_directory", publish_another_dump_first) + with pytest.raises(OSError): + data_dump.dump_data_by_key(dump, [], "second") + assert data_dump.read_row_index(dump)["partition_id"] == "first" + assert [path.name for path in tmp_path.iterdir()] == ["dump"] def test_unit_reads_only_assigned_ranges_and_merges(unit, tmp_path, monkeypatch): diff --git a/tests/test_dump_lock.py b/tests/test_dump_lock.py deleted file mode 100644 index d79d5f9a..00000000 --- a/tests/test_dump_lock.py +++ /dev/null @@ -1,132 +0,0 @@ -# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. -# Copyright 2025 The TransferQueue Team -# -# Licensed under the Apache License, Version 2.0 (the "License"); -# you may not use this file except in compliance with the License. -# You may obtain a copy of the License at -# -# http://www.apache.org/licenses/LICENSE-2.0 -# -# Unless required by applicable law or agreed to in writing, software -# distributed under the License is distributed on an "AS IS" BASIS, -# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. -# See the License for the specific language governing permissions and -# limitations under the License. - -"""Publication, reads, and recovery coordinate across processes.""" - -import multiprocessing -from pathlib import Path -from types import SimpleNamespace - -import pytest - -from transfer_queue import data_dump, interface - - -def _configure(): - interface._TQ_CONTROLLER = object() - interface._maybe_create_tq_client = lambda: SimpleNamespace(validate_dump_schema=lambda *_: None) - - -def _publisher(path, paused, release): - _configure() - rename = Path.rename - - def pause_at_publish(source, target): - if source == path.with_name(path.name + ".tmp"): - paused.set() - if not release.wait(20): - raise RuntimeError("Test publisher was not released") - return rename(source, target) - - Path.rename = pause_at_publish - data_dump.dump_data_by_key(path, [], "new") - - -def _reader(path, started, results): - started.set() - results.put(data_dump.read_row_index(path)["partition_id"]) - - -def _loader(path, paused, release): - _configure() - - def pause_load(*args): - paused.set() - if not release.wait(20): - raise RuntimeError("Test loader was not released") - - data_dump._load_via_kv = pause_load - data_dump.load_data_by_key(path) - - -@pytest.mark.parametrize("kill_writer", [False, True]) -def test_reader_waits_for_publisher_and_recovers_after_exit(tmp_path, monkeypatch, kill_writer): - monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) - monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: object()) - path = tmp_path / "dump" - data_dump.dump_data_by_key(path, [], "old") - context = multiprocessing.get_context("spawn") - paused, release, reader_started = context.Event(), context.Event(), context.Event() - results = context.Queue() - writer = context.Process(target=_publisher, args=(path, paused, release)) - reader = context.Process(target=_reader, args=(path, reader_started, results)) - writer.start() - try: - assert paused.wait(20) - reader.start() - assert reader_started.wait(20) - reader.join(0.2) - assert reader.is_alive(), "Reader rolled back a live publisher" - if kill_writer: - writer.terminate() - else: - release.set() - writer.join(20) - assert not writer.is_alive() - assert results.get(timeout=20) == ("old" if kill_writer else "new") - reader.join(20) - assert reader.exitcode == 0 - if not kill_writer: - assert writer.exitcode == 0 - assert path.with_name("dump.lock").exists() - finally: - # A terminated process may have held the Event's semaphore; do not reuse it. - for process in (reader, writer): - if process.pid and process.is_alive(): - process.terminate() - process.join(10) - results.close() - - -def test_publish_waits_until_load_finishes(tmp_path, monkeypatch): - monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) - monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: object()) - path = tmp_path / "dump" - data_dump.dump_data_by_key(path, [], "old") - context = multiprocessing.get_context("spawn") - load_paused, load_release, publish_paused, publish_release = [context.Event() for _ in range(4)] - loader = context.Process(target=_loader, args=(path, load_paused, load_release)) - writer = context.Process(target=_publisher, args=(path, publish_paused, publish_release)) - loader.start() - try: - assert load_paused.wait(20) - writer.start() - assert not publish_paused.wait(0.5) - assert path.exists() - load_release.set() - loader.join(20) - assert loader.exitcode == 0 - assert publish_paused.wait(20) - publish_release.set() - writer.join(20) - assert writer.exitcode == 0 - assert data_dump.read_row_index(path)["partition_id"] == "new" - finally: - load_release.set() - publish_release.set() - for process in (writer, loader): - if process.pid and process.is_alive(): - process.terminate() - process.join(10) diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index 709c2bf7..e6027cec 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -29,15 +29,14 @@ shard__.pkl # independent {global_index, fields} records """ -import fcntl import json import os import pickle import shutil from collections import defaultdict -from contextlib import contextmanager from pathlib import Path from typing import Any +from uuid import uuid4 import torch from tensordict import TensorDict @@ -77,36 +76,15 @@ def _fsync_directory(path: Path) -> None: os.close(fd) -@contextmanager -def _dump_lock(dump_dir: str | Path): - """Serialize publication and recovery using a stable sibling inode shared by all callers.""" - directory = Path(dump_dir).resolve() - directory.parent.mkdir(parents=True, exist_ok=True) - lock_path = directory.with_name(directory.name + ".lock") - with lock_path.open("a+b") as lock: - fcntl.flock(lock.fileno(), fcntl.LOCK_EX) - try: - yield directory - finally: - fcntl.flock(lock.fileno(), fcntl.LOCK_UN) - - -def _recover_dump(dump_dir: Path) -> None: - old_dir = dump_dir.with_name(dump_dir.name + ".old") - if not dump_dir.exists() and old_dir.exists(): - old_dir.rename(dump_dir) - _fsync_directory(dump_dir.parent) - - def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) -> dict[str, int]: """Dump the rows addressed by ``keys`` into ``dump_dir``. Each storage unit pickles the rows it owns in its own process, so the payload never passes through the caller. The caller writes only a small row index. - The directory is replaced wholesale. The previous dump is retained as ``.old`` - until publication is durable, and recovered on the next access after interruption. - Access to the same directory is serialized with a sibling file lock. + The dump is staged in a uniquely named sibling directory and renamed into place + once durable. An existing ``dump_dir`` is refused, so a published dump is never + replaced: loads need no lock and readers no write access. .. note:: **Multi-node limitation**: dump_dir must reside on a shared network filesystem @@ -125,11 +103,11 @@ def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) - ``{"keys", "rows_with_data", "shards", "bytes"}``. Raises: + FileExistsError: ``dump_dir`` already exists. RuntimeError: TransferQueue is not initialized, the partition or a key does not exist, or a storage unit holds no data for a row it was asked to dump. """ - with _dump_lock(dump_dir) as directory: - return _dump_data_by_key(directory, keys, partition_id) + return _dump_data_by_key(Path(dump_dir).resolve(), keys, partition_id) def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dict[str, int]: @@ -138,14 +116,12 @@ def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dic if _TQ_CONTROLLER is None: raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") + if dump_dir.exists(): + raise FileExistsError(f"{dump_dir} already exists; write each dump to a new directory") unique_keys = list(dict.fromkeys(keys)) - dump_dir = Path(dump_dir).resolve() client = _maybe_create_tq_client() - _recover_dump(dump_dir) - tmp_dir = dump_dir.parent / (dump_dir.name + ".tmp") - old_dir = dump_dir.with_name(dump_dir.name + ".old") - if tmp_dir.exists(): - shutil.rmtree(tmp_dir) + # A unique staging name keeps concurrent dumps to the same path from sharing files. + tmp_dir = dump_dir.with_name(f"{dump_dir.name}.tmp-{uuid4().hex}") tmp_dir.mkdir(parents=True) try: @@ -220,27 +196,14 @@ def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dic # published. Syncing the parent makes the rename itself survive a crash. _fsync_directory(tmp_dir) - if dump_dir.exists(): - if old_dir.exists(): - shutil.rmtree(old_dir) - dump_dir.rename(old_dir) - _fsync_directory(dump_dir.parent) + # rename() refuses a non-empty target, so of two dumps racing to one path only + # the first is published. tmp_dir.rename(dump_dir) _fsync_directory(dump_dir.parent) - except Exception: - _recover_dump(dump_dir) - if tmp_dir.exists(): - shutil.rmtree(tmp_dir) + except BaseException: + shutil.rmtree(tmp_dir, ignore_errors=True) raise - # Publication already succeeded; cleanup must not invalidate the new dump. - if old_dir.exists(): - try: - shutil.rmtree(old_dir) - _fsync_directory(dump_dir.parent) - except OSError: - logger.warning("Could not remove previous dump at %s", old_dir, exc_info=True) - total_bytes = sum(path.stat().st_size for path in dump_dir.rglob("*") if path.is_file()) logger.info(f"Dumped {len(unique_keys)} keys of partition {partition_id} to {dump_dir}") return { @@ -300,13 +263,10 @@ def read_row_index(dump_dir: str | Path) -> dict[str, Any]: Raises: FileNotFoundError: The row index is missing. """ - with _dump_lock(dump_dir) as directory: - return _read_row_index(directory) + return _read_row_index(Path(dump_dir)) def _read_row_index(dump_dir: Path) -> dict[str, Any]: - dump_dir = Path(dump_dir) - _recover_dump(dump_dir) row_index_path = dump_dir / _ROW_INDEX_FILE if not row_index_path.exists(): raise FileNotFoundError(f"{_ROW_INDEX_FILE} not found in {dump_dir}") @@ -334,8 +294,7 @@ def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: FileNotFoundError: The dump is incomplete. ValueError: The manifest or a row disagrees with the row index. """ - with _dump_lock(dump_dir) as directory: - return _load_data_by_key(directory) + return _load_data_by_key(Path(dump_dir).resolve()) def _load_data_by_key(dump_dir: Path) -> dict[str, int]: @@ -344,8 +303,6 @@ def _load_data_by_key(dump_dir: Path) -> dict[str, int]: if _TQ_CONTROLLER is None: raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") - dump_dir = Path(dump_dir).resolve() - _recover_dump(dump_dir) info_path = dump_dir / _DUMP_INFO_FILE if not info_path.exists(): raise FileNotFoundError(f"{_DUMP_INFO_FILE} not found in {dump_dir}") From 094c8446dd5e714adce2e20efd31d480e46eba70 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 8 Oct 2026 20:15:26 +0800 Subject: [PATCH 26/30] refactor: ship a single selective dump format Keep only format version 3 and move field packing back into the storage manager. Signed-off-by: OutstanderWang --- docs/checkpoint.md | 8 +- docs/data_dump.md | 17 ++-- tests/e2e/test_data_dump_e2e.py | 29 ------ tests/test_data_dump.py | 14 ++- transfer_queue/data_dump.py | 97 ++++++++----------- .../managers/simple_storage_manager.py | 51 +++++++++- transfer_queue/storage/simple_storage.py | 27 ++---- transfer_queue/utils/tensor_utils.py | 51 ---------- 8 files changed, 115 insertions(+), 179 deletions(-) diff --git a/docs/checkpoint.md b/docs/checkpoint.md index 2d41c13b..b69f8915 100644 --- a/docs/checkpoint.md +++ b/docs/checkpoint.md @@ -172,8 +172,6 @@ If step (1) partially succeeds and step (2) fails, the system is left in a mixed ## Exporting selected keys Use [selective data dumps](data_dump.md) when restoring selected keys into an -existing system or a different number of storage units. SimpleStorage restores for -dump formats v2 and later read assigned row ranges directly on the current storage -owners and merge values instead of replacing entire unit and controller state. -New dumps use v3, which also preserves field schemas. Version-1 dumps use the -caller-side KV put fallback. +existing system or a different number of storage units. SimpleStorage restores read +assigned row ranges directly on the current storage owners and merge values instead +of replacing entire unit and controller state. diff --git a/docs/data_dump.md b/docs/data_dump.md index 03f5dbdb..b8652b8b 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -26,7 +26,7 @@ owner and concurrently asks those units to write their records. Only units holdi selected rows participate. Tensor storage is compacted during serialization, so a row view cannot include the rest of its original batch, including inside tags. -On version-3 SimpleStorage restore: +On SimpleStorage restore: 1. The caller reads the row index and shard manifest and validates all file ranges. 2. The controller resolves existing keys and allocates indexes for new keys. @@ -48,7 +48,7 @@ storage unit. Local temporary storage suffices for single-node deployments. | Operation | State | Payload I/O | Unit count on restore | | --- | --- | --- | --- | | Checkpoint | Entire controller and storage state | Each unit reads/writes its whole file | Must match | -| Selective v3 dump | Selected fields and tags, merged by key | Each owner unit reads/writes its records | May differ | +| Selective dump | Selected fields and tags, merged by key | Each owner unit reads/writes its records | May differ | `DUMP_ROWS` and `LOAD_ROWS` are included in storage operation metrics. Unit logs record loaded rows and bytes; the manager logs total bytes and participating units. @@ -56,7 +56,7 @@ These count application reads, not filesystem read-ahead or physical disk traffi ## Format and compatibility -New dumps use `format_version: 3`: +Dumps use `format_version: 3`, the only version this build reads: ```text dump_info.json @@ -73,7 +73,7 @@ index and a field/value mapping. `shard_info.json` records each source index's current indexes without controller resolution. `row_index.pt` remains readable with `read_row_index` without opening payload shards. -Version 3 also saves the original field schemas and only the selected nested row +A dump also saves the original field schemas and only the selected nested row shapes. Restore uses that schema regardless of target topology or batch boundaries; non-tensor fields remain non-tensor even when a batch happens to contain only tensors. Destination type conflicts are rejected before payload writes. @@ -92,13 +92,8 @@ New puts also carry tensor shape hints for homogeneous tensor values wrapped in complete; real mixed values make the field non-tensor. A field originally declared non-tensor stays non-tensor when later batches contain only tensors. -Version-1 dumps remain readable using the prior caller-side KV put path. Version-2 -dumps retain direct reads, but lack original schemas and use the older inference -behavior; exact field-type preservation cannot be guaranteed for those files. -Restoring to a backend without direct selective loading uses KV puts. Version-1 -and KV fallback restores do not provide distributed file reads. Old builds that only understand -versions 1 or 2 cannot read version-3 dumps. Export of nonempty dumps currently requires -SimpleStorage. +Restoring to a backend without direct selective loading uses KV puts, which do not +provide distributed file reads. Export of nonempty dumps currently requires SimpleStorage. ## Failure behavior diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index b4bb674d..9d13fa5c 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -471,35 +471,6 @@ def no_payload_open(path, *args, **kwargs): assert snapshot.custom_meta[after[keys[0]]]["idx"] == 0 -def test_version_one_dump_remains_readable(tq_system, dump_dir, controller): - partition = "legacy" - dump_dir.mkdir(parents=True) - (dump_dir / "shards").mkdir() - torch.save( - { - "partition_id": partition, - "rows": { - "k": {"global_index": 100, "fields": ["x"], "tag": {"old": True}}, - "empty": {"global_index": 101, "fields": [], "tag": {}}, - }, - }, - dump_dir / "row_index.pt", - ) - (dump_dir / "dump_info.json").write_text( - json.dumps( - {"format_version": 1, "partition_id": partition, "num_keys": 2, "num_rows_with_data": 1, "num_shards": 1} - ) - ) - (dump_dir / "shards" / "shard_info.json").write_text( - json.dumps([{"position": 0, "storage_unit_id": "old", "rows": 1}]) - ) - with (dump_dir / "shards" / "shard_0_old.pkl").open("wb") as f: - pickle.dump({"global_indexes": [100], "field_data": {"x": {100: torch.tensor([7, 8])}}}, f) - tq.load_data_by_key(dump_dir) - assert sorted(_keys_mapping(controller, partition)) == ["empty", "k"] - _assert_rows_equal(tq.kv_batch_get(["k"], partition, select_fields=["x"])["x"], [torch.tensor([7, 8])]) - - @pytest.mark.parametrize("row_count", [3, 127, 128, 129, 130]) @pytest.mark.parametrize("last_kind", ["tensor", "none", "object"]) @pytest.mark.parametrize("legacy_metadata", [False, True]) diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index 1819adb1..b09eebdf 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -165,7 +165,17 @@ def read(self, size=-1): ZMQMessage.create( request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", - body={"shards": [{"path": str(path), "records": records}]}, + body={ + "shards": [ + { + "path": str(path), + "records": records, + "field_schema": { + "x": {"dtype": torch.float32, "shape": (4096,), "is_nested": False, "is_non_tensor": False} + }, + } + ] + }, ) ) assert loaded.body["success"], loaded.body @@ -220,7 +230,7 @@ def test_unit_rejects_invalid_records(unit, tmp_path, problem): ZMQMessage.create( request_type=ZMQRequestType.LOAD_ROWS, sender_id="test", - body={"shards": [{"path": str(path), "records": [record]}]}, + body={"shards": [{"path": str(path), "records": [record], "field_schema": {"x": {"is_non_tensor": True}}}]}, ) ) assert not reply.body["success"] diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index e6027cec..7d3aa323 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -15,9 +15,9 @@ """Persist selected rows and restore them by key without checkpointing controller state. -Version 3 preserves field schemas alongside independent row records in each shard. -The manifest maps source indexes to byte offsets, so current owner units read only their rows -when restoring into a different topology. Version 1 remains readable via KV puts. +Field schemas are saved alongside independent row records in each shard. The manifest +maps source indexes to byte offsets, so current owner units read only their rows when +restoring into a different topology. Layout:: @@ -31,7 +31,6 @@ import json import os -import pickle import shutil from collections import defaultdict from pathlib import Path @@ -44,7 +43,6 @@ from transfer_queue.storage.dump_io import pack_dump_field, read_dump_row, validate_dump_values from transfer_queue.utils import compact_pickle from transfer_queue.utils.logging_utils import get_logger -from transfer_queue.utils.tensor_utils import pack_field_values logger = get_logger(__name__) @@ -276,8 +274,8 @@ def _read_row_index(dump_dir: Path) -> dict[str, Any]: def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: """Merge selected rows into the running system, preserving existing key indexes. - SimpleStorage units read their assigned indexed records directly and in parallel. - Version-1 dumps and other backends use the compatible KV put path. New keys receive + SimpleStorage units read their assigned indexed records directly and in parallel; + other backends use the KV put path. New keys receive current indexes; unrelated rows and fields remain untouched. Writers and clears for these keys must be paused during restore. As with ``kv_batch_put``, a failure may leave new keys registered and payload partially written; retrying is idempotent. @@ -308,10 +306,10 @@ def _load_data_by_key(dump_dir: Path) -> dict[str, int]: raise FileNotFoundError(f"{_DUMP_INFO_FILE} not found in {dump_dir}") with open(info_path, encoding="utf-8") as f: dump_info = json.load(f) - if dump_info["format_version"] not in (1, 2, DUMP_FORMAT_VERSION): + if dump_info["format_version"] != DUMP_FORMAT_VERSION: raise ValueError( f"Unsupported dump format version {dump_info['format_version']} in {dump_dir}; " - f"this build reads versions 1 through {DUMP_FORMAT_VERSION}" + f"this build reads version {DUMP_FORMAT_VERSION}" ) row_index = _read_row_index(dump_dir) @@ -337,42 +335,37 @@ def _load_data_by_key(dump_dir: Path) -> dict[str, int]: if not path.is_file(): raise FileNotFoundError(f"Missing dump shard: {path}") records = [] - if dump_info["format_version"] >= 2: - size = path.stat().st_size - offsets = record["row_offsets"] - if len(offsets) != record["rows"]: - raise ValueError(f"Dump shard row count mismatch: {path}") - for source_index, (offset, length) in offsets.items(): - source_index = int(source_index) - if source_index not in keys_by_index or source_index in seen: - raise ValueError(f"Unexpected or duplicated row {source_index} in {path}") - if offset < 0 or length <= 0 or offset + length > size: - raise ValueError(f"Invalid row range for {source_index} in {path}") - key = keys_by_index[source_index] - records.append( - { - "key": key, - "source_index": source_index, - "fields": rows[key]["fields"], - "offset": offset, - "length": length, - } - ) - seen.add(source_index) - shard = {"path": str(path), "records": records} - if dump_info["format_version"] >= 3: - shard["field_schema"] = row_index["field_schema"] - shards.append(shard) - if dump_info["format_version"] >= 2 and seen != set(keys_by_index): + size = path.stat().st_size + offsets = record["row_offsets"] + if len(offsets) != record["rows"]: + raise ValueError(f"Dump shard row count mismatch: {path}") + for source_index, (offset, length) in offsets.items(): + source_index = int(source_index) + if source_index not in keys_by_index or source_index in seen: + raise ValueError(f"Unexpected or duplicated row {source_index} in {path}") + if offset < 0 or length <= 0 or offset + length > size: + raise ValueError(f"Invalid row range for {source_index} in {path}") + key = keys_by_index[source_index] + records.append( + { + "key": key, + "source_index": source_index, + "fields": rows[key]["fields"], + "offset": offset, + "length": length, + } + ) + seen.add(source_index) + shards.append({"path": str(path), "records": records, "field_schema": row_index["field_schema"]}) + if seen != set(keys_by_index): raise ValueError("Dump shards do not contain every produced row") client = _maybe_create_tq_client() - if dump_info["format_version"] >= 3: - client.validate_dump_schema(partition_id, row_index["field_schema"]) - if dump_info["format_version"] >= 2 and hasattr(getattr(client, "storage_manager", None), "load_rows_by_index"): + client.validate_dump_schema(partition_id, row_index["field_schema"]) + if hasattr(getattr(client, "storage_manager", None), "load_rows_by_index"): client.load_rows_by_key(partition_id, rows, shards) else: - _load_via_kv(partition_id, rows, shards, dump_info["format_version"]) + _load_via_kv(partition_id, rows, shards) total_bytes = sum(path.stat().st_size for path in dump_dir.rglob("*") if path.is_file()) logger.info(f"Restored {len(rows)} keys into partition {partition_id} from {dump_dir}") @@ -384,41 +377,27 @@ def _load_data_by_key(dump_dir: Path) -> dict[str, int]: } -def _load_via_kv(partition_id: str, rows: dict[str, Any], shards: list[dict], version: int) -> None: - """Retain v1 and non-SimpleStorage compatibility without changing their put contract.""" +def _load_via_kv(partition_id: str, rows: dict[str, Any], shards: list[dict]) -> None: + """Restore into non-SimpleStorage backends without changing their put contract.""" from transfer_queue.interface import kv_batch_put - keys_by_index = {row["global_index"]: key for key, row in rows.items()} restored = set() for shard in shards: with open(shard["path"], "rb") as f: - if version == 1: - saved = pickle.load(f) - records = [ - {"key": keys_by_index[index], "source_index": index, "fields": rows[keys_by_index[index]]["fields"]} - for index in saved["global_indexes"] - ] - else: - records = shard["records"] + records = shard["records"] for start in range(0, len(records), 128): groups = defaultdict(list) for record in records[start : start + 128]: index = record["source_index"] fields = record["fields"] - if version == 1: - values = {name: saved["field_data"][name][index] for name in fields} - else: - values = read_dump_row(f, record["offset"], record["length"], index, fields) - if version >= 3: - validate_dump_values(values, shard["field_schema"], index) + values = read_dump_row(f, record["offset"], record["length"], index, fields) + validate_dump_values(values, shard["field_schema"], index) groups[tuple(fields)].append((record["key"], values)) restored.add(index) for signature, batch in groups.items(): keys = [key for key, _ in batch] packed = { name: pack_dump_field([values[name] for _, values in batch], shard["field_schema"][name]) - if version >= 3 - else pack_field_values([values[name] for _, values in batch]) for name in signature } kv_batch_put( diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 94586226..cea2260f 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -37,7 +37,6 @@ ) from transfer_queue.utils.common import log_heavy_operation from transfer_queue.utils.logging_utils import get_logger -from transfer_queue.utils.tensor_utils import pack_field_values from transfer_queue.utils.zmq_utils import ( TQ_SOCKET_POOL_SIZE, ZMQMessage, @@ -555,7 +554,55 @@ async def _put_to_single_storage_unit( ) raise RuntimeError(f"Error in put to storage unit {target_storage_unit}: {type(e).__name__}: {e}") from e - _pack_field_values = staticmethod(pack_field_values) + @staticmethod + def _pack_field_values(values: list) -> torch.Tensor | NonTensorStack: + """ + Pack a list of per-sample values into a batched container. + + For pure tensor lists (no None), this tries nested tensor + (jagged layout first, then strided fallback), then falls back to + ``NonTensorStack``. Scalar tensors are stacked densely. + Mixed types, non-tensor values, or lists containing None placeholders + are grouped into a ``NonTensorStack``. + + Args: + values: List of per-sample values to pack. May contain None for + unfilled batch positions. + + Returns: + A ``torch.Tensor`` (nested or dense) when all values are tensors, + otherwise a ``NonTensorStack``. + + Raises: + ValueError: If *values* is empty. + """ + if not values: + raise ValueError("_pack_field_values received empty values list; caller should filter empty batches") + non_none = [v for v in values if v is not None] + if non_none and all(isinstance(v, torch.Tensor) for v in non_none): + if len(non_none) == len(values): + # Scalar tensors cannot be represented as jagged nested tensors; + # stack them densely to avoid noisy fallback warnings. + if all(v.dim() == 0 for v in non_none): + return torch.stack(non_none) + # Pure tensor list — try nested tensor + try: + return torch.nested.as_nested_tensor(values, layout=torch.jagged) + except (RuntimeError, TypeError) as e: + logger.warning( + f"Failed to pack nested tensor with jagged layout. " + f"Falling back to strided layout. Detailed error: {e}" + ) + try: + return torch.nested.as_nested_tensor(values, layout=torch.strided) + except (RuntimeError, TypeError) as e2: + logger.warning( + f"Failed to pack nested tensor with strided layout. " + f"Falling back to NonTensorStack. Detailed error: {e2}" + ) + return NonTensorStack(*values) + # Mixed tensor + None — cannot create nested tensor, fall through to NonTensorStack + return NonTensorStack(*values) async def get_data(self, metadata: BatchMeta) -> TensorDict: """ diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index 8bb42843..4a297fb5 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -47,16 +47,13 @@ import ray import torch import zmq -from tensordict import TensorDict -from transfer_queue.metadata import extract_field_schema from transfer_queue.storage.dump_io import read_dump_row, select_dump_schema, validate_dump_values from transfer_queue.utils import compact_pickle from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads, log_heavy_operation from transfer_queue.utils.enum_utils import Role from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.perf_utils import IntervalPerfMonitor -from transfer_queue.utils.tensor_utils import pack_field_values from transfer_queue.utils.zmq_utils import ( STORAGE_CLIENT_IDENTITY_PREFIXES, ZMQMessage, @@ -1377,28 +1374,18 @@ def _handle_load_rows(self, request: ZMQMessage) -> ZMQMessage: fields = read_dump_row( f, row["offset"], row["length"], row["source_index"], row["fields"] ) - if "field_schema" in shard: - validate_dump_values(fields, shard["field_schema"], row["source_index"]) + validate_dump_values(fields, shard["field_schema"], row["source_index"]) groups[tuple(row["fields"])].append((row, fields)) bytes_read += row["length"] for signature, rows in groups.items(): indexes = [row["target_index"] for row, _ in rows] values = {name: [fields[name] for _, fields in rows] for name in signature} - if "field_schema" in shard: - schema = select_dump_schema( - shard["field_schema"], - [row["source_index"] for row, _ in rows], - indexes, - signature, - ) - else: - packed = {name: pack_field_values(items) for name, items in values.items()} - schema = extract_field_schema(TensorDict(packed, batch_size=len(rows))) - for field in schema.values(): - if "per_sample_shapes" in field: - field["per_sample_shapes"] = dict( - zip(indexes, field["per_sample_shapes"], strict=True) - ) + schema = select_dump_schema( + shard["field_schema"], + [row["source_index"] for row, _ in rows], + indexes, + signature, + ) self.storage_data.put_data(values, indexes) updates.append({"global_indexes": indexes, "field_schema": schema}) logger.info( diff --git a/transfer_queue/utils/tensor_utils.py b/transfer_queue/utils/tensor_utils.py index a7a3ed7c..b3b8fa06 100644 --- a/transfer_queue/utils/tensor_utils.py +++ b/transfer_queue/utils/tensor_utils.py @@ -19,7 +19,6 @@ from functools import reduce import torch -from tensordict import NonTensorStack from torch import Tensor logger = logging.getLogger(__name__) @@ -181,53 +180,3 @@ def merge_contiguous_memory(ptrs: list[int], sizes: list[int]) -> tuple[list[int merged_sizes.append(current_size) return merged_ptrs, merged_sizes - - -def pack_field_values(values: list) -> torch.Tensor | NonTensorStack: - """ - Pack a list of per-sample values into a batched container. - - For pure tensor lists (no None), this tries nested tensor - (jagged layout first, then strided fallback), then falls back to - ``NonTensorStack``. Scalar tensors are stacked densely. - Mixed types, non-tensor values, or lists containing None placeholders - are grouped into a ``NonTensorStack``. - - Args: - values: List of per-sample values to pack. May contain None for - unfilled batch positions. - - Returns: - A ``torch.Tensor`` (nested or dense) when all values are tensors, - otherwise a ``NonTensorStack``. - - Raises: - ValueError: If *values* is empty. - """ - if not values: - raise ValueError("_pack_field_values received empty values list; caller should filter empty batches") - non_none = [v for v in values if v is not None] - if non_none and all(isinstance(v, torch.Tensor) for v in non_none): - if len(non_none) == len(values): - # Scalar tensors cannot be represented as jagged nested tensors; - # stack them densely to avoid noisy fallback warnings. - if all(v.dim() == 0 for v in non_none): - return torch.stack(non_none) - # Pure tensor list — try nested tensor - try: - return torch.nested.as_nested_tensor(values, layout=torch.jagged) - except (RuntimeError, TypeError) as e: - logger.warning( - f"Failed to pack nested tensor with jagged layout. " - f"Falling back to strided layout. Detailed error: {e}" - ) - try: - return torch.nested.as_nested_tensor(values, layout=torch.strided) - except (RuntimeError, TypeError) as e2: - logger.warning( - f"Failed to pack nested tensor with strided layout. " - f"Falling back to NonTensorStack. Detailed error: {e2}" - ) - return NonTensorStack(*values) - # Mixed tensor + None — cannot create nested tensor, fall through to NonTensorStack - return NonTensorStack(*values) From 9ccae8b0a629835f199f92ef719bfe8524568190 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 8 Oct 2026 20:42:02 +0800 Subject: [PATCH 27/30] fix: take dumped field types from the stored rows Storage units report each dumped row's dtype and shape, and the dump merges them with the field types the controller declares. This replaces the missing-shape recovery and leaves FieldMeta and put-time schema extraction unchanged. Signed-off-by: OutstanderWang --- docs/data_dump.md | 27 +-- tests/e2e/test_data_dump_e2e.py | 70 ++----- tests/test_data_dump.py | 185 ++++++------------ transfer_queue/client.py | 22 +-- transfer_queue/controller.py | 29 +-- transfer_queue/data_dump.py | 85 ++++---- transfer_queue/metadata.py | 19 +- transfer_queue/storage/managers/base.py | 23 +-- .../managers/simple_storage_manager.py | 18 +- transfer_queue/storage/simple_storage.py | 19 +- 10 files changed, 161 insertions(+), 336 deletions(-) diff --git a/docs/data_dump.md b/docs/data_dump.md index b8652b8b..16c45143 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -73,24 +73,15 @@ index and a field/value mapping. `shard_info.json` records each source index's current indexes without controller resolution. `row_index.pt` remains readable with `read_row_index` without opening payload shards. -A dump also saves the original field schemas and only the selected nested row -shapes. Restore uses that schema regardless of target topology or batch boundaries; -non-tensor fields remain non-tensor even when a batch happens to contain only tensors. -Destination type conflicts are rejected before payload writes. - -Legacy imports can leave nested fields with missing row shapes, for example when a -later checkpoint chunk wraps tensor values in `NonTensorStack`. Export asks the owner -units to inspect those values while writing their records. The caller merges only -shape/type metadata before publishing the dump; payloads still stay on the units. -Missing tensor shapes are filled from the actual values and their dtype is checked. -If a missing-shape row contains `None` or an object, the entire selected field is -saved as non-tensor, with a warning, so every unit restores the same field contract. -Fields already declared non-tensor remain non-tensor. This repairs the exported -schema without mutating live controller metadata or inventing shapes for objects. -New puts also carry tensor shape hints for homogeneous tensor values wrapped in -`NonTensorStack`. An existing tensor field uses those hints to keep its shape map -complete; real mixed values make the field non-tensor. A field originally declared -non-tensor stays non-tensor when later batches contain only tensors. +A dump also saves a schema for each selected field. The controller supplies the +declared type; while writing their records, the owner units report each row's dtype +and shape, which the caller merges without seeing payloads. Controller metadata can +trail the stored values, for example when a later put wraps rows of a tensor field in +`NonTensorStack`, so a field is saved as non-tensor (with a warning) unless every +selected row is a tensor of one dtype, and as nested if row shapes differ. A field +declared non-tensor stays non-tensor. Restore uses that schema regardless of target +topology or batch boundaries; destination type conflicts are rejected before payload +writes. Restoring to a backend without direct selective loading uses KV puts, which do not provide distributed file reads. Export of nonempty dumps currently requires SimpleStorage. diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index 9d13fa5c..9d5aebbb 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -26,7 +26,6 @@ import builtins import json import os -import pickle import shutil from pathlib import Path from unittest.mock import AsyncMock @@ -473,11 +472,8 @@ def no_payload_open(path, *args, **kwargs): @pytest.mark.parametrize("row_count", [3, 127, 128, 129, 130]) @pytest.mark.parametrize("last_kind", ["tensor", "none", "object"]) -@pytest.mark.parametrize("legacy_metadata", [False, True]) -def test_legacy_chunks_with_missing_nested_shapes_roundtrip( - tq_system, dump_dir, controller, row_count, last_kind, legacy_metadata, monkeypatch -): - partition = "legacy_nested" +def test_chunks_with_a_wrapped_last_row_roundtrip(tq_system, dump_dir, row_count, last_kind, monkeypatch): + partition = "wrapped_last_row" keys = [f"k{i}" for i in range(row_count)] tensors = [torch.arange(i % 3 + 1, dtype=torch.int64) for i in range(row_count)] last = tensors[-1] if last_kind == "tensor" else None if last_kind == "none" else {"pixels": tensors[-1]} @@ -509,41 +505,11 @@ def test_legacy_chunks_with_missing_nested_shapes_roundtrip( saved = torch.load(chunk, weights_only=False) tq.kv_batch_put(saved["keys"], partition, saved["fields"], tags=[{"key": key} for key in chunk_keys]) - snapshot = ray.get(controller.get_partition_snapshot.remote(partition)) - last_index = snapshot.keys_mapping[keys[-1]] - assert last_index in snapshot.field_metadata["input_ids"].global_indexes - if last_kind == "tensor": - assert tuple(snapshot.field_metadata["input_ids"].per_sample_shapes[last_index]) == tuple(tensors[-1].shape) - else: - assert snapshot.field_metadata["input_ids"].is_non_tensor - assert not snapshot.field_metadata["input_ids"].is_nested - if legacy_metadata: - from transfer_queue.controller import FieldMeta - - # Existing controller checkpoints can retain the old incomplete nested schema. - checkpoint = dump_dir.parent / "legacy-controller.pkl" - tq.get_client().save_controller_checkpoint(str(checkpoint)) - with checkpoint.open("rb") as file: - state = pickle.load(file) - for name in ["input_ids", "multi_modal_inputs#images"]: - state["partitions"][partition].field_metadata[name] = FieldMeta( - global_indexes=set(snapshot.keys_mapping.values()), - dtype=torch.int64, - is_nested=True, - is_non_tensor=False, - per_sample_shapes={ - snapshot.keys_mapping[key]: tuple(value.shape) - for key, value in zip(keys[:-1], tensors[:-1], strict=True) - }, - ) - with checkpoint.open("wb") as file: - pickle.dump(state, file) - tq.get_client().load_controller_checkpoint(str(checkpoint)) open_file = builtins.open def no_shard_read(path, *args, **kwargs): if isinstance(path, str | Path) and Path(path).name.startswith("shard_") and str(path).endswith(".pkl"): - pytest.fail("Dump schema repair read a payload shard in the caller") + pytest.fail("The dump schema merge read a payload shard in the caller") return open_file(path, *args, **kwargs) for target in [dump_dir, dump_dir.parent / "second-dump"]: @@ -559,7 +525,8 @@ def no_shard_read(path, *args, **kwargs): for key, value in zip(keys, expected, strict=True): assert tuple(schema["per_sample_shapes"][index["rows"][key]["global_index"]]) == tuple(value.shape) assert index["field_schema"]["wrapped"]["is_non_tensor"] - if legacy_metadata and last_kind != "tensor" and target == dump_dir: + # The controller still declares a tensor field, which the dump does not change. + if last_kind != "tensor" and target == dump_dir: with pytest.raises(RuntimeError, match="tensor/non-tensor type mismatch"): tq.load_data_by_key(target) tq.kv_clear(keys, partition) @@ -602,18 +569,21 @@ def test_incompatible_schema_rejected_before_writes(tq_system, dump_dir, control _assert_rows_equal(tq.kv_batch_get(["k"], "schema", ["x"])["x"], [torch.tensor([1.5])]) -@pytest.mark.parametrize("tensor_first", [True, False]) -def test_regular_put_keeps_legacy_tensor_nontensor_acceptance(tq_system, controller, tensor_first): - values = [torch.tensor([[7]]), NonTensorStack("text")] - if not tensor_first: - values.reverse() - for key, value in zip(["first", "second"], values, strict=True): - tq.kv_batch_put([key], "legacy_put", TensorDict({"x": value}, batch_size=1)) - metadata = tq.get_client().kv_retrieve_meta(["first", "second"], "legacy_put") - assert metadata.is_ready - assert metadata.field_names == ["x"] - snapshot = ray.get(controller.get_partition_snapshot.remote("legacy_put")) - assert snapshot.field_metadata["x"].global_indexes == set(metadata.global_indexes) +@pytest.mark.parametrize("last", ["text", torch.tensor([3.0, 4.0, 5.0])], ids=["string", "ragged"]) +def test_dense_field_with_a_wrapped_later_row_roundtrips(tq_system, dump_dir, last): + partition = "dense_then_wrapped" + first = torch.tensor([[1.0, 2.0]]) + tq.kv_batch_put(["first"], partition, TensorDict({"x": first}, batch_size=1)) + tq.kv_batch_put(["last"], partition, TensorDict({"x": NonTensorStack(last)}, batch_size=1)) + tq.dump_data_by_key(dump_dir, ["first", "last"], partition) + tq.kv_clear(["first", "last"], partition) + tq.load_data_by_key(dump_dir) + restored = list(tq.kv_batch_get(["first", "last"], partition, ["x"])["x"]) + torch.testing.assert_close(restored[0], first[0]) + if isinstance(last, str): + assert restored[1] == last + else: + torch.testing.assert_close(restored[1], last) def test_failed_load_publishes_no_metadata_and_retry_is_idempotent(tq_system, dump_dir, controller, monkeypatch): diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index b09eebdf..d0d07ed7 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -259,22 +259,19 @@ def dump(shard_dir, indexes, fields_by_index): ) ) assert response.body["success"] - return [ - { - "position": 0, - "storage_unit_id": "unit", - "rows": len(indexes), - "row_offsets": response.body["row_offsets"], - } - ] + shard = { + "position": 0, + "storage_unit_id": "unit", + "rows": len(indexes), + "row_offsets": response.body["row_offsets"], + } + return {"shards": [shard], "row_schema": response.body["row_schema"]} client = SimpleNamespace( describe_data_dump=lambda *_: { "partition_id": "p", "rows": rows, - "field_schema": { - "x": {"dtype": torch.int64, "shape": (2,), "is_nested": False, "is_non_tensor": False}, - }, + "field_schema": {"x": {"is_nested": False, "is_non_tensor": False}}, }, validate_dump_schema=lambda *_: None, dump_rows_by_index=dump, @@ -333,14 +330,14 @@ async def test_dump_waits_for_writers_before_cleanup_can_start(tmp_path): failed = asyncio.Event() completed = [] - async def dump(path, target_storage_unit, global_indexes, fields_by_index, missing_shapes): + async def dump(path, target_storage_unit, global_indexes, fields_by_index): if target_storage_unit == "u0": failed.set() raise OSError("write failed") await failed.wait() await asyncio.sleep(0) completed.append(target_storage_unit) - return {"row_offsets": {1: [0, 1]}} + return {"row_offsets": {1: [0, 1]}, "row_schema": {}} manager._dump_single_shard = dump with pytest.raises(OSError, match="write failed"): @@ -348,54 +345,63 @@ async def dump(path, target_storage_unit, global_indexes, fields_by_index, missi assert completed == ["u1"] -def test_dump_recovers_shapes_from_units_without_forwarding_payloads(unit, tmp_path): +def test_dump_reports_every_rows_stored_types(unit, tmp_path): unit.storage_data.put_data({"x": [torch.arange(2), torch.arange(3)], "y": [None, {"key": "value"}]}, [9, 10]) response = unit._handle_dump_rows( ZMQMessage.create( request_type=ZMQRequestType.DUMP_ROWS, sender_id="test", - body={ - "path": str(tmp_path / "shard.pkl"), - "global_indexes": [9, 10], - "missing_shapes": {9: ["x", "y"], 10: ["x", "y"]}, - }, + body={"path": str(tmp_path / "shard.pkl"), "global_indexes": [9, 10]}, ) ) assert response.body["success"] - assert response.body["recovered_schema"] == { - 9: {"x": {"shape": (2,), "dtype": torch.int64}, "y": None}, - 10: {"x": {"shape": (3,), "dtype": torch.int64}, "y": None}, + assert response.body["row_schema"] == { + 9: {"x": (torch.int64, (2,)), "y": None}, + 10: {"x": (torch.int64, (3,)), "y": None}, } -@pytest.mark.parametrize("reverse_units", [False, True]) -def test_missing_shapes_merge_consistently_across_units(reverse_units): - schema = { - "x": {"dtype": torch.int64, "is_nested": True, "is_non_tensor": False, "per_sample_shapes": {1: None, 2: None}}, - } - shards = [ - {"recovered_schema": {1: {"x": {"shape": (2,), "dtype": torch.int64}}}}, - {"recovered_schema": {2: {"x": None}}}, - ] - if reverse_units: - shards.reverse() - data_dump._complete_dump_schema(schema, {1: ["x"], 2: ["x"]}, shards) - assert schema["x"] == {"dtype": None, "shape": None, "is_nested": False, "is_non_tensor": True} - assert all("recovered_schema" not in shard for shard in shards) - - -@pytest.mark.parametrize("problem", ["missing", "wrong_field", "wrong_dtype"]) -def test_dump_rejects_incomplete_schema_recovery(problem): - schema = {"x": {"dtype": torch.int64, "is_nested": True, "is_non_tensor": False, "per_sample_shapes": {9: None}}} - recovered = {9: {"x": {"shape": (2,), "dtype": torch.int64}}} - if problem == "missing": - recovered = {} - elif problem == "wrong_field": - recovered[9] = {"other": recovered[9]["x"]} - else: - recovered[9]["x"]["dtype"] = torch.float32 - with pytest.raises(ValueError): - data_dump._complete_dump_schema(schema, {9: ["x"]}, [{"recovered_schema": recovered}]) +_DENSE = {"is_nested": False, "is_non_tensor": False} +_NON_TENSOR = {"dtype": None, "shape": None, "is_nested": False, "is_non_tensor": True} + + +@pytest.mark.parametrize( + ("declared", "rows", "expected"), + [ + # A dense field whose later row a put wrapped as a string. + (_DENSE, [(torch.float32, (2,)), None], _NON_TENSOR), + # ... or as a tensor of another dtype. + (_DENSE, [(torch.float32, (2,)), (torch.int64, (2,))], _NON_TENSOR), + # ... or as a tensor of another shape. + ( + _DENSE, + [(torch.float32, (2,)), (torch.float32, (3,))], + { + "dtype": torch.float32, + "shape": None, + "is_nested": True, + "is_non_tensor": False, + "per_sample_shapes": {0: (2,), 1: (3,)}, + }, + ), + (_DENSE, [(torch.int64, ()), (torch.int64, ())], {**_DENSE, "dtype": torch.int64, "shape": (1,)}), + ( + {"is_nested": True, "is_non_tensor": False}, + [(torch.int64, (2,)), (torch.int64, (2,))], + { + "dtype": torch.int64, + "shape": None, + "is_nested": True, + "is_non_tensor": False, + "per_sample_shapes": {0: (2,), 1: (2,)}, + }, + ), + ({"is_nested": False, "is_non_tensor": True}, [(torch.int64, (2,)), (torch.int64, (2,))], _NON_TENSOR), + ], +) +def test_dump_schema_follows_the_stored_rows(declared, rows, expected): + row_schema = {index: {"x": meta} for index, meta in enumerate(rows)} + assert data_dump._dump_field_schema({"x": declared}, row_schema) == {"x": expected} def test_saved_missing_tensor_shape_is_reported_as_invalid_dump(): @@ -410,82 +416,3 @@ def test_saved_missing_tensor_shape_is_reported_as_invalid_dump(): } with pytest.raises(ValueError, match="has no saved shape at row 1"): validate_dump_values({"x": torch.arange(2)}, schema, 1) - - -@pytest.mark.parametrize("value", [torch.arange(4), None, {"image": torch.arange(4)}]) -@pytest.mark.parametrize("wrapped_first", [False, True]) -def test_field_metadata_merges_wrapped_values_without_incomplete_nested_shapes(value, wrapped_first): - from tensordict import NonTensorStack, TensorDict - - from transfer_queue.controller import DataPartitionStatus - from transfer_queue.metadata import extract_field_schema - - batches = [ - TensorDict( - {"x": torch.nested.as_nested_tensor([torch.arange(2), torch.arange(3)], layout=torch.jagged)}, batch_size=2 - ), - TensorDict({"x": NonTensorStack(value)}, batch_size=1), - ] - if wrapped_first: - batches.reverse() - partition = DataPartitionStatus("p") - offset = 0 - for batch in batches: - indexes = list(range(offset, offset + batch.batch_size[0])) - schema = extract_field_schema(batch) - for field in schema.values(): - for part in [field, field.get("tensor_schema", {})]: - if "per_sample_shapes" in part: - part["per_sample_shapes"] = dict(zip(indexes, part["per_sample_shapes"], strict=True)) - assert partition.update_production_status(indexes, [], schema) - offset += batch.batch_size[0] - meta = partition.field_metadata["x"] - if not wrapped_first and isinstance(value, torch.Tensor): - assert meta.is_nested - assert not meta.is_non_tensor - assert meta.per_sample_shapes == {0: (2,), 1: (3,), 2: (4,)} - else: - assert meta.is_non_tensor - assert not meta.is_nested - assert meta.dtype is None - assert not meta.per_sample_shapes - assert meta.global_indexes == {0, 1, 2} - - -@pytest.mark.parametrize("shape", [(), (1,), (4,)]) -def test_wrapped_tensor_hints_preserve_dense_fields(shape): - from tensordict import NonTensorStack, TensorDict - - from transfer_queue.controller import DataPartitionStatus - from transfer_queue.metadata import extract_field_schema - - partition = DataPartitionStatus("p") - value = torch.ones(shape, dtype=torch.int64) - first = extract_field_schema(TensorDict({"x": value.unsqueeze(0)}, batch_size=1)) - second = extract_field_schema(TensorDict({"x": NonTensorStack(value)}, batch_size=1)) - assert second["x"]["is_non_tensor"] - assert partition.update_production_status([0], [], first) - assert partition.update_production_status([1], [], second) - meta = partition.field_metadata["x"] - assert not meta.is_nested - assert not meta.is_non_tensor - assert tuple(meta.shape) == (shape or (1,)) - assert meta.dtype == torch.int64 - - -def test_schema_hints_do_not_iterate_broadcast_nontensor_data(monkeypatch): - from tensordict import NonTensorData, TensorDict - - from transfer_queue.metadata import extract_field_schema - - data = TensorDict({"x": NonTensorData(data={"kind": "image"}, batch_size=(2,))}, batch_size=2) - original = NonTensorData.__getitem__ - - def bounded_getitem(self, index): - assert index == 0, "Schema extraction tried to iterate broadcast NonTensorData" - return original(self, index) - - monkeypatch.setattr(NonTensorData, "__getitem__", bounded_getitem) - field = extract_field_schema(data)["x"] - assert field["is_non_tensor"] - assert "tensor_schema" not in field diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 66a1c726..0bafde72 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1132,7 +1132,7 @@ async def async_describe_data_dump( keys: list[str], socket: zmq.asyncio.Socket | None = None, ) -> dict[str, Any]: - """Fetch selected rows and their original field schemas without payloads.""" + """Fetch selected rows and their fields' declared types without payloads.""" response = await self._request_controller( socket=socket, request_type=ZMQRequestType.DESCRIBE_ROWS_BY_KEY, @@ -1165,18 +1165,17 @@ async def async_dump_rows_by_index( shard_dir: str, global_indexes: list[int], fields_by_index: dict[int, list[str]] | None = None, - missing_shapes: dict[int, list[str]] | None = None, - ) -> list[dict[str, Any]]: + ) -> dict[str, Any]: """Asynchronously dump the given rows into per-storage-unit shards. Args: shard_dir: Directory to write shard files into. global_indexes: Global indexes to dump. fields_by_index: Produced fields to persist; omitted for a raw storage dump. - missing_shapes: Fields whose row shapes must be recovered from stored values. Returns: - One entry per written shard. + ``{"shards", "row_schema"}``: one entry per written shard, and each row's + stored field types. Raises: RuntimeError: If the storage manager is not initialized, or a unit holds @@ -1190,9 +1189,7 @@ async def async_dump_rows_by_index( ) if not hasattr(self.storage_manager, "dump_rows_by_index"): raise NotImplementedError(f"{type(self.storage_manager).__name__} does not support selective data dump") - return await self.storage_manager.dump_rows_by_index( - shard_dir, global_indexes, fields_by_index, **({"missing_shapes": missing_shapes} if missing_shapes else {}) - ) + return await self.storage_manager.dump_rows_by_index(shard_dir, global_indexes, fields_by_index) async def async_load_rows_by_key( self, @@ -1874,25 +1871,24 @@ def dump_rows_by_index( shard_dir: str, global_indexes: list[int], fields_by_index: dict[int, list[str]] | None = None, - missing_shapes: dict[int, list[str]] | None = None, - ) -> list[dict[str, Any]]: + ) -> dict[str, Any]: """Synchronously dump the given rows into per-storage-unit shards. Args: shard_dir: Directory to write shard files into. global_indexes: Global indexes to dump. fields_by_index: Produced fields to persist; omitted for a raw storage dump. - missing_shapes: Fields whose row shapes must be recovered from stored values. Returns: - One entry per written shard. + ``{"shards", "row_schema"}``: one entry per written shard, and each row's + stored field types. Raises: RuntimeError: If the storage manager is not initialized, or a unit holds no data for a row it was asked to dump. NotImplementedError: If the storage backend does not support dumping. """ - return self._dump_rows_by_index(shard_dir, global_indexes, fields_by_index, missing_shapes=missing_shapes) + return self._dump_rows_by_index(shard_dir, global_indexes, fields_by_index) def load_rows_by_key(self, partition_id: str, rows: dict, shards: list[dict]) -> int: """Restore selected payloads at the indexes their keys resolve to now; return bytes read.""" diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index 1fe06baf..a8373953 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -223,20 +223,6 @@ def update(self, incoming: dict[str, Any], incoming_global_indexes: list[int]) - Raises: ValueError: If incoming dtype conflicts with existing dtype. """ - if self.is_non_tensor: - self.global_indexes.update(incoming_global_indexes) - return - if incoming.get("is_non_tensor"): - tensor_schema = incoming.get("tensor_schema") - if tensor_schema is None or tensor_schema["dtype"] != self.dtype: - self.is_non_tensor = True - self.is_nested = False - self.dtype = self.shape = None - self.per_sample_shapes.clear() - self.global_indexes.update(incoming_global_indexes) - return - incoming = tensor_schema - # dtype consistency check new_dtype = incoming.get("dtype") if new_dtype is not None: @@ -2255,15 +2241,12 @@ def _handle_describe_rows_by_key_request(self, request_msg: ZMQMessage) -> ZMQMe params = request_msg.body rows = self.describe_rows_by_key(params["partition_id"], params["keys"]) partition = self.partitions[params["partition_id"]] - field_schema = {} - for name, meta in partition.field_metadata.items(): - indexes = [row["global_index"] for row in rows.values() if name in row["fields"]] - if not indexes: - continue - schema = meta.to_batch_schema(indexes) - if schema.get("is_nested"): - schema["per_sample_shapes"] = {index: meta.per_sample_shapes.get(index) for index in indexes} - field_schema[name] = schema + # Only the declared types: row dtypes and shapes come from the storage units. + field_schema = { + name: {"is_nested": bool(meta.is_nested), "is_non_tensor": bool(meta.is_non_tensor)} + for name, meta in partition.field_metadata.items() + if any(name in row["fields"] for row in rows.values()) + } return self._make_response( request_msg, ZMQRequestType.DESCRIBE_ROWS_BY_KEY_RESPONSE, diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index 7d3aa323..2c7f6c83 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -133,35 +133,20 @@ def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dic } ) rows = row_index["rows"] - field_schema = row_index["field_schema"] - missing_shapes = {} - for row in rows.values(): - index = row["global_index"] - names = [ - name - for name in row["fields"] - if field_schema[name].get("is_nested") - and not field_schema[name].get("is_non_tensor") - and field_schema[name].get("per_sample_shapes", {}).get(index) is None - ] - if names: - missing_shapes[index] = names # A row whose fields are all still unproduced has nothing for a storage unit to # dump, but it keeps its key and tag so the restore can recreate the row. indexes_with_data = sorted(row["global_index"] for row in rows.values() if row["fields"]) - shard_records = ( - client.dump_rows_by_index( + shard_records = [] + if indexes_with_data: + dumped = client.dump_rows_by_index( str(tmp_dir / _SHARD_SUBDIR), indexes_with_data, {row["global_index"]: row["fields"] for row in rows.values() if row["fields"]}, - **({"missing_shapes": missing_shapes} if missing_shapes else {}), ) - if indexes_with_data - else [] - ) - _complete_dump_schema(field_schema, missing_shapes, shard_records) + shard_records = dumped["shards"] + row_index["field_schema"] = _dump_field_schema(row_index["field_schema"], dumped["row_schema"]) shard_dir = tmp_dir / _SHARD_SUBDIR shard_dir.mkdir(parents=True, exist_ok=True) with open(shard_dir / _SHARD_INFO_FILE, "w", encoding="utf-8") as f: @@ -212,37 +197,37 @@ def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dic } -def _complete_dump_schema(field_schema: dict, missing_shapes: dict[int, list[str]], shards: list[dict]) -> None: - """Complete legacy nested metadata before publishing any dump manifest.""" - recovered = {} - for shard in shards: - for index, fields in shard.pop("recovered_schema", {}).items(): - if index in recovered: - raise ValueError(f"Duplicated schema recovery for row {index}") - recovered[index] = fields - if set(recovered) != set(missing_shapes): - raise ValueError("Storage units did not return every requested missing row schema") - for index, names in missing_shapes.items(): - if set(recovered[index]) != set(names): - raise ValueError(f"Recovered schema fields disagree for row {index}") - - non_tensor_fields = {name for fields in recovered.values() for name, meta in fields.items() if meta is None} - for name in non_tensor_fields: - field_schema[name] = {"is_non_tensor": True, "is_nested": False, "shape": None, "dtype": None} - logger.warning( - "Dump field %r contains non-tensor values missing from legacy nested metadata; preserving as non-tensor", - name, - ) - for index, fields in recovered.items(): +def _dump_field_schema(declared: dict, row_schema: dict[int, dict]) -> dict: + """Combine each field's declared type with the values its owner units hold. + + A put that wraps rows in ``NonTensorStack`` leaves a tensor field's metadata + unchanged, so dtypes and shapes come from the stored rows: a field is saved as + non-tensor unless every row is a tensor of one dtype, and nested if shapes differ. + """ + rows_by_field = defaultdict(dict) + for index, fields in row_schema.items(): for name, meta in fields.items(): - if name in non_tensor_fields: - continue - field = field_schema[name] - if meta["dtype"] != field["dtype"]: - raise ValueError(f"Dump field {name!r} dtype mismatch at row {index}") - field["per_sample_shapes"][index] = tuple(meta["shape"]) - if recovered: - logger.info("Recovered missing dump schemas for %s rows", len(recovered)) + rows_by_field[name][index] = meta + schema = {} + for name, rows in rows_by_field.items(): + metas = list(rows.values()) + if declared[name]["is_non_tensor"] or None in metas or len({dtype for dtype, _ in metas}) > 1: + if not declared[name]["is_non_tensor"]: + logger.warning("Dump field %r holds rows that are not tensors of one dtype; saving as non-tensor", name) + schema[name] = {"dtype": None, "shape": None, "is_nested": False, "is_non_tensor": True} + continue + shapes = {index: tuple(shape) for index, (_, shape) in rows.items()} + nested = declared[name]["is_nested"] or len(set(shapes.values())) > 1 + schema[name] = { + "dtype": metas[0][0], + # A dense field of scalars is declared with shape (1,), as a put records it. + "shape": None if nested else shapes[next(iter(shapes))] or (1,), + "is_nested": nested, + "is_non_tensor": False, + } + if nested: + schema[name]["per_sample_shapes"] = shapes + return schema def read_row_index(dump_dir: str | Path) -> dict[str, Any]: diff --git a/transfer_queue/metadata.py b/transfer_queue/metadata.py index df25de54..e89c1531 100644 --- a/transfer_queue/metadata.py +++ b/transfer_queue/metadata.py @@ -23,7 +23,7 @@ import numpy as np import torch -from tensordict import NonTensorStack, TensorDict +from tensordict import TensorDict from transfer_queue.utils.logging_utils import get_logger @@ -185,23 +185,6 @@ def extract_field_schema(data: TensorDict) -> dict[str, dict[str, Any]]: # For nested tensors, record per-sample shapes if is_nested: field_meta["per_sample_shapes"] = [tuple(t.shape) for t in value.unbind()] - elif isinstance(value, NonTensorStack): - # A legacy chunk may wrap tensor rows without changing their values. Keep - # its non-tensor declaration; existing tensor fields can use this hint. - values = value.tolist() - if values and all(isinstance(item, torch.Tensor) and not item.is_nested for item in values): - dtype = values[0].dtype - if all(item.dtype == dtype for item in values): - shapes = [tuple(item.shape) for item in values] - nested = any(shape != shapes[0] for shape in shapes) - field_meta["tensor_schema"] = { - "dtype": dtype, - "shape": None if nested else shapes[0] or (1,), - "is_nested": nested, - "is_non_tensor": False, - } - if nested: - field_meta["tensor_schema"]["per_sample_shapes"] = shapes field_schema[field_name] = field_meta diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index 785a0849..369f800e 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -244,19 +244,16 @@ async def notify_data_update( normalized_field_schema = {} for field_name, field in field_schema.items(): field_copy = field.copy() - schemas = [field_copy] - if "tensor_schema" in field_copy: - field_copy["tensor_schema"] = dict(field_copy["tensor_schema"]) - schemas.append(field_copy["tensor_schema"]) - for schema in schemas: - per_sample_shapes = schema.get("per_sample_shapes") - if isinstance(per_sample_shapes, list | tuple): - if len(per_sample_shapes) != len(global_indexes): - raise ValueError( - f"per_sample_shapes length ({len(per_sample_shapes)}) does not match " - f"number of global_indexes ({len(global_indexes)}) for field '{field_name}'. " - ) - schema["per_sample_shapes"] = dict(zip(global_indexes, per_sample_shapes, strict=True)) + per_sample_shapes = field_copy.get("per_sample_shapes", None) + if isinstance(per_sample_shapes, list | tuple): + if len(per_sample_shapes) != len(global_indexes): + raise ValueError( + f"per_sample_shapes length ({len(per_sample_shapes)}) does not match " + f"number of global_indexes ({len(global_indexes)}) for field '{field_name}'. " + ) + field_copy["per_sample_shapes"] = { + global_indexes[i]: per_sample_shapes[i] for i in range(len(global_indexes)) + } normalized_field_schema[field_name] = field_copy request_msg = ZMQMessage.create( diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index cea2260f..255283d9 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -836,7 +836,6 @@ async def _dump_single_shard( target_storage_unit: str, global_indexes: list[int], fields_by_index: dict[int, list[str]] | None = None, - missing_shapes: dict[int, list[str]] | None = None, socket: zmq.Socket = None, ) -> dict[str, Any]: """Ask one storage unit to write the rows it owns into a shard file.""" @@ -849,7 +848,6 @@ async def _dump_single_shard( "path": path, "global_indexes": global_indexes, "fields_by_index": fields_by_index, - "missing_shapes": missing_shapes or {}, }, ) await socket.send_multipart(request_msg.serialize(), copy=False) @@ -878,8 +876,7 @@ async def dump_rows_by_index( shard_dir: str, global_indexes: list[int], fields_by_index: dict[int, list[str]] | None = None, - missing_shapes: dict[int, list[str]] | None = None, - ) -> list[dict[str, Any]]: + ) -> dict[str, Any]: """Dump the given rows into one shard per storage unit, in parallel. Each unit pickles its own rows in its own process, so the payload never passes @@ -890,10 +887,11 @@ async def dump_rows_by_index( shard_dir: Directory to write shard files into. global_indexes: Global indexes to dump. fields_by_index: Produced fields to persist; omitted for a raw storage dump. - missing_shapes: Fields whose row shapes must be recovered from stored values. Returns: - One entry per written shard: ``{"position", "storage_unit_id", "rows", "row_offsets"}``. + ``{"shards", "row_schema"}``: one ``{"position", "storage_unit_id", "rows", + "row_offsets"}`` entry per written shard, and each dumped row's + ``{field: (dtype, shape) or None}`` as the owner unit holds it. Raises: RuntimeError: A unit holds no data for a row it was asked to dump. @@ -912,35 +910,33 @@ async def dump_rows_by_index( target_storage_unit=su_id, global_indexes=indexes, fields_by_index={index: fields_by_index[index] for index in indexes} if fields_by_index else None, - missing_shapes={index: missing_shapes[index] for index in indexes if index in missing_shapes} - if missing_shapes - else None, ) for path, (su_id, indexes) in zip(paths, targets, strict=True) ), return_exceptions=True, ) shards = [] + row_schema = {} total_rows = 0 for pos, ((su_id, _), result) in enumerate(zip(targets, results, strict=True)): if isinstance(result, BaseException): raise result offsets = result["row_offsets"] total_rows += len(offsets) + row_schema.update(result["row_schema"]) shards.append( { "position": pos, "storage_unit_id": su_id, "rows": len(offsets), "row_offsets": offsets, - "recovered_schema": result.get("recovered_schema", {}), } ) logger.info( f"[{self.storage_manager_id}]: dumped {total_rows} rows across {len(targets)} shards to {shard_dir_path}" ) - return shards + return {"shards": shards, "row_schema": row_schema} async def load_rows_by_index(self, partition_id: str, shards: list[dict[str, Any]]) -> int: """Have current owner units read assigned byte ranges concurrently, then publish metadata. diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index 4a297fb5..1b953cda 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -1311,7 +1311,7 @@ def _handle_dump_rows(self, request: ZMQMessage) -> ZMQMessage: if missing: raise ValueError(f"Storage holds no data for requested rows: {sorted(missing)[:20]}") row_offsets = {} - recovered_schema = {} + row_schema = {} with open(path, "wb") as f: for index in sorted(indexes): described_fields = request.body.get("fields_by_index") @@ -1324,15 +1324,12 @@ def _handle_dump_rows(self, request: ZMQMessage) -> ZMQMessage: else: # Reused global indexes can retain fields no longer present in metadata. fields = {name: self.storage_data.field_data[name][index] for name in described_fields[index]} - missing_fields = request.body.get("missing_shapes", {}).get(index, []) - if missing_fields: - # Inspect values already being written; only shape/type metadata leaves the unit. - recovered_schema[index] = { - name: {"shape": tuple(fields[name].shape), "dtype": fields[name].dtype} - if isinstance(fields[name], torch.Tensor) - else None - for name in missing_fields - } + # Controller metadata can trail the stored values, so report what each + # row actually holds; only type metadata leaves the unit. + row_schema[index] = { + name: (value.dtype, tuple(value.shape)) if isinstance(value, torch.Tensor) else None + for name, value in fields.items() + } offset = f.tell() compact_pickle.dump({"global_index": index, "fields": fields}, f) row_offsets[index] = [offset, f.tell() - offset] @@ -1347,7 +1344,7 @@ def _handle_dump_rows(self, request: ZMQMessage) -> ZMQMessage: "dumped_rows": len(indexes), "missing_rows": [], "row_offsets": row_offsets, - "recovered_schema": recovered_schema, + "row_schema": row_schema, }, ) except Exception as e: From ca94c92570f3dddcd799c3e24fbaed9703290550 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 8 Oct 2026 20:52:42 +0800 Subject: [PATCH 28/30] refactor: narrow the selective dump API Define dump_data_by_key and load_data_by_key in interface.py, next to the checkpoint API, and keep only the file format in data_dump.py. Make the row index reader private, return {key: tag} from load_data_by_key, declare the row dump/load methods on StorageManager instead of probing with hasattr, drop the unused describe_rows_by_key client methods, and drop the KV put fallback: dump and load both require SimpleStorage. Signed-off-by: OutstanderWang --- docs/data_dump.md | 10 +- tests/e2e/test_data_dump_e2e.py | 19 +- tests/test_data_dump.py | 63 +----- transfer_queue/__init__.py | 4 +- transfer_queue/client.py | 38 +--- transfer_queue/data_dump.py | 276 +++++------------------- transfer_queue/interface.py | 121 +++++++++++ transfer_queue/storage/dump_io.py | 10 - transfer_queue/storage/managers/base.py | 21 ++ 9 files changed, 230 insertions(+), 332 deletions(-) diff --git a/docs/data_dump.md b/docs/data_dump.md index 16c45143..cffb2d81 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -10,8 +10,7 @@ import transfer_queue as tq tq.init() tq.dump_data_by_key("/shared/dumps/selected", ["sample-1", "sample-2"], "train") -index = tq.read_row_index("/shared/dumps/selected") -tq.load_data_by_key("/shared/dumps/selected") +tags = tq.load_data_by_key("/shared/dumps/selected") # {key: tag} of the restored keys ``` Pause writes and clears for these keys during both operations. A dump is not an @@ -70,8 +69,7 @@ shards/ Each shard is a sequence of independent pickle records containing a source global index and a field/value mapping. `shard_info.json` records each source index's `[offset, length]`. Source indexes only locate records; they are never reused as -current indexes without controller resolution. `row_index.pt` remains readable -with `read_row_index` without opening payload shards. +current indexes without controller resolution. A dump also saves a schema for each selected field. The controller supplies the declared type; while writing their records, the owner units report each row's dtype @@ -83,8 +81,8 @@ declared non-tensor stays non-tensor. Restore uses that schema regardless of tar topology or batch boundaries; destination type conflicts are rejected before payload writes. -Restoring to a backend without direct selective loading uses KV puts, which do not -provide distributed file reads. Export of nonempty dumps currently requires SimpleStorage. +Dump and load currently require SimpleStorage. Loading into another backend raises +`NotImplementedError` before any key is registered. ## Failure behavior diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index 9d5aebbb..d013989c 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -37,6 +37,7 @@ from tensordict import NonTensorStack, TensorDict import transfer_queue as tq +from transfer_queue.data_dump import _read_row_index os.environ["RAY_DEDUP_LOGS"] = "0" @@ -179,7 +180,7 @@ def test_tags_survive_the_roundtrip(self, tq_system, dump_dir, controller): # Dump + wipe + load tq.dump_data_by_key(dump_dir, selected, partition_id) ray.get(controller.clear_partition.remote(partition_id)) - tq.load_data_by_key(dump_dir) + assert tq.load_data_by_key(dump_dir) == {key: {"idx": keys.index(key)} for key in selected} # Check restored state snapshot = ray.get(controller.get_partition_snapshot.remote(partition_id)) @@ -234,7 +235,7 @@ def test_heterogeneous_field_sets_are_grouped(self, tq_system, dump_dir, control tq.load_data_by_key(dump_dir) # Check restored state: the extra field came back only on its own row - rows = tq.read_row_index(dump_dir)["rows"] + rows = _read_row_index(dump_dir)["rows"] assert "routed_experts" in rows["h1"]["fields"] assert "routed_experts" not in rows["h0"]["fields"] retrieved = tq.kv_batch_get(keys=["h1"], partition_id=partition_id, select_fields=["routed_experts"]) @@ -298,7 +299,7 @@ def test_empty_key_set_writes_a_readable_dump(self, tq_system, dump_dir, control assert (dump_dir / "dump_info.json").exists() # Check that loading it is a no-op rather than an error - assert tq.load_data_by_key(dump_dir)["keys"] == 0 + assert tq.load_data_by_key(dump_dir) == {} def test_live_partition_survives_a_dump(self, tq_system, dump_dir, controller): # Define test data @@ -322,7 +323,7 @@ def test_dump_refuses_an_existing_directory(self, tq_system, dump_dir): with pytest.raises(FileExistsError): tq.dump_data_by_key(dump_dir, ["s0"], partition_id) assert sorted(path.name for path in dump_dir.parent.iterdir()) == [dump_dir.name] - assert tq.read_row_index(dump_dir)["partition_id"] == partition_id + assert _read_row_index(dump_dir)["partition_id"] == partition_id def test_published_dump_loads_from_a_read_only_directory(self, tq_system, dump_dir): partition_id = "d_read_only" @@ -331,7 +332,7 @@ def test_published_dump_loads_from_a_read_only_directory(self, tq_system, dump_d tq.kv_clear(["s0"], partition_id) dump_dir.parent.chmod(0o555) try: - assert sorted(tq.read_row_index(dump_dir)["rows"]) == ["s0"] + assert sorted(_read_row_index(dump_dir)["rows"]) == ["s0"] tq.load_data_by_key(dump_dir) assert sorted(path.name for path in dump_dir.parent.iterdir()) == [dump_dir.name] finally: @@ -355,17 +356,13 @@ def test_row_index_describes_keys_without_reading_payload(self, tq_system, dump_ tq.dump_data_by_key(dump_dir, keys, partition_id) # Check the index - row_index = tq.read_row_index(dump_dir) + row_index = _read_row_index(dump_dir) assert row_index["partition_id"] == partition_id assert sorted(row_index["rows"]) == keys for row, key in enumerate(keys): assert row_index["rows"][key]["fields"] == ["attention_mask", "input_ids"] assert row_index["rows"][key]["tag"] == {"idx": row} - def test_read_row_index_rejects_a_missing_dump(self, tq_system, dump_dir): - with pytest.raises(FileNotFoundError, match="row_index.pt"): - tq.read_row_index(dump_dir) - # --------------------------------------------------------------------------- # error handling @@ -516,7 +513,7 @@ def no_shard_read(path, *args, **kwargs): with monkeypatch.context() as patcher: patcher.setattr(builtins, "open", no_shard_read) tq.dump_data_by_key(target, keys, partition) - index = tq.read_row_index(target) + index = _read_row_index(target) for name in ["input_ids", "multi_modal_inputs#images"]: schema = index["field_schema"][name] assert schema["is_non_tensor"] == (last_kind != "tensor") diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index d0d07ed7..eeb360b0 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -79,8 +79,8 @@ def test_row_index_compacts_tensors_inside_tags(monkeypatch, tmp_path): ) monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: client) - data_dump.dump_data_by_key(tmp_path / "dump", ["key"], "p") - restored = data_dump.read_row_index(tmp_path / "dump")["rows"]["key"]["tag"]["nested"][0].value + interface.dump_data_by_key(tmp_path / "dump", ["key"], "p") + restored = data_dump._read_row_index(tmp_path / "dump")["rows"]["key"]["tag"]["nested"][0].value torch.testing.assert_close(restored, batch[0]) assert restored.untyped_storage().nbytes() == restored.numel() * restored.element_size() @@ -99,7 +99,7 @@ def fail(path): monkeypatch.setattr(data_dump, "_fsync_directory", fail) with pytest.raises(OSError, match="disk full"): - data_dump.dump_data_by_key(tmp_path / "dump", [], "p") + interface.dump_data_by_key(tmp_path / "dump", [], "p") assert not list(tmp_path.iterdir()) @@ -111,13 +111,13 @@ def test_racing_dumps_to_one_path_publish_only_the_first(empty_dump_client, monk def publish_another_dump_first(path): if path.name.startswith("dump.tmp-") and not raced: raced.append(path) - data_dump.dump_data_by_key(dump, [], "first") + interface.dump_data_by_key(dump, [], "first") fsync(path) monkeypatch.setattr(data_dump, "_fsync_directory", publish_another_dump_first) with pytest.raises(OSError): - data_dump.dump_data_by_key(dump, [], "second") - assert data_dump.read_row_index(dump)["partition_id"] == "first" + interface.dump_data_by_key(dump, [], "second") + assert data_dump._read_row_index(dump)["partition_id"] == "first" assert [path.name for path in tmp_path.iterdir()] == ["dump"] @@ -237,56 +237,15 @@ def test_unit_rejects_invalid_records(unit, tmp_path, problem): assert not unit.storage_data._active_keys -def test_version_two_falls_back_to_kv_for_other_backends(unit, monkeypatch, tmp_path): - unit.storage_data.put_data({"x": [torch.tensor([7, 8])]}, [10]) - rows = { - "k": {"global_index": 10, "fields": ["x"], "tag": {"tag": 1}}, - "empty": {"global_index": 11, "fields": [], "tag": {}}, - } - - def dump(shard_dir, indexes, fields_by_index): - directory = type(tmp_path)(shard_dir) - directory.mkdir(parents=True) - response = unit._handle_dump_rows( - ZMQMessage.create( - request_type=ZMQRequestType.DUMP_ROWS, - sender_id="test", - body={ - "path": str(directory / "shard_0_unit.pkl"), - "global_indexes": indexes, - "fields_by_index": fields_by_index, - }, - ) - ) - assert response.body["success"] - shard = { - "position": 0, - "storage_unit_id": "unit", - "rows": len(indexes), - "row_offsets": response.body["row_offsets"], - } - return {"shards": [shard], "row_schema": response.body["row_schema"]} - +def test_load_refuses_other_backends_before_registering_keys(monkeypatch, tmp_path): client = SimpleNamespace( - describe_data_dump=lambda *_: { - "partition_id": "p", - "rows": rows, - "field_schema": {"x": {"is_nested": False, "is_non_tensor": False}}, - }, - validate_dump_schema=lambda *_: None, - dump_rows_by_index=dump, storage_manager=object(), + kv_retrieve_meta=lambda *_, **__: pytest.fail("registered keys on an unsupported backend"), ) monkeypatch.setattr(interface, "_TQ_CONTROLLER", object()) monkeypatch.setattr(interface, "_maybe_create_tq_client", lambda: client) - data_dump.dump_data_by_key(tmp_path / "dump", list(rows), "p") - calls = [] - monkeypatch.setattr(interface, "kv_batch_put", lambda *args, **kwargs: calls.append((args, kwargs))) - data_dump.load_data_by_key(tmp_path / "dump") - assert calls[0][0][:2] == (["k"], "p") - torch.testing.assert_close(calls[0][0][2]["x"][0], torch.tensor([7, 8])) - assert calls[0][1]["tags"] == [{"tag": 1}] - assert calls[1] == ((["empty"], "p"), {"tags": [{}]}) + with pytest.raises(NotImplementedError, match="does not support selective data load"): + interface.load_data_by_key(tmp_path) @pytest.mark.asyncio @@ -401,7 +360,7 @@ def test_dump_reports_every_rows_stored_types(unit, tmp_path): ) def test_dump_schema_follows_the_stored_rows(declared, rows, expected): row_schema = {index: {"x": meta} for index, meta in enumerate(rows)} - assert data_dump._dump_field_schema({"x": declared}, row_schema) == {"x": expected} + assert data_dump.dump_field_schema({"x": declared}, row_schema) == {"x": expected} def test_saved_missing_tensor_shape_is_reported_as_invalid_dump(): diff --git a/transfer_queue/__init__.py b/transfer_queue/__init__.py index 97c180ce..da7a540c 100644 --- a/transfer_queue/__init__.py +++ b/transfer_queue/__init__.py @@ -16,7 +16,6 @@ import os from .client import TransferQueueClient -from .data_dump import dump_data_by_key, load_data_by_key, read_row_index from .dataloader import StreamingDataLoader, StreamingDataset from .interface import ( async_kv_batch_get, @@ -26,6 +25,7 @@ async_kv_list, async_kv_put, close, + dump_data_by_key, get_client, get_metrics_endpoint, init, @@ -36,6 +36,7 @@ kv_list, kv_put, load_checkpoint, + load_data_by_key, save_checkpoint, ) from .metadata import BatchMeta, KVBatchMeta @@ -75,7 +76,6 @@ # Selective Data Dump Interface "dump_data_by_key", "load_data_by_key", - "read_row_index", ] + [ # High-Level StreamingDataLoader Interface diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 0bafde72..0772ccb9 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -1141,10 +1141,6 @@ async def async_describe_data_dump( ) return {name: response.body[name] for name in ("partition_id", "rows", "field_schema")} - async def async_describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, dict[str, Any]]: - """Fetch key-addressed row metadata, preserving the existing row-only API.""" - return (await self.async_describe_data_dump(partition_id, keys))["rows"] - @with_controller_socket async def async_validate_dump_schema( self, @@ -1187,8 +1183,6 @@ async def async_dump_rows_by_index( f"[{self.client_id}]: Storage manager not initialized. " "Call initialize_storage_manager() before dump operations." ) - if not hasattr(self.storage_manager, "dump_rows_by_index"): - raise NotImplementedError(f"{type(self.storage_manager).__name__} does not support selective data dump") return await self.storage_manager.dump_rows_by_index(shard_dir, global_indexes, fields_by_index) async def async_load_rows_by_key( @@ -1201,10 +1195,16 @@ async def async_load_rows_by_key( Failure semantics match ``kv_batch_put``: new keys stay registered and payload writes may be partial, and retrying is idempotent because keys keep their indexes. + + Raises: + RuntimeError: If the storage manager is not initialized, or a unit fails. + NotImplementedError: If the storage backend does not support selective loads. """ - manager = getattr(self, "storage_manager", None) - if manager is None or not hasattr(manager, "load_rows_by_index"): - raise NotImplementedError("Storage backend does not support direct selective load") + if not hasattr(self, "storage_manager") or self.storage_manager is None: + raise RuntimeError( + f"[{self.client_id}]: Storage manager not initialized. " + "Call initialize_storage_manager() before load operations." + ) if not rows: return 0 metadata = await self.async_kv_retrieve_meta(list(rows), partition_id, create=True) @@ -1212,7 +1212,7 @@ async def async_load_rows_by_key( for shard in shards: for record in shard["records"]: record["target_index"] = target_indexes[record["key"]] - bytes_read = await manager.load_rows_by_index(partition_id, shards) + bytes_read = await self.storage_manager.load_rows_by_index(partition_id, shards) metadata.update_custom_meta([row["tag"] for row in rows.values()]) await self.async_set_custom_meta(metadata) return bytes_read @@ -1401,7 +1401,6 @@ def wrapper(*args, **kwargs): self._kv_retrieve_meta = _make_sync(self.async_kv_retrieve_meta) self._kv_retrieve_keys = _make_sync(self.async_kv_retrieve_keys) self._kv_list = _make_sync(self.async_kv_list) - self._describe_rows_by_key = _make_sync(self.async_describe_rows_by_key) self._describe_data_dump = _make_sync(self.async_describe_data_dump) self._validate_dump_schema = _make_sync(self.async_validate_dump_schema) self._dump_rows_by_index = _make_sync(self.async_dump_rows_by_index) @@ -1843,23 +1842,8 @@ def kv_list( return self._kv_list(partition_id=partition_id) # ==================== Selective Data Dump API ==================== - def describe_rows_by_key(self, partition_id: str, keys: list[str]) -> dict[str, dict[str, Any]]: - """Synchronously fetch the row metadata a selective dump needs, via ZMQ RPC. - - Args: - partition_id: Partition that owns ``keys``. - keys: Keys to describe, already deduplicated by the caller. - - Returns: - ``{key: {"global_index": int, "fields": list[str], "tag": dict}}``. - - Raises: - RuntimeError: If the RPC fails, or the partition or a key is unknown. - """ - return self._describe_rows_by_key(partition_id, keys) - def describe_data_dump(self, partition_id: str, keys: list[str]) -> dict[str, Any]: - """Fetch the row index and selected field schemas for a selective dump.""" + """Fetch the row index and the selected fields' declared types for a selective dump.""" return self._describe_data_dump(partition_id, keys) def validate_dump_schema(self, partition_id: str, field_schema: dict) -> None: diff --git a/transfer_queue/data_dump.py b/transfer_queue/data_dump.py index 2c7f6c83..98d2b10b 100644 --- a/transfer_queue/data_dump.py +++ b/transfer_queue/data_dump.py @@ -13,7 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Persist selected rows and restore them by key without checkpointing controller state. +"""File format of selective dumps written by ``dump_data_by_key``. Field schemas are saved alongside independent row records in each shard. The manifest maps source indexes to byte offsets, so current owner units read only their rows when @@ -31,26 +31,22 @@ import json import os -import shutil from collections import defaultdict from pathlib import Path from typing import Any -from uuid import uuid4 import torch -from tensordict import TensorDict -from transfer_queue.storage.dump_io import pack_dump_field, read_dump_row, validate_dump_values from transfer_queue.utils import compact_pickle from transfer_queue.utils.logging_utils import get_logger logger = get_logger(__name__) DUMP_FORMAT_VERSION = 3 +SHARD_SUBDIR = "shards" _DUMP_INFO_FILE = "dump_info.json" _ROW_INDEX_FILE = "row_index.pt" -_SHARD_SUBDIR = "shards" _SHARD_INFO_FILE = "shard_info.json" @@ -74,130 +70,7 @@ def _fsync_directory(path: Path) -> None: os.close(fd) -def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) -> dict[str, int]: - """Dump the rows addressed by ``keys`` into ``dump_dir``. - - Each storage unit pickles the rows it owns in its own process, so the payload never - passes through the caller. The caller writes only a small row index. - - The dump is staged in a uniquely named sibling directory and renamed into place - once durable. An existing ``dump_dir`` is refused, so a published dump is never - replaced: loads need no lock and readers no write access. - - .. note:: - **Multi-node limitation**: dump_dir must reside on a shared network filesystem - (e.g. NFS, GPFS, Lustre) reachable from every storage unit, because each unit - writes its own shard from its own node. - - Callers must freeze writers for ``keys`` for the duration: TransferQueue has no - atomic snapshot, so a concurrent put can land between the row index and the shards. - - Args: - dump_dir: Directory to write the dump into. - keys: Keys to dump. Duplicates are dropped, first occurrence wins. - partition_id: Partition that owns ``keys``. - - Returns: - ``{"keys", "rows_with_data", "shards", "bytes"}``. - - Raises: - FileExistsError: ``dump_dir`` already exists. - RuntimeError: TransferQueue is not initialized, the partition or a key does not - exist, or a storage unit holds no data for a row it was asked to dump. - """ - return _dump_data_by_key(Path(dump_dir).resolve(), keys, partition_id) - - -def _dump_data_by_key(dump_dir: Path, keys: list[str], partition_id: str) -> dict[str, int]: - from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client - - if _TQ_CONTROLLER is None: - raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") - - if dump_dir.exists(): - raise FileExistsError(f"{dump_dir} already exists; write each dump to a new directory") - unique_keys = list(dict.fromkeys(keys)) - client = _maybe_create_tq_client() - # A unique staging name keeps concurrent dumps to the same path from sharing files. - tmp_dir = dump_dir.with_name(f"{dump_dir.name}.tmp-{uuid4().hex}") - tmp_dir.mkdir(parents=True) - - try: - row_index = ( - client.describe_data_dump(partition_id, unique_keys) - if unique_keys - else { - "partition_id": partition_id, - "rows": {}, - "field_schema": {}, - } - ) - rows = row_index["rows"] - - # A row whose fields are all still unproduced has nothing for a storage unit to - # dump, but it keeps its key and tag so the restore can recreate the row. - indexes_with_data = sorted(row["global_index"] for row in rows.values() if row["fields"]) - - shard_records = [] - if indexes_with_data: - dumped = client.dump_rows_by_index( - str(tmp_dir / _SHARD_SUBDIR), - indexes_with_data, - {row["global_index"]: row["fields"] for row in rows.values() if row["fields"]}, - ) - shard_records = dumped["shards"] - row_index["field_schema"] = _dump_field_schema(row_index["field_schema"], dumped["row_schema"]) - shard_dir = tmp_dir / _SHARD_SUBDIR - shard_dir.mkdir(parents=True, exist_ok=True) - with open(shard_dir / _SHARD_INFO_FILE, "w", encoding="utf-8") as f: - json.dump(shard_records, f) - _fsync_file(f) - # The shards themselves were synced by the units that wrote them, but their - # directory entries were created here, on this node. - _fsync_directory(shard_dir) - - # torch.save rather than json: a tag is an arbitrary picklable dict, and this - # path must not fail on a tag that happens to hold a tensor. - with open(tmp_dir / _ROW_INDEX_FILE, "wb") as f: - torch.save(row_index, f, pickle_module=compact_pickle) - _fsync_file(f) - - with open(tmp_dir / _DUMP_INFO_FILE, "w", encoding="utf-8") as f: - json.dump( - { - "format_version": DUMP_FORMAT_VERSION, - "partition_id": partition_id, - "num_keys": len(unique_keys), - "num_rows_with_data": len(indexes_with_data), - "num_shards": len(shard_records), - }, - f, - indent=2, - ) - _fsync_file(f) - # Everything the dump claims is now durable, so the staging directory can be - # published. Syncing the parent makes the rename itself survive a crash. - _fsync_directory(tmp_dir) - - # rename() refuses a non-empty target, so of two dumps racing to one path only - # the first is published. - tmp_dir.rename(dump_dir) - _fsync_directory(dump_dir.parent) - except BaseException: - shutil.rmtree(tmp_dir, ignore_errors=True) - raise - - total_bytes = sum(path.stat().st_size for path in dump_dir.rglob("*") if path.is_file()) - logger.info(f"Dumped {len(unique_keys)} keys of partition {partition_id} to {dump_dir}") - return { - "keys": len(unique_keys), - "rows_with_data": len(indexes_with_data), - "shards": len(shard_records), - "bytes": total_bytes, - } - - -def _dump_field_schema(declared: dict, row_schema: dict[int, dict]) -> dict: +def dump_field_schema(declared: dict, row_schema: dict[int, dict]) -> dict: """Combine each field's declared type with the values its owner units hold. A put that wraps rows in ``NonTensorStack`` leaves a tensor field's metadata @@ -230,23 +103,49 @@ def _dump_field_schema(declared: dict, row_schema: dict[int, dict]) -> dict: return schema -def read_row_index(dump_dir: str | Path) -> dict[str, Any]: - """Read a dump's row index without touching its payload. - - Lets a caller answer questions about which keys and fields a dump holds, for - example whether a row satisfies the current data contract, before paying to - deserialize any shard. +def publish_dump(tmp_dir: Path, dump_dir: Path, row_index: dict[str, Any], shard_records: list[dict]) -> None: + """Write the manifests next to the unit-written shards, then rename into place. - Args: - dump_dir: Directory previously written by ``dump_data_by_key``. + ``dump_info.json`` is written last and the staging directory is renamed only once + everything is durable, so a published dump is always complete. + """ + shard_dir = tmp_dir / SHARD_SUBDIR + shard_dir.mkdir(parents=True, exist_ok=True) + with open(shard_dir / _SHARD_INFO_FILE, "w", encoding="utf-8") as f: + json.dump(shard_records, f) + _fsync_file(f) + # The shards themselves were synced by the units that wrote them, but their + # directory entries were created here, on this node. + _fsync_directory(shard_dir) + + # torch.save rather than json: a tag is an arbitrary picklable dict, and this + # path must not fail on a tag that happens to hold a tensor. + with open(tmp_dir / _ROW_INDEX_FILE, "wb") as f: + torch.save(row_index, f, pickle_module=compact_pickle) + _fsync_file(f) - Returns: - ``{"partition_id": str, "rows": {key: {"global_index", "fields", "tag"}}}``. + rows = row_index["rows"] + with open(tmp_dir / _DUMP_INFO_FILE, "w", encoding="utf-8") as f: + json.dump( + { + "format_version": DUMP_FORMAT_VERSION, + "partition_id": row_index["partition_id"], + "num_keys": len(rows), + "num_rows_with_data": sum(1 for row in rows.values() if row["fields"]), + "num_shards": len(shard_records), + }, + f, + indent=2, + ) + _fsync_file(f) + # Everything the dump claims is now durable, so the staging directory can be + # published. Syncing the parent makes the rename itself survive a crash. + _fsync_directory(tmp_dir) - Raises: - FileNotFoundError: The row index is missing. - """ - return _read_row_index(Path(dump_dir)) + # rename() refuses a non-empty target, so of two dumps racing to one path only + # the first is published. + tmp_dir.rename(dump_dir) + _fsync_directory(dump_dir.parent) def _read_row_index(dump_dir: Path) -> dict[str, Any]: @@ -256,36 +155,17 @@ def _read_row_index(dump_dir: Path) -> dict[str, Any]: return torch.load(row_index_path, weights_only=False) -def load_data_by_key(dump_dir: str | Path) -> dict[str, int]: - """Merge selected rows into the running system, preserving existing key indexes. - - SimpleStorage units read their assigned indexed records directly and in parallel; - other backends use the KV put path. New keys receive - current indexes; unrelated rows and fields remain untouched. Writers and clears for - these keys must be paused during restore. As with ``kv_batch_put``, a failure may - leave new keys registered and payload partially written; retrying is idempotent. +def read_dump(dump_dir: Path) -> tuple[dict[str, Any], list[dict[str, Any]]]: + """Validate a dump's manifests and return its row index and per-shard records. - Args: - dump_dir: Directory previously written by ``dump_data_by_key``. For direct - distributed loading it must be accessible from every storage unit. - - Returns: - ``{"keys", "rows_with_data", "shards", "bytes"}``. + Only manifests are read; each record locates one row's byte range for the unit + that will load it. Raises: - RuntimeError: TransferQueue is not initialized or a storage unit fails. FileNotFoundError: The dump is incomplete. - ValueError: The manifest or a row disagrees with the row index. + ValueError: The format version is unsupported, or a manifest disagrees with + the row index. """ - return _load_data_by_key(Path(dump_dir).resolve()) - - -def _load_data_by_key(dump_dir: Path) -> dict[str, int]: - from transfer_queue.interface import _TQ_CONTROLLER, _maybe_create_tq_client - - if _TQ_CONTROLLER is None: - raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") - info_path = dump_dir / _DUMP_INFO_FILE if not info_path.exists(): raise FileNotFoundError(f"{_DUMP_INFO_FILE} not found in {dump_dir}") @@ -298,16 +178,15 @@ def _load_data_by_key(dump_dir: Path) -> dict[str, int]: ) row_index = _read_row_index(dump_dir) - partition_id = row_index["partition_id"] rows = row_index["rows"] - shard_dir = dump_dir / _SHARD_SUBDIR + shard_dir = dump_dir / SHARD_SUBDIR with open(shard_dir / _SHARD_INFO_FILE, encoding="utf-8") as f: shard_records = json.load(f) if ( len(rows) != dump_info["num_keys"] or len(shard_records) != dump_info["num_shards"] - or partition_id != dump_info["partition_id"] + or row_index["partition_id"] != dump_info["partition_id"] ): raise ValueError("Dump manifest disagrees with the row index") keys_by_index = {row["global_index"]: key for key, row in rows.items() if row["fields"]} @@ -344,55 +223,4 @@ def _load_data_by_key(dump_dir: Path) -> dict[str, int]: shards.append({"path": str(path), "records": records, "field_schema": row_index["field_schema"]}) if seen != set(keys_by_index): raise ValueError("Dump shards do not contain every produced row") - - client = _maybe_create_tq_client() - client.validate_dump_schema(partition_id, row_index["field_schema"]) - if hasattr(getattr(client, "storage_manager", None), "load_rows_by_index"): - client.load_rows_by_key(partition_id, rows, shards) - else: - _load_via_kv(partition_id, rows, shards) - - total_bytes = sum(path.stat().st_size for path in dump_dir.rglob("*") if path.is_file()) - logger.info(f"Restored {len(rows)} keys into partition {partition_id} from {dump_dir}") - return { - "keys": len(rows), - "rows_with_data": len(keys_by_index), - "shards": len(shard_records), - "bytes": total_bytes, - } - - -def _load_via_kv(partition_id: str, rows: dict[str, Any], shards: list[dict]) -> None: - """Restore into non-SimpleStorage backends without changing their put contract.""" - from transfer_queue.interface import kv_batch_put - - restored = set() - for shard in shards: - with open(shard["path"], "rb") as f: - records = shard["records"] - for start in range(0, len(records), 128): - groups = defaultdict(list) - for record in records[start : start + 128]: - index = record["source_index"] - fields = record["fields"] - values = read_dump_row(f, record["offset"], record["length"], index, fields) - validate_dump_values(values, shard["field_schema"], index) - groups[tuple(fields)].append((record["key"], values)) - restored.add(index) - for signature, batch in groups.items(): - keys = [key for key, _ in batch] - packed = { - name: pack_dump_field([values[name] for _, values in batch], shard["field_schema"][name]) - for name in signature - } - kv_batch_put( - keys, - partition_id, - TensorDict(packed, batch_size=len(keys)), - tags=[rows[key]["tag"] for key in keys], - ) - if restored != {row["global_index"] for row in rows.values() if row["fields"]}: - raise ValueError("Dump restore row count mismatch") - keys = [key for key, row in rows.items() if not row["fields"]] - if keys: - kv_batch_put(keys, partition_id, tags=[rows[key]["tag"] for key in keys]) + return row_index, shards diff --git a/transfer_queue/interface.py b/transfer_queue/interface.py index eca8762e..12acfbf3 100644 --- a/transfer_queue/interface.py +++ b/transfer_queue/interface.py @@ -36,6 +36,7 @@ from importlib import resources from pathlib import Path from typing import Any, Callable +from uuid import uuid4 import ray import torch @@ -43,6 +44,7 @@ from tensordict import TensorDict from tensordict.tensorclass import NonTensorStack +from transfer_queue import data_dump from transfer_queue.client import TransferQueueClient from transfer_queue.controller import TransferQueueController from transfer_queue.metadata import KVBatchMeta @@ -1141,3 +1143,122 @@ def load_checkpoint( client.load_controller_checkpoint(str(controller_path)) logger.info(f"Checkpoint loaded from {checkpoint_dir}") + + +# ==================== Selective Data Dump API ==================== + + +def dump_data_by_key(dump_dir: str | Path, keys: list[str], partition_id: str) -> dict[str, int]: + """Dump the rows addressed by ``keys`` into ``dump_dir``. + + Each storage unit pickles the rows it owns in its own process, so the payload never + passes through the caller. The caller writes only a small row index. + + The dump is staged in a uniquely named sibling directory and renamed into place + once durable. An existing ``dump_dir`` is refused, so a published dump is never + replaced: loads need no lock and readers no write access. + + .. note:: + **Multi-node limitation**: dump_dir must reside on a shared network filesystem + (e.g. NFS, GPFS, Lustre) reachable from every storage unit, because each unit + writes its own shard from its own node. + + Callers must freeze writers for ``keys`` for the duration: TransferQueue has no + atomic snapshot, so a concurrent put can land between the row index and the shards. + + Args: + dump_dir: Directory to write the dump into. + keys: Keys to dump. Duplicates are dropped, first occurrence wins. + partition_id: Partition that owns ``keys``. + + Returns: + ``{"keys", "rows_with_data", "shards", "bytes"}``. + + Raises: + FileExistsError: ``dump_dir`` already exists. + RuntimeError: TransferQueue is not initialized, the partition or a key does not + exist, or a storage unit holds no data for a row it was asked to dump. + NotImplementedError: The storage backend does not support selective dumps. + """ + if _TQ_CONTROLLER is None: + raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") + + dump_dir = Path(dump_dir).resolve() + if dump_dir.exists(): + raise FileExistsError(f"{dump_dir} already exists; write each dump to a new directory") + unique_keys = list(dict.fromkeys(keys)) + client = _maybe_create_tq_client() + # A unique staging name keeps concurrent dumps to the same path from sharing files. + tmp_dir = dump_dir.with_name(f"{dump_dir.name}.tmp-{uuid4().hex}") + tmp_dir.mkdir(parents=True) + + try: + row_index = ( + client.describe_data_dump(partition_id, unique_keys) + if unique_keys + else {"partition_id": partition_id, "rows": {}, "field_schema": {}} + ) + rows = row_index["rows"] + # A row whose fields are all still unproduced has nothing for a storage unit to + # dump, but it keeps its key and tag so the restore can recreate the row. + fields_by_index = {row["global_index"]: row["fields"] for row in rows.values() if row["fields"]} + shard_records = [] + if fields_by_index: + dumped = client.dump_rows_by_index( + str(tmp_dir / data_dump.SHARD_SUBDIR), sorted(fields_by_index), fields_by_index + ) + shard_records = dumped["shards"] + row_index["field_schema"] = data_dump.dump_field_schema(row_index["field_schema"], dumped["row_schema"]) + data_dump.publish_dump(tmp_dir, dump_dir, row_index, shard_records) + except BaseException: + shutil.rmtree(tmp_dir, ignore_errors=True) + raise + + total_bytes = sum(path.stat().st_size for path in dump_dir.rglob("*") if path.is_file()) + logger.info(f"Dumped {len(unique_keys)} keys of partition {partition_id} to {dump_dir}") + return { + "keys": len(unique_keys), + "rows_with_data": len(fields_by_index), + "shards": len(shard_records), + "bytes": total_bytes, + } + + +def load_data_by_key(dump_dir: str | Path) -> dict[str, dict]: + """Merge selected rows into the running system, preserving existing key indexes. + + SimpleStorage units read their assigned indexed records directly and in parallel. + New keys receive current indexes; unrelated rows and fields remain untouched. + Writers and clears for these keys must be paused during restore. As with + ``kv_batch_put``, a failure may leave new keys registered and payload partially + written; retrying is idempotent. + + Args: + dump_dir: Directory previously written by ``dump_data_by_key``. It must be + accessible from every storage unit. + + Returns: + ``{key: tag}`` for every restored key. + + Raises: + RuntimeError: TransferQueue is not initialized or a storage unit fails. + NotImplementedError: The storage backend is not SimpleStorage. + FileNotFoundError: The dump is incomplete. + ValueError: The manifest or a row disagrees with the row index. + """ + if _TQ_CONTROLLER is None: + raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") + + client = _maybe_create_tq_client() + # Checked up front: the client registers the keys before any unit reads a shard. + if not isinstance(client.storage_manager, AsyncSimpleStorageManager): + raise NotImplementedError(f"{type(client.storage_manager).__name__} does not support selective data load") + + dump_dir = Path(dump_dir).resolve() + row_index, shards = data_dump.read_dump(dump_dir) + partition_id = row_index["partition_id"] + rows = row_index["rows"] + client.validate_dump_schema(partition_id, row_index["field_schema"]) + bytes_read = client.load_rows_by_key(partition_id, rows, shards) + logger.info(f"Restored {len(rows)} keys into partition {partition_id} from {dump_dir}, reading {bytes_read} bytes") + return {key: row["tag"] for key, row in rows.items()} diff --git a/transfer_queue/storage/dump_io.py b/transfer_queue/storage/dump_io.py index 8840eaca..d59982f7 100644 --- a/transfer_queue/storage/dump_io.py +++ b/transfer_queue/storage/dump_io.py @@ -18,7 +18,6 @@ import pickle import torch -from tensordict import NonTensorStack def read_dump_row(file, offset: int, length: int, global_index: int, fields: list[str]) -> dict: @@ -65,12 +64,3 @@ def select_dump_schema(schema: dict, source_indexes: list[int], target_indexes: } selected[name] = field return selected - - -def pack_dump_field(values: list, schema: dict): - """Build fallback KV batches according to the original field contract.""" - if schema["is_non_tensor"]: - return NonTensorStack(*values) - if schema["is_nested"]: - return torch.nested.as_nested_tensor(values, layout=torch.jagged) - return torch.stack(values) diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index 369f800e..78a9c449 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -373,6 +373,27 @@ async def load_checkpoint(self, checkpoint_dir: str) -> None: """ raise NotImplementedError(f"{self.__class__.__name__} does not support checkpoint") + async def dump_rows_by_index( + self, + shard_dir: str, + global_indexes: list[int], + fields_by_index: dict[int, list[str]] | None = None, + ) -> dict[str, Any]: + """Have the owner units write the given rows into shards under shard_dir. + + Raises: + NotImplementedError: If this storage backend does not support selective dumps. + """ + raise NotImplementedError(f"{self.__class__.__name__} does not support selective data dump") + + async def load_rows_by_index(self, partition_id: str, shards: list[dict[str, Any]]) -> int: + """Have the current owner units read dumped rows at their target indexes. + + Raises: + NotImplementedError: If this storage backend does not support selective loads. + """ + raise NotImplementedError(f"{self.__class__.__name__} does not support selective data load") + def close(self) -> None: """Close all ZMQ sockets/contexts and stop the notify loop.""" From d2cfa0ef6c83760b4fb5819d6fe800c87827cc7b Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 8 Oct 2026 21:00:05 +0800 Subject: [PATCH 29/30] docs: describe keys left registered by a failed load Signed-off-by: OutstanderWang --- docs/data_dump.md | 6 ++++-- tests/e2e/test_data_dump_e2e.py | 1 + 2 files changed, 5 insertions(+), 2 deletions(-) diff --git a/docs/data_dump.md b/docs/data_dump.md index cffb2d81..78dd114b 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -94,8 +94,10 @@ succeeds. Readers need no lock and no write access. Restore has the failure semantics of `kv_batch_put`: it is not transactional, and payload writes before a failure remain. Metadata is published only after every unit -has succeeded, so a failed load leaves those writes invisible. Existing keys keep -their indexes, so retrying the same load is idempotent; clearing the keys abandons it. +has succeeded, so a failed load leaves those writes invisible. Keys it registered +stay registered: `kv_list` shows them with empty tags and no readable fields. Those +keys keep their indexes, so retrying the same load is idempotent and restores the +tags; clearing the keys abandons it. Like an ordinary put, a load does not fence late writes against indexes that are cleared and reused while it runs, so keep writers and clears for these keys paused. diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index d013989c..3989a1df 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -602,6 +602,7 @@ async def lose_reply(*args, **kwargs): snapshot = ray.get(controller.get_partition_snapshot.remote("retry")) index = snapshot.keys_mapping["key"] assert not snapshot.field_metadata + assert tq.kv_list("retry")["retry"] == {"key": {}} tq.load_data_by_key(dump_dir) snapshot = ray.get(controller.get_partition_snapshot.remote("retry")) From 14c1804d2d5a16ed8ac781ee161d41792b2111ec Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 8 Oct 2026 21:03:16 +0800 Subject: [PATCH 30/30] docs, test: state the pickle trust boundary; build real objects in dump tests Loading a dump or checkpoint unpickles its files, so both documents now say to load only trusted directories. The dump unit tests construct a real storage unit, and the manager tests run against the manager tq.init builds, instead of bypassing __init__. Signed-off-by: OutstanderWang --- docs/checkpoint.md | 3 ++ docs/data_dump.md | 3 ++ tests/e2e/test_data_dump_e2e.py | 86 ++++++++++++++++++++++++++++++ tests/test_data_dump.py | 93 ++------------------------------- 4 files changed, 95 insertions(+), 90 deletions(-) diff --git a/docs/checkpoint.md b/docs/checkpoint.md index b69f8915..38cf0edb 100644 --- a/docs/checkpoint.md +++ b/docs/checkpoint.md @@ -101,6 +101,9 @@ checkpoint_dir/ ] ``` +Loading unpickles the controller state and every storage unit file, which can run +arbitrary code. Load only checkpoints from directories that you trust. + ## Atomic Checkpoint Replacement `save_checkpoint` writes to `.tmp`, renames the existing `checkpoint_dir` to `.old`, renames `.tmp` into `checkpoint_dir`, then deletes `.old`. This keeps the old checkpoint recoverable until the new one is fully in place; a failure partway through restores `.old` automatically, and the directory stays a plain folder (no symlink). diff --git a/docs/data_dump.md b/docs/data_dump.md index 78dd114b..39091804 100644 --- a/docs/data_dump.md +++ b/docs/data_dump.md @@ -71,6 +71,9 @@ index and a field/value mapping. `shard_info.json` records each source index's `[offset, length]`. Source indexes only locate records; they are never reused as current indexes without controller resolution. +Loading unpickles `row_index.pt` and every shard record, which can run arbitrary +code. Load only dumps from directories that you trust. + A dump also saves a schema for each selected field. The controller supplies the declared type; while writing their records, the owner units report each row's dtype and shape, which the caller merges without seeing payloads. Controller metadata can diff --git a/tests/e2e/test_data_dump_e2e.py b/tests/e2e/test_data_dump_e2e.py index 3989a1df..4bc8d1c3 100644 --- a/tests/e2e/test_data_dump_e2e.py +++ b/tests/e2e/test_data_dump_e2e.py @@ -23,6 +23,7 @@ pytest tests/e2e/test_data_dump_e2e.py -v """ +import asyncio import builtins import json import os @@ -611,6 +612,91 @@ async def lose_reply(*args, **kwargs): assert tq.kv_list("retry")["retry"] == {"key": {"idx": 0}} +def test_load_reaches_every_owner_unit_concurrently(tq_system, monkeypatch): + manager = tq.get_client().storage_manager + units = list(manager.storage_unit_infos) + seen = [] + + async def load_all(): + started, ready = set(), asyncio.Event() + + async def load(shards, target_storage_unit): + started.add(target_storage_unit) + if len(started) == len(units): + ready.set() + # Every unit waits for all the others, so this finishes only if they run at once. + await asyncio.wait_for(ready.wait(), timeout=2) + for shard in shards: + for row in shard["records"]: + assert target_storage_unit == units[row["target_index"] % len(units)] + seen.append(row["source_index"]) + return {"updates": [], "bytes_read": 0} + + monkeypatch.setattr(manager, "_load_selected_rows", load) + records = [{"source_index": i, "target_index": 31 - i} for i in range(16)] + return await manager.load_rows_by_index("p", [{"path": "shard.pkl", "records": records}]) + + assert asyncio.run(load_all()) == 0 + assert sorted(seen) == list(range(16)) + + +def test_load_waits_for_other_units_before_raising(tq_system, monkeypatch): + manager = tq.get_client().storage_manager + units = list(manager.storage_unit_infos) + finished, notified = [], [] + + async def load_with_one_failure(): + failed = asyncio.Event() + + async def load(shards, target_storage_unit): + if target_storage_unit == units[0]: + failed.set() + raise RuntimeError("unit failed") + await failed.wait() + await asyncio.sleep(0) + finished.append(target_storage_unit) + return {"bytes_read": 0, "updates": [{"global_indexes": [1], "field_schema": {}}]} + + async def notify(*args): + notified.append(args) + + monkeypatch.setattr(manager, "_load_selected_rows", load) + monkeypatch.setattr(manager, "notify_data_update", notify) + await manager.load_rows_by_index( + "p", [{"path": "shard", "records": [{"target_index": 0}, {"target_index": 1}]}] + ) + + with pytest.raises(RuntimeError, match="unit failed"): + asyncio.run(load_with_one_failure()) + assert finished == [units[1]] + assert not notified + + +def test_dump_waits_for_writers_before_cleanup_can_start(tq_system, dump_dir, monkeypatch): + manager = tq.get_client().storage_manager + units = list(manager.storage_unit_infos) + completed = [] + + async def dump_with_one_failure(): + failed = asyncio.Event() + + async def dump(path, target_storage_unit, global_indexes, fields_by_index): + if target_storage_unit == units[0]: + failed.set() + raise OSError("write failed") + await failed.wait() + await asyncio.sleep(0) + completed.append(target_storage_unit) + return {"row_offsets": {1: [0, 1]}, "row_schema": {}} + + monkeypatch.setattr(manager, "_dump_single_shard", dump) + await manager.dump_rows_by_index(str(dump_dir), [0, 1]) + + with pytest.raises(OSError, match="write failed"): + asyncio.run(dump_with_one_failure()) + assert completed == [units[1]] + + def test_dump_and_load_do_not_use_the_put_get_timeout_pool(tq_system, dump_dir, monkeypatch): _put_rows("own_pool", ["key"]) manager = tq.get_client().storage_manager diff --git a/tests/test_data_dump.py b/tests/test_data_dump.py index eeb360b0..d38b28d6 100644 --- a/tests/test_data_dump.py +++ b/tests/test_data_dump.py @@ -15,7 +15,6 @@ """Selective dump integrity and publication tests.""" -import asyncio import builtins import io import pickle @@ -26,18 +25,15 @@ from transfer_queue import data_dump, interface from transfer_queue.storage.dump_io import validate_dump_values -from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager from transfer_queue.storage.simple_storage import SimpleStorageUnit, StorageUnitData from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType @pytest.fixture def unit(): - cls = SimpleStorageUnit.__ray_metadata__.modified_class - unit = cls.__new__(cls) - unit.storage_unit_id = "test_unit" - unit.storage_data = StorageUnitData() - return unit + unit = SimpleStorageUnit.__ray_metadata__.modified_class({}) + yield unit + unit.shutdown() @pytest.mark.parametrize("row_count", [1, 32]) @@ -188,33 +184,6 @@ def read(self, size=-1): assert sorted(index for update in loaded.body["updates"] for index in update["global_indexes"]) == [101, 106] -@pytest.mark.asyncio -async def test_manager_loads_current_owners_concurrently(): - manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) - manager.storage_manager_id = "test" - manager.storage_unit_infos = dict.fromkeys(["u0", "u1", "u2", "u3"]) - manager.close = lambda: None - started = set() - ready = asyncio.Event() - seen = [] - - async def load(shards, target_storage_unit): - started.add(target_storage_unit) - if len(started) == 4: - ready.set() - await asyncio.wait_for(ready.wait(), timeout=2) - for shard in shards: - for row in shard["records"]: - assert target_storage_unit == f"u{row['target_index'] % 4}" - seen.append(row["source_index"]) - return {"updates": [], "bytes_read": 0} - - manager._load_selected_rows = load - records = [{"source_index": i, "target_index": 31 - i} for i in range(16)] - assert await manager.load_rows_by_index("p", [{"path": "shard.pkl", "records": records}]) == 0 - assert sorted(seen) == list(range(16)) - - @pytest.mark.parametrize("problem", ["wrong_index", "truncated", "missing_field"]) def test_unit_rejects_invalid_records(unit, tmp_path, problem): path = tmp_path / "row.pkl" @@ -248,62 +217,6 @@ def test_load_refuses_other_backends_before_registering_keys(monkeypatch, tmp_pa interface.load_data_by_key(tmp_path) -@pytest.mark.asyncio -async def test_load_waits_for_other_units_before_raising(): - manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) - manager.storage_manager_id = "test" - manager.storage_unit_infos = dict.fromkeys(["u0", "u1"]) - manager.close = lambda: None - failed = asyncio.Event() - finished = [] - - async def load(shards, target_storage_unit): - if target_storage_unit == "u0": - failed.set() - raise RuntimeError("unit failed") - await failed.wait() - await asyncio.sleep(0) - finished.append(target_storage_unit) - return {"bytes_read": 0, "updates": [{"global_indexes": [1], "field_schema": {}}]} - - async def notify(*args): - notified.append(args) - - notified = [] - manager._load_selected_rows = load - manager.notify_data_update = notify - with pytest.raises(RuntimeError, match="unit failed"): - await manager.load_rows_by_index( - "p", [{"path": "shard", "records": [{"target_index": 0}, {"target_index": 1}]}] - ) - assert finished == ["u1"] - assert not notified - - -@pytest.mark.asyncio -async def test_dump_waits_for_writers_before_cleanup_can_start(tmp_path): - manager = AsyncSimpleStorageManager.__new__(AsyncSimpleStorageManager) - manager.storage_manager_id = "test" - manager.storage_unit_infos = dict.fromkeys(["u0", "u1"]) - manager.close = lambda: None - failed = asyncio.Event() - completed = [] - - async def dump(path, target_storage_unit, global_indexes, fields_by_index): - if target_storage_unit == "u0": - failed.set() - raise OSError("write failed") - await failed.wait() - await asyncio.sleep(0) - completed.append(target_storage_unit) - return {"row_offsets": {1: [0, 1]}, "row_schema": {}} - - manager._dump_single_shard = dump - with pytest.raises(OSError, match="write failed"): - await manager.dump_rows_by_index(str(tmp_path), [0, 1]) - assert completed == ["u1"] - - def test_dump_reports_every_rows_stored_types(unit, tmp_path): unit.storage_data.put_data({"x": [torch.arange(2), torch.arange(3)], "y": [None, {"key": "value"}]}, [9, 10]) response = unit._handle_dump_rows(