From 7bc6cf32da6879a3a82c4e75ec79cd32da904710 Mon Sep 17 00:00:00 2001 From: Ginray Date: Wed, 2 Sep 2026 18:22:13 +0800 Subject: [PATCH] feat: add NIXL-UCX payload transfer Signed-off-by: Ginray --- README.md | 5 +- docs/nixl_ucx_payload.md | 188 ++++++ tests/test_payload_transfer.py | 146 +++++ tests/test_serial_utils_batch_on_cpu.py | 22 + transfer_queue/config.yaml | 7 + .../bootstrap/simple_storage_bootstrap.py | 5 + .../managers/simple_storage_manager.py | 96 +--- .../storage/payload_transfer/__init__.py | 32 ++ .../storage/payload_transfer/base.py | 76 +++ .../storage/payload_transfer/factory.py | 65 +++ .../storage/payload_transfer/nixl.py | 534 ++++++++++++++++++ .../payload_transfer/nixl_ucx_runtime.py | 392 +++++++++++++ .../storage/payload_transfer/zmq.py | 173 ++++++ transfer_queue/storage/simple_storage.py | 208 +++---- transfer_queue/utils/serial_utils.py | 30 + transfer_queue/utils/zmq_utils.py | 8 + 16 files changed, 1787 insertions(+), 200 deletions(-) create mode 100644 docs/nixl_ucx_payload.md create mode 100644 tests/test_payload_transfer.py create mode 100644 transfer_queue/storage/payload_transfer/__init__.py create mode 100644 transfer_queue/storage/payload_transfer/base.py create mode 100644 transfer_queue/storage/payload_transfer/factory.py create mode 100644 transfer_queue/storage/payload_transfer/nixl.py create mode 100644 transfer_queue/storage/payload_transfer/nixl_ucx_runtime.py create mode 100644 transfer_queue/storage/payload_transfer/zmq.py diff --git a/README.md b/README.md index 0e026e80..e2395263 100644 --- a/README.md +++ b/README.md @@ -211,6 +211,9 @@ pip install TransferQueue pip install dist/*.whl ``` +For the optional SimpleStorage NIXL-UCX Host payload path, see +[the NIXL-UCX payload guide](docs/nixl_ucx_payload.md). +

📊 Performance

### Simple Case: Regular Tensor @@ -345,4 +348,4 @@ Please kindly cite our paper if you find this repo is useful: journal={arXiv preprint arXiv:2507.01663}, year={2025} } -``` \ No newline at end of file +``` diff --git a/docs/nixl_ucx_payload.md b/docs/nixl_ucx_payload.md new file mode 100644 index 00000000..48beb1f2 --- /dev/null +++ b/docs/nixl_ucx_payload.md @@ -0,0 +1,188 @@ +# SimpleStorage NIXL-UCX Host Payload Transfer + +After setting `payload_transfer` to `nixl-ucx`, all non-empty payloads are transferred through NIXL-UCX; +ZMQ handles control messages only. This document describes the configuration and usage. + +## 1. Check RDMA Devices + +Run the following on every node running TQ or a Ray worker: + +```bash +ls /sys/class/infiniband +rdma link show +ibv_devinfo +``` + +`ls` should list RDMA devices, and the ports shown by `rdma link show` should be `ACTIVE`. +`ibv_devinfo` does not show the provider dynamic library name. If no device is present or a port is +not active, check the driver, `rdma-core`, provider, and container device mappings first. + +## 2. Install NIXL and TQ + +Install the NIXL wheel: + +```bash +python -m pip install nixl +``` + +The NIXL wheel includes the UCX runtime, but the system must still provide `rdma-core`, +`libibverbs`, and the provider for the network adapter. + +Install TQ from the source directory: + +```bash +python -m pip install -e . +``` + +## 3. Check the NIXL-UCX Backend + +Run: + +```bash +python - <<'PY' +from nixl import nixl_agent, nixl_agent_config + +config = nixl_agent_config( + enable_prog_thread=True, + enable_listen_thread=True, + listen_port=0, + backends=["UCX"], +) +agent = nixl_agent("tq-nixl-check", config) +assert "UCX" in agent.backends +print("NIXL UCX backend is available") +PY +``` + +The command should output `NIXL UCX backend is available`. If it fails, see the common issues at the end. + +## 4. Enable SimpleStorage NIXL-UCX Transfer + +Enable NIXL-UCX in the TQ configuration: + +```yaml +backend: + storage_backend: SimpleStorage + SimpleStorage: + payload_transfer: + backend: nixl-ucx + ucx_env_vars: {} +``` + +`ucx_env_vars: {}` means that TQ does not set additional UCX environment variables. TQ and Ray workers +continue to use the `UCX_*` variables inherited when they were started. To specify a transport, device, +or GID, add the corresponding variables to `ucx_env_vars`. + +If NIXL initialization or a transfer fails, TQ reports the error directly and does not fall back to ZMQ. +If the transport is not restricted, or if `UCX_TLS` includes `tcp`, UCX may use TCP. + +### Common UCX Configuration + +| Variable | Purpose | Reference value | +| --- | --- | --- | +| `UCX_TLS` | Restrict the transports available to UCX | `,tcp,sm,self` | +| `UCX_NET_DEVICES` | Specify the RDMA device and port | `:` | +| `UCX_IB_GID_INDEX` | Specify the RoCE GID index | `` | +| `UCX_MODULE_DIR` | Specify the UCX transport module directory in the NIXL wheel | `` | + +Restart TQ/Ray after making changes. Set the device name and GID index for each node. + +### Memory Registration + +Before NIXL registers memory, check the system limit in the current shell: + +```bash +ulimit -l +``` + +If the value is too small, set it to `unlimited` in the shell that starts TQ/Ray: + +```bash +ulimit -l unlimited +``` + +This setting applies only to the current shell and its child processes. + +## 5. Verify SimpleStorage NIXL-UCX Transfer + +After enabling it, the StorageUnit startup log will contain: + +```text +SimpleStorage payload transfer selected: nixl-ucx device=ucx-auto gid_index=ucx-auto tls=ucx-auto +``` + +After a cross-node PUT/GET completes, the GET content should match the PUT content. The log and data +validation only confirm that the NIXL-UCX path is usable; to confirm RDMA, also check the payload lane. +`rc_*` indicates RDMA, while a TCP lane indicates that TCP is being used. + +## Common Issues + +### RDMA Devices Are Ready, but NIXL-UCX Fails to Start + +If `ibv_devinfo` shows RDMA devices and active ports but NIXL initialization fails, the log typically contains: + +```text +no userspace device-specific driver found +failed to open ... libuct_ib ... +NIXL_ERR_BACKEND +``` + +First confirm that the provider for the network adapter is installed. If the provider is installed but the +error persists, use a NIXL wheel compatible with the system `rdma-core/provider`. If no suitable wheel is +available, follow the [official NIXL source build instructions](https://github.com/ai-dynamo/nixl#prerequisites-for-source-build-linux) +to build UCX with multi-thread and verbs enabled: + +```bash +python -m pip install meson ninja pybind11 tomlkit +git clone https://github.com/openucx/ucx.git +cd +git checkout +./autogen.sh +./contrib/configure-release-mt \ + --prefix= \ + --enable-shared \ + --disable-static \ + --with-verbs +make -j"$(nproc)" +make install +``` + +Then configure NIXL to use this UCX: + +```bash +git clone https://github.com/ai-dynamo/nixl.git +cd +python -m pip install . +meson setup build \ + -Ducx_path= \ + -Dprefix= \ + -Dbuildtype=release +ninja -C build +ninja -C build install +python -m pip install build/src/bindings/python/nixl-meta/nixl-*-py3-none-any.whl +``` + +### `libnixl.so` Cannot Be Found + +If `import nixl` reports `libnixl.so: cannot open shared object file`, the NIXL Python extension +cannot find the NIXL shared library; UCX initialization has not started yet. + +First locate `libnixl.so` in the wheel: + +```bash +NIXL_SITE=$(python -c 'import site; print(site.getsitepackages()[0])') +find "${NIXL_SITE}" -name libnixl.so +``` + +If it is found, add its containing directory to `LD_LIBRARY_PATH` in the same shell that starts TQ/Ray: + +```bash +NIXL_LIB_DIR=/path/to/directory/containing/libnixl.so +export LD_LIBRARY_PATH="${NIXL_LIB_DIR}${LD_LIBRARY_PATH:+:${LD_LIBRARY_PATH}}" +``` + +If it is not found, reinstall the NIXL wheel corresponding to the reported error. For example: + +```bash +python -m pip install --no-cache-dir --force-reinstall --no-deps nixl-cu12 +``` diff --git a/tests/test_payload_transfer.py b/tests/test_payload_transfer.py new file mode 100644 index 00000000..8495b56f --- /dev/null +++ b/tests/test_payload_transfer.py @@ -0,0 +1,146 @@ +# 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. + +"""Payload transfer contract and NIXL-UCX configuration tests.""" + +from __future__ import annotations + +import os +from pathlib import Path + +import pytest +from omegaconf import OmegaConf + +from transfer_queue.storage.payload_transfer import ( + PayloadTransferError, + create_payload_transfer, + parse_payload_transfer_config, +) +from transfer_queue.storage.payload_transfer.nixl import PayloadDescriptor +from transfer_queue.storage.payload_transfer.nixl_ucx_runtime import _configure_ucx_environment +from transfer_queue.storage.payload_transfer.zmq import ZmqPayloadTransfer +from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType + + +def test_payload_descriptor_preserves_frame_layout(): + descriptor = PayloadDescriptor("framed", 4 + 8 * 2 + 5, (2, 3)) + descriptor.validate() + assert PayloadDescriptor.from_dict(descriptor.to_dict()) == descriptor + + with pytest.raises(PayloadTransferError, match="packed payload length"): + PayloadDescriptor("framed", 5, (2, 3)).validate() + + +def test_payload_descriptor_requires_frame_layout_and_rejects_negative_lengths(): + with pytest.raises(KeyError, match="frame_sizes"): + PayloadDescriptor.from_dict({"transfer_id": "payload", "payload_bytes": 3}) + + with pytest.raises(PayloadTransferError, match="negative"): + PayloadDescriptor.from_dict({"transfer_id": "payload", "payload_bytes": -1, "frame_sizes": [1]}) + + +def test_yaml_ucx_settings_override_process_environment(monkeypatch): + monkeypatch.setenv("UCX_TLS", "sm") + monkeypatch.delenv("UCX_IB_GID_INDEX", raising=False) + + configured = _configure_ucx_environment( + { + "UCX_TLS": "tcp", + "UCX_IB_GID_INDEX": 3, + } + ) + + assert configured == {"UCX_TLS": "tcp", "UCX_IB_GID_INDEX": "3"} + assert os.environ["UCX_TLS"] == "tcp" + assert os.environ["UCX_IB_GID_INDEX"] == "3" + + +def test_empty_yaml_ucx_settings_preserve_process_environment(monkeypatch): + monkeypatch.setenv("UCX_NET_DEVICES", "custom_hca:2") + + assert _configure_ucx_environment({}) == {} + assert os.environ["UCX_NET_DEVICES"] == "custom_hca:2" + + +def test_payload_transfer_rejects_unsupported_backend(): + with pytest.raises(ValueError, match="expected 'zmq' or 'nixl-ucx'"): + create_payload_transfer({"backend": "unsupported"}) + + +def test_factory_returns_zmq_payload_transfer(): + transfer = create_payload_transfer({"backend": "zmq"}) + + assert isinstance(transfer, ZmqPayloadTransfer) + assert transfer.bootstrap_info() is None + + +def test_nixl_factory_rejects_incomplete_peer_endpoints(): + with pytest.raises(RuntimeError, match="payload transfer endpoints are missing"): + create_payload_transfer( + {"backend": "nixl-ucx"}, + control_peer_infos={"storage": object()}, + ) + + +def test_zmq_payload_transfer_handles_put_and_get_requests(): + stored = {} + transfer = ZmqPayloadTransfer() + + put_request = ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA, + sender_id="manager", + receiver_id="storage", + body={"global_indexes": [1], "data": {"value": [42]}}, + ) + put_response = transfer.handle_request( + put_request, + storage_id="storage", + load_data=lambda fields, indexes: {field: stored[field] for field in fields}, + store_data=lambda indexes, data, parser: stored.update(data), + ) + + assert put_response.request_type == ZMQRequestType.PUT_DATA_RESPONSE + assert stored == {"value": [42]} + + get_request = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA, + sender_id="manager", + receiver_id="storage", + body={"global_indexes": [1], "fields": ["value"]}, + ) + get_response = transfer.handle_request( + get_request, + storage_id="storage", + load_data=lambda fields, indexes: {field: stored[field] for field in fields}, + store_data=lambda indexes, data, parser: stored.update(data), + ) + + assert get_response.request_type == ZMQRequestType.GET_DATA_RESPONSE + assert get_response.body["data"] == {"value": [42]} + + +def test_payload_transfer_config_extracts_ucx_settings(): + assert parse_payload_transfer_config( + { + "backend": "nixl-ucx", + "ucx_env_vars": {"UCX_TLS": "rc", "UCX_IB_GID_INDEX": 3}, + } + ) == ("nixl-ucx", {"ucx_env_vars": {"UCX_TLS": "rc", "UCX_IB_GID_INDEX": 3}}) + + +def test_public_config_defaults_to_zmq_payload_transfer(): + config = OmegaConf.load(Path(__file__).parents[1] / "transfer_queue/config.yaml") + + assert config.backend.SimpleStorage.payload_transfer.backend == "zmq" diff --git a/tests/test_serial_utils_batch_on_cpu.py b/tests/test_serial_utils_batch_on_cpu.py index 7720f4a9..e7e8d3f5 100644 --- a/tests/test_serial_utils_batch_on_cpu.py +++ b/tests/test_serial_utils_batch_on_cpu.py @@ -22,6 +22,8 @@ * ``batch_decode_from`` """ +import struct + import numpy as np import pytest import torch @@ -42,6 +44,15 @@ def test_calc_packed_size_then_pack_unpack_roundtrip(): assert [bytes(mv) for mv in recovered] == items +def test_initialize_packed_frame_table_leaves_payload_for_direct_receive(): + items = [b"hello", b"world!"] + buf = bytearray(serial_utils.calc_packed_size(items)) + serial_utils.initialize_packed_frame_table(buf, [len(item) for item in items]) + payload_start = serial_utils._PACK_HEADER_SIZE + len(items) * serial_utils._PACK_ENTRY_SIZE + buf[payload_start:] = b"helloworld!" + assert [bytes(mv) for mv in serial_utils.unpack_from(buf)] == items + + def test_pack_into_writes_only_within_its_slice(): items = [b"alpha", b"beta", b"gamma"] sz = serial_utils.calc_packed_size(items) @@ -64,6 +75,17 @@ def test_unpack_from_zero_item_buffer(): assert serial_utils.unpack_from(buf) == [] +def test_unpack_from_rejects_invalid_frame_bounds(): + items = [b"payload"] + buf = bytearray(serial_utils.calc_packed_size(items)) + serial_utils.pack_into(buf, items) + + # Corrupt the frame offset so it points into the frame table. + struct.pack_into(" dict[str, Any]: num_data_storage_units = conf.backend.SimpleStorage.num_data_storage_units total_storage_size = conf.backend.SimpleStorage.get("total_storage_size", None) required_node_resource = conf.backend.SimpleStorage.get("required_node_resource", None) + payload_transfer_config = conf.backend.SimpleStorage.get("payload_transfer") scheduling_strategies = get_node_round_robin_scheduling_strategies( num_data_storage_units, required_node_resource=required_node_resource ) @@ -50,6 +52,7 @@ def initialize_simple_storage(conf: DictConfig) -> dict[str, Any]: name=f"TransferQueueStorageUnit#{storage_unit_rank}", ).remote( storage_unit_size=storage_unit_size, + payload_transfer=payload_transfer_config, ) simple_storage_handles[f"TransferQueueStorageUnit#{storage_unit_rank}"] = storage_node logger.info( @@ -60,5 +63,7 @@ def initialize_simple_storage(conf: DictConfig) -> dict[str, Any]: storage_zmq_info = process_zmq_server_info(simple_storage_handles) backend_name = conf.backend.storage_backend conf.backend[backend_name].zmq_info = storage_zmq_info + infos = ray.get([storage.get_payload_transfer_info.remote() for storage in simple_storage_handles.values()]) + conf.backend[backend_name].payload_transfer_infos = {info["id"]: info for info in infos if info is not None} return simple_storage_handles diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 88340006..598795d5 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -24,11 +24,13 @@ import torch import zmq +import zmq.asyncio from omegaconf import DictConfig from tensordict import NonTensorStack, TensorDict from transfer_queue.metadata import BatchMeta, extract_field_schema from transfer_queue.storage.managers.base import StorageManager, StorageManagerFactory +from transfer_queue.storage.payload_transfer import create_payload_transfer from transfer_queue.storage.simple_storage import KEY_NOT_FOUND_MARKER, StorageKeyNotFoundError from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.zmq_utils import ( @@ -97,6 +99,11 @@ def __init__( raise ValueError("AsyncSimpleStorageManager requires non-empty 'zmq_info' in config.") self.storage_unit_infos = self._register_servers(server_infos) + self.payload_transfer = create_payload_transfer( + config.get("payload_transfer"), + peer_infos=config.get("payload_transfer_infos", {}) or {}, + control_peer_infos=self.storage_unit_infos, + ) def _register_servers(self, server_infos: "ZMQServerInfo | dict[Any, ZMQServerInfo]"): """Register and validate server information. @@ -299,41 +306,15 @@ async def _put_to_single_storage_unit( Send data to a specific storage unit. """ - request_msg = ZMQMessage.create( - request_type=ZMQRequestType.PUT_DATA, # type: ignore[arg-type] + await self.payload_transfer.put( + control_socket=socket, sender_id=self.storage_manager_id, - receiver_id=target_storage_unit, - body={"global_indexes": global_indexes, "data": storage_data, "data_parser": data_parser}, + target_id=target_storage_unit, + global_indexes=global_indexes, + data=storage_data, + data_parser=data_parser, ) - try: - data = request_msg.serialize() - await socket.send_multipart(data, copy=False) - messages = await socket.recv_multipart(copy=False) - response_msg = ZMQMessage.deserialize(messages) - - if response_msg.request_type != ZMQRequestType.PUT_DATA_RESPONSE: - raise RuntimeError( - f"Failed to put data to storage unit {target_storage_unit}: " - f"{response_msg.body.get('message', 'Unknown error')}" - ) - except zmq.error.Again as e: - timeout_sec = TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT - logger.error( - f"[{self.storage_manager_id}]: ZMQ recv timeout ({timeout_sec}s) " - f"during put to storage unit {target_storage_unit}. " - f"The storage unit may be overloaded or crashed." - ) - raise RuntimeError( - f"ZMQ recv timeout ({timeout_sec}s) during put to storage unit {target_storage_unit}" - ) from e - except Exception as e: - logger.error( - f"[{self.storage_manager_id}]: Unexpected error during put to storage unit " - f"{target_storage_unit}: {type(e).__name__}: {e}" - ) - 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: """ @@ -445,43 +426,21 @@ async def _get_from_single_storage_unit( socket: zmq.Socket = None, ): """Get data from a single SU by global index keys.""" - request_msg = ZMQMessage.create( - request_type=ZMQRequestType.GET_DATA, # type: ignore[arg-type] - sender_id=self.storage_manager_id, - receiver_id=target_storage_unit, - body={"global_indexes": global_indexes, "fields": fields}, - ) try: - await socket.send_multipart(request_msg.serialize()) - messages = await socket.recv_multipart(copy=False) - response_msg = ZMQMessage.deserialize(messages) - - if response_msg.request_type == ZMQRequestType.GET_DATA_RESPONSE: - storage_unit_data = response_msg.body["data"] - return fields, storage_unit_data - else: - message = response_msg.body.get("message", "Unknown error") - error_type = StorageKeyNotFoundError if KEY_NOT_FOUND_MARKER in message else RuntimeError - raise error_type(f"Failed to get data from storage unit {target_storage_unit}: {message}") - except zmq.error.Again as e: - timeout_sec = TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT - logger.error( - f"[{self.storage_manager_id}]: ZMQ recv timeout ({timeout_sec}s) " - f"from storage unit {target_storage_unit}. " - f"The storage unit may be overloaded or crashed." + data = await self.payload_transfer.get( + control_socket=socket, + sender_id=self.storage_manager_id, + target_id=target_storage_unit, + global_indexes=global_indexes, + fields=fields, ) - raise RuntimeError(f"ZMQ recv timeout ({timeout_sec}s) from storage unit {target_storage_unit}") from e except StorageKeyNotFoundError: - # Already logged at debug by the storage unit; propagate for the caller to classify. raise - except Exception as e: - logger.error( - f"[{self.storage_manager_id}]: Unexpected error from storage unit " - f"{target_storage_unit}: {type(e).__name__}: {e}" - ) - raise RuntimeError( - f"Error getting data from storage unit {target_storage_unit}: {type(e).__name__}: {e}" - ) from e + except Exception as exc: + if KEY_NOT_FOUND_MARKER in str(exc): + raise StorageKeyNotFoundError(str(exc)) from exc + raise + return fields, data async def clear_data(self, metadata: BatchMeta) -> None: """Clear data in remote StorageUnit. @@ -658,5 +617,8 @@ async def load_checkpoint(self, checkpoint_dir: str) -> None: logger.info(f"[{self.storage_manager_id}]: restored {len(su_ids)} storage units from {su_dir}") def close(self) -> None: - """Close all ZMQ sockets and context to prevent resource leaks.""" - super().close() + """Close payload transfer resources before ZMQ ownership is released.""" + try: + self.payload_transfer.close() + finally: + super().close() diff --git a/transfer_queue/storage/payload_transfer/__init__.py b/transfer_queue/storage/payload_transfer/__init__.py new file mode 100644 index 00000000..8df2ac9c --- /dev/null +++ b/transfer_queue/storage/payload_transfer/__init__.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. + +"""Payload transfer strategies for SimpleStorage.""" + +from transfer_queue.storage.payload_transfer.base import ( + PayloadTransfer, + PayloadTransferError, +) +from transfer_queue.storage.payload_transfer.factory import ( + create_payload_transfer, + parse_payload_transfer_config, +) + +__all__ = [ + "PayloadTransfer", + "PayloadTransferError", + "create_payload_transfer", + "parse_payload_transfer_config", +] diff --git a/transfer_queue/storage/payload_transfer/base.py b/transfer_queue/storage/payload_transfer/base.py new file mode 100644 index 00000000..fa2003e0 --- /dev/null +++ b/transfer_queue/storage/payload_transfer/base.py @@ -0,0 +1,76 @@ +# 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. + +"""High-level payload transfer contract used by SimpleStorage.""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from typing import TYPE_CHECKING, Any, Callable + +if TYPE_CHECKING: + from transfer_queue.utils.zmq_utils import ZMQMessage + + +class PayloadTransferError(RuntimeError): + """A payload transfer could not be completed safely.""" + + +class PayloadTransfer(ABC): + """Complete SimpleStorage payload strategy, including its wire protocol.""" + + @abstractmethod + async def put( + self, + *, + control_socket: Any, + sender_id: str, + target_id: str, + global_indexes: list[int], + data: dict[str, Any], + data_parser: Callable[[Any], Any] | None, + ) -> None: + """Put decoded storage data through this strategy.""" + + @abstractmethod + async def get( + self, + *, + control_socket: Any, + sender_id: str, + target_id: str, + global_indexes: list[int], + fields: list[str], + ) -> dict[str, Any]: + """Get storage data through this strategy.""" + + @abstractmethod + def handle_request( + self, + request: ZMQMessage, + *, + storage_id: str, + load_data: Callable[..., dict[str, Any]], + store_data: Callable[..., None], + ) -> ZMQMessage | None: + """Handle a strategy-owned request on a SimpleStorageUnit.""" + + def bootstrap_info(self) -> dict[str, Any] | None: + """Return transport-specific metadata needed by peer instances.""" + return None + + def close(self) -> None: + """Release transport resources.""" + return None diff --git a/transfer_queue/storage/payload_transfer/factory.py b/transfer_queue/storage/payload_transfer/factory.py new file mode 100644 index 00000000..28f9ec4a --- /dev/null +++ b/transfer_queue/storage/payload_transfer/factory.py @@ -0,0 +1,65 @@ +# 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. + +"""Construction helper for SimpleStorage payload transfer strategies.""" + +from collections.abc import Mapping + +from transfer_queue.storage.payload_transfer.base import PayloadTransfer +from transfer_queue.utils.zmq_utils import ZMQServerInfo + +_SUPPORTED_TRANSFERS = frozenset({"nixl-ucx", "zmq"}) + + +def parse_payload_transfer_config(value: object | None = None) -> tuple[str, dict[str, object]]: + """Return the backend name and options from the payload transfer block.""" + config = {"backend": "zmq"} if value is None else value + if not isinstance(config, Mapping): + raise TypeError("SimpleStorage.payload_transfer must be a mapping") + + config = dict(config) + if "backend" not in config: + raise ValueError("SimpleStorage.payload_transfer.backend is required") + backend = str(config.pop("backend")).strip().lower() + if backend not in _SUPPORTED_TRANSFERS: + raise ValueError(f"unsupported SimpleStorage payload transfer: {backend!r}; expected 'zmq' or 'nixl-ucx'") + return backend, config + + +def create_payload_transfer( + value: object | None = None, + *, + peer_infos: Mapping[str, object] | None = None, + control_peer_infos: Mapping[str, ZMQServerInfo] | None = None, +) -> PayloadTransfer: + """Create the configured SimpleStorage payload transfer strategy.""" + backend, options = parse_payload_transfer_config(value) + if backend == "zmq": + from transfer_queue.storage.payload_transfer.zmq import ZmqPayloadTransfer + + return ZmqPayloadTransfer() + if backend == "nixl-ucx": + from transfer_queue.storage.payload_transfer.nixl import NixlPayloadTransfer + + ucx_env_vars = options.get("ucx_env_vars") + if ucx_env_vars is not None and not isinstance(ucx_env_vars, Mapping): + raise TypeError("SimpleStorage.payload_transfer.ucx_env_vars must be a mapping") + return NixlPayloadTransfer( + ucx_env_vars=None if ucx_env_vars is None else dict(ucx_env_vars), + peer_infos=peer_infos, + control_peer_infos=control_peer_infos, + ) + + raise RuntimeError(f"unhandled SimpleStorage payload transfer: {backend!r}") diff --git a/transfer_queue/storage/payload_transfer/nixl.py b/transfer_queue/storage/payload_transfer/nixl.py new file mode 100644 index 00000000..4350905a --- /dev/null +++ b/transfer_queue/storage/payload_transfer/nixl.py @@ -0,0 +1,534 @@ +# 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. + +"""NIXL H2H payload transfer strategy for SimpleStorage.""" + +from __future__ import annotations + +import asyncio +import os +from collections.abc import Mapping +from concurrent.futures import Future +from dataclasses import dataclass +from typing import Any, Callable +from uuid import uuid4 + +import zmq.asyncio + +from transfer_queue.storage.payload_transfer.base import PayloadTransfer, PayloadTransferError +from transfer_queue.storage.payload_transfer.nixl_ucx_runtime import NixlError, NixlRuntime +from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads +from transfer_queue.utils.logging_utils import get_logger +from transfer_queue.utils.serial_utils import calc_packed_size, decode, encode, unpack_from +from transfer_queue.utils.zmq_utils import ( + ZMQMessage, + ZMQRequestType, + ZMQServerInfo, + create_zmq_socket, + format_zmq_address, +) + +logger = get_logger(__name__) +TQ_NUM_THREADS = int(os.environ.get("TQ_NUM_THREADS", 8)) +TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT = int(os.environ.get("TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT", 200)) + + +@dataclass(frozen=True) +class PayloadDescriptor: + """Description of one encoded payload.""" + + transfer_id: str + payload_bytes: int + frame_sizes: tuple[int, ...] + + @staticmethod + def validate_transfer_id(transfer_id: str) -> None: + if not transfer_id or len(transfer_id) > 128: + raise PayloadTransferError("transfer_id must contain 1 to 128 characters") + + def validate(self) -> None: + self.validate_transfer_id(self.transfer_id) + if self.payload_bytes < 0: + raise PayloadTransferError(f"negative payload length for {self.transfer_id}") + if any(size < 0 for size in self.frame_sizes): + raise PayloadTransferError(f"negative frame length for {self.transfer_id}") + packed_size = 4 + 8 * len(self.frame_sizes) + sum(self.frame_sizes) + if packed_size != self.payload_bytes: + raise PayloadTransferError( + f"packed payload length mismatch for {self.transfer_id}: " + f"expected {packed_size}, got {self.payload_bytes}" + ) + + def to_dict(self) -> dict[str, int | str | list[int]]: + return { + "transfer_id": self.transfer_id, + "payload_bytes": self.payload_bytes, + "frame_sizes": list(self.frame_sizes), + } + + @classmethod + def from_dict(cls, value: dict[str, Any]) -> PayloadDescriptor: + descriptor = cls( + transfer_id=str(value["transfer_id"]), + payload_bytes=int(value["payload_bytes"]), + frame_sizes=tuple(int(size) for size in value["frame_sizes"]), + ) + descriptor.validate() + return descriptor + + +@dataclass +class TransferEndpoint: + """Bootstrap metadata for a NIXL endpoint.""" + + transport: str + data: dict[str, Any] + + def to_dict(self) -> dict[str, Any]: + return {"transport": self.transport, "data": self.data} + + @classmethod + def from_dict(cls, value: dict[str, Any]) -> TransferEndpoint: + return cls(transport=str(value["transport"]), data=dict(value["data"])) + + +@dataclass +class ReceiveToken: + """Transport-owned metadata returned after preparing a receive.""" + + data: dict[str, Any] + + def to_dict(self) -> dict[str, Any]: + return {"data": self.data} + + @classmethod + def from_dict(cls, value: dict[str, Any]) -> ReceiveToken: + return cls(data=dict(value["data"])) + + +@dataclass +class _PendingPut: + descriptor: PayloadDescriptor + sender_id: str + global_indexes: tuple[int, ...] + data_parser: Any + + +@dataclass +class _PendingGet: + descriptor: PayloadDescriptor + sender_id: str + frames: tuple[bytes | bytearray | memoryview, ...] + + +class NixlPayloadTransfer(PayloadTransfer): + """Own the NIXL payload protocol and its transfer lifecycle.""" + + transport = "nixl-ucx" + + def __init__( + self, + ucx_env_vars: dict[str, object] | None = None, + peer_infos: Mapping[str, object] | None = None, + control_peer_infos: Mapping[str, ZMQServerInfo] | None = None, + ): + self._peer_infos = dict(peer_infos or {}) + self._control_peer_infos = dict(control_peer_infos or {}) + if self._control_peer_infos and set(self._control_peer_infos) != set(self._peer_infos): + raise RuntimeError("SimpleStorage payload transfer endpoints are missing") + try: + self._runtime = NixlRuntime(ucx_env_vars) + except NixlError: + raise + except Exception as exc: + raise NixlError(f"failed to create NIXL runtime: {exc}") from exc + self._pending_puts: dict[str, _PendingPut] = {} + self._pending_gets: dict[str, _PendingGet] = {} + + def bootstrap_info(self) -> dict[str, Any]: + return {"endpoint": self.endpoint().to_dict()} + + def endpoint(self) -> TransferEndpoint: + return TransferEndpoint( + transport=self.transport, + data={ + "agent_name": self._runtime.agent_name, + "agent_metadata": self._runtime.endpoint_metadata(), + }, + ) + + async def put( + self, + *, + control_socket: zmq.asyncio.Socket, + sender_id: str, + target_id: str, + global_indexes: list[int], + data: dict[str, Any], + data_parser: Callable[[Any], Any] | None, + ) -> None: + frames = tuple(encode(data)) + descriptor = PayloadDescriptor( + transfer_id=uuid4().hex, + payload_bytes=calc_packed_size(frames), + frame_sizes=tuple(memoryview(frame).nbytes for frame in frames), + ) + descriptor.validate() + remote_may_be_prepared = True + try: + prepare = ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_PREPARE, + sender_id=sender_id, + receiver_id=target_id, + body={ + "global_indexes": global_indexes, + "descriptor": descriptor.to_dict(), + "data_parser": data_parser, + }, + ) + await control_socket.send_multipart(prepare.serialize(), copy=False) + ready = ZMQMessage.deserialize(await control_socket.recv_multipart(copy=False)) + self._expect(ready, ZMQRequestType.PUT_DATA_READY, target_id) + if PayloadDescriptor.from_dict(ready.body["descriptor"]) != descriptor: + raise RuntimeError(f"PUT descriptor changed by storage unit {target_id}") + token = ReceiveToken.from_dict(ready.body["receive_token"]) + endpoint = self._peer_endpoint(target_id) + await asyncio.wrap_future(self.send(endpoint, token, descriptor, frames)) + + commit = ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_COMMIT, + sender_id=sender_id, + receiver_id=target_id, + body={"transfer_id": descriptor.transfer_id}, + ) + await control_socket.send_multipart(commit.serialize(), copy=False) + response = ZMQMessage.deserialize(await control_socket.recv_multipart(copy=False)) + self._expect(response, ZMQRequestType.PUT_DATA_RESPONSE, target_id) + remote_may_be_prepared = False + except BaseException: + if remote_may_be_prepared: + await self._cancel(sender_id, target_id, ZMQRequestType.PUT_DATA_CANCEL, descriptor.transfer_id) + raise + + async def get( + self, + *, + control_socket: zmq.asyncio.Socket, + sender_id: str, + target_id: str, + global_indexes: list[int], + fields: list[str], + ) -> dict[str, Any]: + transfer_id = uuid4().hex + remote_prepared = False + receive_prepared = False + descriptor = None + try: + prepare = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_PREPARE, + sender_id=sender_id, + receiver_id=target_id, + body={"global_indexes": global_indexes, "fields": fields, "transfer_id": transfer_id}, + ) + await control_socket.send_multipart(prepare.serialize(), copy=False) + remote_prepared = True + ready = ZMQMessage.deserialize(await control_socket.recv_multipart(copy=False)) + self._expect(ready, ZMQRequestType.GET_DATA_READY, target_id) + descriptor = PayloadDescriptor.from_dict(ready.body["descriptor"]) + if descriptor.transfer_id != transfer_id: + raise RuntimeError(f"GET descriptor identity changed by storage unit {target_id}") + token = self.prepare_receive(descriptor) + receive_prepared = True + commit = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_COMMIT, + sender_id=sender_id, + receiver_id=target_id, + body={ + "transfer_id": descriptor.transfer_id, + "receiver_endpoint": self.endpoint().to_dict(), + "receive_token": token.to_dict(), + }, + ) + await control_socket.send_multipart(commit.serialize(), copy=False) + response = ZMQMessage.deserialize(await control_socket.recv_multipart(copy=False)) + self._expect(response, ZMQRequestType.GET_DATA_RESPONSE, target_id) + remote_prepared = False + payload = await asyncio.wrap_future(self.receive(descriptor)) + return decode(unpack_from(payload)) + except BaseException: + if receive_prepared and descriptor is not None: + self.cancel_receive(descriptor.transfer_id) + if remote_prepared: + await self._cancel(sender_id, target_id, ZMQRequestType.GET_DATA_CANCEL, transfer_id) + raise + + def handle_request( + self, + request: ZMQMessage, + *, + storage_id: str, + load_data: Callable[..., dict[str, Any]], + store_data: Callable[..., None], + ) -> ZMQMessage | None: + if request.request_type == ZMQRequestType.PUT_DATA_PREPARE: + return self._handle_put_prepare(request, storage_id) + if request.request_type == ZMQRequestType.PUT_DATA_COMMIT: + return self._handle_put_commit(request, storage_id, store_data) + if request.request_type == ZMQRequestType.PUT_DATA_CANCEL: + return self._handle_put_cancel(request, storage_id) + if request.request_type == ZMQRequestType.GET_DATA_PREPARE: + return self._handle_get_prepare(request, storage_id, load_data) + if request.request_type == ZMQRequestType.GET_DATA_COMMIT: + return self._handle_get_commit(request, storage_id) + if request.request_type == ZMQRequestType.GET_DATA_CANCEL: + return self._handle_get_cancel(request, storage_id) + return None + + def _handle_put_prepare(self, request: ZMQMessage, storage_id: str) -> ZMQMessage: + descriptor = None + prepared = False + try: + descriptor = PayloadDescriptor.from_dict(request.body["descriptor"]) + if descriptor.transfer_id in self._pending_puts: + raise RuntimeError(f"duplicate PUT transfer_id: {descriptor.transfer_id}") + token = self.prepare_receive(descriptor) + prepared = True + self._pending_puts[descriptor.transfer_id] = _PendingPut( + descriptor, + request.sender_id, + tuple(request.body["global_indexes"]), + request.body.get("data_parser"), + ) + return ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_READY, + sender_id=storage_id, + body={"descriptor": descriptor.to_dict(), "receive_token": token.to_dict()}, + ) + except Exception as exc: + if descriptor is not None and prepared: + self._pending_puts.pop(descriptor.transfer_id, None) + self.cancel_receive(descriptor.transfer_id) + return self._error(storage_id, "PUT prepare", exc) + + def _handle_put_commit(self, request: ZMQMessage, storage_id: str, store_data: Callable[..., None]) -> ZMQMessage: + transfer_id = request.body["transfer_id"] + owns_receive = False + try: + pending = self._pending_puts.get(transfer_id) + if pending is None: + raise RuntimeError(f"unknown or expired PUT transfer_id: {transfer_id}") + if pending.sender_id != request.sender_id: + raise RuntimeError(f"PUT transfer {transfer_id} belongs to another sender") + self._pending_puts.pop(transfer_id) + owns_receive = True + payload = self.receive(pending.descriptor).result() + with limit_pytorch_auto_parallel_threads( + target_num_threads=TQ_NUM_THREADS, info=f"[{storage_id}] PUT commit" + ): + store_data(list(pending.global_indexes), decode(unpack_from(payload)), pending.data_parser) + return ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_RESPONSE, + sender_id=storage_id, + body={"transfer_id": transfer_id}, + ) + except Exception as exc: + if owns_receive: + self.cancel_receive(transfer_id) + return self._error(storage_id, "PUT commit", exc) + + def _handle_put_cancel(self, request: ZMQMessage, storage_id: str) -> ZMQMessage: + transfer_id = request.body["transfer_id"] + pending = self._pending_puts.get(transfer_id) + if pending is not None and pending.sender_id != request.sender_id: + return self._error( + storage_id, + "PUT cancel", + RuntimeError(f"PUT transfer {transfer_id} belongs to another sender"), + ) + self._pending_puts.pop(transfer_id, None) + self.cancel_receive(transfer_id) + return ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_RESPONSE, + sender_id=storage_id, + body={"transfer_id": transfer_id}, + ) + + def _handle_get_prepare( + self, + request: ZMQMessage, + storage_id: str, + load_data: Callable[..., dict[str, Any]], + ) -> ZMQMessage: + try: + transfer_id = str(request.body["transfer_id"]) + PayloadDescriptor.validate_transfer_id(transfer_id) + if transfer_id in self._pending_gets: + raise RuntimeError(f"duplicate GET transfer_id: {transfer_id}") + with limit_pytorch_auto_parallel_threads( + target_num_threads=TQ_NUM_THREADS, info=f"[{storage_id}] GET prepare" + ): + frames = tuple(encode(load_data(request.body["fields"], request.body["global_indexes"]))) + descriptor = PayloadDescriptor( + transfer_id=transfer_id, + payload_bytes=calc_packed_size(frames), + frame_sizes=tuple(memoryview(frame).nbytes for frame in frames), + ) + descriptor.validate() + self._pending_gets[transfer_id] = _PendingGet(descriptor, request.sender_id, frames) + return ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_READY, + sender_id=storage_id, + body={"descriptor": descriptor.to_dict()}, + ) + except Exception as exc: + return self._error(storage_id, "GET prepare", exc) + + def _handle_get_commit(self, request: ZMQMessage, storage_id: str) -> ZMQMessage: + transfer_id = request.body["transfer_id"] + try: + pending = self._pending_gets.get(transfer_id) + if pending is None: + raise RuntimeError(f"unknown or expired GET transfer_id: {transfer_id}") + if pending.sender_id != request.sender_id: + raise RuntimeError(f"GET transfer {transfer_id} belongs to another sender") + self._pending_gets.pop(transfer_id) + endpoint = TransferEndpoint.from_dict(request.body["receiver_endpoint"]) + token = ReceiveToken.from_dict(request.body["receive_token"]) + self.send(endpoint, token, pending.descriptor, pending.frames).result() + return ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_RESPONSE, + sender_id=storage_id, + body={"transfer_id": transfer_id}, + ) + except Exception as exc: + return self._error(storage_id, "GET commit", exc) + + def _handle_get_cancel(self, request: ZMQMessage, storage_id: str) -> ZMQMessage: + transfer_id = request.body["transfer_id"] + pending = self._pending_gets.get(transfer_id) + if pending is not None and pending.sender_id != request.sender_id: + return self._error( + storage_id, + "GET cancel", + RuntimeError(f"GET transfer {transfer_id} belongs to another sender"), + ) + self._pending_gets.pop(transfer_id, None) + return ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_RESPONSE, + sender_id=storage_id, + body={"transfer_id": transfer_id}, + ) + + def prepare_receive(self, descriptor: PayloadDescriptor) -> ReceiveToken: + return ReceiveToken(data=self._runtime.prepare_receive(descriptor)) + + def send( + self, + endpoint: TransferEndpoint, + token: ReceiveToken, + descriptor: PayloadDescriptor, + frames: tuple[bytes | bytearray | memoryview, ...] | list[bytes | bytearray | memoryview], + ) -> Future[None]: + self._validate_metadata(endpoint, token, descriptor) + return self._runtime.send(endpoint.data, token.data, descriptor, tuple(frames)) + + def receive(self, descriptor: PayloadDescriptor) -> Future[memoryview]: + return self._runtime.receive(descriptor) + + def cancel_receive(self, transfer_id: str) -> None: + self._runtime.cancel_receive(transfer_id) + + def close(self) -> None: + self._runtime.close() + + def _peer_endpoint(self, target_id: str) -> TransferEndpoint: + try: + peer_info = self._peer_infos[target_id] + if not isinstance(peer_info, Mapping): + raise TypeError("invalid NIXL peer info") + return TransferEndpoint.from_dict(peer_info["endpoint"]) + except (KeyError, TypeError) as exc: + raise RuntimeError(f"NIXL endpoint is missing for storage unit {target_id}") from exc + + async def _cancel( + self, + sender_id: str, + target_id: str, + request_type: ZMQRequestType, + transfer_id: str, + ) -> None: + """Send cancellation on a fresh DEALER to isolate late responses.""" + cancel_context = zmq.asyncio.Context() + cancel_socket = None + try: + server_info = self._control_peer_infos[target_id] + cancel_socket = create_zmq_socket( + cancel_context, + zmq.DEALER, + server_info.ip, + identity=(f"{sender_id}_cancel_{target_id}_{uuid4().hex[:8]}").encode(), + ) + timeout_ms = min(TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT, 10) * 1000 + cancel_socket.setsockopt(zmq.RCVTIMEO, timeout_ms) + cancel_socket.setsockopt(zmq.SNDTIMEO, timeout_ms) + cancel_socket.connect(format_zmq_address(server_info.ip, server_info.ports["put_get_socket"])) + cancel = ZMQMessage.create( + request_type=request_type, + sender_id=sender_id, + receiver_id=target_id, + body={"transfer_id": transfer_id}, + ) + await cancel_socket.send_multipart(cancel.serialize(), copy=False) + response = ZMQMessage.deserialize(await cancel_socket.recv_multipart(copy=False)) + expected = ( + ZMQRequestType.PUT_DATA_RESPONSE + if request_type == ZMQRequestType.PUT_DATA_CANCEL + else ZMQRequestType.GET_DATA_RESPONSE + ) + self._expect(response, expected, target_id) + except Exception as exc: + logger.warning("failed to cancel %s transfer %s at %s: %s", request_type.value, transfer_id, target_id, exc) + finally: + if cancel_socket is not None: + cancel_socket.close(linger=0) + cancel_context.term() + + @staticmethod + def _expect(response: ZMQMessage, expected: ZMQRequestType, target_id: str) -> None: + if response.request_type != expected: + raise RuntimeError( + f"storage unit {target_id} returned {response.request_type}: " + f"{response.body.get('message', 'unknown error')}" + ) + + @staticmethod + def _error(storage_id: str, operation: str, error: Exception) -> ZMQMessage: + return ZMQMessage.create( + request_type=ZMQRequestType.PUT_GET_ERROR, + sender_id=storage_id, + body={"message": f"{operation} failed in storage unit {storage_id}: {error}"}, + ) + + @staticmethod + def _validate_metadata(endpoint: TransferEndpoint, token: ReceiveToken, descriptor: PayloadDescriptor) -> None: + descriptor.validate() + if endpoint.transport != "nixl-ucx": + raise PayloadTransferError(f"NIXL cannot use endpoint for {endpoint.transport!r}") + if token.data.get("agent_name") != endpoint.data.get("agent_name"): + raise PayloadTransferError("NIXL endpoint and receive token agent names differ") + if not isinstance(token.data.get("frame_remote_descs"), bytes): + raise PayloadTransferError("NIXL receive token is missing frame descriptors") + if not isinstance(token.data.get("agent_metadata"), bytes): + raise PayloadTransferError("NIXL receive token is missing agent metadata") diff --git a/transfer_queue/storage/payload_transfer/nixl_ucx_runtime.py b/transfer_queue/storage/payload_transfer/nixl_ucx_runtime.py new file mode 100644 index 00000000..3c915482 --- /dev/null +++ b/transfer_queue/storage/payload_transfer/nixl_ucx_runtime.py @@ -0,0 +1,392 @@ +# 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. + +"""Small NIXL runtime used by the SimpleStorage H2H payload adapter.""" + +from __future__ import annotations + +import ctypes +import os +import socket +import threading +import time +from concurrent.futures import Future, ThreadPoolExecutor +from dataclasses import dataclass +from typing import Any +from uuid import uuid4 + +from transfer_queue.storage.payload_transfer.base import PayloadTransferError +from transfer_queue.utils.logging_utils import get_logger +from transfer_queue.utils.serial_utils import initialize_packed_frame_table + +logger = get_logger(__name__) + +DEFAULT_NIXL_TRANSFER_TIMEOUT_SECONDS = 180 + + +class NixlError(PayloadTransferError): + """A NIXL runtime operation failed.""" + + +@dataclass +class _RegisteredReceiveBuffer: + buffer: bytearray + registrations: Any + + @property + def capacity(self) -> int: + return len(self.buffer) + + +@dataclass +class _ReceiveState: + scratch: _RegisteredReceiveBuffer + serialized_frame_descs: bytes + + +@dataclass +class _FrameRegistration: + """One registered source frame kept alive for an active transfer.""" + + owner: bytearray | memoryview + registrations: Any + address: int + size: int + + +def _buffer_address(buffer: bytearray | memoryview) -> int: + """Return the address of a writable, contiguous host buffer.""" + if not buffer: + raise NixlError("NIXL does not support an empty payload buffer") + return ctypes.addressof(ctypes.c_ubyte.from_buffer(buffer)) + + +def _configure_ucx_environment(ucx_env_vars: dict[str, object] | None) -> dict[str, str]: + """Apply YAML UCX settings before the NIXL agent reads its environment.""" + config = {str(key): str(value) for key, value in (ucx_env_vars or {}).items()} + os.environ.update(config) + return config + + +def _warn_if_tcp_fallback_possible() -> None: + """Warn when the configured UCX transports allow a TCP payload fallback.""" + ucx_tls = os.environ.get("UCX_TLS") + if not ucx_tls: + return + + transports = {item.strip() for item in ucx_tls.split(",") if item.strip()} + if "tcp" not in transports and "all" not in transports: + return + + logger.warning("NIXL-UCX may fall back to TCP; check UCX logs for the actual transport (UCX_TLS=%s)", ucx_tls) + + +class NixlRuntime: + """Own one NIXL agent and serialize metadata updates safely. + + SimpleStorage's control plane already orders prepare/send/commit. The + runtime therefore only keeps registered receive buffers alive and waits + for the sender-side NIXL request to finish. + """ + + def __init__(self, ucx_env_vars: dict[str, object] | None = None): + _configure_ucx_environment(ucx_env_vars) + try: + from nixl import nixl_agent, nixl_agent_config + except Exception as exc: # pragma: no cover - optional dependency + raise NixlError("NIXL Python bindings are unavailable; install a NIXL build with the UCX backend") from exc + + self._agent_name = self._make_agent_name() + try: + config = nixl_agent_config( + enable_prog_thread=True, + enable_listen_thread=True, + listen_port=0, + backends=["UCX"], + ) + self._agent = nixl_agent(self._agent_name, config) + except Exception as exc: # pragma: no cover - native runtime dependent + raise NixlError(f"failed to initialize the NIXL UCX backend: {exc}") from exc + + if "UCX" not in getattr(self._agent, "backends", {}): + raise NixlError("NIXL UCX backend is not available") + + _warn_if_tcp_fallback_possible() + logger.info( + "SimpleStorage payload transfer selected: nixl-ucx device=%s gid_index=%s tls=%s", + os.environ.get("UCX_NET_DEVICES", "ucx-auto"), + os.environ.get("UCX_IB_GID_INDEX", "ucx-auto"), + os.environ.get("UCX_TLS", "ucx-auto"), + ) + + self._receives: dict[str, _ReceiveState] = {} + self._reusable_receive_buffer: _RegisteredReceiveBuffer | None = None + self._remote_metadata: dict[str, bytes] = {} + self._lock = threading.RLock() + self._closed = False + self._timeout_seconds = DEFAULT_NIXL_TRANSFER_TIMEOUT_SECONDS + self._send_executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="tq-nixl-send") + + @staticmethod + def _make_agent_name() -> str: + return f"tq-{socket.gethostname()}-{os.getpid()}-{uuid4().hex[:12]}" + + @property + def agent_name(self) -> str: + """Return the stable name advertised by this runtime instance.""" + return self._agent_name + + def endpoint_metadata(self) -> bytes: + """Return serialized metadata that peers need to address this agent.""" + with self._lock: + self._ensure_open() + return self._agent.get_agent_metadata() + + def prepare_receive(self, descriptor: Any) -> dict[str, Any]: + """Allocate or reuse registered storage for a scatter receive.""" + descriptor.validate() + if not descriptor.frame_sizes or not any(descriptor.frame_sizes): + raise NixlError("NIXL direct-frame receive requires a non-empty frame") + with self._lock: + self._ensure_open() + if descriptor.transfer_id in self._receives: + raise NixlError(f"duplicate NIXL receive: {descriptor.transfer_id}") + scratch = self._acquire_receive_buffer(descriptor.payload_bytes) + try: + initialize_packed_frame_table(scratch.buffer, descriptor.frame_sizes) + address = _buffer_address(scratch.buffer) + payload_offset = 4 + 8 * len(descriptor.frame_sizes) + regions = [] + for size in descriptor.frame_sizes: + if size: + regions.append((address + payload_offset, size, 0)) + payload_offset += size + frame_descs = self._agent.get_serialized_descs(self._agent.get_xfer_descs(regions, mem_type="DRAM")) + state = _ReceiveState(scratch, frame_descs) + metadata = self._agent.get_agent_metadata() + self._receives[descriptor.transfer_id] = state + except Exception as exc: + self._return_receive_buffer(scratch) + raise NixlError(f"failed to register NIXL receive buffer: {exc}") from exc + result = { + "agent_name": self._agent_name, + "agent_metadata": metadata, + "frame_remote_descs": state.serialized_frame_descs, + "payload_bytes": descriptor.payload_bytes, + } + return result + + def send( + self, + endpoint: dict[str, Any], + token: dict[str, Any], + descriptor: Any, + frames: tuple[bytes | bytearray | memoryview, ...], + ) -> Future[None]: + """Run the NIXL transfer on the dedicated send thread.""" + descriptor.validate() + if not descriptor.frame_sizes or not any(descriptor.frame_sizes): + raise NixlError("NIXL direct-frame send requires a non-empty frame") + if tuple(memoryview(frame).nbytes for frame in frames) != descriptor.frame_sizes: + raise NixlError(f"frame lengths do not match descriptor for {descriptor.transfer_id}") + try: + remote_name, metadata = self._validate_send_metadata(endpoint, token, descriptor) + except Exception as exc: + future: Future[None] = Future() + future.set_exception(exc) + return future + with self._lock: + self._ensure_open() + return self._send_executor.submit( + self._send_scatter, + remote_name, + metadata, + tuple(frames), + token["frame_remote_descs"], + ) + + def _acquire_receive_buffer(self, required_capacity: int) -> _RegisteredReceiveBuffer: + reusable = self._reusable_receive_buffer + if reusable is not None and reusable.capacity >= required_capacity: + self._reusable_receive_buffer = None + return reusable + + buffer = bytearray(required_capacity) + address = _buffer_address(buffer) + registrations = self._agent.register_memory( + [(address, required_capacity, 0, "")], mem_type="DRAM", backends=["UCX"] + ) + if registrations is None: + raise NixlError("failed to register NIXL receive scratch buffer") + allocated = _RegisteredReceiveBuffer(buffer, registrations) + if reusable is not None: + self._reusable_receive_buffer = None + self._deregister_registration(reusable.registrations) + return allocated + + def _return_receive_buffer(self, scratch: _RegisteredReceiveBuffer) -> None: + reusable = self._reusable_receive_buffer + if reusable is None: + self._reusable_receive_buffer = scratch + elif scratch.capacity > reusable.capacity: + self._deregister_registration(reusable.registrations) + self._reusable_receive_buffer = scratch + else: + self._deregister_registration(scratch.registrations) + + def _acquire_frame_registration(self, frame: bytes | bytearray | memoryview) -> _FrameRegistration: + view = memoryview(frame) + if view.readonly or not view.c_contiguous: + owner: bytearray | memoryview = bytearray(view) + else: + owner = view.cast("B") + address = _buffer_address(owner) + registrations = self._agent.register_memory( + [(address, memoryview(owner).nbytes, 0, "")], mem_type="DRAM", backends=["UCX"] + ) + if registrations is None: + raise NixlError("failed to register NIXL source frame") + return _FrameRegistration(owner, registrations, address, memoryview(owner).nbytes) + + def _send_scatter( + self, + remote_name: str, + metadata: bytes, + frames: tuple[bytes | bytearray | memoryview, ...], + serialized_remote_descs: bytes, + ) -> None: + """Send frames directly to matching registered remote regions.""" + frame_registrations: list[_FrameRegistration] = [] + handle = None + try: + with self._lock: + self._ensure_open() + previous = self._remote_metadata.get(remote_name) + if previous != metadata: + if previous is not None: + self._agent.remove_remote_agent(remote_name) + loaded_name = self._agent.add_remote_agent(metadata) + if isinstance(loaded_name, bytes): + loaded_name = loaded_name.decode() + if loaded_name != remote_name: + raise NixlError( + f"NIXL remote agent name mismatch: expected {remote_name!r}, got {loaded_name!r}" + ) + self._remote_metadata[remote_name] = metadata + + for frame in frames: + if memoryview(frame).nbytes: + frame_registrations.append(self._acquire_frame_registration(frame)) + local_descs = self._agent.get_xfer_descs( + [(registration.address, registration.size, 0) for registration in frame_registrations], + mem_type="DRAM", + ) + remote_descs = self._agent.deserialize_descs(serialized_remote_descs) + handle = self._agent.initialize_xfer("WRITE", local_descs, remote_descs, remote_name, backends=["UCX"]) + status = self._agent.transfer(handle) + + deadline = time.monotonic() + self._timeout_seconds + while status == "PROC": + if time.monotonic() >= deadline: + raise NixlError(f"NIXL WRITE timed out after {self._timeout_seconds:g}s") + with self._lock: + self._ensure_open() + status = self._agent.check_xfer_state(handle) + if status == "PROC": + time.sleep(0.0005) + if status != "DONE": + raise NixlError(f"NIXL WRITE failed with status {status!r}") + except NixlError: + raise + except Exception as exc: + raise NixlError(f"NIXL H2H scatter WRITE failed: {exc}") from exc + finally: + with self._lock: + if handle is not None: + try: + self._agent.release_xfer_handle(handle) + except Exception as exc: + logger.warning("failed to release NIXL transfer handle: %s", exc) + for registration in frame_registrations: + self._deregister_registration(registration.registrations) + + @staticmethod + def _validate_send_metadata( + endpoint: dict[str, Any], + token: dict[str, Any], + descriptor: Any, + ) -> tuple[str, bytes]: + remote_name = str(token.get("agent_name") or endpoint.get("agent_name") or "") + metadata = token.get("agent_metadata") or endpoint.get("agent_metadata") + if not remote_name or not isinstance(metadata, bytes): + raise NixlError("NIXL endpoint is missing remote agent metadata") + if int(token.get("payload_bytes", -1)) != descriptor.payload_bytes: + raise NixlError("NIXL receive token length does not match descriptor") + if not isinstance(token.get("frame_remote_descs"), bytes): + raise NixlError("NIXL receive token is missing frame descriptors") + return remote_name, metadata + + def receive(self, descriptor: Any) -> Future[memoryview]: + """Complete a prepared receive and return detached packed payload bytes.""" + descriptor.validate() + future: Future[memoryview] = Future() + with self._lock: + self._ensure_open() + state = self._receives.pop(descriptor.transfer_id, None) + if state is None: + future.set_exception(NixlError(f"no prepared NIXL receive for {descriptor.transfer_id}")) + return future + try: + payload = state.scratch.buffer[: descriptor.payload_bytes] + future.set_result(memoryview(payload)) + except Exception as exc: + future.set_exception(NixlError(f"failed to copy NIXL receive payload: {exc}")) + finally: + self._return_receive_buffer(state.scratch) + return future + + def cancel_receive(self, transfer_id: str) -> None: + """Cancel a prepared receive and deregister its scratch buffer.""" + with self._lock: + state = self._receives.pop(transfer_id, None) + if state is None: + return + self._deregister_registration(state.scratch.registrations) + + def close(self) -> None: + """Stop the send executor and release all NIXL registrations.""" + with self._lock: + if self._closed: + return + self._closed = True + self._send_executor.shutdown(wait=True, cancel_futures=True) + with self._lock: + for state in self._receives.values(): + self._deregister_registration(state.scratch.registrations) + self._receives.clear() + if self._reusable_receive_buffer is not None: + self._deregister_registration(self._reusable_receive_buffer.registrations) + self._reusable_receive_buffer = None + self._agent = None + + def _ensure_open(self) -> None: + if self._closed: + raise NixlError("NIXL runtime is closed") + + def _deregister_registration(self, registrations: Any) -> None: + try: + self._agent.deregister_memory(registrations, backends=["UCX"]) + except Exception as exc: + logger.warning("failed to deregister NIXL memory: %s", exc) diff --git a/transfer_queue/storage/payload_transfer/zmq.py b/transfer_queue/storage/payload_transfer/zmq.py new file mode 100644 index 00000000..8abb3cf3 --- /dev/null +++ b/transfer_queue/storage/payload_transfer/zmq.py @@ -0,0 +1,173 @@ +# 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. + +"""Legacy inline ZMQ payload transfer strategy for SimpleStorage.""" + +from __future__ import annotations + +import os +from typing import Any, Callable, NoReturn + +import zmq +import zmq.asyncio + +from transfer_queue.storage.payload_transfer.base import PayloadTransfer +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 ZMQMessage, ZMQRequestType + +logger = get_logger(__name__) +TQ_NUM_THREADS = int(os.environ.get("TQ_NUM_THREADS", 8)) +TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT = int(os.environ.get("TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT", 200)) + + +class ZmqPayloadTransfer(PayloadTransfer): + """Keep the original decoded-data ZMQ request protocol behind the contract.""" + + async def put( + self, + *, + control_socket: zmq.asyncio.Socket, + sender_id: str, + target_id: str, + global_indexes: list[int], + data: dict[str, Any], + data_parser: Callable[[Any], Any] | None, + ) -> None: + request = ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA, + sender_id=sender_id, + receiver_id=target_id, + body={"global_indexes": global_indexes, "data": data, "data_parser": data_parser}, + ) + try: + await control_socket.send_multipart(request.serialize(), copy=False) + response = ZMQMessage.deserialize(await control_socket.recv_multipart(copy=False)) + if response.request_type != ZMQRequestType.PUT_DATA_RESPONSE: + raise RuntimeError( + f"Failed to put data to storage unit {target_id}: {response.body.get('message', 'Unknown error')}" + ) + except zmq.error.Again as exc: + self._raise_timeout(sender_id, target_id, "put", exc) + except Exception as exc: + logger.error( + f"[{sender_id}]: Unexpected error during put to storage unit {target_id}: {type(exc).__name__}: {exc}" + ) + raise RuntimeError(f"Error in put to storage unit {target_id}: {type(exc).__name__}: {exc}") from exc + + async def get( + self, + *, + control_socket: zmq.asyncio.Socket, + sender_id: str, + target_id: str, + global_indexes: list[int], + fields: list[str], + ) -> dict[str, Any]: + request = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA, + sender_id=sender_id, + receiver_id=target_id, + body={"global_indexes": global_indexes, "fields": fields}, + ) + try: + await control_socket.send_multipart(request.serialize()) + response = ZMQMessage.deserialize(await control_socket.recv_multipart(copy=False)) + if response.request_type != ZMQRequestType.GET_DATA_RESPONSE: + raise RuntimeError( + f"Failed to get data from storage unit {target_id}: {response.body.get('message', 'Unknown error')}" + ) + return response.body["data"] + except zmq.error.Again as exc: + self._raise_timeout(sender_id, target_id, "get", exc) + except Exception as exc: + logger.error(f"[{sender_id}]: Unexpected error from storage unit {target_id}: {type(exc).__name__}: {exc}") + raise RuntimeError( + f"Error getting data from storage unit {target_id}: {type(exc).__name__}: {exc}" + ) from exc + + def handle_request( + self, + request: ZMQMessage, + *, + storage_id: str, + load_data: Callable[..., dict[str, Any]], + store_data: Callable[..., None], + ) -> ZMQMessage | None: + if request.request_type == ZMQRequestType.PUT_DATA: + try: + with limit_pytorch_auto_parallel_threads( + target_num_threads=TQ_NUM_THREADS, info=f"[{storage_id}] _handle_put" + ): + store_data( + request.body["global_indexes"], + request.body["data"], + request.body.get("data_parser"), + ) + return self._response(ZMQRequestType.PUT_DATA_RESPONSE, storage_id) + except Exception as exc: + return ZMQMessage.create( + request_type=ZMQRequestType.PUT_ERROR, + sender_id=storage_id, + body={ + "message": f"Failed to put data into storage unit id " + f"#{storage_id}, detail error message: {str(exc)}" + }, + ) + + if request.request_type == ZMQRequestType.GET_DATA: + try: + fields = request.body["fields"] + global_indexes = request.body["global_indexes"] + with limit_pytorch_auto_parallel_threads( + target_num_threads=TQ_NUM_THREADS, info=f"[{storage_id}] _handle_get" + ): + data = load_data(fields, global_indexes) + return self._response(ZMQRequestType.GET_DATA_RESPONSE, storage_id, {"data": data}) + except Exception as exc: + logger.error( + f"[{storage_id}]: _handle_get error, " + f"fields={fields}, global_indexes={global_indexes}: {type(exc).__name__}: {exc}" + ) + return ZMQMessage.create( + request_type=ZMQRequestType.GET_ERROR, + sender_id=storage_id, + body={ + "message": f"Failed to get data from storage unit id #{storage_id}, " + f"detail error message: {str(exc)}" + }, + ) + + return None + + @staticmethod + def _response(request_type: ZMQRequestType, sender_id: str, body: dict[str, Any] | None = None) -> ZMQMessage: + return ZMQMessage.create(request_type=request_type, sender_id=sender_id, body=body or {}) + + @staticmethod + def _raise_timeout(sender_id: str, target_id: str, operation: str, error: Exception) -> NoReturn: + timeout = TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT + if operation == "put": + logger.error( + f"[{sender_id}]: ZMQ recv timeout ({timeout}s) during put to storage unit {target_id}. " + "The storage unit may be overloaded or crashed." + ) + raise RuntimeError(f"ZMQ recv timeout ({timeout}s) during put to storage unit {target_id}") from error + + logger.error( + f"[{sender_id}]: ZMQ recv timeout ({timeout}s) from storage unit {target_id}. " + "The storage unit may be overloaded or crashed." + ) + raise RuntimeError(f"ZMQ recv timeout ({timeout}s) from storage unit {target_id}") from error diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index f80f905e..a7f8bb73 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -17,6 +17,7 @@ import pickle import time import weakref +from collections.abc import Mapping from threading import Event, Thread from typing import TYPE_CHECKING, Any from uuid import uuid4 @@ -25,6 +26,7 @@ import ray import zmq +from transfer_queue.storage.payload_transfer import PayloadTransfer, create_payload_transfer from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads from transfer_queue.utils.enum_utils import Role from transfer_queue.utils.logging_utils import get_logger @@ -166,17 +168,23 @@ class SimpleStorageUnit: zmq_server_info: ZMQ connection information for clients. """ - def __init__(self, storage_unit_size: int | None = None): + def __init__( + self, + storage_unit_size: int | None = None, + payload_transfer: Mapping[str, object] | None = None, + ): """Initialize a SimpleStorageUnit with the specified size. Args: storage_unit_size: Maximum number of elements that can be stored in this storage unit. If None, the storage unit has unlimited capacity. + payload_transfer: Backend and optional backend-specific settings. """ self.storage_unit_id = f"TQ_STORAGE_UNIT_{uuid4().hex[:8]}" self.storage_unit_size = storage_unit_size self.storage_data = StorageUnitData(self.storage_unit_size) + self.payload_transfer = create_payload_transfer(payload_transfer) # Internal communication address for proxy and workers self._inproc_addr = f"inproc://simple_storage_workers_{self.storage_unit_id}" @@ -204,6 +212,7 @@ def __init__(self, storage_unit_size: int | None = None): self.proxy_thread, self.zmq_context, self.put_get_socket, + self.payload_transfer, ) def _init_zmq_socket(self) -> None: @@ -311,23 +320,32 @@ def _worker_routine(self) -> None: try: logger.debug(f"[{self.storage_unit_id}]: worker received operation: {operation}") - # Process request - 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.GET_DATA: # type: ignore[arg-type] - with monitor.measure(op_type="GET_DATA"): - response_msg = self._handle_get(request_msg) - elif operation == ZMQRequestType.CLEAR_DATA: # type: ignore[arg-type] + if operation in (ZMQRequestType.PUT_DATA, ZMQRequestType.GET_DATA): + with monitor.measure(op_type=operation.name): + response_msg = self.payload_transfer.handle_request( + request_msg, + storage_id=self.storage_unit_id, + load_data=self._load_data, + store_data=self._put_decoded_data, + ) + else: + response_msg = self.payload_transfer.handle_request( + request_msg, + storage_id=self.storage_unit_id, + load_data=self._load_data, + store_data=self._put_decoded_data, + ) + + if response_msg is None and operation == ZMQRequestType.CLEAR_DATA: # type: ignore[arg-type] with monitor.measure(op_type="CLEAR_DATA"): response_msg = self._handle_clear(request_msg) - elif operation == ZMQRequestType.GET_METRICS: # type: ignore[arg-type] + elif response_msg is None and operation == ZMQRequestType.GET_METRICS: # type: ignore[arg-type] response_msg = self._handle_get_metrics() - elif operation == ZMQRequestType.SAVE_STORAGE_CHECKPOINT: # type: ignore[arg-type] + elif response_msg is None and operation == ZMQRequestType.SAVE_STORAGE_CHECKPOINT: # type: ignore[arg-type] response_msg = self._handle_save_checkpoint(request_msg) - elif operation == ZMQRequestType.LOAD_STORAGE_CHECKPOINT: # type: ignore[arg-type] + elif response_msg is None and operation == ZMQRequestType.LOAD_STORAGE_CHECKPOINT: # type: ignore[arg-type] response_msg = self._handle_load_checkpoint(request_msg) - else: + elif response_msg is None: response_msg = ZMQMessage.create( request_type=ZMQRequestType.PUT_GET_OPERATION_ERROR, # type: ignore[arg-type] sender_id=self.storage_unit_id, @@ -357,129 +375,45 @@ def _worker_routine(self) -> None: poller.unregister(worker_socket) worker_socket.close(linger=0) - def _handle_put(self, data_parts: ZMQMessage) -> ZMQMessage: - """ - Handle put request, add or update data into storage unit. - - Args: - data_parts: ZMQMessage from client. + def _put_decoded_data(self, global_indexes: list[int], field_data: dict[str, Any], data_parser: Any) -> None: + """Validate parsed data and store it.""" + if data_parser is not None: + if not callable(data_parser): + raise TypeError(f"data_parser must be callable, got {type(data_parser).__name__}") + original_keys = set(field_data) + original_lengths = {key: self._field_length(value) for key, value in field_data.items()} + field_data = data_parser(field_data) + if not isinstance(field_data, dict): + raise TypeError(f"data_parser must return a dict, got {type(field_data).__name__}") + if set(field_data) != original_keys: + raise ValueError( + f"data_parser must not change dict keys. Original keys: {sorted(original_keys)}, " + f"got: {sorted(field_data)}" + ) + for key, value in field_data.items(): + original_length = original_lengths[key] + new_length = self._field_length(value) + if original_length is not None and new_length is not None and original_length != new_length: + raise ValueError( + f"data_parser changed the number of elements for key '{key}': " + f"expected {original_length}, got {new_length}" + ) + self.storage_data.put_data(field_data, global_indexes) - Returns: - Put data success response ZMQMessage. - """ + @staticmethod + def _field_length(value: Any) -> int | None: + if hasattr(value, "shape") and isinstance(value.shape, tuple | list) and len(value.shape) > 0: + return value.shape[0] try: - global_indexes = data_parts.body["global_indexes"] - field_data = data_parts.body["data"] # field_data should be a dict. - data_parser = data_parts.body.get("data_parser", None) - - with limit_pytorch_auto_parallel_threads( - target_num_threads=TQ_NUM_THREADS, info=f"[{self.storage_unit_id}] _handle_put" - ): - if data_parser is not None: - if not callable(data_parser): - raise TypeError(f"data_parser must be callable, got {type(data_parser).__name__}") - - original_keys = set(field_data.keys()) - original_lengths = {} - for k, v in field_data.items(): - if hasattr(v, "shape") and isinstance(v.shape, tuple | list) and len(v.shape) > 0: - original_lengths[k] = v.shape[0] - else: - try: - original_lengths[k] = len(v) - except Exception: - original_lengths[k] = None - - field_data = data_parser(field_data) - - if not isinstance(field_data, dict): - raise TypeError(f"data_parser must return a dict, got {type(field_data).__name__}") - - new_keys = set(field_data.keys()) - if new_keys != original_keys: - raise ValueError( - f"data_parser must not change dict keys. " - f"Original keys: {sorted(original_keys)}, got: {sorted(new_keys)}" - ) - - for k, v in field_data.items(): - if hasattr(v, "shape") and isinstance(v.shape, tuple | list) and len(v.shape) > 0: - new_len = v.shape[0] - else: - try: - new_len = len(v) - except Exception: - new_len = None - - orig_len = original_lengths[k] - if orig_len is not None and new_len is not None and orig_len != new_len: - raise ValueError( - f"data_parser changed the number of elements for key '{k}': " - f"expected {orig_len}, got {new_len}" - ) - self.storage_data.put_data(field_data, global_indexes) - - # After put operation finish, send a message to the client - response_msg = ZMQMessage.create( - request_type=ZMQRequestType.PUT_DATA_RESPONSE, # type: ignore[arg-type] - sender_id=self.storage_unit_id, - body={}, - ) - - return response_msg - except Exception as e: - return ZMQMessage.create( - request_type=ZMQRequestType.PUT_ERROR, # type: ignore[arg-type] - sender_id=self.storage_unit_id, - body={ - "message": f"Failed to put data into 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. - - Args: - data_parts: ZMQMessage from client. + return len(value) + except Exception: + return None - Returns: - Get data success response ZMQMessage, containing target data. - """ + def _load_data(self, fields: list[str], global_indexes: list[int]) -> dict[str, list]: try: - fields = data_parts.body["fields"] - global_indexes = data_parts.body["global_indexes"] - - with limit_pytorch_auto_parallel_threads( - target_num_threads=TQ_NUM_THREADS, info=f"[{self.storage_unit_id}] _handle_get" - ): - result_data = self.storage_data.get_data(fields, global_indexes) - - response_msg = ZMQMessage.create( - request_type=ZMQRequestType.GET_DATA_RESPONSE, # type: ignore[arg-type] - sender_id=self.storage_unit_id, - body={ - "data": result_data, - }, - ) - except Exception as e: - key_not_found = isinstance(e, StorageKeyNotFoundError) - log = logger.debug if key_not_found else logger.error - log( - f"[{self.storage_unit_id}]: _handle_get error, " - f"fields={fields}, global_indexes={global_indexes}: {type(e).__name__}: {e}" - ) - marker = f"[{KEY_NOT_FOUND_MARKER}] " if key_not_found else "" - response_msg = ZMQMessage.create( - request_type=ZMQRequestType.GET_ERROR, # type: ignore[arg-type] - sender_id=self.storage_unit_id, - body={ - "message": f"{marker}Failed to get data from storage unit id #{self.storage_unit_id}, " - f"detail error message: {str(e)}" - }, - ) - return response_msg + return self.storage_data.get_data(fields, global_indexes) + except StorageKeyNotFoundError as exc: + raise StorageKeyNotFoundError(f"{KEY_NOT_FOUND_MARKER}: {exc}") from exc def _handle_clear(self, data_parts: ZMQMessage) -> ZMQMessage: """ @@ -691,6 +625,7 @@ def _shutdown_resources( proxy_thread: Thread | None, zmq_context: zmq.Context | None, put_get_socket: zmq.Socket | None, + payload_transfer: PayloadTransfer, ) -> None: """Clean up resources on garbage collection.""" logger.info("Shutting down SimpleStorageUnit resources...") @@ -712,6 +647,8 @@ def _shutdown_resources( if proxy_thread and proxy_thread.is_alive(): proxy_thread.join(timeout=5) + payload_transfer.close() + logger.info("SimpleStorageUnit resources shutdown complete.") def start_metrics(self, port: int = 0) -> str: @@ -742,3 +679,10 @@ def get_zmq_server_info(self) -> ZMQServerInfo: ZMQServerInfo containing connection details for this storage unit. """ return self.zmq_server_info + + def get_payload_transfer_info(self) -> dict[str, Any] | None: + """Return bootstrap-safe transfer endpoint metadata.""" + info = self.payload_transfer.bootstrap_info() + if info is None: + return None + return {"id": self.storage_unit_id, **info} diff --git a/transfer_queue/utils/serial_utils.py b/transfer_queue/utils/serial_utils.py index 2ba8e7d1..54c68259 100644 --- a/transfer_queue/utils/serial_utils.py +++ b/transfer_queue/utils/serial_utils.py @@ -462,14 +462,44 @@ def pack_into(target_buffer: bytestr, items: Sequence[bytestr]) -> None: def unpack_from(source_buffer: bytestr) -> list[memoryview]: """Split a packed buffer back into N memoryview slices over ``source_buffer``.""" mv = memoryview(source_buffer) + if mv.nbytes < _PACK_HEADER_SIZE: + raise ValueError("unpack_from: payload is shorter than its header") item_count = struct.unpack_from(_PACK_HEADER_FMT, mv, 0)[0] + payload_start = _PACK_HEADER_SIZE + item_count * _PACK_ENTRY_SIZE + if payload_start > mv.nbytes: + raise ValueError("unpack_from: frame table exceeds payload size") + result: list[memoryview] = [] for i in range(item_count): offset, length = struct.unpack_from(_PACK_ENTRY_FMT, mv, _PACK_HEADER_SIZE + i * _PACK_ENTRY_SIZE) + if offset < payload_start or length > mv.nbytes - offset: + raise ValueError(f"unpack_from: frame {i} is outside the payload") result.append(mv[offset : offset + length]) return result +def initialize_packed_frame_table(target_buffer: bytestr, frame_sizes: Sequence[int]) -> None: + """Write only the packed frame header/table into a receive buffer. + + The frame payload bytes are filled directly by a scatter/gather receive, + so this avoids copying them into a second contiguous buffer. + """ + if any(size < 0 for size in frame_sizes): + raise ValueError("frame sizes must be non-negative") + target_mv = memoryview(target_buffer) + payload_start = _PACK_HEADER_SIZE + len(frame_sizes) * _PACK_ENTRY_SIZE + required = payload_start + sum(frame_sizes) + if target_mv.nbytes < required: + raise ValueError(f"frame table target has {target_mv.nbytes} bytes, requires {required}") + struct.pack_into(_PACK_HEADER_FMT, target_mv, 0, len(frame_sizes)) + entry_offset = _PACK_HEADER_SIZE + payload_offset = payload_start + for size in frame_sizes: + struct.pack_into(_PACK_ENTRY_FMT, target_mv, entry_offset, payload_offset, size) + entry_offset += _PACK_ENTRY_SIZE + payload_offset += size + + def batch_encode_into( objs: list[Any], alloc_buff_func: Callable[[list[int]], list[Any]], diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index 2ca63d7f..f63c39e7 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -53,6 +53,14 @@ class ZMQRequestType(ExplicitEnum): PUT_DATA = "PUT" GET_DATA_RESPONSE = "GET_DATA_RESPONSE" PUT_DATA_RESPONSE = "PUT_DATA_RESPONSE" + PUT_DATA_PREPARE = "PUT_DATA_PREPARE" + PUT_DATA_READY = "PUT_DATA_READY" + PUT_DATA_COMMIT = "PUT_DATA_COMMIT" + PUT_DATA_CANCEL = "PUT_DATA_CANCEL" + GET_DATA_PREPARE = "GET_DATA_PREPARE" + GET_DATA_READY = "GET_DATA_READY" + GET_DATA_COMMIT = "GET_DATA_COMMIT" + GET_DATA_CANCEL = "GET_DATA_CANCEL" CLEAR_DATA = "CLEAR_DATA" CLEAR_DATA_RESPONSE = "CLEAR_DATA_RESPONSE"