From 0adb7b6c71989c0436d67bbc1c0644127b88cb4d Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Sun, 20 Sep 2026 15:31:59 +0800 Subject: [PATCH 1/5] [test] Add unit tests for kv_update and empty Cover argument normalization, unit-side parser/empty application, cross-unit field_schema assembly, controller non-tensor marking, and KV backend rejection. Signed-off-by: OutstanderWang --- tests/test_kv_update.py | 203 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 203 insertions(+) create mode 100644 tests/test_kv_update.py diff --git a/tests/test_kv_update.py b/tests/test_kv_update.py new file mode 100644 index 00000000..28a6694d --- /dev/null +++ b/tests/test_kv_update.py @@ -0,0 +1,203 @@ +# 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 unittest.mock import MagicMock, patch + +import numpy as np +import pytest +import torch + +from transfer_queue.controller import DataPartitionStatus +from transfer_queue.interface import _normalize_kv_update_args +from transfer_queue.metadata import BatchMeta +from transfer_queue.storage.managers.base import KVStorageManager +from transfer_queue.storage.managers.simple_storage_manager import _build_update_field_schema +from transfer_queue.storage.simple_storage import StorageUnitData + + +def _concat(old, new): + if old is None: + return new + return torch.cat([old, new]) + + +def test_normalize_empty_rejects_values_and_parser(): + with pytest.raises(ValueError, match="must not specify values"): + _normalize_kv_update_args("tokens", torch.tensor([1]), None, empty=True) + with pytest.raises(ValueError, match="must not specify parser"): + _normalize_kv_update_args("tokens", None, _concat, empty=True) + names, batch, parser, use_empty = _normalize_kv_update_args("tokens", None, None, empty=True) + assert names == ["tokens"] + assert batch is None + assert parser is None + assert use_empty is True + + +def test_normalize_custom_parser_requires_values(): + with pytest.raises(ValueError, match="requires values"): + _normalize_kv_update_args("tokens", None, _concat) + + +def test_normalize_rejects_non_callable_parser(): + with pytest.raises(TypeError, match="parser must be callable unless empty=True"): + _normalize_kv_update_args("tokens", torch.tensor([1]), None) + + +def test_normalize_single_field_wraps_a_batch(): + names, batch, parser, use_empty = _normalize_kv_update_args("tokens", torch.tensor([4, 5]), _concat) + assert names == ["tokens"] + assert use_empty is False + assert parser is _concat + assert batch is not None + assert batch.batch_size == torch.Size([1]) + assert torch.equal(batch["tokens"][0], torch.tensor([4, 5])) + + +def test_normalize_multi_field_requires_matching_dict(): + with pytest.raises(TypeError, match="must be a dict"): + _normalize_kv_update_args(["a", "b"], torch.tensor([1]), _concat) + with pytest.raises(ValueError, match="same columns"): + _normalize_kv_update_args(["a", "b"], {"a": 1}, _concat) + + +def test_apply_update_concat_prompt_and_response_keeps_field_name(): + """Stored prompt_ids plus new response_ids become the sequence; the field name is unchanged.""" + prompt_ids = torch.tensor([10, 11, 12]) + response_ids = torch.tensor([20, 21]) + data = StorageUnitData() + data.put_data({"sequence_ids": [prompt_ids.clone()]}, [0]) + + described = data.apply_update([0], ["sequence_ids"], {"sequence_ids": [response_ids]}, _concat, False) + + assert list(described) == ["sequence_ids"] + assert torch.equal(data.field_data["sequence_ids"][0], torch.tensor([10, 11, 12, 20, 21])) + assert "prompt_ids" not in data.field_data + assert "response_ids" not in data.field_data + + +def test_apply_update_concatenates_and_is_atomic_on_parser_error(): + data = StorageUnitData() + data.put_data({"tokens": [torch.tensor([1, 2, 3])]}, [7]) + + described = data.apply_update([7], ["tokens"], {"tokens": [torch.tensor([4, 5])]}, _concat, False) + assert described["tokens"] == {"dtype": torch.int64, "shapes": [(5,)]} + assert torch.equal(data.field_data["tokens"][7], torch.tensor([1, 2, 3, 4, 5])) + + def boom(old, new): + raise RuntimeError("parser failed") + + with pytest.raises(RuntimeError, match="parser failed"): + data.apply_update([7], ["tokens"], {"tokens": [torch.tensor([9])]}, boom, False) + assert torch.equal(data.field_data["tokens"][7], torch.tensor([1, 2, 3, 4, 5])) + + +def test_apply_update_empty_stores_none(): + data = StorageUnitData() + data.put_data({"tokens": [torch.tensor([1])], "keep": [torch.tensor([2])]}, [3]) + described = data.apply_update([3], ["tokens"], None, None, True) + assert described["tokens"]["shapes"] is None + assert data.field_data["tokens"][3] is None + assert torch.equal(data.field_data["keep"][3], torch.tensor([2])) + + +def test_apply_update_missing_field_passes_none_as_old(): + data = StorageUnitData() + seen = [] + + def record(old, new): + seen.append(old) + return new + + data.apply_update([1], ["fresh"], {"fresh": [torch.tensor([8])]}, record, False) + assert seen == [None] + assert torch.equal(data.field_data["fresh"][1], torch.tensor([8])) + + +def test_build_update_field_schema_orders_shapes_across_units(): + """Units describe only their own rows; the batch schema must follow metadata order.""" + described = _build_update_field_schema( + [0, 1, 2, 3], + [ + ([0, 2], {"tokens": {"dtype": torch.int64, "shapes": [(4,), (4,)]}}), + ([1, 3], {"tokens": {"dtype": torch.int64, "shapes": [(9,), (9,)]}}), + ], + ) + + assert described["tokens"]["is_nested"] is True + assert described["tokens"]["shape"] is None + assert described["tokens"]["per_sample_shapes"] == [(4,), (9,), (4,), (9,)] + + +def test_build_update_field_schema_keeps_uniform_column_flat(): + described = _build_update_field_schema( + [0, 1], + [ + ([0], {"tokens": {"dtype": torch.int64, "shapes": [(4,)]}}), + ([1], {"tokens": {"dtype": torch.int64, "shapes": [(4,)]}}), + ], + ) + + assert described["tokens"] == { + "dtype": torch.int64, + "shape": (4,), + "is_nested": False, + "is_non_tensor": False, + } + + +def test_build_update_field_schema_marks_column_non_tensor_if_any_unit_is(): + described = _build_update_field_schema( + [0, 1], + [ + ([0], {"tokens": {"dtype": torch.int64, "shapes": [(4,)]}}), + ([1], {"tokens": {"dtype": None, "shapes": None}}), + ], + ) + + assert described["tokens"]["is_non_tensor"] is True + assert described["tokens"]["shape"] is None + + +def test_empty_marks_the_controller_field_non_tensor(): + """tq.kv_empty stores None, so the controller must stop describing the column as a tensor.""" + partition = DataPartitionStatus(partition_id="p") + partition._update_field_metadata( + [0], {"tokens": {"dtype": torch.int64, "shape": (5,), "is_nested": False, "is_non_tensor": False}} + ) + + partition._update_field_metadata( + [0], {"tokens": {"dtype": None, "shape": None, "is_nested": False, "is_non_tensor": True}} + ) + + tokens = partition.field_metadata["tokens"] + assert tokens.is_non_tensor is True + assert tokens.shape is None + assert tokens.to_batch_schema([0])["is_non_tensor"] is True + + +@pytest.mark.asyncio +@patch("transfer_queue.storage.managers.base.StorageClientFactory.create") +@patch.object(KVStorageManager, "_connect_to_controller", lambda self: None) +async def test_kv_backend_rejects_update(mock_create): + mock_create.return_value = MagicMock() + manager = KVStorageManager(controller_info=MagicMock(), config={"client_name": "YuanrongStorageClient"}) + meta = BatchMeta( + global_indexes=[0], + partition_ids=["p"], + field_schema={"x": {"dtype": torch.int64, "shape": (1,), "is_nested": False, "is_non_tensor": False}}, + production_status=np.ones(1, dtype=np.int8), + ) + with pytest.raises(NotImplementedError, match="kv_update is not supported for KV-based backends"): + await manager.update_data(meta, ["x"], empty=True) From b5f6d2f8bd9612bc508c3c11626610ae0d92a770 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Sun, 20 Sep 2026 15:34:52 +0800 Subject: [PATCH 2/5] [feat] Add kv_update for read-modify-write on an existing key kv_update rewrites selected fields of an existing key by running a custom parser(old, new) on the SimpleStorage unit that holds the row, so the payload never travels to the client and back. kv_empty stores None instead, keeping the key and its production status. The unit computes every value before writing, so a failing parser leaves the row unchanged, and it reports only the shape it stored per sample; the manager assembles one batch-level field_schema ordered to match the notified indexes. FieldMeta now accepts a column turning non-tensor, so an emptied field stops being described as a tensor. KV-based backends reject the call. Signed-off-by: OutstanderWang --- README.md | 1 + tests/e2e/test_kv_interface_e2e.py | 127 ++++++++++++ tests/test_kv_update.py | 17 +- transfer_queue/__init__.py | 8 + transfer_queue/client.py | 51 +++++ transfer_queue/controller.py | 50 +++-- transfer_queue/interface.py | 147 ++++++++++++++ transfer_queue/metadata.py | 16 +- transfer_queue/storage/managers/base.py | 35 ++++ .../managers/simple_storage_manager.py | 186 ++++++++++++++++++ transfer_queue/storage/simple_storage.py | 87 +++++++- transfer_queue/utils/zmq_utils.py | 3 + tutorial/02_kv_interface.py | 41 ++-- 13 files changed, 729 insertions(+), 40 deletions(-) diff --git a/README.md b/README.md index 3c08a6ad..62610857 100644 --- a/README.md +++ b/README.md @@ -116,6 +116,7 @@ To simplify the usage of TransferQueue, we provide a Redis-style high-level API - **(async_)kv_put**: Insert/Update a multi-column sample by key, with an optional metadata tag. - **(async_)kv_batch_put**: Put multiple key-value pairs efficiently in batches. - **(async_)kv_batch_get**: Retrieve samples (by keys), supporting column selection (by fields). +- **(async_)kv_update**: Rewrite selected fields of an existing key with `parser(old, new)`, or `empty=True` to store `None`. `kv_empty()` is the same as `kv_update(..., empty=True)`. SimpleStorage only. - **(async_)kv_list**: List keys and tags (metadata) in a partition. - **(async_)kv_clear**: Remove key-value pairs from storage. diff --git a/tests/e2e/test_kv_interface_e2e.py b/tests/e2e/test_kv_interface_e2e.py index 271758ee..c3f9c447 100644 --- a/tests/e2e/test_kv_interface_e2e.py +++ b/tests/e2e/test_kv_interface_e2e.py @@ -1110,6 +1110,133 @@ def test_field_expansion_across_samples(self, controller, tq_api): tq_api.kv_clear(keys=keys, partition_id=partition_id) +class TestKVUpdateE2E: + """kv_update: SimpleStorage runs parsers on the unit; KV backends reject.""" + + def test_concat_then_empty(self, controller, tq_api, backend_name): + if backend_name != "SimpleStorage": + pytest.skip("parser-backed kv_update is implemented only for SimpleStorage") + + partition_id = "test_partition" + key = "sample_update" + + tq_api.kv_put(key=key, partition_id=partition_id, fields={"tokens": torch.tensor([1, 2, 3])}, tag={"step": 1}) + + def concat(old, new): + return torch.cat([old, new]) + + meta = tq_api.kv_update( + key=key, partition_id=partition_id, fields="tokens", values=torch.tensor([4, 5]), parser=concat + ) + assert "tokens" in meta.fields + + retrieved = tq_api.kv_batch_get(keys=key, partition_id=partition_id, select_fields="tokens") + assert_tensor_equal(retrieved["tokens"][0], torch.tensor([1, 2, 3, 4, 5])) + + tq_api.kv_empty(key=key, partition_id=partition_id, fields="tokens") + emptied = tq_api.kv_batch_get(keys=key, partition_id=partition_id, select_fields="tokens") + assert emptied["tokens"][0] is None + + partition = get_controller_partition(controller, partition_id) + col = partition.field_name_mapping["tokens"] + global_idx = partition.keys_mapping[key] + assert partition.production_status[global_idx, col] == 1 + # The column now holds None, so the controller must not still describe it as a tensor. + assert partition.field_metadata["tokens"].is_non_tensor is True + + tq_api.kv_clear(keys=key, partition_id=partition_id) + + def test_concat_prompt_and_response_keeps_field_name(self, tq_api, backend_name): + """prompt_ids already stored; update appends response_ids; field stays sequence_ids.""" + if backend_name != "SimpleStorage": + pytest.skip("parser-backed kv_update is implemented only for SimpleStorage") + + partition_id = "test_partition" + key = "sample_seq" + prompt_ids = torch.tensor([10, 11, 12]) + response_ids = torch.tensor([20, 21]) + + tq_api.kv_put(key=key, partition_id=partition_id, fields={"sequence_ids": prompt_ids}) + tq_api.kv_update( + key=key, + partition_id=partition_id, + fields="sequence_ids", + values=response_ids, + parser=lambda old, new: torch.cat([old, new]), + ) + + retrieved = tq_api.kv_batch_get(keys=key, partition_id=partition_id) + assert "sequence_ids" in retrieved + assert "prompt_ids" not in retrieved + assert "response_ids" not in retrieved + assert_tensor_equal(retrieved["sequence_ids"][0], torch.tensor([10, 11, 12, 20, 21])) + + tq_api.kv_clear(keys=key, partition_id=partition_id) + + def test_update_multiple_fields_at_once(self, tq_api, backend_name): + if backend_name != "SimpleStorage": + pytest.skip("parser-backed kv_update is implemented only for SimpleStorage") + + partition_id = "test_partition" + key = "sample_multi" + + tq_api.kv_put( + key=key, + partition_id=partition_id, + fields={"a": torch.tensor([1, 2]), "b": torch.tensor([10])}, + ) + tq_api.kv_update( + key=key, + partition_id=partition_id, + fields=["a", "b"], + values={"a": torch.tensor([3]), "b": torch.tensor([20, 30])}, + parser=lambda old, new: torch.cat([old, new]), + ) + + retrieved = tq_api.kv_batch_get(keys=key, partition_id=partition_id) + assert_tensor_equal(retrieved["a"][0], torch.tensor([1, 2, 3])) + assert_tensor_equal(retrieved["b"][0], torch.tensor([10, 20, 30])) + + tq_api.kv_empty(key=key, partition_id=partition_id, fields=["a", "b"]) + emptied = tq_api.kv_batch_get(keys=key, partition_id=partition_id) + assert emptied["a"][0] is None + assert emptied["b"][0] is None + + tq_api.kv_clear(keys=key, partition_id=partition_id) + + def test_missing_key_and_empty_with_values(self, tq_api, backend_name): + if backend_name != "SimpleStorage": + pytest.skip("parser-backed kv_update is implemented only for SimpleStorage") + + with pytest.raises(ValueError, match="must not specify values"): + tq_api.kv_update( + key="no_such", partition_id="test_partition", fields="tokens", values=torch.tensor([1]), empty=True + ) + with pytest.raises(ValueError, match="not found"): + tq_api.kv_update( + key="no_such", + partition_id="test_partition", + fields="tokens", + values=torch.tensor([1]), + parser=lambda old, new: new, + ) + + def test_kv_backend_rejects(self, tq_api, backend_name): + if backend_name == "SimpleStorage": + pytest.skip("rejection is the KV-backend contract") + + tq_api.kv_put(key="k", partition_id="test_partition", fields={"tokens": torch.tensor([1])}) + with pytest.raises(NotImplementedError, match="kv_update is not supported"): + tq_api.kv_update( + key="k", + partition_id="test_partition", + fields="tokens", + values=torch.tensor([2]), + parser=lambda old, new: new, + ) + tq_api.kv_clear(keys="k", partition_id="test_partition") + + def run_tests(): """Run all e2e tests manually for debugging.""" pytest.main([__file__, "-v", "-s"]) diff --git a/tests/test_kv_update.py b/tests/test_kv_update.py index 28a6694d..aed9d4a8 100644 --- a/tests/test_kv_update.py +++ b/tests/test_kv_update.py @@ -24,7 +24,7 @@ from transfer_queue.metadata import BatchMeta from transfer_queue.storage.managers.base import KVStorageManager from transfer_queue.storage.managers.simple_storage_manager import _build_update_field_schema -from transfer_queue.storage.simple_storage import StorageUnitData +from transfer_queue.storage.simple_storage import HybridStorageUnitData, StorageUnitData def _concat(old, new): @@ -125,6 +125,21 @@ def record(old, new): assert torch.equal(data.field_data["fresh"][1], torch.tensor([8])) +def test_apply_update_decodes_ssd_offloaded_old_value(tmp_path): + """With SSD offload on, the parser must see the stored tensor, not its file reference.""" + data = HybridStorageUnitData( + storage_size=4, threshold_bytes=64, ssd_path=str(tmp_path), run_id="run", unit_id="unit" + ) + prompt = torch.arange(32) + data.put_data({"tokens": [prompt]}, [0]) + assert data.ssd_active_values == 1, "precondition: the prompt must live on SSD" + + data.apply_update([0], ["tokens"], {"tokens": [torch.tensor([99])]}, _concat, False) + + assert torch.equal(data.get_data(["tokens"], [0])["tokens"][0], torch.cat([prompt, torch.tensor([99])])) + assert data.ssd_active_values == 1, "the replaced SSD file must be released, not leaked" + + def test_build_update_field_schema_orders_shapes_across_units(): """Units describe only their own rows; the batch schema must follow metadata order.""" described = _build_update_field_schema( diff --git a/transfer_queue/__init__.py b/transfer_queue/__init__.py index 754bb4d8..e0b59b9b 100644 --- a/transfer_queue/__init__.py +++ b/transfer_queue/__init__.py @@ -22,8 +22,10 @@ async_kv_batch_get_by_meta, async_kv_batch_put, async_kv_clear, + async_kv_empty, async_kv_list, async_kv_put, + async_kv_update, close, get_client, get_metrics_endpoint, @@ -32,8 +34,10 @@ kv_batch_get_by_meta, kv_batch_put, kv_clear, + kv_empty, kv_list, kv_put, + kv_update, load_checkpoint, save_checkpoint, ) @@ -57,12 +61,16 @@ "kv_batch_get_by_meta", "kv_list", "kv_clear", + "kv_update", + "kv_empty", "async_kv_put", "async_kv_batch_put", "async_kv_batch_get", "async_kv_batch_get_by_meta", "async_kv_list", "async_kv_clear", + "async_kv_update", + "async_kv_empty", "KVBatchMeta", ] + [ diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 0f7f13f1..63234827 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -492,6 +492,42 @@ async def async_put( return metadata + async def async_update( + self, + metadata: BatchMeta, + field_names: list[str], + values: TensorDict | None = None, + parser: Callable[[Any, Any], Any] | None = None, + empty: bool = False, + ) -> BatchMeta: + """Apply empty or parser(old, new) on SimpleStorage units for the named fields. + + Args: + metadata: Samples to update. The key must already exist. + field_names: Fields to rewrite. + values: New values aligned with ``metadata``, or None when empty. + parser: ``parser(old, new) -> stored``. Ignored when empty. + empty: Store None for each named field. + + Returns: + The same metadata with field_schema and production status updated. + + Raises: + NotImplementedError: If the storage backend is not SimpleStorage. + """ + 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 performing storage operations." + ) + if not metadata or metadata.size == 0: + raise ValueError("metadata cannot be none or empty") + + field_schema = await self.storage_manager.update_data( + metadata, field_names, values=values, parser=parser, empty=empty + ) + return metadata.apply_field_schema(field_schema) + async def async_get_data(self, metadata: BatchMeta) -> TensorDict: """Asynchronously fetch data from storage units and organize into TensorDict. @@ -1294,6 +1330,7 @@ def wrapper(*args, **kwargs): # Bind internal sync wrappers. Public methods are defined explicitly below # to ensure proper type hints and documentation. self._put = _make_sync(self.async_put) + self._update = _make_sync(self.async_update) self._get_meta = _make_sync(self.async_get_meta) self._get_data = _make_sync(self.async_get_data) self._clear_partition = _make_sync(self.async_clear_partition) @@ -1480,6 +1517,20 @@ def put( """ return self._put(data=data, metadata=metadata, partition_id=partition_id, data_parser=data_parser) + def update( + self, + metadata: BatchMeta, + field_names: list[str], + values: TensorDict | None = None, + parser: Callable[[Any, Any], Any] | None = None, + empty: bool = False, + ) -> BatchMeta: + """Synchronously apply empty or parser(old, new) on SimpleStorage units. + + See ``async_update``. + """ + return self._update(metadata, field_names, values=values, parser=parser, empty=empty) + def get_data(self, metadata: BatchMeta) -> TensorDict: """Synchronously fetch data from storage units and organize into TensorDict. diff --git a/transfer_queue/controller.py b/transfer_queue/controller.py index b45a85a6..23ddaa3e 100644 --- a/transfer_queue/controller.py +++ b/transfer_queue/controller.py @@ -253,28 +253,36 @@ def update(self, incoming: dict[str, Any], incoming_global_indexes: list[int]) - # Update newly provided per_sample_shapes self.per_sample_shapes.update(new_per_sample_shapes) + elif new_is_non_tensor: + # Some samples now hold plain objects (e.g. kv_update empty stores None), so no + # single dtype/shape describes the column any more. dtype is kept so a later + # tensor write still has to agree with the one this field was created with. + self.is_non_tensor = True + self.is_nested = False + self.shape = None + self.per_sample_shapes.clear() + else: - if not new_is_non_tensor: - # newly input is regular tensor - new_shape = incoming.get("shape", None) - if new_shape is None: - raise ValueError("Receiving a regular tensor without 'shape'!") - if self.is_nested: - # we need to update incoming shape into per_sample_shapes - for gi in incoming_global_indexes: - self.per_sample_shapes[gi] = new_shape - else: - if self.is_non_tensor is not None and not self.is_non_tensor: - # original data is also regular tensor - assert self.shape is not None - if self.shape != new_shape: - for gi in self.global_indexes: - self.per_sample_shapes[gi] = self.shape - for gi in incoming_global_indexes: - self.per_sample_shapes[gi] = new_shape - - self.shape = None - self.is_nested = True + # newly input is regular tensor + new_shape = incoming.get("shape", None) + if new_shape is None: + raise ValueError("Receiving a regular tensor without 'shape'!") + if self.is_nested: + # we need to update incoming shape into per_sample_shapes + for gi in incoming_global_indexes: + self.per_sample_shapes[gi] = new_shape + else: + if self.is_non_tensor is not None and not self.is_non_tensor: + # original data is also regular tensor + assert self.shape is not None + if self.shape != new_shape: + for gi in self.global_indexes: + self.per_sample_shapes[gi] = self.shape + for gi in incoming_global_indexes: + self.per_sample_shapes[gi] = new_shape + + self.shape = None + self.is_nested = True self.global_indexes.update(incoming_global_indexes) diff --git a/transfer_queue/interface.py b/transfer_queue/interface.py index eca8762e..8920cfaf 100644 --- a/transfer_queue/interface.py +++ b/transfer_queue/interface.py @@ -606,6 +606,61 @@ def kv_clear(keys: list[str] | str, partition_id: str) -> None: tq_client._run_coroutine(async_kv_clear(keys=keys, partition_id=partition_id)) +def kv_update( + key: str, + partition_id: str, + fields: str | list[str], + *, + values: Any | dict[str, Any] | None = None, + parser: Callable[[Any, Any], Any] | None = None, + empty: bool = False, +) -> KVBatchMeta: + """Update fields of an existing key by combining the stored value with a new one. + + A custom ``parser(old, new)`` runs once per sample per field on the SimpleStorage + unit that holds the row, then writes the return value back. ``empty=True`` + stores ``None`` instead and must not be given ``values`` or ``parser``. + ``kv_empty()`` is the same as ``kv_update(..., empty=True)``. + + Args: + key: Existing user-specified key. The key is not created if it is missing. + partition_id: Partition that holds the key. + fields: One field name, or a list of field names, to update. + values: New value for a single field, or ``{field: new}`` for several. + Must be omitted when ``empty=True``. + parser: ``parser(old, new) -> stored``. Required unless ``empty=True``. + Only SimpleStorage executes parsers; KV backends reject the call. + ``old`` is ``None`` when the field has never been written. + empty: If True, store None for each named field. + + Returns: + KVBatchMeta for the updated sample, including every field stored for it. + + Raises: + ValueError: If the key is missing, ``empty=True`` is given values or a + parser, a custom parser is given no values, or ``fields`` / ``values`` + do not line up. + TypeError: If ``parser`` is missing or not callable when ``empty`` is False. + NotImplementedError: If the storage backend is not SimpleStorage. + """ + tq_client = _maybe_create_tq_client() + return tq_client._run_coroutine( + async_kv_update( + key=key, + partition_id=partition_id, + fields=fields, + values=values, + parser=parser, + empty=empty, + ) + ) + + +def kv_empty(key: str, partition_id: str, fields: str | list[str]) -> KVBatchMeta: + """Store None for the named fields of an existing key. Same as kv_update(..., empty=True).""" + return kv_update(key=key, partition_id=partition_id, fields=fields, empty=True) + + # ==================== KV Interface API ==================== async def async_kv_put( key: str, @@ -987,6 +1042,98 @@ async def async_kv_clear(keys: list[str] | str, partition_id: str) -> None: await tq_client.async_clear_samples(batch_meta) +def _normalize_kv_update_args( + fields: str | list[str], + values: Any | dict[str, Any] | None, + parser: Callable[[Any, Any], Any] | None, + empty: bool = False, +) -> tuple[list[str], TensorDict | None, Callable[[Any, Any], Any] | None, bool]: + """Validate kv_update arguments and wrap new values as a one-row TensorDict.""" + if isinstance(fields, str): + field_names = [fields] + elif isinstance(fields, list) and fields and all(isinstance(f, str) for f in fields): + field_names = list(fields) + else: + raise TypeError("fields must be a field name or a non-empty list of field names") + + if empty: + if values is not None: + raise ValueError("kv_update with empty=True must not specify values") + if parser is not None: + raise ValueError("kv_update with empty=True must not specify parser") + return field_names, None, None, True + + if not callable(parser): + raise TypeError("parser must be callable unless empty=True") + if values is None: + raise ValueError("kv_update with a custom parser requires values") + + if isinstance(fields, str): + value_map = {fields: values} + else: + if not isinstance(values, dict): + raise TypeError("values must be a dict mapping field name to new value when updating multiple fields") + missing = [name for name in field_names if name not in values] + extra = [name for name in values if name not in field_names] + if missing or extra: + raise ValueError( + f"fields and values must name the same columns; fields={field_names}, extra={extra}, missing={missing}" + ) + value_map = values + + batch: dict[str, Any] = {} + for field_name, value in value_map.items(): + if isinstance(value, torch.Tensor): + if value.is_nested: + raise ValueError("nested tensors are not supported for single-key kv_update") + batch[field_name] = value.unsqueeze(0) + else: + batch[field_name] = NonTensorStack(value) + return field_names, TensorDict(batch, batch_size=[1]), parser, False + + +async def async_kv_update( + key: str, + partition_id: str, + fields: str | list[str], + *, + values: Any | dict[str, Any] | None = None, + parser: Callable[[Any, Any], Any] | None = None, + empty: bool = False, +) -> KVBatchMeta: + """Asynchronously update fields of an existing key. See ``kv_update``.""" + field_names, value_batch, bound_parser, use_empty = _normalize_kv_update_args(fields, values, parser, empty) + + tq_client = _maybe_create_tq_client() + batch_meta = await tq_client.async_kv_retrieve_meta(keys=[key], partition_id=partition_id, create=False) + + if batch_meta.size == 0: + raise ValueError("keys or partition were not found!") + if batch_meta.size != 1: + raise RuntimeError(f"Retrieved BatchMeta size {batch_meta.size} does not match with input `key` size of 1!") + + batch_meta = await tq_client.async_update( + metadata=batch_meta, + field_names=field_names, + values=value_batch, + parser=bound_parser, + empty=use_empty, + ) + + return KVBatchMeta( + keys=[key], + tags=batch_meta.custom_meta, + partition_id=partition_id, + fields=batch_meta.field_names, + extra_info=batch_meta.extra_info, + ) + + +async def async_kv_empty(key: str, partition_id: str, fields: str | list[str]) -> KVBatchMeta: + """Store None for the named fields of an existing key. Same as async_kv_update(..., empty=True).""" + return await async_kv_update(key=key, partition_id=partition_id, fields=fields, empty=True) + + # ==================== Low-Level Native API ==================== # For low-level API support, please refer to transfer_queue/client.py for details. def get_client(): diff --git a/transfer_queue/metadata.py b/transfer_queue/metadata.py index e89c1531..ab395e07 100644 --- a/transfer_queue/metadata.py +++ b/transfer_queue/metadata.py @@ -461,10 +461,20 @@ def add_fields(self, tensor_dict: TensorDict, set_all_ready: bool = True) -> "Ba if batch_size != self.size: raise ValueError(f"add_fields batch size mismatch: self.size={self.size} vs tensor_dict={batch_size}") - field_schema = extract_field_schema(tensor_dict) + return self.apply_field_schema(extract_field_schema(tensor_dict), set_all_ready=set_all_ready) - for key, value in field_schema.items(): - self.field_schema[key] = value + def apply_field_schema(self, field_schema: dict[str, dict[str, Any]], set_all_ready: bool = True) -> "BatchMeta": + """Merge a field_schema into this batch in place. + + The only place that keeps ``field_names`` and ``is_ready`` in step with + ``field_schema`` and ``production_status``; callers holding a schema from storage + should go through here rather than writing the derived state themselves. + + Args: + field_schema: Per-field metadata to add or overwrite. + set_all_ready (bool): If True, set all production_status to READY_FOR_CONSUME. Default is True. + """ + self.field_schema.update(field_schema) if set_all_ready: self.production_status[:] = 1 diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index 369f800e..310ab823 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -354,6 +354,28 @@ async def clear_data(self, metadata: BatchMeta) -> None: """ raise NotImplementedError("Subclasses must implement clear_data") + async def update_data( + self, + metadata: BatchMeta, + field_names: list[str], + values: TensorDict | None = None, + parser: Callable[[Any, Any], Any] | None = None, + empty: bool = False, + ) -> dict[str, dict[str, Any]]: + """Apply empty or parser(old, new) on the unit that holds each sample. + + Args: + metadata: Samples to update. + field_names: Fields to rewrite. + values: New values as a TensorDict aligned with ``metadata``, or None for empty. + parser: Called per sample per field as ``parser(old, new)``. Ignored when empty. + empty: If True, store None for each named field and ignore values/parser. + + Returns: + field_schema of the stored values, keyed by field name. + """ + raise NotImplementedError(f"{self.__class__.__name__} does not support update_data") + async def save_checkpoint(self, checkpoint_dir: str) -> None: """Save storage state into checkpoint_dir. @@ -784,6 +806,19 @@ async def put_data( per_field_custom_backend_meta, ) + async def update_data( + self, + metadata: BatchMeta, + field_names: list[str], + values: TensorDict | None = None, + parser: Callable[[Any, Any], Any] | None = None, + empty: bool = False, + ) -> dict[str, dict[str, Any]]: + """kv_update is only implemented for SimpleStorage.""" + raise NotImplementedError( + "kv_update is not supported for KV-based backends (MooncakeStore, Yuanrong, RayStore)." + ) + async def get_data(self, metadata: BatchMeta) -> TensorDict: """ Retrieve tensor data from the backend storage. diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index c89d1cfb..cfe9dca6 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -87,6 +87,48 @@ def _describe_unit_state(body: dict[str, Any]) -> str: return f"verdict=unit_serving_again ({' '.join(parts)})" +def _build_update_field_schema( + global_indexes: list[int], + per_unit: list[tuple[list[int], dict[str, dict[str, Any]]]], +) -> dict[str, dict[str, Any]]: + """Assemble one batch-level field_schema from each unit's per-sample description. + + A unit only describes its own rows, so shapes are collected per global index and + then ordered to match ``global_indexes``: the controller reads per_sample_shapes + positionally against the indexes it is notified about. + """ + dtypes: dict[str, Any] = {} + shapes: dict[str, dict[int, tuple]] = defaultdict(dict) + non_tensor: set[str] = set() + for unit_indexes, described in per_unit: + for field_name, description in described.items(): + unit_shapes = description["shapes"] + if unit_shapes is None: + non_tensor.add(field_name) + continue + dtypes[field_name] = description["dtype"] + # Serialization turns the unit's tuples into lists; shapes are compared and hashed here. + shapes[field_name].update(zip(unit_indexes, (tuple(s) for s in unit_shapes), strict=True)) + + field_schema: dict[str, dict[str, Any]] = {} + for field_name in non_tensor | set(shapes): + # A column is only a tensor column when every unit stored tensors for it. + if field_name in non_tensor: + field_schema[field_name] = {"dtype": None, "shape": None, "is_nested": False, "is_non_tensor": True} + continue + ordered = [shapes[field_name][gi] for gi in global_indexes] + is_nested = len(set(ordered)) > 1 + field_schema[field_name] = { + "dtype": dtypes[field_name], + "shape": None if is_nested else ordered[0], + "is_nested": is_nested, + "is_non_tensor": False, + } + if is_nested: + field_schema[field_name]["per_sample_shapes"] = ordered + return field_schema + + _SU_SUBDIR = "simple_storage" _SU_INFO_FILE = "storage_unit_info.json" @@ -481,6 +523,86 @@ async def put_data( field_schema, ) + async def update_data( + self, + metadata: BatchMeta, + field_names: list[str], + values: TensorDict | None = None, + parser: Callable[[Any, Any], Any] | None = None, + empty: bool = False, + ) -> dict[str, dict[str, Any]]: + """Send an update to each hashed storage unit and notify the controller. + + The unit reads the old value, applies empty or parser(old, new), and reports + the shape it stored per sample; this assembles one batch-level field_schema + so production status stays ready. + """ + logger.debug(f"[{self.storage_manager_id}]: receive update_data request, updating {metadata.size} samples.") + + if metadata.size == 0: + return {} + if empty: + if values is not None or parser is not None: + raise ValueError("empty update must not include values or a parser") + else: + if values is None or parser is None: + raise ValueError("update requires values and a parser unless empty is set") + if values.batch_size[0] != metadata.size: + raise ValueError( + f"Batch size of values ({values.batch_size[0]}) does not match metadata size ({metadata.size})" + ) + + routing = self._group_by_hash(metadata.global_indexes) + # Parser-backed updates are not replayed: parser(old, new) reads the stored value, so a + # second attempt after a lost answer would fold the new value in twice. Empty is a + # plain overwrite and stays retryable. + max_attempts = 1 if parser is not None else None + tasks = [] + for su_id, group in routing.items(): + storage_data = None + if values is not None: + storage_data = {f: self._select_by_positions(values[f], group.batch_positions) for f in field_names} + tasks.append( + self._request_with_retry( + "update", + su_id, + f"samples={len(group.global_indexes)} fields={field_names}", + partial( + self._update_to_single_storage_unit, + group.global_indexes, + field_names, + storage_data, + parser, + empty, + target_storage_unit=su_id, + ), + max_attempts=max_attempts, + ) + ) + + try: + described = await asyncio.gather(*tasks) + except Exception as e: + logger.error( + f"[{self.storage_manager_id}]: update_data failed. " + f"partition_id={metadata.partition_ids[0]}, " + f"num_samples={metadata.size}, " + f"num_storage_units={len(routing)}, " + f"error={type(e).__name__}: {e}" + ) + raise + + field_schema = _build_update_field_schema( + metadata.global_indexes, list(zip([g.global_indexes for g in routing.values()], described, strict=True)) + ) + + await self.notify_data_update( + metadata.partition_ids[0], + metadata.global_indexes, + field_schema, + ) + return field_schema + @with_storage_unit_socket async def _put_to_single_storage_unit( self, @@ -538,6 +660,70 @@ 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 + @with_storage_unit_socket + async def _update_to_single_storage_unit( + self, + global_indexes: list[int], + field_names: list[str], + storage_data: dict[str, Any] | None, + parser: Callable[[Any, Any], Any] | None, + empty: bool, + target_storage_unit: str, + socket: zmq.Socket = None, + ) -> dict[str, dict[str, Any]]: + """Send an update to one storage unit and return what it stored, per field.""" + body: dict[str, Any] = { + "global_indexes": global_indexes, + "fields": field_names, + "empty": empty, + } + if not empty: + body["data"] = storage_data + body["parser"] = parser + + request_msg = ZMQMessage.create( + request_type=ZMQRequestType.UPDATE_DATA, # type: ignore[arg-type] + sender_id=self.storage_manager_id, + receiver_id=target_storage_unit, + body=body, + ) + + serialized_bytes = 0 + started = time.perf_counter() + try: + frames = request_msg.serialize() + serialized_bytes = sum(frame_nbytes(frame) or 0 for frame in frames) + await socket.send_multipart(frames, copy=False) + messages = await socket.recv_multipart(copy=False) + response_msg = ZMQMessage.deserialize(messages) + + if response_msg.request_type != ZMQRequestType.UPDATE_DATA_RESPONSE: + raise RuntimeError( + f"Failed to update data on storage unit {target_storage_unit}: " + f"{response_msg.body.get('message', 'Unknown error')}" + ) + log_heavy_operation( + self.storage_manager_id, + "update", + time.perf_counter() - started, + serialized_bytes, + f"to {target_storage_unit} at {self._describe_storage_unit(target_storage_unit)} " + f"samples={len(global_indexes)} fields={field_names}", + ) + return response_msg.body.get("stored_shapes", {}) + except zmq.error.Again as e: + raise StorageUnitTimeout( + f"no answer in {TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT}s during update to storage unit " + f"{target_storage_unit} at {self._describe_storage_unit(target_storage_unit)}; " + f"serialized_mb={serialized_bytes / 2**20:.1f}" + ) from e + except Exception as e: + logger.error( + f"[{self.storage_manager_id}]: Unexpected error during update to storage unit " + f"{target_storage_unit}: {type(e).__name__}: {e}" + ) + raise RuntimeError(f"Error in update to storage unit {target_storage_unit}: {type(e).__name__}: {e}") from e + @staticmethod def _pack_field_values(values: list) -> torch.Tensor | NonTensorStack: """ diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index df28521e..90d69a52 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -38,7 +38,7 @@ from dataclasses import dataclass from pathlib import Path from threading import Event, Thread -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Callable from uuid import uuid4 import numpy as np @@ -110,6 +110,18 @@ class _SSDValueRef: TQ_ACCEPT_PROBE_INTERVAL = float(os.environ.get("TQ_ACCEPT_PROBE_INTERVAL", 0)) +def _describe_stored_values(values: list) -> dict[str, Any]: + """Describe per-sample stored values for the manager. + + A unit only holds part of a batch, so it reports one shape per sample and + leaves ``shapes=None`` for a column it cannot describe as tensors; only the + manager sees the whole batch and can turn this into a field_schema. + """ + if not all(isinstance(v, torch.Tensor) for v in values): + return {"dtype": None, "shapes": None} + return {"dtype": values[0].dtype, "shapes": [tuple(v.shape) for v in values]} + + class StorageKeyNotFoundError(KeyError): """Raised when a requested global index is absent from a storage unit. @@ -215,6 +227,40 @@ def put_data(self, field_data: dict[str, Any], global_indexes: list) -> None: field_dict[key] = val self._active_keys.update(global_indexes) + def apply_update( + self, + global_indexes: list[int], + fields: list[str], + new_data: dict[str, Any] | None, + parser: Callable[[Any, Any], Any] | None, + use_empty: bool, + ) -> dict[str, dict[str, Any]]: + """Read each (field, index), compute the stored value, then write once. + + Missing fields yield ``old=None``. All parser calls finish before any write, + so a raise leaves storage unchanged. + + Returns: + Per-sample description of the stored values, keyed by field name. + """ + computed: dict[str, list] = {} + if use_empty: + computed = {field: [None] * len(global_indexes) for field in fields} + else: + if parser is None or new_data is None: + raise TypeError("apply_update requires a parser and new values unless use_empty is set") + for field in fields: + # Read through get_data so SSD-offloaded values reach the parser decoded, + # not as file references; indexes never written keep old=None. + stored = self.field_data.get(field, {}) + present = [idx for idx in global_indexes if idx in stored] + old_values = dict(zip(present, self.get_data([field], present)[field], strict=True)) if present else {} + computed[field] = [ + parser(old_values.get(idx), new_data[field][i]) for i, idx in enumerate(global_indexes) + ] + self.put_data(computed, global_indexes) + return {field: _describe_stored_values(values) for field, values in computed.items()} + def clear(self, keys: list[int]) -> None: """Remove data at given global index keys, immediately freeing memory. @@ -993,6 +1039,9 @@ def _process_one_worker_request(self, worker_socket: zmq.Socket, monitor: Any) - if operation == ZMQRequestType.PUT_DATA: # type: ignore[arg-type] with monitor.measure(op_type="PUT_DATA"): response_msg = self._handle_put(request_msg) + elif operation == ZMQRequestType.UPDATE_DATA: # type: ignore[arg-type] + with monitor.measure(op_type="UPDATE_DATA"): + response_msg = self._handle_update(request_msg) elif operation == ZMQRequestType.GET_DATA: # type: ignore[arg-type] with monitor.measure(op_type="GET_DATA"): response_msg = self._handle_get(request_msg) @@ -1120,6 +1169,40 @@ def _handle_put(self, data_parts: ZMQMessage) -> ZMQMessage: }, ) + def _handle_update(self, data_parts: ZMQMessage) -> ZMQMessage: + """Read each field, apply empty or parser(old, new), write once, describe what was stored.""" + try: + global_indexes = data_parts.body["global_indexes"] + fields = data_parts.body["fields"] + use_empty = bool(data_parts.body.get("empty", False)) + parser = data_parts.body.get("parser") + new_data = data_parts.body.get("data") + + with limit_pytorch_auto_parallel_threads( + target_num_threads=TQ_NUM_THREADS, info=f"[{self.storage_unit_id}] _handle_update" + ): + if not use_empty: + if not callable(parser): + raise TypeError(f"update parser must be callable, got {type(parser).__name__}") + if not isinstance(new_data, dict): + raise TypeError("update data must be a dict of new field values") + stored_shapes = self.storage_data.apply_update(global_indexes, fields, new_data, parser, use_empty) + + return ZMQMessage.create( + request_type=ZMQRequestType.UPDATE_DATA_RESPONSE, # type: ignore[arg-type] + sender_id=self.storage_unit_id, + body={"stored_shapes": stored_shapes}, + ) + except Exception as e: + return ZMQMessage.create( + request_type=ZMQRequestType.UPDATE_ERROR, # type: ignore[arg-type] + sender_id=self.storage_unit_id, + body={ + "message": f"Failed to update data in storage unit id " + f"#{self.storage_unit_id}, detail error message: {str(e)}" + }, + ) + def _handle_get(self, data_parts: ZMQMessage) -> ZMQMessage: """ Handle get request, return data from storage unit. @@ -1238,7 +1321,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", "UPDATE_DATA", "GET_DATA", "CLEAR_DATA"): try: hist = self._metrics.request_duration.labels(op_type=op_type) counter = self._metrics.request_total.labels(op_type=op_type) diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index a94cd3ab..80742d8b 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -69,8 +69,10 @@ class ZMQRequestType(ExplicitEnum): # DATA_OPERATION GET_DATA = "GET" PUT_DATA = "PUT" + UPDATE_DATA = "UPDATE" GET_DATA_RESPONSE = "GET_DATA_RESPONSE" PUT_DATA_RESPONSE = "PUT_DATA_RESPONSE" + UPDATE_DATA_RESPONSE = "UPDATE_DATA_RESPONSE" CLEAR_DATA = "CLEAR_DATA" CLEAR_DATA_RESPONSE = "CLEAR_DATA_RESPONSE" @@ -78,6 +80,7 @@ class ZMQRequestType(ExplicitEnum): PUT_GET_ERROR = "PUT_GET_ERROR" PUT_ERROR = "PUT_ERROR" GET_ERROR = "GET_ERROR" + UPDATE_ERROR = "UPDATE_ERROR" CLEAR_DATA_ERROR = "CLEAR_DATA_ERROR" # META_OPERATION diff --git a/tutorial/02_kv_interface.py b/tutorial/02_kv_interface.py index 778739c2..601f80d8 100644 --- a/tutorial/02_kv_interface.py +++ b/tutorial/02_kv_interface.py @@ -56,10 +56,10 @@ def demonstrate_kv_api(): """ Demonstrate the Key-Value (KV) semantic API: - kv_put & kv_batch_put -> kv_list -> kv_batch_get -> kv_clear + kv_put & kv_batch_put -> kv_update -> kv_list -> kv_batch_get -> kv_clear """ print("=" * 80) - print("Key-Value Semantic API Demo: kv_put/kv_batch_put → kv_list → kv_batch_get → kv_clear") + print("Key-Value Semantic API Demo: kv_put/kv_batch_put → kv_update → kv_list → kv_batch_get → kv_clear") print("=" * 80) # Step 1: Put a single key-value pair with kv_put @@ -153,8 +153,22 @@ def demonstrate_kv_api(): tq.kv_put(key=key_for_update_tags, partition_id=partition_id, fields=None, tag=tag_update) print(f" ✓ Update success: Samples '0_0' now has tag as {tag_update}.") - # Step 5: List all keys and tags in a partition - print("\n[Step 5] Listing all keys and tags in partition...") + # Step 5: Concatenate a field with kv_update, then empty it + print("\n[Step 5] Updating a field with kv_update (concat, then empty)...") + tq.kv_put(key=key, partition_id=partition_id, fields={"scratch": torch.tensor([1, 2])}) + tq.kv_update( + key=key, + partition_id=partition_id, + fields="scratch", + values=torch.tensor([3, 4]), + parser=lambda old, new: torch.cat([old, new]), + ) + print(" ✓ kv_update concat: scratch of '0_0' is now [1, 2, 3, 4].") + tq.kv_empty(key=key, partition_id=partition_id, fields="scratch") + print(" ✓ tq.kv_empty: scratch of '0_0' is stored as None (key remains).") + + # Step 6: List all keys and tags in a partition + print("\n[Step 6] Listing all keys and tags in partition...") partition_info = tq.kv_list() print(f" Found {len(partition_info.keys())} partitions: '{list(partition_info.keys())}'") @@ -162,16 +176,16 @@ def demonstrate_kv_api(): for k, t in keys_and_tags.items(): print(f"Partition: {pid}, - key='{k}' | tag={t}") - # Step 6: Retrieve specific fields using kv_batch_get - print("\n[Step 6] Retrieving specific fields (Column) with kv_batch_get...") + # Step 7: Retrieve specific fields using kv_batch_get + print("\n[Step 7] Retrieving specific fields (Column) with kv_batch_get...") print(" Fetching only 'input_ids' to save bandwidth (ignoring 'attention_mask' and 'response').") all_keys = list(partition_info[partition_id].keys()) retrieved_input_ids = tq.kv_batch_get(keys=all_keys, partition_id=partition_id, select_fields="input_ids") print(f" ✓ Successfully retrieved only {list(retrieved_input_ids.keys())} field for all samples.") - # # Step 7: Retrieve all fields using kv_batch_get - print("\n[Step 7] Retrieving all fields with kv_batch_get...") + # Step 8: Retrieve all fields using kv_batch_get + print("\n[Step 8] Retrieving all fields with kv_batch_get...") retrieved_all = tq.kv_batch_get(keys=all_keys, partition_id=partition_id) print(f" Retrieved all fields for {all_keys}:") print(f" Fields: {list(retrieved_all.keys())}") @@ -179,8 +193,8 @@ def demonstrate_kv_api(): f" Note: We cannot retrieve fields {list(response_batch.keys())}, since they only available in {append_keys}" ) - # Step 8: Clear specific keys - print("\n[Step 8] Clearing keys from partition...") + # Step 9: Clear specific keys + print("\n[Step 9] Clearing keys from partition...") keys_to_clear = all_keys[:2] # Delete the first 2 keys tq.kv_clear(keys=keys_to_clear, partition_id=partition_id) print(f" ✓ Cleared keys: {keys_to_clear}") @@ -202,9 +216,10 @@ def main(): Key Methods: 1. (async_)kv_put - Insert/Update a multi-column sample by key, with optional metadata tag 2. (async_)kv_batch_put - Put multiple key-value pairs efficiently in batch - 3. (async_)kv_batch_get - Retrieve samples (by keys), supporting column selection (by fields) - 4. (async_)kv_list - List keys and tags (metadata) in a partition - 5. (async_)kv_clear - Remove key-value pairs from storage + 3. (async_)kv_update - Rewrite fields with parser(old, new), or empty=True / tq.kv_empty to store None + 4. (async_)kv_batch_get - Retrieve samples (by keys), supporting column selection (by fields) + 5. (async_)kv_list - List keys and tags (metadata) in a partition + 6. (async_)kv_clear - Remove key-value pairs from storage Key Features: ✓ Redis-style Semantics - Familiar KV interface (Put/Get/List) for zero learning curve From bbd1372f6484a75dc87f1a3e8df4282c9f88eebe Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 23 Sep 2026 14:55:08 +0800 Subject: [PATCH 3/5] [fix] Name the unsupported backend when kv_update is rejected kv_update needs to run parser(old, new) where the value lives, which only SimpleStorage can do. The refusal used to list backend names in a string that would go stale as soon as another KV backend was added, and it never told the caller which backend they were actually on. Report the concrete manager class instead, and cover Mooncake, Yuanrong and Ray individually rather than only the shared base class. Signed-off-by: OutstanderWang --- tests/test_kv_update.py | 18 +++++++++++++++--- transfer_queue/storage/managers/base.py | 8 ++++++-- 2 files changed, 21 insertions(+), 5 deletions(-) diff --git a/tests/test_kv_update.py b/tests/test_kv_update.py index aed9d4a8..6fd4e741 100644 --- a/tests/test_kv_update.py +++ b/tests/test_kv_update.py @@ -23,7 +23,10 @@ from transfer_queue.interface import _normalize_kv_update_args from transfer_queue.metadata import BatchMeta from transfer_queue.storage.managers.base import KVStorageManager +from transfer_queue.storage.managers.mooncake_manager import MooncakeStorageManager +from transfer_queue.storage.managers.ray_storage_manager import RayStorageManager from transfer_queue.storage.managers.simple_storage_manager import _build_update_field_schema +from transfer_queue.storage.managers.yuanrong_manager import YuanrongStorageManager from transfer_queue.storage.simple_storage import HybridStorageUnitData, StorageUnitData @@ -203,16 +206,25 @@ def test_empty_marks_the_controller_field_non_tensor(): @pytest.mark.asyncio +@pytest.mark.parametrize( + "manager_cls", [MooncakeStorageManager, YuanrongStorageManager, RayStorageManager, KVStorageManager] +) @patch("transfer_queue.storage.managers.base.StorageClientFactory.create") @patch.object(KVStorageManager, "_connect_to_controller", lambda self: None) -async def test_kv_backend_rejects_update(mock_create): +async def test_kv_backends_reject_update(mock_create, manager_cls): + """Every KV backend must refuse kv_update by name; only SimpleStorage implements it.""" mock_create.return_value = MagicMock() - manager = KVStorageManager(controller_info=MagicMock(), config={"client_name": "YuanrongStorageClient"}) + # Each manager validates its own config before reaching update_data. + config = { + KVStorageManager: {"client_name": "YuanrongStorageClient"}, + YuanrongStorageManager: {"worker_port": 31501}, + }.get(manager_cls, {}) + manager = manager_cls(controller_info=MagicMock(), config=config) meta = BatchMeta( global_indexes=[0], partition_ids=["p"], field_schema={"x": {"dtype": torch.int64, "shape": (1,), "is_nested": False, "is_non_tensor": False}}, production_status=np.ones(1, dtype=np.int8), ) - with pytest.raises(NotImplementedError, match="kv_update is not supported for KV-based backends"): + with pytest.raises(NotImplementedError, match=f"not supported by {manager_cls.__name__}"): await manager.update_data(meta, ["x"], empty=True) diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index 310ab823..97b7aa21 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -814,9 +814,13 @@ async def update_data( parser: Callable[[Any, Any], Any] | None = None, empty: bool = False, ) -> dict[str, dict[str, Any]]: - """kv_update is only implemented for SimpleStorage.""" + """kv_update is only implemented for SimpleStorage. + + A KV backend stores each sample-field under its own key and offers no hook to run + parser(old, new) where the value lives, so read-modify-write cannot be made atomic. + """ raise NotImplementedError( - "kv_update is not supported for KV-based backends (MooncakeStore, Yuanrong, RayStore)." + f"kv_update is not supported by {type(self).__name__}; it requires the SimpleStorage backend." ) async def get_data(self, metadata: BatchMeta) -> TensorDict: From c72af0b7f8b6345139b86bce993fb774628565d4 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 23 Sep 2026 14:55:20 +0800 Subject: [PATCH 4/5] [test] Add a cluster end-to-end example for kv_update The unit tests stub out the transport, so they cannot show that a parser really travels to the storage unit or that a batch spanning several units reports its per-sample shapes back to the controller in the right order. This script exercises both against a live Ray cluster, alongside kv_empty, multi-field updates, non-tensor payloads, parser failure atomicity, argument validation, the async entry points, and updates issued from remote workers. It refuses to run when a controller already exists, so it can never attach to somebody else's deployment and clear their keys on the way out. Signed-off-by: OutstanderWang --- recipe/kv_update_e2e.py | 389 ++++++++++++++++++++++++++++++++++++++++ 1 file changed, 389 insertions(+) create mode 100644 recipe/kv_update_e2e.py diff --git a/recipe/kv_update_e2e.py b/recipe/kv_update_e2e.py new file mode 100644 index 00000000..c25c6bc9 --- /dev/null +++ b/recipe/kv_update_e2e.py @@ -0,0 +1,389 @@ +# 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 exercise of kv_update / kv_empty against a real Ray cluster. + +kv_update rewrites fields of an existing key by running ``parser(old, new)`` on the +SimpleStorage unit that holds the row, so nothing here works unless the parser +actually travels to the unit and the resulting schema travels back to the +controller. Every check below reads its result back through the public API, and +the multi-unit ones also inspect controller metadata, which is the part a +data-only assertion cannot see. + +Run it as a Ray job so the driver lives on the cluster: + + ray job submit --address http://: --working-dir . \ + -- python recipe/kv_update_e2e.py + +Locally, ``python recipe/kv_update_e2e.py`` starts its own cluster instead. +""" + +import asyncio +import os +import traceback + +# Disable Ray's cross-worker log deduplication before importing Ray itself, +# otherwise worker-side prints get folded into "[repeated Nx across cluster]". +os.environ.setdefault("RAY_DEDUP_LOGS", "0") + +import ray +import torch +from omegaconf import OmegaConf +from tensordict import TensorDict + +import transfer_queue as tq + +PARTITION = "kv_update_e2e" + +# Several units so the hash routing actually splits a batch; the cross-unit checks +# below are meaningless with a single unit. +CONFIG = { + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": {"total_storage_size": 512, "num_data_storage_units": 4}, + }, +} + + +def concat(old, new): + """Append new tokens to what is already stored, keeping the field name.""" + return new if old is None else torch.cat([old, new]) + + +# ==================== checks ==================== + + +def check_concat_keeps_field_name(): + """The rollout case: prompt_ids in place, response_ids appended, one field.""" + key = "concat" + tq.kv_put(key=key, partition_id=PARTITION, fields={"sequence_ids": torch.tensor([10, 11, 12])}) + + tq.kv_update( + key=key, + partition_id=PARTITION, + fields="sequence_ids", + values=torch.tensor([20, 21]), + parser=concat, + ) + + got = tq.kv_batch_get(keys=key, partition_id=PARTITION) + assert torch.equal(got["sequence_ids"][0], torch.tensor([10, 11, 12, 20, 21])), got["sequence_ids"][0] + assert list(got.keys()) == ["sequence_ids"], f"update must not add columns: {list(got.keys())}" + + +def check_parser_sees_none_for_a_new_field(): + """A field the key never held reaches the parser as old=None. + + The parser runs in the storage unit's process, so what it observed can only + come back inside the value it returns; a closure variable would stay untouched + here in the driver. + """ + key = "fresh_field" + tq.kv_put(key=key, partition_id=PARTITION, fields={"anchor": torch.tensor([1])}) + + def report_old(old, new): + return {"old_was_none": old is None, "new": new.tolist()} + + tq.kv_update(key=key, partition_id=PARTITION, fields="logprobs", values=torch.tensor([7, 8]), parser=report_old) + + got = tq.kv_batch_get(keys=key, partition_id=PARTITION, select_fields="logprobs") + assert got["logprobs"][0] == {"old_was_none": True, "new": [7, 8]}, got["logprobs"][0] + + +def check_multiple_fields_in_one_call(): + key = "multi_field" + tq.kv_put( + key=key, + partition_id=PARTITION, + fields={"a": torch.tensor([1, 2]), "b": torch.tensor([10])}, + ) + + tq.kv_update( + key=key, + partition_id=PARTITION, + fields=["a", "b"], + values={"a": torch.tensor([3]), "b": torch.tensor([20, 30])}, + parser=concat, + ) + + got = tq.kv_batch_get(keys=key, partition_id=PARTITION) + assert torch.equal(got["a"][0], torch.tensor([1, 2, 3])), got["a"][0] + assert torch.equal(got["b"][0], torch.tensor([10, 20, 30])), got["b"][0] + + +def check_non_tensor_payload(): + """Parsers are not tensor-only: a plain object round-trips the same way.""" + key = "non_tensor" + tq.kv_put(key=key, partition_id=PARTITION, fields={"stats": {"step": 1}}) + + tq.kv_update( + key=key, + partition_id=PARTITION, + fields="stats", + values={"reward": 0.5}, + parser=lambda old, new: {**old, **new}, + ) + + got = tq.kv_batch_get(keys=key, partition_id=PARTITION, select_fields="stats") + assert got["stats"][0] == {"step": 1, "reward": 0.5}, got["stats"][0] + + +def check_empty_keeps_the_key_and_updates_the_controller(): + """kv_empty stores None; the key stays produced and stops being a tensor column.""" + key = "to_empty" + tq.kv_put(key=key, partition_id=PARTITION, fields={"tokens": torch.tensor([1, 2, 3])}) + + tq.kv_empty(key=key, partition_id=PARTITION, fields="tokens") + + got = tq.kv_batch_get(keys=key, partition_id=PARTITION, select_fields="tokens") + assert got["tokens"][0] is None, got["tokens"][0] + assert key in tq.kv_list(partition_id=PARTITION)[PARTITION], "kv_empty must not delete the key" + + partition = partition_snapshot() + col = partition.field_name_mapping["tokens"] + row = partition.keys_mapping[key] + assert partition.production_status[row, col] == 1, "an emptied field stays produced" + assert partition.field_metadata["tokens"].is_non_tensor is True, "controller still calls the column a tensor" + + +def check_failed_parser_leaves_the_row_unchanged(): + """The unit computes every value before writing, so a raise is a no-op.""" + key = "atomic" + tq.kv_put(key=key, partition_id=PARTITION, fields={"tokens": torch.tensor([1, 2, 3])}) + + def boom(old, new): + raise RuntimeError("parser failed on purpose") + + try: + tq.kv_update(key=key, partition_id=PARTITION, fields="tokens", values=torch.tensor([4]), parser=boom) + except Exception as e: + assert "parser failed on purpose" in str(e), f"unexpected error: {e}" + else: + raise AssertionError("a raising parser must not succeed") + + got = tq.kv_batch_get(keys=key, partition_id=PARTITION, select_fields="tokens") + assert torch.equal(got["tokens"][0], torch.tensor([1, 2, 3])), got["tokens"][0] + + +def check_rejected_arguments(): + key = "rejects" + tq.kv_put(key=key, partition_id=PARTITION, fields={"tokens": torch.tensor([1])}) + + def expect(exc_type, match, fn): + try: + fn() + except exc_type as e: + assert match in str(e), f"expected {match!r} in {e!r}" + else: + raise AssertionError(f"expected {exc_type.__name__} containing {match!r}") + + expect( + ValueError, + "must not specify values", + lambda: tq.kv_update(key=key, partition_id=PARTITION, fields="tokens", values=torch.tensor([1]), empty=True), + ) + expect( + ValueError, + "must not specify parser", + lambda: tq.kv_update(key=key, partition_id=PARTITION, fields="tokens", parser=concat, empty=True), + ) + expect( + TypeError, + "parser must be callable", + lambda: tq.kv_update(key=key, partition_id=PARTITION, fields="tokens", values=torch.tensor([1])), + ) + expect( + ValueError, + "requires values", + lambda: tq.kv_update(key=key, partition_id=PARTITION, fields="tokens", parser=concat), + ) + expect( + ValueError, + "same columns", + lambda: tq.kv_update( + key=key, partition_id=PARTITION, fields=["tokens", "other"], values={"tokens": 1}, parser=concat + ), + ) + expect( + ValueError, + "not found", + lambda: tq.kv_update( + key="never_written", partition_id=PARTITION, fields="tokens", values=torch.tensor([1]), parser=concat + ), + ) + + +def check_async_variants(): + key = "async_key" + tq.kv_put(key=key, partition_id=PARTITION, fields={"tokens": torch.tensor([1, 2])}) + + asyncio.run( + tq.async_kv_update(key=key, partition_id=PARTITION, fields="tokens", values=torch.tensor([3]), parser=concat) + ) + got = tq.kv_batch_get(keys=key, partition_id=PARTITION, select_fields="tokens") + assert torch.equal(got["tokens"][0], torch.tensor([1, 2, 3])), got["tokens"][0] + + asyncio.run(tq.async_kv_empty(key=key, partition_id=PARTITION, fields="tokens")) + got = tq.kv_batch_get(keys=key, partition_id=PARTITION, select_fields="tokens") + assert got["tokens"][0] is None, got["tokens"][0] + + +def check_multi_sample_update_across_units(): + """One update spanning several units must report per-sample shapes in batch order. + + Each key starts with a different length, so concatenating turns the column + nested. The controller's per_sample_shapes is the only place a mis-ordered or + overwritten merge shows up; the payload alone would still look correct. + """ + keys = [f"span_{i}" for i in range(8)] + lengths = [i + 1 for i in range(len(keys))] + for key, length in zip(keys, lengths, strict=True): + tq.kv_put(key=key, partition_id=PARTITION, fields={"tokens": torch.ones(length, dtype=torch.int64)}) + + client = tq.get_client() + metadata = client.kv_retrieve_meta(keys=keys, partition_id=PARTITION, create=False) + # Same addition for every row, so this holds whichever order the metadata came back in. + values = TensorDict({"tokens": torch.full((len(keys), 1), 9, dtype=torch.int64)}, batch_size=[len(keys)]) + client.update(metadata, ["tokens"], values=values, parser=concat) + + got = tq.kv_batch_get(keys=keys, partition_id=PARTITION, select_fields="tokens") + for i, (key, length) in enumerate(zip(keys, lengths, strict=True)): + row = got["tokens"][i] + assert row.shape == (length + 1,), f"{key}: expected len {length + 1}, got {tuple(row.shape)}" + assert row[-1].item() == 9, f"{key}: appended value missing" + + partition = partition_snapshot() + tokens = partition.field_metadata["tokens"] + assert tokens.is_nested is True, "mixed lengths must promote the column to nested" + for key, length in zip(keys, lengths, strict=True): + recorded = tuple(tokens.per_sample_shapes[partition.keys_mapping[key]]) + assert recorded == (length + 1,), f"{key}: controller recorded {recorded}, expected {(length + 1,)}" + + +@ray.remote(num_cpus=0) +class Updater: + """A worker that attaches to the running TransferQueue and updates its own key.""" + + def __init__(self): + tq.init() + + def run(self, key: str, addition: int) -> str: + tq.kv_update( + key=key, + partition_id=PARTITION, + fields="tokens", + values=torch.tensor([addition]), + parser=concat, + ) + return key + + +def check_updates_from_remote_workers(): + """Updates issued by other processes on the cluster, not just the driver.""" + keys = [f"worker_{i}" for i in range(4)] + for i, key in enumerate(keys): + tq.kv_put(key=key, partition_id=PARTITION, fields={"tokens": torch.tensor([i])}) + + updaters = [Updater.remote() for _ in keys] + ray.get([u.run.remote(key, 100 + i) for i, (u, key) in enumerate(zip(updaters, keys, strict=True))]) + for u in updaters: + ray.kill(u) + + got = tq.kv_batch_get(keys=keys, partition_id=PARTITION, select_fields="tokens") + for i, key in enumerate(keys): + expected = torch.tensor([i, 100 + i]) + assert torch.equal(got["tokens"][i], expected), f"{key}: {got['tokens'][i]} != {expected}" + + +# ==================== harness ==================== + + +def partition_snapshot(): + controller = ray.get_actor("TransferQueueController", namespace="transfer_queue") + return ray.get(controller.get_partition_snapshot.remote(PARTITION)) + + +CHECKS = [ + ("concat keeps the field name", check_concat_keeps_field_name), + ("parser sees old=None for a new field", check_parser_sees_none_for_a_new_field), + ("several fields in one call", check_multiple_fields_in_one_call), + ("non-tensor payload", check_non_tensor_payload), + ("kv_empty keeps the key, updates the controller", check_empty_keeps_the_key_and_updates_the_controller), + ("a failed parser leaves the row unchanged", check_failed_parser_leaves_the_row_unchanged), + ("rejected arguments", check_rejected_arguments), + ("async_kv_update / async_kv_empty", check_async_variants), + ("multi-sample update across units", check_multi_sample_update_across_units), + ("updates from remote workers", check_updates_from_remote_workers), +] + + +def refuse_if_already_deployed() -> None: + """Never attach to a TransferQueue somebody else is using. + + tq.init() silently joins an existing controller by name, so on a shared cluster + this script would otherwise run its checks against a live deployment and clear + its keys on the way out. + """ + try: + ray.get_actor("TransferQueueController", namespace="transfer_queue") + except ValueError: + return + raise SystemExit( + "A TransferQueueController is already running on this cluster. This script needs a " + "deployment of its own; stop the owning job or run it on an idle cluster." + ) + + +def main() -> int: + if not ray.is_initialized(): + ray.init(namespace="transfer_queue") + + print("=" * 78) + print("kv_update end-to-end checks") + print(f"cluster: {ray.get_runtime_context().gcs_address} nodes: {len(ray.nodes())}") + print("=" * 78) + + refuse_if_already_deployed() + tq.init(OmegaConf.create(CONFIG)) + failures = [] + try: + for name, check in CHECKS: + try: + check() + except Exception as e: + failures.append(name) + print(f" FAIL {name}: {type(e).__name__}: {e}") + traceback.print_exc() + else: + print(f" ok {name}") + finally: + # Each check owns its keys; clear between them so one failure cannot cascade. + leftover = list(tq.kv_list(partition_id=PARTITION).get(PARTITION, {})) + if leftover: + tq.kv_clear(keys=leftover, partition_id=PARTITION) + finally: + tq.close() + + print("=" * 78) + print(f"{len(CHECKS) - len(failures)}/{len(CHECKS)} checks passed") + if failures: + print("failed: " + ", ".join(failures)) + print("=" * 78) + return 1 if failures else 0 + + +if __name__ == "__main__": + raise SystemExit(main()) From caf3753011a72d1e579f5e8cd98c7cb2f7ea83b1 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Sat, 3 Oct 2026 13:32:20 +0800 Subject: [PATCH 5/5] [fix] Make SSD-offloaded put all-or-nothing across fields With SSD offload on, put_data wrote, committed and cleaned up one field at a time. If writing a later field's files failed (for example a full disk), the earlier fields were already replaced and their old files deleted, yet the caller got an error. A plain put converges on retry, but kv_update parsers such as concat are not idempotent and the manager does not retry updates, so retrying applied the earlier fields twice. Write every field's files first and only then swap all references, update the SSD accounting and delete superseded files; any failure removes just the files this call wrote. The base put_data now checks every field's length before writing any, so the single commit cannot fail halfway either. apply_update's "a raise leaves storage unchanged" now holds for SSD write errors as well as parser errors. Signed-off-by: OutstanderWang --- tests/test_kv_update.py | 34 ++++++++++++++ tests/test_simple_storage_unit.py | 34 ++++++++++++++ transfer_queue/storage/simple_storage.py | 59 ++++++++++++++---------- 3 files changed, 103 insertions(+), 24 deletions(-) diff --git a/tests/test_kv_update.py b/tests/test_kv_update.py index 6fd4e741..0ab2f83f 100644 --- a/tests/test_kv_update.py +++ b/tests/test_kv_update.py @@ -13,6 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import errno from unittest.mock import MagicMock, patch import numpy as np @@ -143,6 +144,39 @@ def test_apply_update_decodes_ssd_offloaded_old_value(tmp_path): assert data.ssd_active_values == 1, "the replaced SSD file must be released, not leaked" +def test_apply_update_is_atomic_when_a_later_field_fails_to_reach_ssd(tmp_path, monkeypatch): + """Concat is not idempotent, so a partly applied update would double-apply on retry.""" + data = HybridStorageUnitData( + storage_size=4, threshold_bytes=64, ssd_path=str(tmp_path), run_id="run", unit_id="unit" + ) + try: + prompt, mask = torch.arange(32), torch.ones(32, dtype=torch.int64) + data.put_data({"tokens": [prompt], "mask": [mask]}, [0]) + old_files = set(tmp_path.rglob("*.bin")) + + write_values = data._ssd_store.write_values + calls = [] + + def fail_second_field(encoded_values): + calls.append(encoded_values) + if len(calls) == 2: + raise OSError(errno.ENOSPC, "No space left on device") + return write_values(encoded_values) + + monkeypatch.setattr(data._ssd_store, "write_values", fail_second_field) + new_data = {"tokens": [torch.tensor([99])], "mask": [torch.tensor([1])]} + with pytest.raises(OSError, match="No space left"): + data.apply_update([0], ["tokens", "mask"], new_data, _concat, False) + + stored = data.get_data(["tokens", "mask"], [0]) + assert torch.equal(stored["tokens"][0], prompt) + assert torch.equal(stored["mask"][0], mask) + assert set(tmp_path.rglob("*.bin")) == old_files + assert (data.ssd_active_values, data.ssd_active_bytes) == (2, prompt.nbytes + mask.nbytes) + finally: + data.close() + + def test_build_update_field_schema_orders_shapes_across_units(): """Units describe only their own rows; the batch schema must follow metadata order.""" described = _build_update_field_schema( diff --git a/tests/test_simple_storage_unit.py b/tests/test_simple_storage_unit.py index 99ceb4ec..fe52f175 100644 --- a/tests/test_simple_storage_unit.py +++ b/tests/test_simple_storage_unit.py @@ -28,6 +28,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import errno import pickle import time from pathlib import Path @@ -603,6 +604,39 @@ def fail_write(_fd, _data): storage.close() +def test_failed_ssd_write_on_second_field_leaves_every_field_unchanged(tmp_path, monkeypatch): + storage = HybridStorageUnitData( + storage_size=10, + threshold_bytes=64, + ssd_path=str(tmp_path), + run_id="test-run", + unit_id="test-unit", + ) + try: + storage.put_data({"a": [b"a" * 100], "b": [b"b" * 100]}, [1]) + old_files = set(tmp_path.rglob("*.bin")) + + write_values = storage._ssd_store.write_values + calls = [] + + def fail_second_field(encoded_values): + calls.append(encoded_values) + if len(calls) == 2: + raise OSError(errno.ENOSPC, "No space left on device") + return write_values(encoded_values) + + monkeypatch.setattr(storage._ssd_store, "write_values", fail_second_field) + with pytest.raises(OSError, match="No space left"): + storage.put_data({"a": [b"A" * 100], "b": [b"B" * 100]}, [1]) + + assert storage.get_data(["a", "b"], [1]) == {"a": [b"a" * 100], "b": [b"b" * 100]} + assert set(tmp_path.rglob("*.bin")) == old_files + assert storage.ssd_active_values == 2 + assert storage.ssd_active_bytes == 200 + finally: + storage.close() + + @pytest.mark.parametrize("threshold_bytes", [0, -1]) def test_hybrid_storage_requires_positive_threshold(tmp_path, threshold_bytes): with pytest.raises(ValueError, match="threshold must be greater than zero"): diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index 90d69a52..5b77dd21 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -220,6 +220,7 @@ def put_data(self, field_data: dict[str, Any], global_indexes: list) -> None: f"StorageUnitData put_data: field '{f}' values length {len(values)} " f"!= global_indexes length {len(global_indexes)}, length mismatch" ) + for f, values in field_data.items(): if f not in self.field_data: self.field_data[f] = {} field_dict = self.field_data[f] @@ -237,8 +238,9 @@ def apply_update( ) -> dict[str, dict[str, Any]]: """Read each (field, index), compute the stored value, then write once. - Missing fields yield ``old=None``. All parser calls finish before any write, - so a raise leaves storage unchanged. + Missing fields yield ``old=None``. All parser calls finish before one + ``put_data`` that writes every field or none, so a raise from a parser or + an SSD write leaves storage unchanged. Returns: Per-sample description of the stored values, keyed by field name. @@ -595,38 +597,47 @@ def put_data( unique_global_indexes = set(global_indexes) has_duplicate_indexes = len(unique_global_indexes) != len(global_indexes) - for field, values in field_data.items(): - prepared_values, entries, fallback_values = self._prepare_field_values(values) - + old_ssd_values = [] + for field in field_data: stored_field = self.field_data.get(field, {}) - old_ssd_values = [] for global_index in unique_global_indexes: old_value = stored_field.get(global_index) if isinstance(old_value, _SSDValueRef): old_ssd_values.append(old_value) - try: - super().put_data({field: prepared_values}, global_indexes) - except Exception: - for entry in entries: - self._ssd_store.unlink(entry) - raise - - obsolete_ssd_values = old_ssd_values - if has_duplicate_indexes: - retained_ssd_paths = set() + + # Write every field's files before replacing any value: kv_update parsers are not + # idempotent, so a failure on a later field must not leave earlier fields committed. + prepared_data = {} + entries: list[_SSDValueRef] = [] + fallback_values = 0 + try: + for field, values in field_data.items(): + prepared_data[field], field_entries, field_fallback_values = self._prepare_field_values(values) + entries.extend(field_entries) + fallback_values += field_fallback_values + super().put_data(prepared_data, global_indexes) + except Exception: + for entry in entries: + self._ssd_store.unlink(entry) + raise + + obsolete_ssd_values = old_ssd_values + if has_duplicate_indexes: + retained_ssd_paths = set() + for field in field_data: for global_index in unique_global_indexes: retained_value = self.field_data[field][global_index] if isinstance(retained_value, _SSDValueRef): retained_ssd_paths.add(retained_value.path) - obsolete_ssd_values.extend(entry for entry in entries if entry.path not in retained_ssd_paths) + obsolete_ssd_values.extend(entry for entry in entries if entry.path not in retained_ssd_paths) - self._ssd_active_values += len(entries) - len(obsolete_ssd_values) - self._ssd_active_bytes += sum(entry.size_bytes for entry in entries) - sum( - value.size_bytes for value in obsolete_ssd_values - ) - self._ssd_fallback_values_total += fallback_values - for obsolete_value in obsolete_ssd_values: - self._ssd_store.unlink(obsolete_value) + self._ssd_active_values += len(entries) - len(obsolete_ssd_values) + self._ssd_active_bytes += sum(entry.size_bytes for entry in entries) - sum( + value.size_bytes for value in obsolete_ssd_values + ) + self._ssd_fallback_values_total += fallback_values + for obsolete_value in obsolete_ssd_values: + self._ssd_store.unlink(obsolete_value) def get_data(self, fields: list[str], global_indexes: list) -> dict[str, list]: """Read mixed memory- and SSD-backed samples in request order."""