diff --git a/docs/checkpoint.md b/docs/checkpoint.md index 8c0a2829..2d41c13b 100644 --- a/docs/checkpoint.md +++ b/docs/checkpoint.md @@ -167,4 +167,13 @@ 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. 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. diff --git a/docs/data_dump.md b/docs/data_dump.md new file mode 100644 index 00000000..462d1e31 --- /dev/null +++ b/docs/data_dump.md @@ -0,0 +1,187 @@ +# 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. 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 + +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-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. +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. 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 +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 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. +These count application reads, not filesystem read-ahead or physical disk traffic. + +## Format and compatibility + +New dumps use `format_version: 3`: + +```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 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. + +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. + +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. + +## 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; 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 +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 +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. + +## 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/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 new file mode 100644 index 00000000..1fd0d0ee --- /dev/null +++ b/tests/e2e/test_data_dump_cross_topology_e2e.py @@ -0,0 +1,162 @@ +# 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 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" + + +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(dump_test_root, request): + return dump_test_root / request.node.name / "dump" + + +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), (1, 4)], +) +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: + _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"]) + _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() + + +@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 new file mode 100644 index 00000000..13b7cf82 --- /dev/null +++ b/tests/e2e/test_data_dump_e2e.py @@ -0,0 +1,706 @@ +# 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. +Single-node runs default to pytest-managed temporary storage. + +Run with: + 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 +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 + + +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 +def dump_dir(dump_test_root, request): + case = dump_test_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) + + +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])]) + + +@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" + 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 + 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"]: + 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] + 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"] + 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): + 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"]) + 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 + + +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])]) + + +@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("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") + 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..51e95861 --- /dev/null +++ b/tests/e2e/test_restore_timeout_e2e.py @@ -0,0 +1,226 @@ +# 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 new file mode 100644 index 00000000..02ac5f29 --- /dev/null +++ b/tests/test_data_dump.py @@ -0,0 +1,504 @@ +# 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 asyncio +import builtins +import io +import pickle +from types import SimpleNamespace + +import pytest +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 + + +@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"] + assert set(reply.body["row_offsets"]) == set(indexes) + with path.open("rb") as f: + 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() + + +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_data_dump=lambda *_: { + "partition_id": "p", + "rows": {"key": {"global_index": 0, "fields": [], "tag": tag}}, + "field_schema": {}, + } + ) + 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() + + +@pytest.fixture +def empty_dump_client(monkeypatch): + monkeypatch.setattr(interface, "_TQ_CONTROLLER", 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): + 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" + + +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._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.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._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_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") + 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_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, 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 {"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) + + +@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/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/tests/test_restore_lifecycle.py b/tests/test_restore_lifecycle.py new file mode 100644 index 00000000..78e4613d --- /dev/null +++ b/tests/test_restore_lifecycle.py @@ -0,0 +1,552 @@ +# 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 754bb4d8..a0fa4eee 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, recover_data_load from .dataloader import StreamingDataLoader, StreamingDataset from .interface import ( async_kv_batch_get, @@ -44,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__ = ( [ @@ -70,6 +72,14 @@ "save_checkpoint", "load_checkpoint", ] + + [ + # Selective Data Dump Interface + "dump_data_by_key", + "load_data_by_key", + "read_row_index", + "recover_data_load", + "RestorePendingError", + ] + [ # High-Level StreamingDataLoader Interface "StreamingDataset", diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 0f7f13f1..f5d4c98b 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 ( @@ -1124,6 +1126,181 @@ def _can_destroy_zmq_context(self) -> bool: return False return True + # ==================== Selective Data Dump API ==================== + @with_controller_socket + async def async_describe_data_dump( + self, + partition_id: str, + keys: list[str], + socket: zmq.asyncio.Socket | None = None, + ) -> 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")} + + 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, + 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, + 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. + + 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. + + 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, 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.""" + 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]) + # ==================== Checkpoint API ==================== @with_controller_socket async def async_save_controller_checkpoint( @@ -1308,6 +1485,13 @@ 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) + 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) @@ -1744,6 +1928,69 @@ 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.""" + 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, + 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. + + 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. + + 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) + + 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) + # ==================== 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..4f393baf 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 @@ -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, @@ -223,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: @@ -557,6 +572,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], @@ -997,6 +1024,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._restore_outcomes: dict[str, bool] = {} + self._clearing_indexes: set[int] = set() # Connected storage managers tracking self._connected_storage_managers: set[str] = set() @@ -1479,30 +1510,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): """ @@ -1512,21 +1547,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): """ @@ -1563,54 +1600,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, @@ -1630,64 +1670,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] + + # 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, @@ -1725,6 +1767,164 @@ 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, 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() @@ -1948,6 +2148,12 @@ 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.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, } @@ -2181,6 +2387,48 @@ def _handle_load_controller_checkpoint_request(self, request_msg: ZMQMessage) -> {"success": True}, ) + 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.get(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_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"]) + 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 @@ -2205,23 +2453,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. @@ -2232,24 +2482,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 new file mode 100644 index 00000000..bf3daaff --- /dev/null +++ b/transfer_queue/data_dump.py @@ -0,0 +1,531 @@ +# 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. + +"""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. + +Layout:: + + / + dump_info.json # version, partition, counts + row_index.pt # key -> {global_index, fields, tag} + shards/ + shard_info.json # unit, row count, source index -> [offset, length] + 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 + +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 + +logger = get_logger(__name__) + +DUMP_FORMAT_VERSION = 3 + +_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) + + +@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. + + .. 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. + """ + 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: + raise RuntimeError("TransferQueue is not initialized. Call tq.init() first.") + + 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") + if tmp_dir.exists(): + shutil.rmtree(tmp_dir) + 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"] + 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( + 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: + 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) + + if dump_dir.exists(): + 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 { + "keys": len(unique_keys), + "rows_with_data": len(indexes_with_data), + "shards": len(shard_records), + "bytes": total_bytes, + } + + +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. + + 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. + """ + 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 + 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]: + """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 + 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. + + 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"}``. + + 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. + """ + 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: + 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(): + 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): + raise ValueError( + f"Unsupported dump format version {dump_info['format_version']} in {dump_dir}; " + f"this build reads versions 1 through {DUMP_FORMAT_VERSION}" + ) + + row_index = _read_row_index(dump_dir) + partition_id = row_index["partition_id"] + rows = row_index["rows"] + + 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"] + ): + 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: + 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) + 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"] >= 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) + 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(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) + 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_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, + 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]) + + +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 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/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/dump_io.py b/transfer_queue/storage/dump_io.py new file mode 100644 index 00000000..11772fef --- /dev/null +++ b/transfer_queue/storage/dump_io.py @@ -0,0 +1,113 @@ +# 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 + +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.""" + 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"] + + +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.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"] + 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) + + +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/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( diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index c89d1cfb..7ebaea3a 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, NamedTuple +from typing import Any, Callable import torch import zmq @@ -37,6 +37,8 @@ ) 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, ZMQMessage, @@ -105,13 +107,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. @@ -217,15 +212,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. @@ -538,55 +525,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: """ @@ -813,6 +752,211 @@ 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], + 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.""" + 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, + "missing_shapes": missing_shapes or {}, + }, + ) + 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 + 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], + 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. + + 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. + 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"}``. + + 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)] + + 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) + ), + return_exceptions=True, + ) + shards = [] + 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) + 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 + + 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: + rows = shard["records"] + for unit_id, group in self._group_by_hash([row["target_index"] for row in rows]).items(): + assignments[unit_id].append( + { + **shard, + "records": [rows[pos] for pos in group.batch_positions], + } + ) + 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() + ), + 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, + bytes_read, + len(assignments), + ) + return updates + + @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}, + ) + 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 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 df28521e..fbfc7a67 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,11 +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, 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, @@ -792,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() @@ -1003,6 +1010,14 @@ 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] + 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.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: @@ -1238,7 +1253,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) @@ -1293,6 +1308,202 @@ def _handle_save_checkpoint(self, data_parts) -> ZMQMessage: body={"success": False, "message": str(e)}, ) + 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: + missing = indexes - self.storage_data._active_keys + 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") + 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]} + 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] + f.flush() + os.fsync(f.fileno()) + logger.info("[%s]: dumped %s rows to %s", self.storage_unit_id, len(indexes), path) + 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, + "recovered_schema": recovered_schema, + }, + ) + 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 _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 + 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"] + ) + 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 = [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) + ) + 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("[%s]: load rows failed: %s", self.storage_unit_id, e) + return ZMQMessage.create( + request_type=ZMQRequestType.LOAD_ROWS_RESPONSE, + 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/compact_pickle.py b/transfer_queue/utils/compact_pickle.py new file mode 100644 index 00000000..f83a5d60 --- /dev/null +++ b/transfer_queue/utils/compact_pickle.py @@ -0,0 +1,37 @@ +# 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): + """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): + return value.clone().__reduce_ex__(pickle.HIGHEST_PROTOCOL) + return NotImplemented + + +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/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} 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 a94cd3ab..1ef84c28 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -135,6 +135,27 @@ 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" + LOAD_ROWS = "LOAD_ROWS" + LOAD_ROWS_RESPONSE = "LOAD_ROWS_RESPONSE" + 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: """