diff --git a/README.md b/README.md index 0e026e80..44e21730 100644 --- a/README.md +++ b/README.md @@ -211,6 +211,10 @@ pip install TransferQueue pip install dist/*.whl ``` +For the optional SimpleStorage UCX Host RDMA path, including official UCX installation, +native extension build, configuration, and lane verification, see the +[TQ UCX Host RDMA developer guide](docs/ucx_rdma_developer_guide.md). +

📊 Performance

### Simple Case: Regular Tensor @@ -345,4 +349,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/hixl_a2_rh2h_experiment_and_selection.md b/docs/hixl_a2_rh2h_experiment_and_selection.md new file mode 100644 index 00000000..65a85fb9 --- /dev/null +++ b/docs/hixl_a2_rh2h_experiment_and_selection.md @@ -0,0 +1,271 @@ +# A2-26 <-> A2-27 HIXL rH2H experiment and TQ selection note + +> This file is a historical validation record for one test cluster. Paths, +> addresses, devices, versions, and commands are not TransferQueue product configuration. + +Date: 2026-08-04 +Scope: decide whether HIXL can be the RDMA backend for TransferQueue +`SimpleStorage` Host payloads on the current A2-26/27/28 cluster. + +## Executive conclusion + +HIXL `rH2H` is documented as supported for **A2 RoCE**, but it is not a +CPU-only Host-NIC RDMA abstraction on the exercised implementations. Both +available HIXL connection paths require Ascend device communication resources: + +| HIXL path | How selected | Actual dependency | A2-26 <-> 27 result | +| --- | --- | --- | --- | +| legacy/default | no `LocalCommRes`, which is the benchmark default | ADXL -> HCCL communicator -> HCCL one-sided Get/Put | fails because generated rank table has an empty NPU `device_ip` | +| 1.3 / recommended | `LocalCommRes={"version":"1.3"}` | HixlCS endpoint generation; reads NPU HCCN IP from `hccn.conf` or `hccn_tool` | fails before connection because NPU 0 has no HCCN IP | + +Therefore HIXL is **not yet a Go** for TQ SimpleStorage on this cluster. This +is an environment-precondition result, not a claim that A2 `rH2H` is absent. +No successful transfer or bandwidth number has been measured. + +For a TQ backend that must work on CPU-only nodes or must use only the Host +RoCE NIC, HIXL has the wrong dependency boundary. It can be reconsidered for +an Ascend deployment after HCCN is configured and an end-to-end HIXL test +passes. A Host RDMA backend should otherwise use an independent Host transport +(UCX or a small native-Verbs service). + +## What the official 9.1.0 benchmark states + +Source inspected: HIXL branch `9.1.0`, commit +`c12f9a56ab66299f62adb7f1f3f34d92e725e856`. + +`benchmarks/README.md` defines: + +| direction | meaning | operation | +| --- | --- | --- | +| `H2rH` | Host writes remote Host | write | +| `rH2H` | Host reads remote Host | read | + +The same document's support table says the A2 RoCE path supports all eight +memory directions, including these two. Its documented dual-host invocation +is: + +```bash +# target host +python3 benchmarks/comm_benchmark/scripts/run_comm_benchmark.py \ + --role=target --transport=rdma + +# initiator host: run the command printed by target +``` + +This proves the intended product capability. It does not by itself establish +that a bare Host NIC can be used without an Ascend communication environment. + +## Environment actually used + +| item | A2-26 | A2-27 | +| --- | --- | --- | +| host RoCE test address | `178.123.4.4` | `178.123.4.3` | +| NPU | Ascend 910B3 (A2) | Ascend 910B3 (A2) | +| driver / HDK | 25.5.1 | 25.5.1 | +| installed Toolkit and 910B ops | 9.1.0-beta.1 | 9.1.0-beta.1 | +| default `/usr/local/Ascend/cann` | **9.0.0** | 9.1.0-beta.1 | + +All benchmark commands explicitly sourced: + +```bash +source /usr/local/Ascend/cann-9.1.0-beta.1/set_env.sh +``` + +On A2-27, the official source branch was built exactly as documented: + +```bash +cd /tmp/hixl-v910-benchmark +bash build.sh --examples +``` + +The build completed and produced `hixl_comm_bench`. The generated HIXL package +was **not installed**; system CANN, driver, HCCN, and network configuration +were not changed. A2-26 lacks `cmake`, so it used the same temporary benchmark +binary and temporary source-built `libcann_hixl.so` copied from A2-27. + +## Experiment A: benchmark default path + +Test intent: + +```text +target: A2-27, 178.123.4.3:16000 +initiator: A2-26, 178.123.4.4:16001 +direction: rH2H +transport: rdma +block: 64 KiB +``` + +Observed before failure: + +1. Both processes initialized HIXL. +2. Both allocated and registered Host memory with `aclrtMallocHost`. +3. TCP coordination connected and exchanged the remote Host address. +4. The initiator failed during HIXL `Connect`. + +Exact failure: + +```text +Config_Error_Ranktable(EI0014): Value [] for ranktable variable [device_ip] is invalid +HcclCommInitClusterInfoMemConfig ... fail +``` + +### Why this default used HCCL + +This is source behavior, not an inference from the error: + +```text +EngineFactory::CreateEngine + no LocalCommRes option + -> CommEngine + -> adxl::AdxlInnerEngine + -> CommChannel::InitializeHcclComm + -> HcclCommInitClusterInfoMemConfig + +CommChannel::TransferSync(read) + -> HcclBatchGet +``` + +Relevant source files: + +- `src/hixl/engine/engine_factory.cc` +- `src/hixl/engine/comm_engine.cc` +- `src/llm_datadist/adxl/comm_channel.cc` +- `src/llm_datadist/hccl/hccl_adapter.cc` + +The default benchmark supplies `BufferPool=0:0` but no `LocalCommRes`; it +therefore intentionally selects this compatibility/collective-communicator +backend. + +## Experiment B: documented 1.3 path + +The current `docs/cpp/HIXL接口.md` distinguishes the paths: + +```text +LocalCommRes empty / 1.0 / 1.2 -> collective-communication-domain connection +LocalCommRes = {"version":"1.3"} -> HixlCS connection (recommended) + requires HDK >= 25.5 and Toolkit >= 9.1 +``` + +We reran the official target command with: + +```bash +-H='LocalCommRes={"version":"1.3"}' +``` + +Important benchmark syntax: the executable requires the `-H=KEY=VALUE` form; +`-H KEY=VALUE` is rejected by its parser. + +The target accepted the option and its effective configuration showed: + +```text +LocalCommRes={"version":"1.3"} +``` + +It then failed **before** TCP/HIXL connection: + +```text +Failed to get device ip from hccn.conf and hccn_tool, phy_device_id:0 + -> EndpointGenerator::BuildRoceEndpoint + -> EndpointGenerator::BuildEndpointList + -> HixlEngine::Initialize +``` + +This agrees with the generic HIXL source: `LocalCommRes` version 1.3 selects +`HixlEngine`; its `DirectClientHandler` calls HixlCS APIs. This run did not +reach a data transfer, so it does not prove the eventual on-wire data path. +It does prove that HixlCS needs an NPU HCCN IP in order to construct its RoCE +endpoint. + +## Live HCCN check + +On both hosts, NPU 0 and NPU 1 currently return: + +```text +Get ipconf failed, because no ip was preset there! +``` + +Thus the failures are expected from the current device network state: + +```text +default path: HCCL rank table needs device_ip +1.3 path: HixlCS endpoint generator needs device_ip +``` + +The Host NIC addresses (`178.123.4.4` and `178.123.4.3`) do not substitute for +per-NPU HCCN IPs. + +## Implications for TransferQueue + +### HIXL option + +Use HIXL only if all of the following are acceptable: + +- each TQ worker has an Ascend device and an ACL runtime context; +- the participating NPU HCCN IPs and links are configured and reachable; +- a benchmark verifies the exact HIXL route selected by the application; +- TQ accepts that a Host payload transport has an NPU/HCCN lifecycle and + operational dependency. + +After HCCN is made available, rerun the 1.3 test first, then validate 64 KiB, +1 MiB, and 16 MiB in both directions, concurrency, reconnect, and peer exit. + +### Independent Host RDMA option + +Choose an independent Host transport if SimpleStorage must run without NPU or +HCCN setup. On this cluster, ordinary Host Verbs bandwidth testing has already +worked, whereas UCX `rc_verbs` is separately blocked by an HNS UD-QP bootstrap +issue. That is a UCX/provider compatibility problem, not a test of HIXL. + +Do not claim a TQ speedup until the selected transport has passed the full +payload lifecycle and has measured end-to-end timings. + +## Reference implementations and corrected comparison + +This section separates verified implementation facts from an adoption +recommendation. A framework that can register an accelerator buffer does not +automatically provide every H2H/H2D/D2H/D2D direction as a portable backend. + +| system / option | confirmed implementation boundary | useful TQ lesson | qualification | +| --- | --- | --- | --- | +| Mooncake | Its C++ RDMA transport directly manages Verbs resources: PD, CQ, QP, MR and `ibv_reg_mr`. Its CUDA build distinguishes Host and device pointers and can use dma-buf or NVIDIA peer-memory registration for device RDMA. Recent releases also contain Ascend HIXL RoCE samples and Ascend-direct work. | Reuse the design idea: a uniform transfer interface above native, accelerator-specific transports; cache registrations and make teardown explicit. | Do **not** summarize this as “Mooncake Ascend = HIXL with every memory direction proven”. Exact Ascend behavior is release-, CANN-, hardware-, and route-dependent. | +| YuanRong DataSystem TransferEngine | Python API is pybind11 (`yr.datasystem.TransferEngine`). It exposes initialization, registration, synchronous read and batch synchronous read. Its Python API requires `npu:${id}` and does not expose write or async operations. The current source tree has recent HIXL D2D backend work. | Keep the Python surface narrow and isolate the native backend behind one engine interface; preserve a fallback path. | It is **not** evidence of a generic Python Host H2H RDMA backend, nor evidence that all NPU H2H/H2D/D2H/D2D directions are exposed. | +| UCX/UCP | UCP provides C/C++ high-level communication APIs (Tag, AM, RMA) over multiple transports and supports CUDA memory when built with CUDA support. It may select GPUDirect RDMA when the build, NIC, driver, GPU topology and selected transport permit it. | Best initial portable candidate for a Host/GPU C++ backend when the target environment is known-good. | Tag or AM use does not itself guarantee zero-copy/GDR. On these A2 hosts, UCP `rc_verbs` is not currently usable because the HNS UD-QP bootstrap fails; do not declare UCX a current Go without fixing or bypassing that compatibility issue. | +| HIXL | On A2, the benchmark documents `rH2H`, but the tested default and 1.3 routes both require NPU communication configuration. | Use only as the Ascend-specific backend, with lifecycle tied to ACL and HCCN. | Not a CPU-only Host-NIC RDMA backend on the current cluster. | +| native Verbs | Direct `libibverbs` gives complete control of Host RC QPs, registrations, completions, and connection metadata exchange. Basic Host Verbs traffic has passed in this environment. | Keep as a targeted fallback when UCX cannot support a provider/topology. | Highest implementation and operational cost; do not add it merely as a speculative performance optimization. | + +### TQ backend selection, revised + +```text +TQ Python layer + thin pybind11 interface only; it does not poll CQs or own RDMA lifecycle + +Host H2H + candidate 1: UCX/UCP after provider-specific acceptance passes + candidate 2: a small C++ Verbs backend if UCX remains incompatible + fallback: existing TCP/ZMQ path + +NVIDIA GPU + CUDA-aware UCX backend; enable/measure GDR only when runtime capability + detection proves it is active + +Ascend NPU + HIXL backend; require compatible CANN/HIXL, ACL context, HCCN IP and links +``` + +The initial TQ proposal must **not** be described as a low-cost simultaneous +implementation of `UcpBackend + HixlBackend + TcpBackend`. These are three +independent lifecycle, registration and failure-recovery systems. A practical +sequence is: preserve TCP/ZMQ fallback, validate one selected Host backend, +then add the Ascend-specific HIXL backend only after its HCCN acceptance gate +passes. + +## Decision status + +| decision question | status | +| --- | --- | +| Does 9.1.0 benchmark document A2 `rH2H`? | Yes | +| Does the current 26/27 environment meet the HCCN prerequisite? | No | +| Did default HIXL `rH2H` transfer pass? | No, HCCL `device_ip` absent | +| Did HixlCS 1.3 transfer pass? | No, HixlCS endpoint `device_ip` absent | +| Is HIXL a CPU-only SimpleStorage RDMA backend on this cluster? | No | +| Is HIXL permanently ruled out for A2? | No; configure HCCN and rerun 1.3 | diff --git a/docs/rdma_official_ucx_20260819.md b/docs/rdma_official_ucx_20260819.md new file mode 100644 index 00000000..1018f889 --- /dev/null +++ b/docs/rdma_official_ucx_20260819.md @@ -0,0 +1,98 @@ +# 官方 UCX v1.22.0:A2 Host H2H 验证记录 + +日期:2026-08-19 + +本记录只描述未修改的官方 UCX release,不代表 TQ 的默认部署配置。当前 A2 机器只保留 +`/opt/tq-ucx/1.22.0-official`;历史测试目录 `/opt/tq-ucx/1.18.1` 和 +`/opt/tq-ucx/11307` 已删除。 + +开发者要复现安装和 TQ 使用流程,请先看 +[`ucx_rdma_developer_guide.md`](ucx_rdma_developer_guide.md);本文只保存本次 A2 实验的 +具体结果。 + +## 环境 + +| 项目 | 值 | +| --- | --- | +| UCX | 1.22.0,revision `8a6b06f` | +| A2-26 | `178.123.4.4`, `hns_0:1`, `enp189s0f0` | +| A2-27 | `178.123.4.3`, `hns_0:1`, `enp189s0f0` | +| GID | index `3`,IPv4-mapped RoCE GID | + +UCX runtime 使用官方 AArch64 RPM 包安装,未修改 UCX 源码。测试时设置: + +```bash +export TQ_UCX_HOME=/opt/tq-ucx/1.22.0-official +export PATH="$TQ_UCX_HOME/bin:$PATH" +export LD_LIBRARY_PATH="$TQ_UCX_HOME/lib64:$TQ_UCX_HOME/lib${LD_LIBRARY_PATH:+:$LD_LIBRARY_PATH}" +export UCX_IB_GID_INDEX=3 +export UCX_IB_ADDR_TYPE=ib_global +``` + +## 原生 UCX 结果 + +### 纯 RC 失败 + +```bash +UCX_TLS=rc_verbs +UCX_NET_DEVICES=hns_0:1 +ucx_perftest -t tag_bw ... +``` + +结果: + +```text +no auxiliary transport ... Unsupported operation +ucp_ep_create() failed: Destination is unreachable +``` + +这确认当前 HNS 环境的 UD auxiliary 仍不可用;官方 v1.22.0 没有自动消除这个硬件/驱动 +组合上的限制。 + +### RC 数据 + TCP 辅助通道成功 + +```bash +export UCX_TLS=rc_verbs,tcp,sm,self +export UCX_NET_DEVICES=hns_0:1,enp189s0f0 +``` + +`ucx_perftest -t tag_bw -s 1048576 -n 20 -w 2` 跨 A2-26/A2-27 成功,日志为: + +```text +tag(rc_verbs/hns_0:1) ka(tcp/enp189s0f0) +Final: 20 ... 110.33 MB/s +``` + +因此实际数据通道仍是 RC,TCP 只承担 wireup/keepalive 辅助通道。之前只配置 +`hns_0:1` 时 TCP 不在候选设备中,不能形成辅助通道。 + +## TQ SimpleStorage 结果 + +`transfer_queue._ucx` 使用官方 v1.22.0 头文件和库重新构建,在相同环境下完成跨节点 +SimpleStorage 操作: + +```text +UCX Version 1.22.0 +ucp_context ... tag(rc_verbs/hns_0:1) ka(tcp/enp189s0f0) +cross-node manager PASS ... bytes=1048576 +cross-node CLEAR PASS +``` + +本次测得 1 MiB PUT 约 20.97 MiB/s,GET 约 62.03 MiB/s;该数据只是功能验证样本,不是 +吞吐承诺。 + +## 对 TQ 代码的结论 + +当前 `ucx_discovery.py` 已有正确的基本逻辑: + +1. 根据本机 RoCE GID 找到 `hns_0:1` 和 `enp189s0f0`; +2. 通过 `ucx_info -d` 检查 runtime 是否同时提供 RC 和 TCP; +3. 未显式设置 `UCX_TLS` 时生成 `rc_verbs,tcp,sm,self`; +4. 自动生成 `UCX_NET_DEVICES=hns_0:1,enp189s0f0`。 + +本次只补充了官方 RPM 常见的 `lib64` 目录发现,避免 `ucx_info` 依赖库位于 `lib64` 时 +自动能力探测失败。显式设置 `UCX_TLS=rc_verbs` 仍然会被视为用户的强制覆盖,TQ 不会 +替用户偷偷追加 TCP;部署时不要用纯 RC 配置。 + +因此,支持本次方式不需要新增 transport 组件,也不需要修改 UCX;使用当前 TQ 代码并 +确保 UCX runtime 可发现、TCP 网卡包含在 `UCX_NET_DEVICES` 即可。 diff --git a/docs/rfcs/simple_storage_payload_transport_guide.md b/docs/rfcs/simple_storage_payload_transport_guide.md new file mode 100644 index 00000000..08896465 --- /dev/null +++ b/docs/rfcs/simple_storage_payload_transport_guide.md @@ -0,0 +1,168 @@ +# SimpleStorage Payload Transfer + +开发者安装 UCX、构建 TQ native extension、启用配置和验证实际 RDMA lane 的流程见 +[`TQ UCX Host RDMA Developer Guide`](../ucx_rdma_developer_guide.md)。本文保留设计边界、 +协议和组件责任,不替代环境安装手册。 + +## Scope + +`PayloadTransfer` is a narrow, optional byte-transfer boundary for +SimpleStorage. It does not replace a storage backend and does not own routing, +serialization, object semantics, or durability. + +```text +KV API / routing / data_parser / StorageUnitData + │ + SimpleStorage protocol + (ZMQ control plane) + │ + optional PayloadTransfer + │ + UCX Host byte transfer +``` + +- `payload_transfer: zmq` is the default. It follows the original ZMQ path and + does not create a `PayloadTransfer` object. +- `payload_transfer: ucx` keeps ZMQ for control and small payloads, and uses + UCX Tagged send/receive for contiguous Host payloads of at least 128 KiB. +- A UCX failure is returned to the caller; it never silently falls back to + ZMQ after a transfer has started. +- GPU/NPU direct transfer, RMA and chunking are not implemented here. + +## Storage backend boundary + +| Backend | Relationship to `PayloadTransfer` | +| --- | --- | +| SimpleStorage | The only current consumer. It owns the PREPARE/READY/COMMIT/CANCEL protocol and invokes the transfer for payload bytes. | +| MooncakeStorage | Unchanged. It already owns its data movement and memory lifecycle; wrapping it would duplicate its backend contract. | +| YuanRongStorage | Unchanged. Its client/backend protocol remains its transport boundary. | +| RayStore | Unchanged. Ray object transfer remains owned by Ray. | + +A future backend should use this abstraction only if it shares +SimpleStorage's split control/data-plane protocol. `PayloadTransfer` is not a +mandatory layer below every storage implementation. + +HIXL can later implement the same contract when the payload is already in a +supported device buffer. Device discovery, memory registration and completion +events belong in that implementation, not in SimpleStorage or the generic +descriptor. + +## Contract + +The generic descriptor contains only: + +```text +transfer_id protocol identity and correlation +payload_bytes receive allocation and length validation +``` + +Frame count and UCX tag are not control-plane fields. The frame table is +inside the packed payload, and UCX derives its tag from `transfer_id` on both +peers. `ReceiveToken` remains transport-owned because a future transport may +need receiver-generated metadata; UCX currently returns an empty token. + +The transport API is deliberately small: + +```text +endpoint() +prepare_receive(descriptor) -> token +send(endpoint, token, descriptor, payload) -> Future +receive(descriptor) -> Future +cancel_receive(transfer_id) +close() +``` + +All UCX objects and progress are owned by one dedicated thread. Requests have +a finite timeout, and the native receive result exposes UCX's actual received +length rather than the allocation capacity. + +## PUT sequence + +```mermaid +sequenceDiagram + participant M as SimpleStorageManager + participant MT as PayloadTransfer (Manager) + participant Z as ZMQ control + participant S as SimpleStorageUnit + participant ST as PayloadTransfer (StorageUnit) + participant D as StorageUnitData + + M->>M: encode + pack frames + M->>Z: PUT_PREPARE(descriptor) + Z->>S: PUT_PREPARE + S->>ST: prepare_receive(descriptor) + S-->>Z: PUT_READY(token) + Z-->>M: PUT_READY + M->>MT: send(endpoint, token, payload) + MT-->>ST: payload bytes + M->>Z: PUT_COMMIT(transfer_id) + Z->>S: PUT_COMMIT + S->>ST: receive(descriptor) + S->>S: unpack + decode + data_parser + S->>D: put_data + S-->>M: PUT_RESPONSE +``` + +## GET sequence + +```mermaid +sequenceDiagram + participant M as SimpleStorageManager + participant MT as PayloadTransfer (Manager) + participant Z as ZMQ control + participant S as SimpleStorageUnit + participant ST as PayloadTransfer (StorageUnit) + participant D as StorageUnitData + + M->>Z: GET_PREPARE(fields, indexes, transfer_id) + Z->>S: GET_PREPARE + S->>D: get_data + S->>S: encode + pack frames + S-->>M: GET_READY(descriptor) + M->>MT: prepare_receive(descriptor) + M->>Z: GET_COMMIT(endpoint, token) + Z->>S: GET_COMMIT + S->>ST: send(endpoint, token, payload) + ST-->>MT: payload bytes + S-->>M: GET_RESPONSE + M->>MT: receive(descriptor) + M->>M: unpack + decode +``` + +PREPARE state is bounded by count, aggregate bytes and a TTL. Failed handshakes +send best-effort CANCEL, while transfer timeout remains the final guard against +an unresponsive peer. + +## Configuration and build + +```yaml +backend: + SimpleStorage: + payload_transfer: zmq # default; use ucx to opt in +``` + +The native extension is also an explicit build choice: + +```bash +TQ_BUILD_UCX=1 TQ_UCX_HOME=/path/to/ucx python -m build +``` + +UCX selection must contain a reliable-connection RDMA transport. Device and +GID discovery remains internal; initialization fails if no matching RoCE-v2 +path is found. + +## Code ownership + +| Path | Responsibility | +| --- | --- | +| `storage/payload_transfer/base.py` | Generic descriptor, endpoint, token and lifecycle contract. | +| `storage/payload_transfer/ucx.py` | UCX adapter only. | +| `storage/payload_transfer/ucx_runtime.py` | Owner thread, requests, endpoint cache, timeout and cancellation. | +| `storage/payload_transfer/ucx_discovery.py` | UCX/RoCE device, GID and transport capability discovery. | +| `csrc/ucx/ucx_bindings.cpp` | Minimal UCP Tagged binding and actual receive length. | +| `storage/simple_storage.py` | StorageUnit protocol state, decode/parser/store and resource bounds. | +| `storage/managers/simple_storage_manager.py` | Manager-side handshake and ZMQ/UCX selection. | + +Current automated tests cover the protocol and lifecycle without RDMA +hardware. Hardware end-to-end and HIXL support require separate validation and +are not implied by those tests. diff --git a/docs/ucx_rdma_developer_guide.md b/docs/ucx_rdma_developer_guide.md new file mode 100644 index 00000000..d7672b8e --- /dev/null +++ b/docs/ucx_rdma_developer_guide.md @@ -0,0 +1,304 @@ +# TQ UCX Host RDMA 开发者指南 + +本文面向需要在开发机或集群上构建、启用和验证 TQ UCX Host RDMA 的开发者。 +当前实现只覆盖 **SimpleStorage 的 Host payload transfer**:ZMQ 仍负责控制协议和小 +payload,UCX UCP Tagged Send/Receive 负责达到阈值的大块 Host buffer。 + +本文不是 UCX 通用调优手册,也不覆盖 Mooncake、openYuanrong 或 HIXL 的独立传输配置。 + +## 1. 当前能力和边界 + +```text +SimpleStorage KV/control protocol ── ZMQ(始终保留) + │ +large contiguous Host payload ── UCX Tagged + RC data lane + └─ TCP auxiliary/wireup(需要时) +``` + +- 默认配置仍为 `payload_transfer: zmq`,不需要安装 UCX。 +- `payload_transfer: ucx` 只改变 SimpleStorage 的大 payload 数据面,不替换存储后端、 + 路由、序列化或 ZMQ 控制面。 +- 当前 UCX payload 路径要求 Host memory、可靠连接 RDMA transport 和 RoCE-v2 GID。 +- 小于 `128 KiB` 的 payload 仍走 inline/ZMQ;具体阈值由当前 SimpleStorage payload + transfer 实现定义,不应当作 UCX 的通用阈值。 +- GPU/NPU 直连、GDR、RMA、自动分块不是当前 TQ UCX payload 的能力。 +- Mooncake、openYuanrong 和 RayStore 继续使用各自 backend 的 transport,不要因为安装 + UCX 就替这些 backend 设置 TQ UCX payload。 + +权威设计说明见 [SimpleStorage Payload Transfer RFC](rfcs/simple_storage_payload_transport_guide.md)。 + +## 2. 节点准入 + +每个运行 Controller、SimpleStorageUnit 或 TQ client 的节点都需要满足: + +1. Linux、可用的 `rdma-core`/libibverbs 和 RoCE/InfiniBand NIC; +2. 所有参与节点能够通过同一网络平面互通; +3. 每个节点的 Python 环境包含 TQ、PyTorch(如业务 payload 使用 Torch)、pyzmq、Ray + 和构建 native extension 所需的 `pybind11`; +4. UCX runtime、UCX native extension 和 Python 环境在每个 Ray worker 节点可见; +5. `memlock`、容器 capability 和防火墙策略允许 RDMA memory registration 与对应的 + auxiliary TCP 连接。 + +先检查底层设备,不要先把失败归因到 TQ: + +```bash +ibv_devices +ibv_devinfo +ip -br addr +``` + +多网卡机器不应根据网卡名称猜测配置。TQ 会遍历本机 RDMA sysfs,使用本机控制 IP 对应 +的 RoCE-v2 GID 选择 RDMA device、port、netdev 和 GID index;多个候选无法唯一确定时会 +拒绝初始化,而不是随机选一张卡。 + +## 3. 安装官方 UCX + +不要使用开发中的 UCX 分支或为某台机器维护私有 UCX patch。固定一个官方 release,并在 +所有节点使用同一个版本。当前验证版本是官方 [UCX v1.22.0](https://github.com/openucx/ucx/releases/tag/v1.22.0)。 + +### 3.1 官方二进制/RPM 包 + +从 release 页面选择与操作系统、架构和 verbs 栈匹配的 asset。下面以 AArch64/CentOS 兼容 +RPM 包为例;asset 名称需要按目标平台替换: + +```bash +UCX_VERSION=1.22.0 +UCX_PREFIX=/opt/tq-ucx/${UCX_VERSION}-official +UCX_ASSET=ucx-1.22.0-centos8-mofed5-cuda11-aarch64.tar.bz2 + +curl -fL --retry 2 \ + -o "/tmp/${UCX_ASSET}" \ + "https://github.com/openucx/ucx/releases/download/v${UCX_VERSION}/${UCX_ASSET}" + +mkdir -p /tmp/ucx-rpms +tar -xjf "/tmp/${UCX_ASSET}" -C /tmp/ucx-rpms + +rpm -Uvh --replacepkgs --nodeps --prefix="${UCX_PREFIX}" \ + /tmp/ucx-rpms/ucx-${UCX_VERSION}-*.rpm \ + /tmp/ucx-rpms/ucx-ib-${UCX_VERSION}-*.rpm \ + /tmp/ucx-rpms/ucx-rdmacm-${UCX_VERSION}-*.rpm \ + /tmp/ucx-rpms/ucx-devel-${UCX_VERSION}-*.rpm +``` + +Host payload 只需要 UCX core、verbs/RDMA-CM 和 devel headers。CUDA/ROCm/GDR package +只有在对应 device-memory backend 已经单独准入时才安装,不要为当前 Host payload 默认引入。 + +### 3.2 官方源码构建 + +当 release asset 与目标发行版不匹配时,使用同一个官方 tag 构建,不要切换到带有私有 +patch 的分支: + +```bash +git clone --depth 1 --branch v1.22.0 https://github.com/openucx/ucx.git /tmp/ucx-v1.22.0 +cd /tmp/ucx-v1.22.0 +./autogen.sh +./contrib/configure-release \ + --prefix=/opt/tq-ucx/1.22.0-official \ + --with-verbs \ + --with-rdmacm \ + --without-cuda \ + --without-rocm +make -j"$(nproc)" +sudo make install +``` + +### 3.3 验证安装 + +下面的环境变量只指向当前进程使用的 UCX,不应写死到 TQ 源码: + +```bash +export TQ_UCX_HOME=/opt/tq-ucx/1.22.0-official +export PATH="${TQ_UCX_HOME}/bin:${PATH}" +export LD_LIBRARY_PATH="${TQ_UCX_HOME}/lib64:${TQ_UCX_HOME}/lib${LD_LIBRARY_PATH:+:${LD_LIBRARY_PATH}}" + +ucx_info -v +ucx_info -d +``` + +确认: + +- `ucx_info -v` 的 version/revision 是预期的官方 release; +- `ucx_info -d` 同时能看到可靠连接 RC transport 和本机实际 TCP transport; +- `Device` 来自当前节点,不要复制另一台机器的 device 名称。 + +如果只看到 TCP、看不到 RC,先修复 rdma-core/NIC/UCX 安装,不要继续 TQ 集成。 + +## 4. 构建 TQ native UCX extension + +TQ 的 UCX 接入是可选 native extension,必须显式构建: + +```bash +cd /path/to/TransferQueue +TQ_BUILD_UCX=1 TQ_UCX_HOME="${TQ_UCX_HOME}" \ + python setup.py build_ext --inplace +``` + +也可以在构建 wheel 时使用相同的变量: + +```bash +TQ_BUILD_UCX=1 TQ_UCX_HOME="${TQ_UCX_HOME}" \ + python -m build +``` + +验证 native extension 使用的是目标 UCX: + +```bash +ldd transfer_queue/_ucx*.so | grep -E 'libucp|libuct|libucs|libucm' +``` + +所有节点都要重复构建或部署同一 ABI 的 extension,并确保 Ray worker 能加载对应的 +`libucp.so`。官方 RPM 可能把库放在 `lib64`,TQ discovery 已支持 `lib` 和 `lib64`。 + +## 5. 启用 TQ UCX payload + +使用 OmegaConf 覆盖 SimpleStorage 配置: + +```python +from omegaconf import OmegaConf +import transfer_queue as tq + +conf = OmegaConf.create({ + "backend": { + "SimpleStorage": { + "payload_transfer": "ucx", + }, + }, +}) + +tq.init(conf) +``` + +或者在配置文件中设置: + +```yaml +backend: + storage_backend: SimpleStorage + SimpleStorage: + payload_transfer: ucx +``` + +必须在启动 Ray/TQ 进程前,把 UCX 的 `PATH`、`LD_LIBRARY_PATH` 和 TQ Python/native +extension 配置到每个参与节点。推荐让 TQ 自动发现本机 RDMA device、netdev 和 GID: + +```bash +unset UCX_TLS UCX_NET_DEVICES UCX_IB_GID_INDEX UCX_IB_ADDR_TYPE +``` + +如果部署平台必须显式设置 UCX 变量,配置要表达完整的 RC + TCP auxiliary 候选: + +```bash +export UCX_TLS=rc_verbs,tcp,sm,self +export UCX_NET_DEVICES=':,' +export UCX_IB_GID_INDEX='' +export UCX_IB_ADDR_TYPE=ib_global +``` + +``、`` 和 `` 必须在每个节点按本机实际设备填写。 +不要只设置 `UCX_TLS=rc_verbs`:某些 verbs 环境的 RC endpoint 需要 TCP auxiliary 完成 +wireup,纯 RC 会出现 `no auxiliary transport` 或 `Destination is unreachable`。 + +显式设置会覆盖 TQ 自动选择;TQ 不会把用户写的纯 RC 配置偷偷改成 TCP。若不确定,清空 +这些变量并使用 TQ discovery。 + +## 6. 验证 TQ 真正在走 RDMA + +仅看到 PUT/GET 成功不能证明使用了 RDMA。启动测试时打开 UCX 日志: + +```bash +export UCX_LOG_LEVEL=info +``` + +有效的 UCX Host RDMA 日志应包含类似: + +```text +tag(rc_verbs/:) ka(tcp/) +``` + +判定规则: + +- `tag(rc_verbs/...)` 或等价 RC transport 是数据 lane,才算 Host RDMA 数据路径; +- `ka(tcp/...)`、wireup TCP 是辅助通道,不等于 TCP data fallback; +- 如果日志是 `tag(tcp/...)`,只能算 TCP 测试,不能计入 RDMA 结果; +- 如果出现 `no auxiliary transport`,优先检查 TCP netdev 是否包含在 + `UCX_NET_DEVICES`,以及对应端口/GID 是否在节点间可达。 + +仓库内的验证工具: + +```bash +# standalone native payload path +python tools/test_ucx_payload_transfer.py --help + +# real SimpleStorage actor/manager path +python tools/test_simplestorage_ucx_integration.py --mode ucx + +# 在已有 Ray 集群上运行时,把 native extension 和 UCX 库路径传给 Ray worker +export TQ_RAY_WORKER_PYTHONPATH="$(pwd)" +export TQ_RAY_WORKER_LD_LIBRARY_PATH="${LD_LIBRARY_PATH}" +python tools/test_simplestorage_ucx_integration.py --mode ucx --ray-address=auto + +# native payload throughput (不替代 SimpleStorage E2E 性能) +python tools/bench_ucx_payload_transfer.py --help +``` + +SimpleStorage 验证至少应包含: + +1. 大 payload PUT 后远端数据校验; +2. GET 后内容和长度校验; +3. CLEAR 后再次 GET 确认数据已删除; +4. UCX 日志确认实际 data lane; +5. 与同拓扑、同 payload 的默认 ZMQ 结果对照。 + +## 7. 故障定位顺序 + +### UCX 初始化失败 + +```text +no RoCE-v2 device and GID match local IP +``` + +检查本机控制 IP 是否确实落在目标 RoCE netdev,检查 `/sys/class/infiniband`、GID 类型 +和 `ip -br addr`。多候选设备时不要在代码里硬编码选择,先明确部署的本机网络拓扑。 + +### RC endpoint 无法建立 + +```text +no auxiliary transport +ucp_ep_create() failed: Destination is unreachable +``` + +确认 `ucx_info -d` 能看到 TCP,并且 TCP netdev 和 RDMA device 同属可互通网络平面。纯 +`UCX_TLS=rc_verbs` 不是通用的修复方式。 + +### TQ 能初始化但数据走 TCP + +检查日志中的 `tag(...)`: + +- `tag(tcp/...)`:UCX 没有选到 RC,检查 UCX capabilities、`UCX_TLS` 和设备过滤; +- `tag(rc_verbs/...) ka(tcp/...)`:数据已走 RC,TCP 只是 auxiliary; +- 不要只根据吞吐判断 transport,必须以 UCX lane 日志和数据校验为准。 + +### Ray worker 找不到 UCX + +Controller 节点能执行 `ucx_info` 不代表远端 SimpleStorageUnit 能加载 UCX。确认每个 +Ray 节点都有: + +- 相同 TQ native extension ABI; +- 相同官方 UCX runtime 或兼容的 ABI; +- `PATH`、`LD_LIBRARY_PATH` 和 Python package 可见; +- 正确的本机网络和 GID。 + +## 8. 代码入口 + +| 文件 | 责任 | +| --- | --- | +| `setup.py` | `TQ_BUILD_UCX`/`TQ_UCX_HOME` 控制 native extension 构建 | +| `transfer_queue/storage/payload_transfer/factory.py` | 选择 `zmq` 或 `ucx` | +| `transfer_queue/storage/payload_transfer/ucx_discovery.py` | 本机 RDMA/GID/netdev 和 UCX 能力探查 | +| `transfer_queue/storage/payload_transfer/ucx.py` | TQ payload contract 到 UCX 的适配 | +| `transfer_queue/storage/payload_transfer/ucx_runtime.py` | UCP worker、endpoint、request 和生命周期 | +| `transfer_queue/csrc/ucx/ucx_bindings.cpp` | 最小 pybind11/UCP Tagged binding | +| `transfer_queue/storage/bootstrap/simple_storage_bootstrap.py` | 为 SimpleStorage actor 发布 payload transfer metadata | +| `transfer_queue/storage/simple_storage.py` | PREPARE/READY/COMMIT/CANCEL 和数据校验 | + +新增 transport 时先确认它是否真的共享 SimpleStorage 的控制/数据面边界;Mooncake、 +openYuanrong 和 HIXL 不应被强行改造成 `PayloadTransfer` 的实现。 diff --git a/pyproject.toml b/pyproject.toml index 85e5514e..42bb3676 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,8 @@ [build-system] requires = [ "setuptools>=61.0", - "wheel" + "wheel", + "pybind11>=2.12" ] build-backend = "setuptools.build_meta" @@ -132,5 +133,6 @@ mooncake = [ [tool.setuptools.package-data] transfer_queue = [ "version/*", - "*.yaml" -] \ No newline at end of file + "*.yaml", + "csrc/ucx/*.cpp" +] diff --git a/setup.py b/setup.py new file mode 100644 index 00000000..7afa4b01 --- /dev/null +++ b/setup.py @@ -0,0 +1,87 @@ +"""Build the optional UCX extension when explicitly requested.""" + +from __future__ import annotations + +import os +import subprocess +from pathlib import Path + +from setuptools import Extension, setup + + +def _valid_ucx_installation(include_dir: Path, libdir: Path) -> tuple[Path, Path] | None: + if (include_dir / "ucp/api/ucp.h").is_file() and any(libdir.glob("libucp.so*")): + return include_dir, libdir + return None + + +def _pkg_config_ucx() -> tuple[Path, Path] | None: + try: + values = [ + subprocess.run( + ["pkg-config", f"--variable={variable}", "ucx"], + check=True, + capture_output=True, + text=True, + ).stdout.strip() + for variable in ("includedir", "libdir") + ] + except (FileNotFoundError, subprocess.CalledProcessError): + return None + if not all(values): + return None + return _valid_ucx_installation(Path(values[0]), Path(values[1])) + + +def _prefix_ucx(prefix: Path) -> tuple[Path, Path] | None: + include_dir = prefix / "include" + libdirs = [prefix / "lib", prefix / "lib64", *sorted((prefix / "lib").glob("*-linux-gnu"))] + for libdir in libdirs: + installation = _valid_ucx_installation(include_dir, libdir) + if installation is not None: + return installation + return None + + +def find_ucx() -> tuple[Path, Path] | None: + mode = os.environ.get("TQ_BUILD_UCX", "0").lower() + if mode in {"0", "false", "off"}: + return None + if mode not in {"1", "true", "on"}: + raise RuntimeError("TQ_BUILD_UCX must be one of: 0, false, off, 1, true, on") + configured = os.environ.get("TQ_UCX_HOME") + if configured: + installation = _prefix_ucx(Path(configured)) + if installation is not None: + return installation + installation = _pkg_config_ucx() + if installation is not None: + return installation + for prefix in (Path("/usr/local"), Path("/usr")): + installation = _prefix_ucx(prefix) + if installation is not None: + return installation + raise RuntimeError("TQ_BUILD_UCX requested a UCX build, but no UCX installation was found") + + +def ucx_extension() -> list[Extension]: + installation = find_ucx() + if installation is None: + return [] + import pybind11 + + include_dir, libdir = installation + return [ + Extension( + "transfer_queue._ucx", + ["transfer_queue/csrc/ucx/ucx_bindings.cpp"], + include_dirs=[pybind11.get_include(), str(include_dir)], + library_dirs=[str(libdir)], + libraries=["ucp", "uct", "ucs", "ucm"], + language="c++", + extra_compile_args=["-std=c++17"], + ) + ] + + +setup(ext_modules=ucx_extension()) diff --git a/tests/test_payload_transfer.py b/tests/test_payload_transfer.py new file mode 100644 index 00000000..0e046ebf --- /dev/null +++ b/tests/test_payload_transfer.py @@ -0,0 +1,474 @@ +# Copyright 2026 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. + +"""Payload transfer contract, UCX runtime and lifecycle tests.""" + +from __future__ import annotations + +from pathlib import Path +from threading import Event + +import pytest +from omegaconf import OmegaConf + +from transfer_queue.storage.payload_transfer import ( + PayloadDescriptor, + PayloadTransferError, + ReceiveToken, + TransferEndpoint, + create_payload_transfer, +) +from transfer_queue.storage.payload_transfer.ucx import UcxPayloadTransfer +from transfer_queue.storage.payload_transfer.ucx_discovery import ( + UcxDeviceSelection, + _parse_ucx_info_devices, + discover_ucx_device, + gid_matches_ip, +) +from transfer_queue.storage.payload_transfer.ucx_runtime import ( + DEFAULT_TRANSFER_TIMEOUT_SECONDS, + UcxError, + UcxRuntime, + UcxTransfer, + _UcxOwnerThread, + address_digest, + transfer_tag, +) + + +class _FakeRequest: + def __init__(self, pending_polls: int, result: str = "done"): + self.pending_polls = pending_polls + self.result = result + self.cancelled = False + self.started = Event() + + def test(self): + self.started.set() + if self.pending_polls: + self.pending_polls -= 1 + return None + return self.result + + def cancel(self) -> None: + self.cancelled = True + + def start_cancel(self) -> None: + self.cancelled = True + + def test_cancel(self): + return True if self.cancelled else None + + +def test_owner_progresses_multiple_requests_without_blocking_each_other(): + owner = _UcxOwnerThread(lambda: None, lambda: None) + try: + first = _FakeRequest(pending_polls=2, result="first") + second = _FakeRequest(pending_polls=0, result="second") + + first_future = owner.submit_request(lambda: first, lambda request: request.test(), None, 1) + second_future = owner.submit_request(lambda: second, lambda request: request.test(), None, 1) + + assert second_future.result(timeout=1) == "second" + assert first_future.result(timeout=1) == "first" + finally: + owner.stop() + + +def test_default_transfer_timeout_is_finite(): + owner = _UcxOwnerThread(lambda: None, lambda: None) + try: + request = _FakeRequest(pending_polls=1, result="complete") + future = owner.submit_request( + lambda: request, + lambda pending: pending.test(), + None, + DEFAULT_TRANSFER_TIMEOUT_SECONDS, + ) + + assert future.result(timeout=1) == "complete" + assert DEFAULT_TRANSFER_TIMEOUT_SECONDS > 0 + finally: + owner.stop() + + +def test_progress_failure_is_reported_to_active_request(): + failed = False + + def progress(): + nonlocal failed + if not failed: + failed = True + raise RuntimeError("worker failed") + + owner = _UcxOwnerThread(lambda: None, progress) + try: + request = _FakeRequest(pending_polls=1_000_000) + future = owner.submit_request(lambda: request, lambda pending: pending.test(), None, 10) + + with pytest.raises(UcxError, match="UCX progress failed"): + future.result(timeout=1) + assert request.cancelled + finally: + owner.stop() + + +def test_owner_cancels_active_request_before_shutdown(): + owner = _UcxOwnerThread(lambda: None, lambda: None) + try: + request = _FakeRequest(pending_polls=1_000_000) + future = owner.submit_request(lambda: request, lambda pending: pending.test(), None, 10) + assert request.started.wait(timeout=1) + + owner.cancel_active_requests() + + assert request.cancelled + with pytest.raises(UcxError, match="canceled during shutdown"): + future.result(timeout=1) + finally: + owner.stop() + + +def test_cancelled_future_does_not_block_later_owner_tasks(): + owner = _UcxOwnerThread(lambda: None, lambda: None) + try: + request = _FakeRequest(pending_polls=1_000_000) + future = owner.submit_request(lambda: request, lambda pending: pending.test(), None, 10) + assert request.started.wait(timeout=1) + + assert future.cancel() + assert owner.submit(lambda: "still-responsive").result(timeout=1) == "still-responsive" + assert request.cancelled + finally: + owner.stop() + + +def test_ucx_transfer_rejects_inconsistent_control_metadata(): + transfer_id = "descriptor-validation" + tag = transfer_tag(transfer_id) + + UcxTransfer(transfer_id, tag, 3).validate() + large = UcxTransfer(transfer_id, tag, 64 * 1024 * 1024 * 1024 + 1) + large.validate() + + invalid = ( + UcxTransfer(transfer_id, tag + 1, 3), + UcxTransfer(transfer_id, tag, -1), + ) + for descriptor in invalid: + with pytest.raises(UcxError): + descriptor.validate() + + +def test_payload_descriptor_is_always_complete_when_deserialized(): + descriptor = PayloadDescriptor.from_dict({"transfer_id": "payload", "payload_bytes": 3}) + assert descriptor.payload_bytes == 3 + + with pytest.raises(PayloadTransferError, match="negative"): + PayloadDescriptor.from_dict({"transfer_id": "payload", "payload_bytes": -1}) + + +def test_ucx_adapter_derives_tag_instead_of_sending_it_in_control_metadata(): + class Runtime: + address = b"a" * 32 + + def prepare_receive(self, transfer): + self.prepared = transfer + + def send_async(self, address, transfer, payload, digest): + self.sent = address, transfer, payload, digest + return "future" + + adapter = UcxPayloadTransfer.__new__(UcxPayloadTransfer) + adapter._runtime = Runtime() + descriptor = PayloadDescriptor("payload", 3) + + token = adapter.prepare_receive(descriptor) + assert token == ReceiveToken(data={}) + endpoint = TransferEndpoint("ucx", {"address": b"b" * 32}) + assert adapter.send(endpoint, token, descriptor, b"abc") == "future" + assert adapter._runtime.prepared.tag == transfer_tag("payload") + assert adapter._runtime.sent[1].tag == transfer_tag("payload") + + +def test_ucx_runtime_direct_send_receive_roundtrip(): + mailbox = {} + + class Receive: + def __init__(self, tag, target): + self.tag = tag + self.target = target + + def wait(self, _timeout): + payload = mailbox[self.tag] + assert len(payload) == len(self.target) + self.target[:] = payload + return self.target + + class Worker: + def post_receive(self, tag, size): + return Receive(tag, bytearray(size)) + + class Endpoint: + def post_send(self, tag, payload): + mailbox[tag] = bytes(payload) + return _FakeRequest(0) + + plane = UcxRuntime.__new__(UcxRuntime) + plane._timeout_seconds = 1 + plane._receives = {} + plane._worker = Worker() + plane._closed = False + plane._call = lambda operation: operation() + peer_address = b"p" * 32 + plane._endpoints = {peer_address: Endpoint()} + transfer_id = "direct-roundtrip" + descriptor = UcxTransfer(transfer_id, transfer_tag(transfer_id), 10) + + plane.prepare_receive(descriptor) + send = plane._post_send(peer_address, None, descriptor, b"abcdefghij") + assert send.test() == "done" + + assert bytes(plane.finish_receive(descriptor)) == b"abcdefghij" + + +def test_peer_address_digest_is_checked_before_ucx_submission(): + address = b"a" * 32 + UcxRuntime._validate_peer_address(address, address_digest(address)) + + with pytest.raises(UcxError, match="digest"): + UcxRuntime._validate_peer_address(address, address_digest(b"b" * 32)) + + with pytest.raises(UcxError, match="length"): + UcxRuntime._validate_peer_address(b"short") + + +def test_receive_uses_actual_length_reported_by_binding(): + descriptor = UcxTransfer("short-receive", transfer_tag("short-receive"), 8) + + with pytest.raises(UcxError, match="received length mismatch"): + UcxRuntime._received_payload(descriptor, memoryview(b"short")) + + +def test_roce_device_is_discovered_from_local_ip_and_gid(tmp_path): + port = tmp_path / "rdma0" / "ports" / "1" + for relative, value in ( + ("gid_attrs/ndevs/3", "eth1\n"), + ("gid_attrs/types/3", "RoCE v2\n"), + ("gids/3", "::ffff:192.0.2.10\n"), + ): + path = port / relative + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(value) + + selection = discover_ucx_device( + "192.0.2.10", + infiniband_root=tmp_path, + interface_addresses={"eth1": {"192.0.2.10"}}, + ) + + assert selection == UcxDeviceSelection("rdma0", 1, "eth1", 3) + assert selection.net_devices == "rdma0:1,eth1" + assert gid_matches_ip("::ffff:192.0.2.10", "192.0.2.10") + + +def test_unique_roce_device_is_discovered_when_control_ip_is_separate(tmp_path): + port = tmp_path / "rdma0" / "ports" / "1" + for relative, value in ( + ("gid_attrs/ndevs/3", "eth1\n"), + ("gid_attrs/types/3", "RoCE v2\n"), + ("gids/3", "::ffff:192.0.2.10\n"), + ): + path = port / relative + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(value) + + selection = discover_ucx_device( + "198.51.100.10", + infiniband_root=tmp_path, + interface_addresses={"eth1": {"192.0.2.10"}}, + ) + + assert selection == UcxDeviceSelection("rdma0", 1, "eth1", 3) + + +def test_ambiguous_roce_devices_are_not_guessed(tmp_path): + for device, netdev, address in ( + ("rdma0", "eth1", "192.0.2.10"), + ("rdma1", "eth2", "192.0.2.11"), + ): + port = tmp_path / device / "ports" / "1" + for relative, value in ( + ("gid_attrs/ndevs/3", f"{netdev}\n"), + ("gid_attrs/types/3", "RoCE v2\n"), + ("gids/3", f"::ffff:{address}\n"), + ): + path = port / relative + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(value) + + selection = discover_ucx_device( + "198.51.100.10", + infiniband_root=tmp_path, + interface_addresses={"eth1": {"192.0.2.10"}, "eth2": {"192.0.2.11"}}, + ) + + assert selection == UcxDeviceSelection() + + +def test_ucx_config_is_derived_from_local_discovery(): + # The actual transport selection is runtime capability based. Keep this + # test focused on the node-local device/GID settings. + selection = UcxDeviceSelection("rdma0", 1, "eth1", 3) + assert selection.ucx_config == { + "NET_DEVICES": "rdma0:1,eth1", + "IB_GID_INDEX": "3", + "IB_ADDR_TYPE": "ib_global", + } + + +def test_ucx_info_devices_are_parsed_without_vendor_assumptions(): + output = """ +# Transport: rc_verbs +# Device: hca0:1 +# Transport: tcp +# Device: eth9 +# Transport: self +# Device: memory +""" + + assert _parse_ucx_info_devices(output) == { + ("rc_verbs", "hca0:1"), + ("tcp", "eth9"), + ("self", "memory"), + } + + +def test_ucx_config_selects_runtime_supported_transports(monkeypatch): + from transfer_queue.storage.payload_transfer import ucx_discovery + + monkeypatch.setattr( + ucx_discovery, + "_discover_ucx_transports", + lambda: frozenset( + { + ("rc_verbs", "hca0:1"), + ("tcp", "eth9"), + ("sysv", "memory"), + ("self", "memory"), + } + ), + ) + + selection = UcxDeviceSelection("hca0", 1, "eth9", 3) + + assert selection.ucx_config["TLS"] == "rc_verbs,tcp,sm,self" + + +def test_ucx_info_discovery_supports_lib64_installations(monkeypatch, tmp_path): + from transfer_queue.storage.payload_transfer import ucx_discovery + + prefix = tmp_path / "ucx" + executable = prefix / "bin" / "ucx_info" + executable.parent.mkdir(parents=True) + executable.touch() + (prefix / "lib64").mkdir() + captured = {} + + class Result: + stdout = "# Transport: tcp\n# Device: eth9\n" + + def run(_args, *, env, **_kwargs): + captured.update(env) + return Result() + + monkeypatch.setattr(ucx_discovery, "_find_ucx_info", lambda: executable) + monkeypatch.setattr(ucx_discovery.subprocess, "run", run) + monkeypatch.setenv("LD_LIBRARY_PATH", "/existing/lib") + ucx_discovery._discover_ucx_transports.cache_clear() + + try: + assert ucx_discovery._discover_ucx_transports() == frozenset({("tcp", "eth9")}) + assert captured["LD_LIBRARY_PATH"] == f"{prefix / 'lib64'}:/existing/lib" + finally: + ucx_discovery._discover_ucx_transports.cache_clear() + + +def test_explicit_ucx_tls_overrides_runtime_selection(monkeypatch): + monkeypatch.setenv("UCX_TLS", "tcp,self") + monkeypatch.setenv("UCX_NET_DEVICES", "custom_hca:2,custom_eth") + monkeypatch.setenv("UCX_IB_GID_INDEX", "7") + monkeypatch.setenv("UCX_IB_ADDR_TYPE", "ib_global") + + selection = UcxDeviceSelection("rdma0", 1, "eth1", 3) + assert selection.ucx_config == { + "TLS": "tcp,self", + "NET_DEVICES": "custom_hca:2,custom_eth", + "IB_GID_INDEX": "7", + "IB_ADDR_TYPE": "ib_global", + } + + +def test_ucx_payload_transfer_fails_when_ucx_is_unavailable(monkeypatch): + monkeypatch.setenv("UCX_TLS", "rc_verbs") + monkeypatch.setattr( + "transfer_queue.storage.payload_transfer.ucx_runtime.discover_ucx_device", + lambda _ip: UcxDeviceSelection("rdma0", 1, "eth0", 3), + ) + + class UnavailableRuntime: + def __init__(self, **_kwargs): + raise RuntimeError("UCX unavailable") + + monkeypatch.setattr("transfer_queue.storage.payload_transfer.ucx_runtime.UcxRuntime", UnavailableRuntime) + + assert create_payload_transfer(local_ip="192.0.2.1") is None + assert create_payload_transfer("zmq", local_ip="192.0.2.1") is None + with pytest.raises(UcxError, match="UCX unavailable"): + create_payload_transfer("ucx", local_ip="192.0.2.1") + + +def test_ucx_payload_transfer_requires_an_rdma_device(monkeypatch): + monkeypatch.delenv("UCX_NET_DEVICES", raising=False) + monkeypatch.setattr( + "transfer_queue.storage.payload_transfer.ucx_runtime.discover_ucx_device", + lambda _ip: UcxDeviceSelection(), + ) + + constructed = False + + class UnexpectedRuntime: + def __init__(self, **_kwargs): + nonlocal constructed + constructed = True + + monkeypatch.setattr("transfer_queue.storage.payload_transfer.ucx_runtime.UcxRuntime", UnexpectedRuntime) + + with pytest.raises(UcxError, match="no RoCE-v2 device"): + create_payload_transfer("ucx", local_ip="192.0.2.1") + assert not constructed + + +def test_ucx_payload_transfer_rejects_tcp_only_tls(monkeypatch): + monkeypatch.setenv("UCX_TLS", "tcp,sm,self") + monkeypatch.setattr( + "transfer_queue.storage.payload_transfer.ucx_runtime.discover_ucx_device", + lambda _ip: UcxDeviceSelection("rdma0", 1, "eth0", 3), + ) + + with pytest.raises(UcxError, match="no reliable-connection RDMA transport"): + create_payload_transfer("ucx", local_ip="192.0.2.1") + + +def test_payload_transfer_rejects_unknown_implementation(): + with pytest.raises(ValueError, match="expected 'zmq' or 'ucx'"): + create_payload_transfer("hixl") + + +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 == "zmq" + assert "payload_transport" not in config.backend.SimpleStorage + assert "data_plane" not in config.backend.SimpleStorage diff --git a/tests/test_serial_utils_batch_on_cpu.py b/tests/test_serial_utils_batch_on_cpu.py index 7720f4a9..4c004a11 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 @@ -64,6 +66,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(" bytearray: + seed = hashlib.sha256(f"tq-bench:{size}:{iteration}".encode()).digest() + return bytearray((seed * ((size // len(seed)) + 1))[:size]) + + +def descriptor(size: int, iteration: int) -> UcxTransfer: + transfer_id = f"tq-bench-{size}-{iteration}" + return UcxTransfer(transfer_id, transfer_tag(transfer_id), size) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("role", choices=("server", "client")) + parser.add_argument("address_file", type=Path) + parser.add_argument("--size", type=int, default=16 * 1024 * 1024) + parser.add_argument("--repetitions", type=int, default=5) + args = parser.parse_args() + + dp = UcxRuntime(180) + try: + if args.role == "server": + # Post the first receive before publishing the address. Otherwise + # the client can start the first rendezvous before the server has + # installed its matching receive, which is not representative of + # SimpleStorage's PREPARE handshake. + first = descriptor(args.size, 0) + dp.prepare_receive(first) + args.address_file.write_text(base64.b64encode(dp.address).decode()) + for iteration in range(args.repetitions): + item = descriptor(args.size, iteration) + if iteration > 0: + dp.prepare_receive(item) + actual = dp.finish_receive(item) + if len(actual) != args.size: + raise AssertionError(f"payload length mismatch at iteration {iteration}") + print(f"bench server PASS size={args.size} repetitions={args.repetitions}", flush=True) + return + + deadline = time.monotonic() + 60 + while not args.address_file.exists(): + if time.monotonic() > deadline: + raise TimeoutError("server address file was not created") + time.sleep(0.1) + peer = base64.b64decode(args.address_file.read_text()) + peer_digest = address_digest(peer) + dp.warmup(peer, peer_digest) + elapsed = [] + for iteration in range(args.repetitions): + item = descriptor(args.size, iteration) + start = time.perf_counter() + dp.send(peer, item, make_payload(args.size, iteration), peer_address_digest=peer_digest) + elapsed.append(time.perf_counter() - start) + measured = elapsed[1:] if len(elapsed) > 1 else elapsed + median_seconds = statistics.median(measured) + mib_s = args.size / median_seconds / 2**20 + print( + f"bench client PASS size={args.size} repetitions={args.repetitions} " + f"median_seconds={median_seconds:.6f} throughput_mib_s={mib_s:.2f}", + flush=True, + ) + finally: + dp.close() + + +if __name__ == "__main__": + main() diff --git a/tools/bench_zmq_multipart.py b/tools/bench_zmq_multipart.py new file mode 100644 index 00000000..94f13310 --- /dev/null +++ b/tools/bench_zmq_multipart.py @@ -0,0 +1,77 @@ +#!/usr/bin/env python3 +"""Measure a ZMQ multipart payload with a receive-side acknowledgement.""" + +from __future__ import annotations + +import argparse +import statistics +import time + +import torch +import zmq + +from transfer_queue.utils.serial_utils import encode + + +def make_frames(elements: int) -> list: + return encode({"tensor": [torch.arange(elements, dtype=torch.int64)]}) + + +def server(bind_address: str, elements: int, repetitions: int) -> None: + context = zmq.Context() + socket = context.socket(zmq.REP) + socket.bind(bind_address) + try: + for _ in range(repetitions): + frames = socket.recv_multipart(copy=False) + socket.send(b"ok") + if not frames or sum(memoryview(frame).nbytes for frame in frames) == 0: + raise AssertionError("empty multipart payload") + print(f"zmq server PASS elements={elements} repetitions={repetitions}", flush=True) + finally: + socket.close(0) + context.term() + + +def client(address: str, elements: int, repetitions: int) -> None: + context = zmq.Context() + socket = context.socket(zmq.REQ) + socket.connect(address) + frames = make_frames(elements) + elapsed = [] + try: + for _ in range(repetitions): + start = time.perf_counter() + socket.send_multipart(frames, copy=False) + if socket.recv() != b"ok": + raise AssertionError("invalid server acknowledgement") + elapsed.append(time.perf_counter() - start) + measured = elapsed[1:] if len(elapsed) > 1 else elapsed + median_seconds = statistics.median(measured) + payload_bytes = sum(memoryview(frame).nbytes for frame in frames) + print( + f"zmq client PASS elements={elements} repetitions={repetitions} " + f"median_seconds={median_seconds:.6f} " + f"throughput_mib_s={payload_bytes / median_seconds / 2**20:.2f}", + flush=True, + ) + finally: + socket.close(0) + context.term() + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("role", choices=("server", "client")) + parser.add_argument("address") + parser.add_argument("--elements", type=int, default=1_000_000) + parser.add_argument("--repetitions", type=int, default=5) + args = parser.parse_args() + if args.role == "server": + server(args.address, args.elements, args.repetitions) + else: + client(args.address, args.elements, args.repetitions) + + +if __name__ == "__main__": + main() diff --git a/tools/phase1_host_rdma_bench.sh b/tools/phase1_host_rdma_bench.sh new file mode 100755 index 00000000..d73a6355 --- /dev/null +++ b/tools/phase1_host_rdma_bench.sh @@ -0,0 +1,53 @@ +#!/usr/bin/env bash +set -euo pipefail + +# Run on a control machine after both remote shells have been prepared with +# the intended UCX runtime and node-local UCX configuration. + +: "${TQ_RDMA_SERVER_SSH:?set TQ_RDMA_SERVER_SSH}" +: "${TQ_RDMA_CLIENT_SSH:?set TQ_RDMA_CLIENT_SSH}" +: "${TQ_RDMA_SERVER_IP:?set TQ_RDMA_SERVER_IP}" +: "${TQ_RDMA_SERVER_DEVICE:?set TQ_RDMA_SERVER_DEVICE}" +: "${TQ_RDMA_CLIENT_DEVICE:?set TQ_RDMA_CLIENT_DEVICE}" +: "${TQ_RDMA_GID_INDEX:?set TQ_RDMA_GID_INDEX}" + +transport=${1:-all} +case "$transport" in + ucx|write|read|all) ;; + *) echo "usage: $0 {ucx|write|read|all}" >&2; exit 2 ;; +esac + +sizes=(65536 1048576 16777216) +run_id="$$" + +run_ucx() { + local size=$1 port=$2 log="/tmp/tq-phase1-ucx-${run_id}-${size}.log" + ssh "$TQ_RDMA_SERVER_SSH" "rm -f '$log'; nohup ucx_perftest -p '$port' -t tag_bw -s '$size' -n 200 >'$log' 2>&1 &" + sleep 3 + ssh "$TQ_RDMA_CLIENT_SSH" "timeout 90 ucx_perftest '$TQ_RDMA_SERVER_IP' -p '$port' -t tag_bw -s '$size' -n 200" + ssh "$TQ_RDMA_SERVER_SSH" "grep -E 'Version|cfg#|Final:|ERROR' '$log' || true" +} + +run_verbs() { + local kind=$1 size=$2 port=$3 log="/tmp/tq-phase1-${kind}-${run_id}-${size}.log" + ssh -tt "$TQ_RDMA_SERVER_SSH" "/usr/bin/ib_${kind}_bw -d '$TQ_RDMA_SERVER_DEVICE' -i 1 -x '$TQ_RDMA_GID_INDEX' -s '$size' -n 200 -p '$port'" >"$log" 2>&1 & + local server_ssh_pid=$! + sleep 3 + ssh "$TQ_RDMA_CLIENT_SSH" "timeout 90 /usr/bin/ib_${kind}_bw '$TQ_RDMA_SERVER_IP' -d '$TQ_RDMA_CLIENT_DEVICE' -i 1 -x '$TQ_RDMA_GID_INDEX' -s '$size' -n 200 -p '$port'" + wait "$server_ssh_pid" || true + grep -E 'Connection type|Link type|GID index|BW average' "$log" || true +} + +idx=0 +for size in "${sizes[@]}"; do + if [[ "$transport" == ucx || "$transport" == all ]]; then + run_ucx "$size" $((29600 + idx)) + fi + if [[ "$transport" == write || "$transport" == all ]]; then + run_verbs write "$size" $((29700 + idx)) + fi + if [[ "$transport" == read || "$transport" == all ]]; then + run_verbs read "$size" $((29800 + idx)) + fi + idx=$((idx + 1)) +done diff --git a/tools/test_controller_reinit_after_crash.py b/tools/test_controller_reinit_after_crash.py new file mode 100644 index 00000000..cf2fc9b9 --- /dev/null +++ b/tools/test_controller_reinit_after_crash.py @@ -0,0 +1,60 @@ +#!/usr/bin/env python3 +"""Check same-process TQ re-entry after the controller actor is killed.""" + +from __future__ import annotations + +import os + +import ray +from omegaconf import OmegaConf + +import transfer_queue as tq +from transfer_queue import interface + + +def config(): + return OmegaConf.create( + { + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "total_storage_size": 64, + "num_data_storage_units": 1, + "payload_transfer": "zmq", + }, + }, + }, + flags={"allow_objects": True}, + ) + + +def main() -> None: + ray.init(address=os.environ["RAY_ADDRESS"], include_dashboard=False) + try: + tq.init(config()) + controller = interface._TQ_CONTROLLER + assert controller is not None + print("controller initial init PASS", flush=True) + + ray.kill(controller) + try: + ray.get(controller.get_config.remote(), timeout=2) + except Exception: + print("controller crash observed PASS", flush=True) + else: + raise AssertionError("controller remained available after ray.kill") + + try: + tq.init(config()) + except Exception as exc: + print(f"same-process tq.init after crash expected error={type(exc).__name__}", flush=True) + else: + raise AssertionError("same-process tq.init unexpectedly recovered a dead controller") + finally: + tq.close() + ray.shutdown() + + +if __name__ == "__main__": + main() diff --git a/tools/test_simplestorage_bootstrap_ucx.py b/tools/test_simplestorage_bootstrap_ucx.py new file mode 100644 index 00000000..e28ddd7a --- /dev/null +++ b/tools/test_simplestorage_bootstrap_ucx.py @@ -0,0 +1,74 @@ +#!/usr/bin/env python3 +"""Exercise the production TQ bootstrap and manager wiring with UCX enabled.""" + +from __future__ import annotations + +import asyncio +import os + +import ray +import torch +from omegaconf import OmegaConf + +import transfer_queue as tq + + +async def run_roundtrip(label: str, index: int) -> None: + client = tq.get_client() + manager = client.storage_manager + assert manager.payload_transfer is not None + assert manager.payload_transfer_infos + storage_id, storage_info = next(iter(manager.storage_unit_infos.items())) + endpoint = manager.payload_transfer_infos[storage_id]["endpoint"] + assert endpoint["transport"] == "ucx" + print( + f"bootstrap metadata PASS storage={storage_id} host={storage_info.ip} " + f"data_address_len={len(endpoint['data']['address'])}", + flush=True, + ) + + value = torch.arange(1_000_000, dtype=torch.int64) + await manager._put_to_single_storage_unit( + [index], {"tensor": [value]}, target_storage_unit=storage_id + ) + _, result = await manager._get_from_single_storage_unit( + [index], ["tensor"], target_storage_unit=storage_id + ) + assert torch.equal(result["tensor"][0], value) + print(f"{label} manager roundtrip PASS", flush=True) + + +def main() -> None: + conf = OmegaConf.create( + { + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "total_storage_size": 64 * 1024 * 1024, + "num_data_storage_units": 1, + "payload_transfer": "ucx", + }, + }, + }, + flags={"allow_objects": True}, + ) + ray.init(address=os.environ["RAY_ADDRESS"], include_dashboard=False) + try: + final_conf = tq.init(conf) + assert final_conf.backend.SimpleStorage.payload_transfer_infos + asyncio.run(run_roundtrip("bootstrap initial", 301)) + tq.close() + print("controller close PASS", flush=True) + + final_conf = tq.init(conf) + assert final_conf.backend.SimpleStorage.payload_transfer_infos + asyncio.run(run_roundtrip("bootstrap replacement", 302)) + print("controller replacement PASS", flush=True) + finally: + tq.close() + ray.shutdown() + + +if __name__ == "__main__": + main() diff --git a/tools/test_simplestorage_bootstrap_ucx_multinode.py b/tools/test_simplestorage_bootstrap_ucx_multinode.py new file mode 100644 index 00000000..1e7119c3 --- /dev/null +++ b/tools/test_simplestorage_bootstrap_ucx_multinode.py @@ -0,0 +1,72 @@ +#!/usr/bin/env python3 +"""Validate formal TQ bootstrap with one UCX StorageUnit on each Ray node.""" + +from __future__ import annotations + +import asyncio +import os + +import ray +import torch +from omegaconf import OmegaConf + +import transfer_queue as tq + + +async def roundtrip_all_units() -> None: + manager = tq.get_client().storage_manager + assert manager.payload_transfer is not None + assert len(manager.storage_unit_infos) == 2 + assert set(manager.payload_transfer_infos) == set(manager.storage_unit_infos) + + for offset, (storage_id, storage_info) in enumerate(manager.storage_unit_infos.items()): + data_info = manager.payload_transfer_infos[storage_id] + print( + f"bootstrap unit metadata PASS id={storage_id} host={storage_info.ip} " + f"address_len={len(data_info['endpoint']['data']['address'])}", + flush=True, + ) + value = torch.arange(1_000_000, dtype=torch.int64) + offset + index = [700 + offset] + await manager._put_to_single_storage_unit( + index, {"tensor": [value]}, target_storage_unit=storage_id + ) + _, result = await manager._get_from_single_storage_unit( + index, ["tensor"], target_storage_unit=storage_id + ) + assert torch.equal(result["tensor"][0], value) + print(f"bootstrap unit roundtrip PASS id={storage_id} host={storage_info.ip}", flush=True) + + +def main() -> None: + conf = OmegaConf.create( + { + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "total_storage_size": 128 * 1024 * 1024, + "num_data_storage_units": 2, + "payload_transfer": "ucx", + }, + }, + }, + flags={"allow_objects": True}, + ) + ray.init(address=os.environ["RAY_ADDRESS"], include_dashboard=False) + try: + tq.init(conf) + asyncio.run(roundtrip_all_units()) + print("bootstrap multinode PASS", flush=True) + tq.close() + print("bootstrap multinode close PASS", flush=True) + tq.init(conf) + asyncio.run(roundtrip_all_units()) + print("bootstrap multinode replacement PASS", flush=True) + finally: + tq.close() + ray.shutdown() + + +if __name__ == "__main__": + main() diff --git a/tools/test_simplestorage_controller_crash.py b/tools/test_simplestorage_controller_crash.py new file mode 100644 index 00000000..b856ba10 --- /dev/null +++ b/tools/test_simplestorage_controller_crash.py @@ -0,0 +1,87 @@ +#!/usr/bin/env python3 +"""Check the behavior of an existing Manager after the controller actor dies.""" + +from __future__ import annotations + +import asyncio +import os + +import ray +import torch +from omegaconf import OmegaConf + +import transfer_queue as tq +from transfer_queue import interface + + +def make_config(): + return OmegaConf.create( + { + "controller": {"polling_mode": True}, + "backend": { + "storage_backend": "SimpleStorage", + "SimpleStorage": { + "total_storage_size": 64 * 1024 * 1024, + "num_data_storage_units": 1, + "payload_transfer": "ucx", + }, + }, + }, + flags={"allow_objects": True}, + ) + + +async def roundtrip(index: int) -> None: + manager = tq.get_client().storage_manager + storage_id = next(iter(manager.storage_unit_infos)) + value = torch.arange(1_000_000, dtype=torch.int64) + index + await manager._put_to_single_storage_unit( + [index], {"tensor": [value]}, target_storage_unit=storage_id + ) + _, result = await manager._get_from_single_storage_unit( + [index], ["tensor"], target_storage_unit=storage_id + ) + assert torch.equal(result["tensor"][0], value) + + +async def observe_control_plane_failure() -> None: + manager = tq.get_client().storage_manager + try: + await manager.notify_data_update( + "controller_crash", + [402], + {"tensor": {"dtype": "int64", "shape": [1_000_000]}}, + ) + except (TimeoutError, RuntimeError) as exc: + print(f"existing manager notify failure propagated PASS type={type(exc).__name__}", flush=True) + else: + raise AssertionError("notify_data_update unexpectedly succeeded without a controller") + + +def main() -> None: + ray.init(address=os.environ["RAY_ADDRESS"], include_dashboard=False) + try: + tq.init(make_config()) + asyncio.run(roundtrip(401)) + controller = interface._TQ_CONTROLLER + assert controller is not None + ray.kill(controller) + try: + ray.get(controller.get_config.remote(), timeout=2) + except Exception: + print("controller unavailable PASS", flush=True) + else: + raise AssertionError("controller actor still answered after ray.kill") + print("controller kill PASS", flush=True) + + # Keep the original Manager and StorageUnit; do not call tq.init/close here. + asyncio.run(observe_control_plane_failure()) + asyncio.run(roundtrip(402)) + print("existing manager after controller crash PASS", flush=True) + finally: + tq.close() + ray.shutdown() + + +if __name__ == "__main__": + main() diff --git a/tools/test_simplestorage_ucx_concurrent.py b/tools/test_simplestorage_ucx_concurrent.py new file mode 100644 index 00000000..379a5e47 --- /dev/null +++ b/tools/test_simplestorage_ucx_concurrent.py @@ -0,0 +1,106 @@ +#!/usr/bin/env python3 +"""Concurrent SimpleStorage UCX PUT/GET validation on one remote actor.""" + +from __future__ import annotations + +import argparse +import asyncio +import os +import time + +import ray +import torch +from test_simplestorage_ucx_integration import make_manager + +from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager +from transfer_queue.storage.simple_storage import SimpleStorageUnit + + +async def main_case(concurrency: int, elements: int, resource: str | None, enabled: bool) -> None: + threshold = 64 * 1024 + actor_options = {"num_cpus": 1} + if resource: + actor_options["resources"] = {resource: 1} + actor = SimpleStorageUnit.options(**actor_options).remote( + storage_unit_size=None, + payload_transfer="ucx" if enabled else "zmq", + ) + info = ray.get(actor.get_zmq_server_info.remote()) + data_info = ray.get(actor.get_payload_transfer_info.remote()) + managers: list[AsyncSimpleStorageManager] = [] + try: + values = [torch.arange(elements, dtype=torch.int64) + i for i in range(concurrency)] + for i in range(concurrency): + manager = make_manager(info, data_info, enabled, threshold) + manager.storage_manager_id = f"TQ_CONCURRENT_MANAGER_{i}" + managers.append(manager) + + async def put_one(i: int) -> None: + await managers[i]._put_to_single_storage_unit( + [10_000 + i], {"tensor": [values[i]]}, target_storage_unit=info.id + ) + + start = time.perf_counter() + await asyncio.gather(*(put_one(i) for i in range(concurrency))) + put_seconds = time.perf_counter() - start + + async def get_one(i: int) -> None: + _, result = await managers[i]._get_from_single_storage_unit( + [10_000 + i], ["tensor"], target_storage_unit=info.id + ) + assert torch.equal(result["tensor"][0], values[i]) + + start = time.perf_counter() + await asyncio.gather(*(get_one(i) for i in range(concurrency))) + get_seconds = time.perf_counter() - start + total_bytes = concurrency * elements * 8 + print( + f"concurrent PASS mode={'ucx' if enabled else 'legacy'} " + f"concurrency={concurrency} elements={elements} " + f"put_seconds={put_seconds:.6f} get_seconds={get_seconds:.6f} " + f"put_mib_s={total_bytes / put_seconds / 2**20:.2f} " + f"get_mib_s={total_bytes / get_seconds / 2**20:.2f}", + flush=True, + ) + finally: + for manager in managers: + if manager.payload_transfer is not None: + manager.payload_transfer.close() + ray.kill(actor) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--ray-address", required=True) + parser.add_argument("--concurrency", type=int, default=4) + parser.add_argument("--elements", type=int, default=1_000_000) + parser.add_argument("--resource", default=None) + parser.add_argument("--mode", choices=("legacy", "ucx"), default="ucx") + args = parser.parse_args() + runtime_env_vars = { + "TORCH_DEVICE_BACKEND_AUTOLOAD": "0", + } + for name in ("TQ_RAY_WORKER_PYTHONPATH", "TQ_RAY_WORKER_LD_LIBRARY_PATH"): + value = os.environ.get(name) + if value: + runtime_env_vars[name.removeprefix("TQ_RAY_WORKER_")] = value + ray.init( + address=args.ray_address, + include_dashboard=False, + runtime_env={"env_vars": runtime_env_vars} if runtime_env_vars else None, + ) + try: + asyncio.run( + main_case( + args.concurrency, + args.elements, + args.resource, + enabled=args.mode == "ucx", + ) + ) + finally: + ray.shutdown() + + +if __name__ == "__main__": + main() diff --git a/tools/test_simplestorage_ucx_failed_put.py b/tools/test_simplestorage_ucx_failed_put.py new file mode 100644 index 00000000..53d6f361 --- /dev/null +++ b/tools/test_simplestorage_ucx_failed_put.py @@ -0,0 +1,139 @@ +#!/usr/bin/env python3 +"""Verify that a failed UCX PUT does not commit a partial object.""" + +from __future__ import annotations + +import argparse +import asyncio +import os +from unittest.mock import PropertyMock, patch + +import ray +import torch +from test_simplestorage_ucx_integration import make_manager + +from transfer_queue.storage.payload_transfer.ucx_runtime import UcxRuntime, address_digest +from transfer_queue.storage.simple_storage import SimpleStorageUnit + +THRESHOLD = 64 * 1024 + + +async def run_case(resource: str | None) -> None: + options = {"num_cpus": 1} + if resource: + options["resources"] = {resource: 1} + actor = SimpleStorageUnit.options(**options).remote( + storage_unit_size=None, + payload_transfer="ucx", + ) + manager = None + dead_peer = None + try: + info = ray.get(actor.get_zmq_server_info.remote()) + data_info = ray.get(actor.get_payload_transfer_info.remote()) + manager = make_manager(info, data_info, True, THRESHOLD) + value = torch.arange(1_000_000, dtype=torch.int64) + endpoint_data = manager.payload_transfer_infos[info.id]["endpoint"]["data"] + valid_address = endpoint_data["address"] + + # Production bootstrap metadata must reject a same-length mutation + # before ucp_ep_create(). Keep the original digest intentionally. + corrupt_address = bytearray(valid_address) + corrupt_address[0] ^= 0xFF + endpoint_data["address"] = bytes(corrupt_address) + try: + await manager._put_to_single_storage_unit([500], {"tensor": [value]}, target_storage_unit=info.id) + except Exception as exc: + print(f"corrupt bootstrap address expected error={type(exc).__name__}", flush=True) + else: + raise AssertionError("corrupt bootstrap address unexpectedly succeeded") + pending = ray.get(actor.get_payload_transfer_pending_counts.remote()) + assert pending == {"pending_puts": 0, "pending_gets": 0, "pending_receives": 0}, pending + print(f"corrupt bootstrap address guard PASS state={pending}", flush=True) + endpoint_data["address"] = valid_address + + dead_peer = UcxRuntime(10) + dead_address = dead_peer.address + dead_peer.close() + endpoint_data["address"] = dead_address + endpoint_data["address_digest"] = address_digest(dead_address) + try: + await manager._put_to_single_storage_unit([501], {"tensor": [value]}, target_storage_unit=info.id) + except Exception as exc: + print(f"failed PUT expected error={type(exc).__name__}", flush=True) + else: + raise AssertionError("failed UCX PUT unexpectedly succeeded") + + pending = ray.get(actor.get_payload_transfer_pending_counts.remote()) + assert pending == {"pending_puts": 0, "pending_gets": 0, "pending_receives": 0}, pending + print(f"failed PUT remote cleanup PASS state={pending}", flush=True) + + endpoint_data["address"] = valid_address + endpoint_data["address_digest"] = address_digest(valid_address) + try: + await manager._get_from_single_storage_unit([501], ["tensor"], target_storage_unit=info.id) + except Exception: + print("failed PUT not committed PASS", flush=True) + else: + raise AssertionError("failed UCX PUT left a readable object") + + await manager._put_to_single_storage_unit([502], {"tensor": [value + 1]}, target_storage_unit=info.id) + + # Exercise the reverse GET direction with an address altered in the + # control message while retaining the original digest. + receiver_address = manager.payload_transfer.endpoint().data["address"] + corrupt_receiver = bytearray(receiver_address) + corrupt_receiver[0] ^= 0xFF + with ( + patch.object(UcxRuntime, "address", new_callable=PropertyMock, return_value=bytes(corrupt_receiver)), + patch( + "transfer_queue.storage.payload_transfer.ucx.address_digest", + return_value=address_digest(receiver_address), + ), + ): + try: + await manager._get_from_single_storage_unit([502], ["tensor"], target_storage_unit=info.id) + except Exception as exc: + print(f"corrupt GET receiver expected error={type(exc).__name__}", flush=True) + else: + raise AssertionError("corrupt GET receiver unexpectedly succeeded") + pending = ray.get(actor.get_payload_transfer_pending_counts.remote()) + assert pending == {"pending_puts": 0, "pending_gets": 0, "pending_receives": 0}, pending + print(f"corrupt GET receiver guard PASS state={pending}", flush=True) + + _, result = await manager._get_from_single_storage_unit([502], ["tensor"], target_storage_unit=info.id) + assert torch.equal(result["tensor"][0], value + 1) + print("subsequent valid PUT/GET PASS", flush=True) + finally: + if manager is not None and manager.payload_transfer is not None: + manager.payload_transfer.close() + if dead_peer is not None: + dead_peer.close() + ray.kill(actor) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--ray-address", required=True) + parser.add_argument("--resource", default=None) + args = parser.parse_args() + runtime_env_vars = { + "TORCH_DEVICE_BACKEND_AUTOLOAD": "0", + } + for name in ("TQ_RAY_WORKER_PYTHONPATH", "TQ_RAY_WORKER_LD_LIBRARY_PATH"): + value = os.environ.get(name) + if value: + runtime_env_vars[name.removeprefix("TQ_RAY_WORKER_")] = value + ray.init( + address=args.ray_address, + include_dashboard=False, + runtime_env={"env_vars": runtime_env_vars}, + ) + try: + asyncio.run(run_case(args.resource)) + finally: + ray.shutdown() + + +if __name__ == "__main__": + main() diff --git a/tools/test_simplestorage_ucx_integration.py b/tools/test_simplestorage_ucx_integration.py new file mode 100644 index 00000000..f5b1163a --- /dev/null +++ b/tools/test_simplestorage_ucx_integration.py @@ -0,0 +1,190 @@ +#!/usr/bin/env python3 +"""Minimal real SimpleStorage ZMQ/UCX integration test. + +This intentionally bypasses the controller metadata path. It exercises the +same manager methods and the real SimpleStorageUnit Ray actor, while keeping +the test focused on the storage data plane handshake. +""" + +from __future__ import annotations + +import argparse +import asyncio +import os +import time + +import numpy as np +import ray +import torch +import zmq.asyncio +from tensordict import NonTensorStack +from tq_test_types import PickleValue + +from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager +from transfer_queue.storage.payload_transfer import create_payload_transfer +from transfer_queue.storage.simple_storage import SimpleStorageUnit + + +def parser(field_data): + field_data["reference"] = [torch.full((3,), 7, dtype=torch.int64) for _ in field_data["reference"]] + return field_data + + +def make_manager( + server_info, + data_info, + enabled: bool, + threshold: int, +): + manager = object.__new__(AsyncSimpleStorageManager) + manager.storage_manager_id = "TQ_INTEGRATION_MANAGER" + # The focused integration tool bypasses StorageManager.__init__, but its + # normal close/__del__ path still expects these base lifecycle fields. + manager.controller_handshake_socket = None + manager.zmq_context = zmq.asyncio.Context() + manager.storage_unit_infos = {server_info.id: server_info} + manager.payload_transfer = create_payload_transfer( + "ucx" if enabled else "zmq", local_ip=ray.util.get_node_ip_address() + ) + manager.payload_transfer_infos = {server_info.id: data_info} if enabled else {} + manager.inline_threshold_bytes = threshold + return manager + + +async def run_case(enabled: bool) -> None: + threshold = 64 * 1024 + repeated_gets = max(1, int(os.environ.get("TQ_INTEGRATION_REPEATED_GETS", "1"))) + use_parser = os.environ.get("TQ_INTEGRATION_USE_PARSER", "1") == "1" + actor_options = {"num_cpus": 1} + actor_resource = os.environ.get("TQ_STORAGE_NODE_RESOURCE") + if actor_resource: + actor_options["resources"] = {actor_resource: 1} + actor = SimpleStorageUnit.options(**actor_options).remote( + storage_unit_size=None, + payload_transfer="ucx" if enabled else "zmq", + ) + info = ray.get(actor.get_zmq_server_info.remote()) + data_info = ray.get(actor.get_payload_transfer_info.remote()) + address = data_info["endpoint"]["data"]["address"] if data_info else None + data_address_len = len(address) if address else None + data_address_type = type(address).__name__ if address else None + print( + f"enabled={enabled} actor ready on {info.ip} data_address_len={data_address_len} type={data_address_type}", + flush=True, + ) + manager = make_manager(info, data_info, enabled, threshold) + try: + large_elements = int(os.environ.get("TQ_INTEGRATION_ELEMENTS", "1000000")) + large = torch.arange(large_elements, dtype=torch.int64) + values = { + "tensor": [large], + "numpy": [np.arange(32, dtype=np.float32)], + "nested": [torch.nested.as_nested_tensor([torch.arange(3), torch.arange(5)], layout=torch.jagged)], + "non_tensor_stack": [NonTensorStack("left", "right")], + "pickle": [PickleValue("fallback")], + "reference": ["shape:3"], + } + payload_bytes = large.numel() * large.element_size() + repetitions = int(os.environ.get("TQ_INTEGRATION_REPETITIONS", "1")) + for repetition in range(repetitions): + key = 101 + repetition + put_start = time.perf_counter() + await manager._put_to_single_storage_unit( + [key], values, target_storage_unit=info.id, data_parser=parser if use_parser else None + ) + put_seconds = time.perf_counter() - put_start + print( + f"enabled={enabled} repetition={repetition} large PUT done seconds={put_seconds:.6f} " + f"payload_bytes={payload_bytes} throughput_mib_s={payload_bytes / put_seconds / 2**20:.2f}", + flush=True, + ) + for get_repeat in range(repeated_gets): + get_start = time.perf_counter() + _, result = await manager._get_from_single_storage_unit( + [key], list(values), target_storage_unit=info.id + ) + get_seconds = time.perf_counter() - get_start + print( + f"enabled={enabled} repetition={repetition} get_repeat={get_repeat} " + f"large GET done seconds={get_seconds:.6f} payload_bytes={payload_bytes} " + f"throughput_mib_s={payload_bytes / get_seconds / 2**20}", + flush=True, + ) + assert torch.equal(result["tensor"][0], large) + np.testing.assert_array_equal(result["numpy"][0], values["numpy"][0]) + if use_parser: + assert result["reference"][0].tolist() == [7, 7, 7] + else: + assert result["reference"][0] == "shape:3" + assert isinstance(result["pickle"][0], PickleValue) + assert result["pickle"][0].value == "fallback" + + # A small payload must use the legacy ZMQ path even when the UCX plane is enabled. + small = {"small": [torch.tensor([1, 2, 3], dtype=torch.int64)]} + await manager._put_to_single_storage_unit([102], small, target_storage_unit=info.id) + print(f"enabled={enabled} small PUT done", flush=True) + _, small_result = await manager._get_from_single_storage_unit([102], ["small"], target_storage_unit=info.id) + print(f"enabled={enabled} small GET done", flush=True) + assert torch.equal(small_result["small"][0], small["small"][0]) + + await manager._clear_single_storage_unit( + list(range(101, 101 + repetitions)) + [102], target_storage_unit=info.id + ) + try: + await manager._get_from_single_storage_unit([101 + repetitions], ["tensor"], target_storage_unit=info.id) + except Exception: + pass + else: + raise AssertionError("CLEAR did not remove stored data") + print(f"simple_storage enabled={enabled} PASS", flush=True) + finally: + if manager.payload_transfer is not None: + manager.payload_transfer.close() + ray.kill(actor) + + +def main() -> None: + parser_args = argparse.ArgumentParser() + parser_args.add_argument("--mode", choices=("legacy", "ucx", "both"), default="both") + parser_args.add_argument("--ray-address", default=None) + args = parser_args.parse_args() + runtime_env_vars = {} + for name in ("TQ_RAY_WORKER_PYTHONPATH", "TQ_RAY_WORKER_LD_LIBRARY_PATH"): + value = os.environ.get(name) + if value: + runtime_env_vars[name.removeprefix("TQ_RAY_WORKER_")] = value + # Standard UCX diagnostics may be propagated by this validation tool; they + # are not TransferQueue user configuration. + for name in ( + "UCX_LOG_LEVEL", + "UCX_LOG_FILE", + "UCX_RNDV_SCHEME", + "UCX_ZCOPY_THRESH", + "UCX_RNDV_THRESH", + "UCX_RNDV_FRAG_SIZE", + ): + value = os.environ.get(name) + if value: + runtime_env_vars[name] = value + runtime_env = {"env_vars": runtime_env_vars} if runtime_env_vars else None + ray.init( + address=args.ray_address, + ignore_reinit_error=True, + include_dashboard=False, + num_cpus=4 if args.ray_address is None else None, + runtime_env=runtime_env, + ) + modes = (False, True) if args.mode == "both" else (args.mode == "ucx",) + try: + asyncio.run(_run_modes(modes)) + finally: + ray.shutdown() + + +async def _run_modes(modes): + for enabled in modes: + await run_case(enabled) + + +if __name__ == "__main__": + main() diff --git a/tools/test_simplestorage_ucx_protocol_state.py b/tools/test_simplestorage_ucx_protocol_state.py new file mode 100644 index 00000000..e6005849 --- /dev/null +++ b/tools/test_simplestorage_ucx_protocol_state.py @@ -0,0 +1,171 @@ +#!/usr/bin/env python3 +"""Exercise UCX control-state ownership and cancellation.""" + +from __future__ import annotations + +import argparse +import asyncio +import os +from uuid import uuid4 + +import ray +import torch +import zmq +import zmq.asyncio +from test_simplestorage_ucx_integration import make_manager + +from transfer_queue.storage.simple_storage import SimpleStorageUnit +from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType, create_zmq_socket + + +async def request(info, message: ZMQMessage) -> ZMQMessage: + context = zmq.asyncio.Context() + socket = create_zmq_socket( + context, + zmq.DEALER, + info.ip, + identity=f"protocol-state-{uuid4().hex[:8]}".encode(), + ) + socket.setsockopt(zmq.RCVTIMEO, 10_000) + socket.setsockopt(zmq.SNDTIMEO, 10_000) + socket.connect(info.to_addr("put_get_socket")) + try: + await socket.send_multipart(message.serialize(), copy=False) + return ZMQMessage.deserialize(await socket.recv_multipart(copy=False)) + finally: + socket.close(linger=0) + context.term() + + +def control_message( + request_type: ZMQRequestType, + sender_id: str, + receiver_id: str, + body: dict, +) -> ZMQMessage: + return ZMQMessage.create( + request_type=request_type, + sender_id=sender_id, + receiver_id=receiver_id, + body=body, + ) + + +async def run_case(resource: str | None) -> None: + options = {"num_cpus": 1} + if resource: + options["resources"] = {resource: 1} + actor = SimpleStorageUnit.options(**options).remote( + storage_unit_size=None, + payload_transfer="ucx", + ) + manager = None + try: + info = ray.get(actor.get_zmq_server_info.remote()) + data_info = ray.get(actor.get_payload_transfer_info.remote()) + manager = make_manager(info, data_info, True, 64 * 1024) + + # Prepare a GET without posting the Manager receive or committing it, + # then cancel it through the production dedicated cancellation socket. + value = torch.arange(1_000_000, dtype=torch.int64) + await manager._put_to_single_storage_unit([700], {"tensor": [value]}, target_storage_unit=info.id) + transfer_id = uuid4().hex + response = await request( + info, + control_message( + ZMQRequestType.GET_DATA_PREPARE, + manager.storage_manager_id, + info.id, + { + "global_indexes": [700], + "fields": ["tensor"], + "transfer_id": transfer_id, + }, + ), + ) + assert response.request_type == ZMQRequestType.GET_DATA_READY + assert ray.get(actor.get_payload_transfer_pending_counts.remote())["pending_gets"] == 1 + response = await request( + info, + control_message( + ZMQRequestType.GET_DATA_CANCEL, + "intruder-B", + info.id, + {"transfer_id": transfer_id}, + ), + ) + assert response.request_type == ZMQRequestType.PUT_GET_ERROR + assert ray.get(actor.get_payload_transfer_pending_counts.remote())["pending_gets"] == 1 + await manager._cancel_payload_get(transfer_id, info.id) + assert ray.get(actor.get_payload_transfer_pending_counts.remote())["pending_gets"] == 0 + print("GET cancel PASS", flush=True) + + # Bind a PUT state to its logical sender. A different sender must not + # cancel the posted receive. + put_descriptor = manager._new_descriptor(1024, 1) + prepare_body = { + "global_indexes": [701], + "descriptor": put_descriptor.to_dict(), + "data_parser": None, + } + response = await request( + info, + control_message(ZMQRequestType.PUT_DATA_PREPARE, "owner-A", info.id, prepare_body), + ) + assert response.request_type == ZMQRequestType.PUT_DATA_READY + response = await request( + info, + control_message( + ZMQRequestType.PUT_DATA_CANCEL, + "intruder-B", + info.id, + {"transfer_id": put_descriptor.transfer_id}, + ), + ) + assert response.request_type == ZMQRequestType.PUT_GET_ERROR + assert ray.get(actor.get_payload_transfer_pending_counts.remote())["pending_puts"] == 1 + response = await request( + info, + control_message( + ZMQRequestType.PUT_DATA_CANCEL, + "owner-A", + info.id, + {"transfer_id": put_descriptor.transfer_id}, + ), + ) + assert response.request_type == ZMQRequestType.PUT_DATA_RESPONSE + assert ray.get(actor.get_payload_transfer_pending_counts.remote())["pending_puts"] == 0 + print("PUT sender binding PASS", flush=True) + + finally: + if manager is not None: + manager.close() + ray.kill(actor) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--ray-address", required=True) + parser.add_argument("--resource", default=None) + args = parser.parse_args() + runtime_env_vars = { + "TORCH_DEVICE_BACKEND_AUTOLOAD": "0", + "TQ_STORAGE_POLLER_TIMEOUT": "1", + } + for name in ("TQ_RAY_WORKER_PYTHONPATH", "TQ_RAY_WORKER_LD_LIBRARY_PATH"): + value = os.environ.get(name) + if value: + runtime_env_vars[name.removeprefix("TQ_RAY_WORKER_")] = value + ray.init( + address=args.ray_address, + include_dashboard=False, + runtime_env={"env_vars": runtime_env_vars}, + ) + try: + asyncio.run(run_case(args.resource)) + finally: + ray.shutdown() + + +if __name__ == "__main__": + main() diff --git a/tools/test_simplestorage_ucx_restart.py b/tools/test_simplestorage_ucx_restart.py new file mode 100644 index 00000000..c5c978b3 --- /dev/null +++ b/tools/test_simplestorage_ucx_restart.py @@ -0,0 +1,88 @@ +#!/usr/bin/env python3 +"""Verify SimpleStorage UCX recovery after replacing the StorageUnit actor.""" + +from __future__ import annotations + +import argparse +import asyncio +import os + +import ray +import torch +from test_simplestorage_ucx_integration import make_manager + +from transfer_queue.storage.simple_storage import SimpleStorageUnit + +THRESHOLD = 64 * 1024 + + +def actor_options(resource: str | None) -> dict: + options = {"num_cpus": 1} + if resource: + options["resources"] = {resource: 1} + return options + + +def create_actor(resource: str | None): + return SimpleStorageUnit.options(**actor_options(resource)).remote( + storage_unit_size=None, + payload_transfer="ucx", + ) + + +async def roundtrip(manager, info, index: int, value: torch.Tensor) -> None: + await manager._put_to_single_storage_unit([index], {"tensor": [value]}, target_storage_unit=info.id) + _, result = await manager._get_from_single_storage_unit([index], ["tensor"], target_storage_unit=info.id) + assert torch.equal(result["tensor"][0], value) + + +async def run_case(resource: str | None) -> None: + actor = create_actor(resource) + manager = None + try: + info = ray.get(actor.get_zmq_server_info.remote()) + data_info = ray.get(actor.get_payload_transfer_info.remote()) + manager = make_manager(info, data_info, True, THRESHOLD) + await roundtrip(manager, info, 201, torch.arange(1_000_000, dtype=torch.int64)) + print(f"restart initial PASS actor={info.id} host={info.ip}", flush=True) + manager.payload_transfer.close() + manager = None + + ray.kill(actor) + actor = create_actor(resource) + new_info = ray.get(actor.get_zmq_server_info.remote()) + new_data_info = ray.get(actor.get_payload_transfer_info.remote()) + manager = make_manager(new_info, new_data_info, True, THRESHOLD) + await roundtrip(manager, new_info, 202, torch.arange(1_000_000, dtype=torch.int64) + 1) + print(f"restart replacement PASS actor={new_info.id} host={new_info.ip}", flush=True) + finally: + if manager is not None and manager.payload_transfer is not None: + manager.payload_transfer.close() + ray.kill(actor) + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("--ray-address", required=True) + parser.add_argument("--resource", default=None) + args = parser.parse_args() + runtime_env_vars = { + "TORCH_DEVICE_BACKEND_AUTOLOAD": "0", + } + for name in ("TQ_RAY_WORKER_PYTHONPATH", "TQ_RAY_WORKER_LD_LIBRARY_PATH"): + value = os.environ.get(name) + if value: + runtime_env_vars[name.removeprefix("TQ_RAY_WORKER_")] = value + ray.init( + address=args.ray_address, + include_dashboard=False, + runtime_env={"env_vars": runtime_env_vars}, + ) + try: + asyncio.run(run_case(args.resource)) + finally: + ray.shutdown() + + +if __name__ == "__main__": + main() diff --git a/tools/test_ucx_binding.py b/tools/test_ucx_binding.py new file mode 100644 index 00000000..777e3775 --- /dev/null +++ b/tools/test_ucx_binding.py @@ -0,0 +1,103 @@ +#!/usr/bin/env python3 +"""Small two-process test for the TQ native UCX Tagged binding.""" + +from __future__ import annotations + +import argparse +import base64 +import hashlib +import os +import sys +import time +from pathlib import Path + +from transfer_queue import _ucx + +SIZES = (64 * 1024, 1024 * 1024, 16 * 1024 * 1024) +REPETITIONS = 10 +REQUEST_TIMEOUT = float(os.environ.get("TQ_TEST_TIMEOUT", "30")) + + +def tag(size: int, iteration: int) -> int: + return ((size << 16) ^ iteration) & ((1 << 63) - 1) + + +def payload(size: int, iteration: int) -> bytes: + seed = hashlib.sha256(f"{size}:{iteration}".encode()).digest() + return (seed * ((size // len(seed)) + 1))[:size] + + +def server(address_file: Path) -> None: + worker = _ucx.Worker() + address_file.write_text(base64.b64encode(worker.address()).decode()) + try: + for size in SIZES: + for iteration in range(REPETITIONS): + expected = payload(size, iteration) + request = worker.post_receive(tag(size, iteration), size) + actual = bytes(request.wait(REQUEST_TIMEOUT)) + if actual != expected: + raise AssertionError(f"payload mismatch: size={size}, iteration={iteration}") + print("server functional PASS", flush=True) + finally: + worker.close() + + +def client(address_file: Path) -> None: + deadline = time.monotonic() + 30 + while not address_file.exists(): + if time.monotonic() > deadline: + raise TimeoutError("server address file was not created") + time.sleep(0.1) + + address = base64.b64decode(address_file.read_text()) + worker = _ucx.Worker() + endpoint = worker.connect(address) + try: + for size in SIZES: + for iteration in range(REPETITIONS): + request = endpoint.post_send(tag(size, iteration), payload(size, iteration)) + request.wait(30.0) + print("client functional PASS", flush=True) + finally: + endpoint.close(30.0) + worker.close() + + +def cancel_and_timeout() -> None: + worker = _ucx.Worker() + try: + request = worker.post_receive(0x1234, 64) + try: + request.wait(0.05) + except RuntimeError: + pass + else: + raise AssertionError("receive timeout did not fail") + + request = worker.post_receive(0x1235, 64) + request.cancel() + print("cancel/timeout PASS", flush=True) + finally: + worker.close() + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("role", choices=("server", "client", "cancel")) + parser.add_argument("address_file", type=Path) + args = parser.parse_args() + if args.role == "server": + server(args.address_file) + elif args.role == "client": + client(args.address_file) + elif args.role == "cancel": + cancel_and_timeout() + + +if __name__ == "__main__": + try: + main() + except Exception as exc: + print(f"FAIL: {type(exc).__name__}: {exc}", file=sys.stderr, flush=True) + raise diff --git a/tools/test_ucx_payload_transfer.py b/tools/test_ucx_payload_transfer.py new file mode 100644 index 00000000..e600825f --- /dev/null +++ b/tools/test_ucx_payload_transfer.py @@ -0,0 +1,76 @@ +#!/usr/bin/env python3 +from __future__ import annotations + +import argparse +import base64 +import hashlib +import time +from pathlib import Path + +from transfer_queue.storage.payload_transfer.ucx_runtime import ( + UcxRuntime, + UcxTransfer, + address_digest, + create_ucx_runtime, + transfer_tag, +) + +TRANSFER_ID = "tq-standalone-data-plane" + + +def payload(size: int) -> bytes: + seed = hashlib.sha256(TRANSFER_ID.encode()).digest() + return (seed * ((size // len(seed)) + 1))[:size] + + +def main() -> None: + p = argparse.ArgumentParser() + p.add_argument("role", choices=("server", "client", "server_exit", "peer_exit")) + p.add_argument("address_file", type=Path) + p.add_argument("--size", type=int, default=8 * 1024 * 1024) + p.add_argument("--timeout", type=float, default=30.0) + p.add_argument("--local-ip", help="use TQ UCX device and transport discovery") + args = p.parse_args() + if args.local_ip: + dp = create_ucx_runtime(local_ip=args.local_ip) + else: + dp = UcxRuntime(args.timeout) + descriptor = UcxTransfer(TRANSFER_ID, transfer_tag(TRANSFER_ID), args.size) + try: + if args.role == "server": + dp.prepare_receive(descriptor) + args.address_file.write_text(base64.b64encode(dp.address).decode()) + actual = dp.finish_receive(descriptor) + assert actual == payload(args.size) + print("payload_transfer server PASS", flush=True) + elif args.role == "server_exit": + args.address_file.write_text(base64.b64encode(dp.address).decode()) + print("payload_transfer server_exit READY", flush=True) + elif args.role == "client": + deadline = time.monotonic() + 30 + while not args.address_file.exists(): + if time.monotonic() > deadline: + raise TimeoutError("address file missing") + time.sleep(0.1) + peer = base64.b64decode(args.address_file.read_text()) + dp.send(peer, descriptor, payload(args.size), peer_address_digest=address_digest(peer)) + print("payload_transfer client PASS", flush=True) + else: + deadline = time.monotonic() + 30 + while not args.address_file.exists(): + if time.monotonic() > deadline: + raise TimeoutError("address file missing") + time.sleep(0.1) + try: + peer = base64.b64decode(args.address_file.read_text()) + dp.send(peer, descriptor, payload(args.size), peer_address_digest=address_digest(peer)) + except (RuntimeError, TimeoutError): + print("payload_transfer peer_exit PASS", flush=True) + else: + raise AssertionError("send unexpectedly succeeded after peer exit") + finally: + dp.close() + + +if __name__ == "__main__": + main() diff --git a/tools/test_ucx_raw_host.py b/tools/test_ucx_raw_host.py new file mode 100644 index 00000000..5d72d50b --- /dev/null +++ b/tools/test_ucx_raw_host.py @@ -0,0 +1,57 @@ +#!/usr/bin/env python3 +"""Minimal raw UCP Tagged host test without torch/TQ imports.""" + +from __future__ import annotations + +import argparse +import base64 +import hashlib +import time +from pathlib import Path + +import _ucx + +TRANSFER_ID = "tq-raw-host-transfer" + + +def tag() -> int: + return int.from_bytes(hashlib.blake2b(TRANSFER_ID.encode(), digest_size=8).digest(), "big") & ((1 << 63) - 1) + + +def payload(size: int) -> bytes: + seed = hashlib.sha256(TRANSFER_ID.encode()).digest() + return (seed * ((size // len(seed)) + 1))[:size] + + +def main() -> None: + parser = argparse.ArgumentParser() + parser.add_argument("role", choices=("server", "client")) + parser.add_argument("address_file", type=Path) + parser.add_argument("--size", type=int, default=16 * 1024 * 1024) + args = parser.parse_args() + worker = _ucx.Worker() + data = payload(args.size) + try: + if args.role == "server": + args.address_file.write_text(base64.b64encode(worker.address()).decode()) + request = worker.post_receive(tag(), args.size) + received = request.wait(60) + assert bytes(received) == data + print(f"raw server PASS size={args.size}", flush=True) + else: + deadline = time.monotonic() + 30 + while not args.address_file.exists(): + if time.monotonic() > deadline: + raise TimeoutError("address file missing") + time.sleep(0.1) + address = base64.b64decode(args.address_file.read_text()) + endpoint = worker.connect(address) + request = endpoint.post_send(tag(), data) + request.wait(60) + print(f"raw client PASS size={args.size}", flush=True) + finally: + worker.close() + + +if __name__ == "__main__": + main() diff --git a/tools/test_ucx_same_length_corrupt.py b/tools/test_ucx_same_length_corrupt.py new file mode 100644 index 00000000..ca2fd9e0 --- /dev/null +++ b/tools/test_ucx_same_length_corrupt.py @@ -0,0 +1,34 @@ +#!/usr/bin/env python3 +"""Probe the remaining same-length-corrupt UCX address boundary.""" + +from __future__ import annotations + +from transfer_queue.storage.payload_transfer.ucx_runtime import UcxError, UcxRuntime, UcxTransfer, transfer_tag + + +def main() -> None: + runtime = UcxRuntime(timeout_seconds=2) + try: + address = bytearray(runtime.address) + if len(address) < 16: + raise AssertionError(f"unexpected worker address length: {len(address)}") + # The first byte is the UCP address version in the current UCX wire + # format. Use the invalid value from the previously observed abort. + address[0] = 9 + descriptor = UcxTransfer( + transfer_id="same-length-corrupt-address", + tag=transfer_tag("same-length-corrupt-address"), + payload_bytes=1, + ) + try: + runtime.send(bytes(address), descriptor, b"x") + except UcxError as exc: + print(f"same-length address guard PASS: {exc}", flush=True) + else: + raise AssertionError("corrupt address unexpectedly connected") + finally: + runtime.close() + + +if __name__ == "__main__": + main() diff --git a/tools/tq_test_types.py b/tools/tq_test_types.py new file mode 100644 index 00000000..7f617a29 --- /dev/null +++ b/tools/tq_test_types.py @@ -0,0 +1,3 @@ +class PickleValue: + def __init__(self, value: str): + self.value = value diff --git a/transfer_queue/config.yaml b/transfer_queue/config.yaml index b7356b60..31805699 100644 --- a/transfer_queue/config.yaml +++ b/transfer_queue/config.yaml @@ -32,6 +32,9 @@ backend: num_data_storage_units: 2 # ZMQ Server IP & Ports (automatically generated during init) zmq_info: null + # Payload transfer for large values. ZMQ keeps the existing path; UCX uses + # ZMQ for control/inline values and UCX for the payload data plane. + payload_transfer: zmq # MooncakeStore: high-performance KV-based hierarchical storage # that supports RDMA transport between GPU and DRAM. diff --git a/transfer_queue/csrc/ucx/ucx_bindings.cpp b/transfer_queue/csrc/ucx/ucx_bindings.cpp new file mode 100644 index 00000000..bb04638b --- /dev/null +++ b/transfer_queue/csrc/ucx/ucx_bindings.cpp @@ -0,0 +1,467 @@ +// Copyright 2026 The TransferQueue Team +// Licensed under the Apache License, Version 2.0 (the "License"); + +#include +#include + +#include + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace py = pybind11; + +class Worker; + +class Endpoint { + public: + Endpoint(std::shared_ptr worker, ucp_ep_h endpoint) : worker_(std::move(worker)), endpoint_(endpoint) {} + ~Endpoint(); + + py::object post_send(uint64_t tag, py::buffer payload); + void flush(double timeout_seconds); + void close(double timeout_seconds); + + private: + std::shared_ptr worker_; + ucp_ep_h endpoint_; +}; + +class ReceiveBuffer; + +struct ReceiveState { + size_t length = 0; +}; + +static void receive_callback(void*, ucs_status_t, const ucp_tag_recv_info_t* info, void* user_data) { + auto* state = static_cast(user_data); + if (info != nullptr) state->length = info->length; +} + +class Request : public std::enable_shared_from_this { + public: + enum class Kind { kSend, kReceive }; + + Request(std::shared_ptr worker, void* request, std::unique_ptr buffer, + std::shared_ptr receive_state) + : worker_(std::move(worker)), + request_(request), + kind_(Kind::kReceive), + buffer_(std::move(buffer)), + receive_state_(std::move(receive_state)), + receive_data_(buffer_.get()) {} + Request(std::shared_ptr worker, void* request, py::object owner) + : worker_(std::move(worker)), request_(request), kind_(Kind::kSend), owner_(std::move(owner)) {} + ~Request(); + + py::object wait(std::optional timeout_seconds); + // Non-blocking completion check. Returns None while the request is in + // progress, True for a completed send, or a receive buffer/list for a + // completed receive. It must be called by the UCX worker owner thread. + py::object test(); + py::object test_cancel(); + void start_cancel(); + void cancel(); + + const uint8_t* data() const { return receive_data_; } + size_t size() const { return receive_state_ == nullptr ? 0 : receive_state_->length; } + + private: + void complete(); + py::object receive_result(); + + std::shared_ptr worker_; + void* request_; + Kind kind_; + // The receive target is overwritten completely by UCX. Allocate it + // without value-initializing/zeroing every byte; zeroing an 8 MiB payload + // before each receive is pure CPU overhead and is not part of the wire + // transfer. The Request owns this allocation until wait() returns. + std::unique_ptr buffer_; + std::shared_ptr receive_state_; + uint8_t* receive_data_ = nullptr; + // Keep the Python send buffer alive until UCX completes; no intermediate + // native copy is needed for sends. + py::object owner_ = py::none(); + bool complete_ = false; +}; + +class ReceiveBuffer { + public: + explicit ReceiveBuffer(std::shared_ptr request) : request_(std::move(request)) {} + + py::buffer_info buffer() { + const uint8_t* data = request_->data(); + const size_t size = request_->size(); + return py::buffer_info( + const_cast(data), + sizeof(uint8_t), + py::format_descriptor::format(), + {static_cast(size)}, + {static_cast(sizeof(uint8_t))}); + } + + private: + std::shared_ptr request_; +}; + +class Worker : public std::enable_shared_from_this { + public: + explicit Worker(const py::dict& options) { + ucp_config_t* config = nullptr; + ucs_status_t status = ucp_config_read(nullptr, nullptr, &config); + check(status, "ucp_config_read"); + for (const auto& item : options) { + const std::string name = py::cast(item.first); + const std::string value = py::cast(item.second); + status = ucp_config_modify(config, name.c_str(), value.c_str()); + if (status != UCS_OK) { + ucp_config_release(config); + check(status, ("ucp_config_modify(" + name + ")").c_str()); + } + } + + ucp_params_t params{}; + params.field_mask = UCP_PARAM_FIELD_FEATURES; + params.features = UCP_FEATURE_TAG; + status = ucp_init(¶ms, config, &context_); + ucp_config_release(config); + check(status, "ucp_init"); + + ucp_worker_params_t worker_params{}; + worker_params.field_mask = UCP_WORKER_PARAM_FIELD_THREAD_MODE; + // All UCP calls are serialized by mutex_. SINGLE avoids relying on + // cross-thread cancellation semantics while retaining safe Python use. + worker_params.thread_mode = UCS_THREAD_MODE_SINGLE; + status = ucp_worker_create(context_, &worker_params, &worker_); + if (status != UCS_OK) { + ucp_cleanup(context_); + context_ = nullptr; + check(status, "ucp_worker_create"); + } + } + + ~Worker() { close(); } + + py::bytes address() { + std::lock_guard lock(mutex_); + ucp_address_t* address = nullptr; + size_t length = 0; + check(ucp_worker_get_address(worker_, &address, &length), "ucp_worker_get_address"); + py::bytes result(reinterpret_cast(address), length); + ucp_worker_release_address(worker_, address); + return result; + } + + std::shared_ptr connect(py::bytes remote_address) { + std::string address = remote_address; + ucp_ep_params_t params{}; + params.field_mask = UCP_EP_PARAM_FIELD_REMOTE_ADDRESS | UCP_EP_PARAM_FIELD_ERR_HANDLING_MODE; + params.address = reinterpret_cast(address.data()); + params.err_mode = UCP_ERR_HANDLING_MODE_PEER; + ucp_ep_h endpoint = nullptr; + { + std::lock_guard lock(mutex_); + check_open(); + check(ucp_ep_create(worker_, ¶ms, &endpoint), "ucp_ep_create"); + } + return std::make_shared(shared_from_this(), endpoint); + } + + std::shared_ptr post_receive(uint64_t tag, size_t length) { + auto buffer = std::make_unique(length); + auto receive_state = std::make_shared(); + ucp_request_param_t params{}; + params.op_attr_mask = UCP_OP_ATTR_FIELD_CALLBACK | UCP_OP_ATTR_FIELD_USER_DATA | + UCP_OP_ATTR_FLAG_NO_IMM_CMPL; + params.cb.recv = receive_callback; + params.user_data = receive_state.get(); + void* request = nullptr; + { + std::lock_guard lock(mutex_); + check_open(); + request = ucp_tag_recv_nbx(worker_, buffer.get(), length, tag, UINT64_MAX, ¶ms); + } + return make_receive_request(request, std::move(buffer), std::move(receive_state), "ucp_tag_recv_nbx"); + } + + std::shared_ptr post_send(ucp_ep_h endpoint, uint64_t tag, py::buffer payload) { + py::buffer_info info = payload.request(); + if (info.ndim != 1 || info.itemsize != 1 || info.strides[0] != 1) { + throw std::runtime_error("UCX send payload must be a contiguous byte buffer"); + } + py::object owner = payload; + ucp_request_param_t params{}; + void* request = nullptr; + { + std::lock_guard lock(mutex_); + check_open(); + request = ucp_tag_send_nbx(endpoint, info.ptr, static_cast(info.size), tag, ¶ms); + } + return make_send_request(request, std::move(owner), "ucp_tag_send_nbx"); + } + + unsigned progress() { + std::lock_guard lock(mutex_); + return worker_ != nullptr ? ucp_worker_progress(worker_) : 0; + } + + void release_request(void* request) { + std::lock_guard lock(mutex_); + if (worker_ != nullptr && request != nullptr) ucp_request_free(request); + } + + void cancel_request(void* request) { + std::lock_guard lock(mutex_); + if (worker_ != nullptr && request != nullptr) ucp_request_cancel(worker_, request); + } + + void close_endpoint(ucp_ep_h endpoint, double timeout_seconds) { + if (endpoint == nullptr) return; + ucp_request_param_t params{}; + params.op_attr_mask = UCP_OP_ATTR_FIELD_FLAGS; + params.flags = UCP_EP_CLOSE_FLAG_FORCE; + void* request = nullptr; + { + std::lock_guard lock(mutex_); + if (worker_ == nullptr) return; + request = ucp_ep_close_nbx(endpoint, ¶ms); + } + wait_native(request, timeout_seconds, "ucp_ep_close_nbx"); + } + + void flush_endpoint(ucp_ep_h endpoint, double timeout_seconds) { + if (endpoint == nullptr) return; + ucp_request_param_t params{}; + void* request = nullptr; + { + std::lock_guard lock(mutex_); + if (worker_ == nullptr) return; + request = ucp_ep_flush_nbx(endpoint, ¶ms); + } + wait_native(request, timeout_seconds, "ucp_ep_flush_nbx"); + } + + void close() { + std::lock_guard lock(mutex_); + if (worker_ != nullptr) { + ucp_worker_destroy(worker_); + worker_ = nullptr; + } + if (context_ != nullptr) { + ucp_cleanup(context_); + context_ = nullptr; + } + } + + static void check(ucs_status_t status, const char* operation) { + if (status != UCS_OK) throw std::runtime_error(std::string(operation) + ": " + ucs_status_string(status)); + } + + void wait_native(void* request, std::optional timeout_seconds, const char* operation, + bool allow_canceled = false) { + if (request == nullptr) return; + if (UCS_PTR_IS_ERR(request)) check(UCS_PTR_STATUS(request), operation); + std::optional deadline; + if (timeout_seconds.has_value()) { + deadline = std::chrono::steady_clock::now() + + std::chrono::duration_cast( + std::chrono::duration(*timeout_seconds)); + } + constexpr unsigned long sleep_us = 50; + while (ucp_request_check_status(request) == UCS_INPROGRESS) { + progress(); + if (deadline.has_value() && std::chrono::steady_clock::now() > *deadline) { + cancel_request(request); + while (ucp_request_check_status(request) == UCS_INPROGRESS) progress(); + ucp_request_free(request); + throw std::runtime_error(std::string(operation) + ": timed out"); + } + if (sleep_us != 0) std::this_thread::sleep_for(std::chrono::microseconds(sleep_us)); + } + const ucs_status_t status = ucp_request_check_status(request); + ucp_request_free(request); + if (status != UCS_OK && !(allow_canceled && status == UCS_ERR_CANCELED)) { + check(status, operation); + } + } + + private: + std::shared_ptr make_receive_request(void* request, std::unique_ptr buffer, + std::shared_ptr receive_state, + const char* operation) { + if (UCS_PTR_IS_ERR(request)) check(UCS_PTR_STATUS(request), operation); + return std::make_shared(shared_from_this(), request, std::move(buffer), std::move(receive_state)); + } + + std::shared_ptr make_send_request(void* request, py::object owner, const char* operation) { + if (UCS_PTR_IS_ERR(request)) check(UCS_PTR_STATUS(request), operation); + return std::make_shared(shared_from_this(), request, std::move(owner)); + } + + void check_open() const { + if (worker_ == nullptr) throw std::runtime_error("UCX worker is closed"); + } + + ucp_context_h context_ = nullptr; + ucp_worker_h worker_ = nullptr; + std::mutex mutex_; +}; + +Endpoint::~Endpoint() { + try { close(0.0); } catch (...) {} +} + +py::object Endpoint::post_send(uint64_t tag, py::buffer payload) { + if (endpoint_ == nullptr) throw std::runtime_error("UCX endpoint is closed"); + return py::cast(worker_->post_send(endpoint_, tag, std::move(payload))); +} + +void Endpoint::close(double timeout_seconds) { + if (endpoint_ == nullptr) return; + worker_->close_endpoint(endpoint_, timeout_seconds); + endpoint_ = nullptr; +} + +void Endpoint::flush(double timeout_seconds) { + if (endpoint_ == nullptr) throw std::runtime_error("UCX endpoint is closed"); + worker_->flush_endpoint(endpoint_, timeout_seconds); +} + +Request::~Request() { + try { cancel(); } catch (...) {} +} + +void Request::complete() { + if (complete_) return; + if (request_ != nullptr) worker_->release_request(request_); + request_ = nullptr; + complete_ = true; +} + +py::object Request::receive_result() { + return py::cast(std::make_shared(shared_from_this())); +} + +py::object Request::wait(std::optional timeout_seconds) { + if (complete_) { + if (kind_ == Kind::kReceive) return receive_result(); + return py::none(); + } + if (request_ != nullptr) { + py::gil_scoped_release release; + try { + worker_->wait_native(request_, timeout_seconds, "UCX request"); + } catch (...) { + // wait_native owns timeout cleanup; do not let the destructor cancel a + // request that has already been freed. + request_ = nullptr; + complete_ = true; + throw; + } + } + request_ = nullptr; + complete(); + if (kind_ == Kind::kReceive) return receive_result(); + return py::none(); +} + +py::object Request::test() { + if (complete_) { + if (kind_ == Kind::kReceive) return receive_result(); + return py::bool_(true); + } + if (request_ == nullptr) { + // UCP may report an immediate completion with a null request handle. + complete_ = true; + if (kind_ == Kind::kReceive) return receive_result(); + return py::bool_(true); + } + + const ucs_status_t status = ucp_request_check_status(request_); + if (status == UCS_INPROGRESS) return py::none(); + if (status != UCS_OK) { + worker_->release_request(request_); + request_ = nullptr; + complete_ = true; + throw std::runtime_error(std::string("UCX request: ") + ucs_status_string(status)); + } + + worker_->release_request(request_); + request_ = nullptr; + complete_ = true; + if (kind_ == Kind::kReceive) return receive_result(); + return py::bool_(true); +} + +void Request::start_cancel() { + if (!complete_ && request_ != nullptr) { + // Some transports may spend measurable time in ucp_request_cancel(). + // Do not hold Python's process-wide GIL while the owner thread enters UCX: + // the independent ZMQ control thread must remain able to acknowledge the + // cancellation and serve unrelated requests. + py::gil_scoped_release release; + worker_->cancel_request(request_); + } +} + +py::object Request::test_cancel() { + if (complete_ || request_ == nullptr) return py::bool_(true); + const ucs_status_t status = ucp_request_check_status(request_); + if (status == UCS_INPROGRESS) return py::none(); + worker_->release_request(request_); + request_ = nullptr; + complete_ = true; + if (status != UCS_OK && status != UCS_ERR_CANCELED) { + throw std::runtime_error(std::string("UCX request cancellation: ") + + ucs_status_string(status)); + } + return py::bool_(true); +} + +void Request::cancel() { + if (complete_) return; + if (request_ != nullptr) { + void* request = request_; + request_ = nullptr; + try { + worker_->cancel_request(request); + worker_->wait_native(request, 1.0, "UCX request cancellation", true); + } catch (...) { + complete_ = true; + throw; + } + } + complete(); +} + +PYBIND11_MODULE(_ucx, m) { + m.doc() = "TransferQueue's narrow UCX UCP Tagged binding"; + py::class_>(m, "Worker") + .def(py::init(), py::arg("config") = py::dict()) + .def("address", &Worker::address) + .def("connect", &Worker::connect) + .def("post_receive", &Worker::post_receive) + .def("progress", &Worker::progress) + .def("close", &Worker::close); + py::class_>(m, "Endpoint") + .def("post_send", &Endpoint::post_send) + .def("flush", &Endpoint::flush) + .def("close", &Endpoint::close); + py::class_>(m, "ReceiveBuffer", py::buffer_protocol()) + .def_buffer(&ReceiveBuffer::buffer); + py::class_>(m, "Request") + .def("wait", &Request::wait) + .def("test", &Request::test) + .def("start_cancel", &Request::start_cancel) + .def("test_cancel", &Request::test_cancel) + .def("cancel", &Request::cancel); +} diff --git a/transfer_queue/storage/bootstrap/simple_storage_bootstrap.py b/transfer_queue/storage/bootstrap/simple_storage_bootstrap.py index 3082120f..fda52312 100644 --- a/transfer_queue/storage/bootstrap/simple_storage_bootstrap.py +++ b/transfer_queue/storage/bootstrap/simple_storage_bootstrap.py @@ -16,9 +16,11 @@ import math from typing import Any +import ray from omegaconf import DictConfig from transfer_queue.storage.bootstrap.provider import StorageBootstrapProvider +from transfer_queue.storage.payload_transfer import normalize_payload_transfer from transfer_queue.storage.simple_storage import SimpleStorageUnit from transfer_queue.utils.common import get_node_round_robin_scheduling_strategies from transfer_queue.utils.logging_utils import get_logger @@ -34,6 +36,7 @@ def initialize_simple_storage(conf: DictConfig) -> dict[str, Any]: simple_storage_handles = {} num_data_storage_units = conf.backend.SimpleStorage.num_data_storage_units total_storage_size = conf.backend.SimpleStorage.get("total_storage_size", None) + payload_transfer = normalize_payload_transfer(conf.backend.SimpleStorage.get("payload_transfer", "zmq")) scheduling_strategies = get_node_round_robin_scheduling_strategies(num_data_storage_units) # Compute per-unit capacity: None means unlimited @@ -47,6 +50,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, ) simple_storage_handles[f"TransferQueueStorageUnit#{storage_unit_rank}"] = storage_node logger.info( @@ -57,5 +61,12 @@ 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 + if payload_transfer != "zmq": + transfer_infos = ray.get( + [storage.get_payload_transfer_info.remote() for storage in simple_storage_handles.values()] + ) + if not all(transfer_infos): + raise RuntimeError("SimpleStorage payload transfer did not initialize on every StorageUnit") + conf.backend[backend_name].payload_transfer_infos = {info["id"]: info for info in transfer_infos} 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 0b6777fb..3458ae42 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -21,19 +21,33 @@ from operator import itemgetter from pathlib import Path from typing import Any, Callable, NamedTuple +from uuid import uuid4 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 ( + DEFAULT_INLINE_PAYLOAD_BYTES, + PayloadDescriptor, + ReceiveToken, + TransferEndpoint, + create_payload_transfer, + normalize_payload_transfer, +) from transfer_queue.utils.logging_utils import get_logger +from transfer_queue.utils.serial_utils import decode, encode, pack_frames, unpack_frames from transfer_queue.utils.zmq_utils import ( ZMQMessage, ZMQRequestType, ZMQServerInfo, + create_zmq_socket, + format_zmq_address, + get_node_ip_address, with_zmq_socket, ) @@ -88,6 +102,13 @@ def __init__(self, controller_info: ZMQServerInfo, config: DictConfig): raise ValueError("AsyncSimpleStorageManager requires non-empty 'zmq_info' in config.") self.storage_unit_infos = self._register_servers(server_infos) + self.payload_transfer_name = normalize_payload_transfer(config.get("payload_transfer", "zmq")) + self.payload_transfer_infos = config.get("payload_transfer_infos", {}) or {} + endpoints_ready = set(self.storage_unit_infos) == set(self.payload_transfer_infos) + if self.payload_transfer_name != "zmq" and not endpoints_ready: + raise RuntimeError("SimpleStorage payload transfer endpoints are missing") + self.payload_transfer = create_payload_transfer(self.payload_transfer_name, local_ip=get_node_ip_address()) + self.inline_threshold_bytes = DEFAULT_INLINE_PAYLOAD_BYTES def _register_servers(self, server_infos: "ZMQServerInfo | dict[Any, ZMQServerInfo]"): """Register and validate server information. @@ -290,6 +311,17 @@ async def _put_to_single_storage_unit( Send data to a specific storage unit. """ + # Preserve the legacy path exactly when optional payload transport is + # disabled. Encoding here would otherwise duplicate the serialization + # performed by ZMQMessage.serialize() and add an unnecessary full + # payload pack/copy to every default PUT. + if self.payload_transfer is not None: + frames = encode(storage_data) + payload = pack_frames(frames) + if len(payload) >= self.inline_threshold_bytes: + await self._put_via_payload_transfer(global_indexes, payload, target_storage_unit, data_parser, socket) + return + request_msg = ZMQMessage.create( request_type=ZMQRequestType.PUT_DATA, # type: ignore[arg-type] sender_id=self.storage_manager_id, @@ -325,6 +357,65 @@ 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 + async def _put_via_payload_transfer( + self, + global_indexes: list[int], + payload: bytes | bytearray | memoryview, + target_storage_unit: str, + data_parser: Callable[[Any], Any] | None, + socket: zmq.Socket, + ) -> None: + """PUT handshake: prepare receive -> payload send -> commit/store.""" + descriptor = self._new_descriptor(len(payload)) + # Once PREPARE is sent, the remote unit may have posted a receive even + # if the control response is lost. Keep cancellation enabled for all + # subsequent failures, including a malformed/late READY response. + remote_may_be_prepared = True + prepare = ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_PREPARE, + sender_id=self.storage_manager_id, + receiver_id=target_storage_unit, + body={ + "global_indexes": global_indexes, + "descriptor": descriptor.to_dict(), + "data_parser": data_parser, + }, + ) + try: + await socket.send_multipart(prepare.serialize(), copy=False) + ready = ZMQMessage.deserialize(await socket.recv_multipart(copy=False)) + self._expect(ready, ZMQRequestType.PUT_DATA_READY, target_storage_unit) + ready_descriptor = PayloadDescriptor.from_dict(ready.body["descriptor"]) + if ready_descriptor != descriptor: + raise RuntimeError(f"PUT descriptor changed by storage unit {target_storage_unit}") + token = ReceiveToken.from_dict(ready.body["receive_token"]) + endpoint = self._peer_endpoint(target_storage_unit) + assert self.payload_transfer is not None + await asyncio.wrap_future(self.payload_transfer.send(endpoint, token, descriptor, payload)) + commit = ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_COMMIT, + sender_id=self.storage_manager_id, + receiver_id=target_storage_unit, + body={"transfer_id": descriptor.transfer_id}, + ) + await socket.send_multipart(commit.serialize(), copy=False) + response = ZMQMessage.deserialize(await socket.recv_multipart(copy=False)) + self._expect(response, ZMQRequestType.PUT_DATA_RESPONSE, target_storage_unit) + remote_may_be_prepared = False + except BaseException: + if remote_may_be_prepared: + await self._cancel_payload_put(descriptor.transfer_id, target_storage_unit) + raise + + async def _cancel_payload_put(self, transfer_id: str, target_storage_unit: str) -> None: + """Best-effort release of a remote receive after a failed PUT.""" + await self._cancel_payload_transfer( + ZMQRequestType.PUT_DATA_CANCEL, + transfer_id, + target_storage_unit, + ZMQRequestType.PUT_DATA_RESPONSE, + ) + @staticmethod def _pack_field_values(values: list) -> torch.Tensor | NonTensorStack: """ @@ -433,6 +524,8 @@ async def _get_from_single_storage_unit( socket: zmq.Socket = None, ): """Get data from a single SU by global index keys.""" + if self.payload_transfer is not None: + return await self._get_via_payload_transfer(global_indexes, fields, target_storage_unit, socket) request_msg = ZMQMessage.create( request_type=ZMQRequestType.GET_DATA, # type: ignore[arg-type] sender_id=self.storage_manager_id, @@ -469,6 +562,146 @@ async def _get_from_single_storage_unit( f"Error getting data from storage unit {target_storage_unit}: {type(e).__name__}: {e}" ) from e + async def _get_via_payload_transfer( + self, global_indexes: list[int], fields: list[str], target_storage_unit: str, socket: zmq.Socket + ): + """GET handshake: encode -> post receive -> commit/start send -> decode.""" + transfer_id = uuid4().hex + remote_prepared = False + receive_prepared = False + prepare = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_PREPARE, + sender_id=self.storage_manager_id, + receiver_id=target_storage_unit, + body={ + "global_indexes": global_indexes, + "fields": fields, + "transfer_id": transfer_id, + }, + ) + try: + await socket.send_multipart(prepare.serialize(), copy=False) + # The StorageUnit may have accepted PREPARE even if READY is lost. + # Use the initial transfer id for best-effort cleanup in that case. + remote_prepared = True + ready = ZMQMessage.deserialize(await socket.recv_multipart(copy=False)) + self._expect(ready, ZMQRequestType.GET_DATA_READY, target_storage_unit) + if ready.body.get("route") == "zmq_inline": + remote_prepared = False + return fields, ready.body["data"] + descriptor = PayloadDescriptor.from_dict(ready.body["descriptor"]) + if descriptor.transfer_id != transfer_id: + raise RuntimeError(f"GET descriptor identity changed by storage unit {target_storage_unit}") + assert self.payload_transfer is not None + token = self.payload_transfer.prepare_receive(descriptor) + receive_prepared = True + commit = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_COMMIT, + sender_id=self.storage_manager_id, + receiver_id=target_storage_unit, + body={ + "transfer_id": descriptor.transfer_id, + "receiver_endpoint": self.payload_transfer.endpoint().to_dict(), + "receive_token": token.to_dict(), + }, + ) + await socket.send_multipart(commit.serialize(), copy=False) + # GET_DATA_RESPONSE acknowledges that the StorageUnit accepted the + # commit and retired its pending protocol entry; data completion is + # represented by the local receive Future. Wait for both so neither + # control nor data-plane completion is left unobserved. + receive_future = self.payload_transfer.receive(descriptor) + receive_task = asyncio.wrap_future(receive_future) + try: + response = ZMQMessage.deserialize(await socket.recv_multipart(copy=False)) + self._expect(response, ZMQRequestType.GET_DATA_RESPONSE, target_storage_unit) + remote_prepared = False + payload = await receive_task + except BaseException: + if not receive_task.done(): + receive_task.cancel() + await asyncio.gather(receive_task, return_exceptions=True) + raise + frames = unpack_frames(payload) + result = decode(frames) + return fields, result + except BaseException: + if receive_prepared: + self.payload_transfer.cancel_receive(descriptor.transfer_id) + if remote_prepared: + await self._cancel_payload_get(transfer_id, target_storage_unit) + raise + + async def _cancel_payload_get(self, transfer_id: str, target_storage_unit: str) -> None: + """Best-effort cancellation of a prepared, uncommitted GET.""" + await self._cancel_payload_transfer( + ZMQRequestType.GET_DATA_CANCEL, + transfer_id, + target_storage_unit, + ZMQRequestType.GET_DATA_RESPONSE, + ) + + async def _cancel_payload_transfer( + self, + request_type: ZMQRequestType, + transfer_id: str, + target_storage_unit: str, + expected_response: ZMQRequestType, + ) -> None: + """Send cancellation on a fresh DEALER to avoid consuming a late response.""" + server_info = self.storage_unit_infos[target_storage_unit] + context = zmq.asyncio.Context() + cancel_socket = create_zmq_socket( + context, + zmq.DEALER, + server_info.ip, + identity=(f"{self.storage_manager_id}_cancel_{target_storage_unit}_{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=self.storage_manager_id, + receiver_id=target_storage_unit, + body={"transfer_id": transfer_id}, + ) + try: + await cancel_socket.send_multipart(cancel.serialize(), copy=False) + response = ZMQMessage.deserialize(await cancel_socket.recv_multipart(copy=False)) + self._expect(response, expected_response, target_storage_unit) + except Exception as exc: + logger.warning( + "[%s]: failed to cancel %s transfer_id=%s at storage unit %s: %s", + self.storage_manager_id, + request_type.value, + transfer_id, + target_storage_unit, + exc, + ) + finally: + cancel_socket.close(linger=0) + context.term() + + def _new_descriptor(self, payload_bytes: int) -> PayloadDescriptor: + return PayloadDescriptor( + transfer_id=uuid4().hex, + payload_bytes=payload_bytes, + ) + + def _peer_endpoint(self, storage_unit_id: str) -> TransferEndpoint: + info = self.payload_transfer_infos[storage_unit_id] + return TransferEndpoint.from_dict(info["endpoint"]) + + @staticmethod + def _expect(response: ZMQMessage, expected: ZMQRequestType, storage_unit_id: str) -> None: + if response.request_type != expected: + raise RuntimeError( + f"storage unit {storage_unit_id} returned {response.request_type}: " + f"{response.body.get('message', 'unknown error')}" + ) + async def clear_data(self, metadata: BatchMeta) -> None: """Clear data in remote StorageUnit. @@ -644,5 +877,7 @@ 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.""" + """Close payload transfer resources before ZMQ ownership is released.""" + if self.payload_transfer is not None: + self.payload_transfer.close() 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..757492b1 --- /dev/null +++ b/transfer_queue/storage/payload_transfer/__init__.py @@ -0,0 +1,30 @@ +# Copyright 2026 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. + +"""Optional out-of-band payload transfer for SimpleStorage.""" + +from transfer_queue.storage.payload_transfer.base import ( + DEFAULT_INLINE_PAYLOAD_BYTES, + PayloadDescriptor, + PayloadTransfer, + PayloadTransferError, + ReceiveToken, + TransferEndpoint, +) +from transfer_queue.storage.payload_transfer.factory import ( + create_payload_transfer, + normalize_payload_transfer, +) + +__all__ = [ + "DEFAULT_INLINE_PAYLOAD_BYTES", + "PayloadDescriptor", + "PayloadTransfer", + "PayloadTransferError", + "ReceiveToken", + "TransferEndpoint", + "create_payload_transfer", + "normalize_payload_transfer", +] diff --git a/transfer_queue/storage/payload_transfer/base.py b/transfer_queue/storage/payload_transfer/base.py new file mode 100644 index 00000000..b1853feb --- /dev/null +++ b/transfer_queue/storage/payload_transfer/base.py @@ -0,0 +1,122 @@ +# Copyright 2026 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. + +"""Optional payload transfer contract used by SimpleStorage.""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from concurrent.futures import Future +from dataclasses import dataclass +from typing import Any + +DEFAULT_INLINE_PAYLOAD_BYTES = 128 * 1024 + + +class PayloadTransferError(RuntimeError): + """A payload transfer could not be completed safely.""" + + +@dataclass(frozen=True) +class PayloadDescriptor: + """Description of one encoded payload.""" + + transfer_id: str + payload_bytes: 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}") + + def to_dict(self) -> dict[str, int | str]: + return { + "transfer_id": self.transfer_id, + "payload_bytes": self.payload_bytes, + } + + @classmethod + def from_dict(cls, value: dict[str, Any]) -> PayloadDescriptor: + descriptor = cls( + transfer_id=str(value["transfer_id"]), + payload_bytes=int(value["payload_bytes"]), + ) + descriptor.validate() + return descriptor + + +@dataclass +class TransferEndpoint: + """Bootstrap metadata for one payload transfer 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"])) + + +class PayloadTransfer(ABC): + """Optional payload data plane; ZMQ remains the control and inline path.""" + + transport: str + + @abstractmethod + def endpoint(self) -> TransferEndpoint: + """Return bootstrap-safe endpoint metadata.""" + + @abstractmethod + def prepare_receive(self, descriptor: PayloadDescriptor) -> ReceiveToken: + """Prepare a receive and return the token required by the sender.""" + + @abstractmethod + def send( + self, + endpoint: TransferEndpoint, + token: ReceiveToken, + descriptor: PayloadDescriptor, + payload: bytes | bytearray | memoryview, + ) -> Future[None]: + """Start sending a payload and return its completion future.""" + + @abstractmethod + def receive(self, descriptor: PayloadDescriptor) -> Future[memoryview]: + """Return a future for a prepared receive.""" + + @abstractmethod + def cancel_receive(self, transfer_id: str) -> None: + """Start best-effort cancellation of a prepared receive.""" + + @property + @abstractmethod + def pending_receive_count(self) -> int: + """Return the number of prepared receives not yet completed.""" + + @abstractmethod + def close(self) -> None: + """Release all transport resources.""" diff --git a/transfer_queue/storage/payload_transfer/factory.py b/transfer_queue/storage/payload_transfer/factory.py new file mode 100644 index 00000000..3f8d6a23 --- /dev/null +++ b/transfer_queue/storage/payload_transfer/factory.py @@ -0,0 +1,29 @@ +# Copyright 2026 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. + +"""Construction helper for optional SimpleStorage payload transfer.""" + +from transfer_queue.storage.payload_transfer.base import PayloadTransfer + +_SUPPORTED_TRANSFERS = frozenset({"zmq", "ucx"}) + + +def normalize_payload_transfer(value: object = "zmq") -> str: + """Return a validated SimpleStorage payload transfer name.""" + normalized = str(value).strip().lower() + if normalized not in _SUPPORTED_TRANSFERS: + raise ValueError(f"unsupported SimpleStorage payload transfer: {normalized!r}; expected 'zmq' or 'ucx'") + return normalized + + +def create_payload_transfer(value: object = "zmq", local_ip: str | None = None) -> PayloadTransfer | None: + """Create the optional data plane; ``zmq`` keeps the existing path.""" + normalized = normalize_payload_transfer(value) + if normalized == "zmq": + return None + + from transfer_queue.storage.payload_transfer.ucx import UcxPayloadTransfer + + return UcxPayloadTransfer(local_ip=local_ip) diff --git a/transfer_queue/storage/payload_transfer/ucx.py b/transfer_queue/storage/payload_transfer/ucx.py new file mode 100644 index 00000000..a2957d27 --- /dev/null +++ b/transfer_queue/storage/payload_transfer/ucx.py @@ -0,0 +1,96 @@ +# Copyright 2026 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. + +"""UCX implementation of the SimpleStorage payload transfer contract.""" + +from __future__ import annotations + +from concurrent.futures import Future + +from transfer_queue.storage.payload_transfer.base import ( + PayloadDescriptor, + PayloadTransfer, + PayloadTransferError, + ReceiveToken, + TransferEndpoint, +) +from transfer_queue.storage.payload_transfer.ucx_runtime import ( + UcxTransfer, + address_digest, + create_ucx_runtime, + transfer_tag, +) + + +class UcxPayloadTransfer(PayloadTransfer): + """Adapt UCX Tagged operations to the payload transfer contract.""" + + transport = "ucx" + + def __init__(self, local_ip: str | None = None): + self._runtime = create_ucx_runtime(local_ip) + + def endpoint(self) -> TransferEndpoint: + address = self._runtime.address + return TransferEndpoint( + transport=self.transport, + data={"address": address, "address_digest": address_digest(address)}, + ) + + def prepare_receive(self, descriptor: PayloadDescriptor) -> ReceiveToken: + self._runtime.prepare_receive(self._ucx_transfer(descriptor)) + return ReceiveToken(data={}) + + def send( + self, + endpoint: TransferEndpoint, + token: ReceiveToken, + descriptor: PayloadDescriptor, + payload: bytes | bytearray | memoryview, + ) -> Future[None]: + self._validate_peer_metadata(endpoint, token, descriptor) + transfer = self._ucx_transfer(descriptor) + return self._runtime.send_async( + endpoint.data["address"], + transfer, + payload, + endpoint.data.get("address_digest"), + ) + + def receive(self, descriptor: PayloadDescriptor) -> Future[memoryview]: + return self._runtime.finish_receive_future(self._ucx_transfer(descriptor)) + + def cancel_receive(self, transfer_id: str) -> None: + self._runtime.cancel_receive(transfer_id) + + @property + def pending_receive_count(self) -> int: + return self._runtime.pending_receive_count + + def close(self) -> None: + self._runtime.close() + + @staticmethod + def _ucx_transfer(descriptor: PayloadDescriptor) -> UcxTransfer: + descriptor.validate() + return UcxTransfer( + transfer_id=descriptor.transfer_id, + tag=transfer_tag(descriptor.transfer_id), + payload_bytes=descriptor.payload_bytes, + ) + + def _validate_peer_metadata( + self, + endpoint: TransferEndpoint, + token: ReceiveToken, + descriptor: PayloadDescriptor, + ) -> None: + descriptor.validate() + if endpoint.transport != self.transport: + raise PayloadTransferError(f"UCX cannot use endpoint for {endpoint.transport!r}") + if token.data: + raise PayloadTransferError("UCX receive token must be empty") + if not isinstance(endpoint.data.get("address"), bytes): + raise PayloadTransferError("UCX endpoint address must be bytes") diff --git a/transfer_queue/storage/payload_transfer/ucx_discovery.py b/transfer_queue/storage/payload_transfer/ucx_discovery.py new file mode 100644 index 00000000..ca533ae0 --- /dev/null +++ b/transfer_queue/storage/payload_transfer/ucx_discovery.py @@ -0,0 +1,252 @@ +# Copyright 2026 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. + +"""Node-local UCX device selection for the UCX payload transfer.""" + +from __future__ import annotations + +import importlib +import ipaddress +import os +import re +import shutil +import subprocess +from dataclasses import dataclass +from functools import lru_cache +from pathlib import Path + +import psutil + +_RC_TRANSPORTS = ("rc_mlx5", "rc_verbs", "rc_v", "rc_x") +_LOCAL_TRANSPORTS = {"sm", "sysv", "posix", "xpmem", "cma", "knem"} +_TRANSPORT_LINE = re.compile(r"^\s*#\s+Transport:\s+(\S+)") +_DEVICE_LINE = re.compile(r"^\s*#\s+Device:\s+(\S+)") + + +@dataclass(frozen=True) +class UcxDeviceSelection: + """RDMA device, port and GID selected for one local IP.""" + + rdma_device: str | None = None + port: int | None = None + netdev: str | None = None + gid_index: int | None = None + + @property + def net_devices(self) -> str | None: + if self.rdma_device is None or self.port is None: + return None + rdma = f"{self.rdma_device}:{self.port}" + return f"{rdma},{self.netdev}" if self.netdev else rdma + + @property + def ucx_config(self) -> dict[str, str]: + """Return UCP context settings for this node-local selection.""" + config: dict[str, str] = {} + if self.rdma_device: + # Do not infer a transport from the vendor or device name. Select + # only from transports advertised by this UCX runtime. An + # explicit process setting remains authoritative. + tls = _select_host_ucx_tls(self.rdma_device, self.port or 0, self.netdev) + if tls: + config["TLS"] = tls + if self.net_devices: + config["NET_DEVICES"] = os.environ.get("UCX_NET_DEVICES", self.net_devices) + if self.gid_index is not None: + config["IB_GID_INDEX"] = os.environ.get("UCX_IB_GID_INDEX", str(self.gid_index)) + config["IB_ADDR_TYPE"] = os.environ.get("UCX_IB_ADDR_TYPE", "ib_global") + elif os.environ.get("UCX_IB_GID_INDEX") is not None: + config["IB_GID_INDEX"] = os.environ["UCX_IB_GID_INDEX"] + config["IB_ADDR_TYPE"] = os.environ.get("UCX_IB_ADDR_TYPE", "ib_global") + return config + + +def discover_ucx_device( + local_ip: str | None, + infiniband_root: Path = Path("/sys/class/infiniband"), + interface_addresses: dict[str, set[str]] | None = None, +) -> UcxDeviceSelection: + """Find the RoCE-v2 GID that carries ``local_ip``. + + An exact local-IP match wins. If the Ray/control IP is on another network, + a single unambiguous RoCE-v2 candidate is accepted; multiple candidates are + left unresolved rather than guessed. + """ + if not infiniband_root.is_dir(): + return UcxDeviceSelection() + try: + normalized_ip = str(ipaddress.ip_address(local_ip.split("%", 1)[0])) if local_ip else None + except ValueError: + normalized_ip = None + interface_addresses = interface_addresses or _interface_addresses() + candidates: list[UcxDeviceSelection] = [] + + for rdma_path in sorted(infiniband_root.iterdir()): + ports_path = rdma_path / "ports" + if not ports_path.is_dir(): + continue + for port_path in sorted(ports_path.iterdir()): + try: + port = int(port_path.name) + except ValueError: + continue + ndevs_path = port_path / "gid_attrs" / "ndevs" + gids_path = port_path / "gids" + types_path = port_path / "gid_attrs" / "types" + if not ndevs_path.is_dir() or not gids_path.is_dir(): + continue + for ndev_path in sorted(ndevs_path.iterdir(), key=lambda path: int(path.name)): + netdev = _read_text(ndev_path) + gid = _read_text(gids_path / ndev_path.name) + gid_type = _read_text(types_path / ndev_path.name) + if not netdev or (gid_type and "v2" not in gid_type.lower()): + continue + assigned_addresses = interface_addresses.get(netdev, set()) + if not any(gid_matches_ip(gid, address) for address in assigned_addresses): + continue + selection = UcxDeviceSelection(rdma_path.name, port, netdev, int(ndev_path.name)) + if normalized_ip and gid_matches_ip(gid, normalized_ip): + return selection + candidates.append(selection) + unique_candidates = list(dict.fromkeys(candidates)) + return unique_candidates[0] if len(unique_candidates) == 1 else UcxDeviceSelection() + + +def _select_host_ucx_tls(rdma_device: str, port: int, netdev: str | None) -> str | None: + explicit = os.environ.get("UCX_TLS") + if explicit is not None: + return explicit + + capabilities = _discover_ucx_transports() + rdma_address = f"{rdma_device}:{port}" + rc_transport = next( + ( + transport + for transport in _RC_TRANSPORTS + if (transport, rdma_address) in capabilities or (transport, rdma_device) in capabilities + ), + None, + ) + if rc_transport is None: + return None + + selected = [rc_transport] + if ("tcp", netdev) in capabilities: + selected.append("tcp") + elif ("ud_verbs", rdma_address) in capabilities: + # RC may use UD only as its auxiliary wireup transport. Do not add + # UD to the data lanes unless UCX advertises no TCP alternative. + selected.append("ud_verbs:aux") + + if capabilities.intersection({(name, "memory") for name in _LOCAL_TRANSPORTS}): + selected.append("sm") + if ("self", "memory") in capabilities: + selected.append("self") + return ",".join(selected) + + +@lru_cache(maxsize=1) +def _discover_ucx_transports() -> frozenset[tuple[str, str]]: + executable = _find_ucx_info() + if executable is None: + return frozenset() + + environment = os.environ.copy() + for name in ("UCX_TLS", "UCX_NET_DEVICES", "UCX_IB_GID_INDEX", "UCX_IB_ADDR_TYPE"): + environment.pop(name, None) + prefix = executable.parent.parent + library_paths = [path for path in (prefix / "lib", prefix / "lib64") if path.is_dir()] + if library_paths: + existing = environment.get("LD_LIBRARY_PATH") + paths = [str(path) for path in library_paths] + if existing: + paths.append(existing) + environment["LD_LIBRARY_PATH"] = ":".join(paths) + try: + result = subprocess.run( + [str(executable), "-d"], + env=environment, + capture_output=True, + check=True, + text=True, + timeout=5, + ) + except (OSError, subprocess.SubprocessError): + return frozenset() + return frozenset(_parse_ucx_info_devices(result.stdout)) + + +def _find_ucx_info() -> Path | None: + configured = os.environ.get("TQ_UCX_INFO") + if configured: + path = Path(configured) + return path if path.is_file() and os.access(path, os.X_OK) else None + + # Prefer the UCX runtime loaded by the native extension. A system + # ``ucx_info`` in PATH may describe a different UCX installation. + try: + importlib.import_module("transfer_queue._ucx") + + for line in Path("/proc/self/maps").read_text().splitlines(): + fields = line.split() + if not fields or "libucp.so" not in fields[-1]: + continue + candidate = Path(fields[-1]).parent.parent / "bin" / "ucx_info" + if candidate.is_file() and os.access(candidate, os.X_OK): + return candidate + except (ImportError, OSError): + pass + + path = shutil.which("ucx_info") + if path: + return Path(path) + return None + + +def _parse_ucx_info_devices(output: str) -> set[tuple[str, str]]: + capabilities: set[tuple[str, str]] = set() + transport: str | None = None + for line in output.splitlines(): + transport_match = _TRANSPORT_LINE.match(line) + if transport_match: + transport = transport_match.group(1) + continue + device_match = _DEVICE_LINE.match(line) + if transport and device_match: + capabilities.add((transport, device_match.group(1))) + return capabilities + + +def gid_matches_ip(gid: str | None, local_ip: str) -> bool: + if not gid: + return False + try: + gid_address = ipaddress.ip_address(gid) + local_address = ipaddress.ip_address(local_ip) + except ValueError: + return False + return gid_address == local_address or ( + isinstance(gid_address, ipaddress.IPv6Address) and gid_address.ipv4_mapped == local_address + ) + + +def _interface_addresses() -> dict[str, set[str]]: + result: dict[str, set[str]] = {} + for interface, addresses in psutil.net_if_addrs().items(): + normalized = set() + for address in addresses: + try: + normalized.add(str(ipaddress.ip_address(address.address.split("%", 1)[0]))) + except ValueError: + continue + result[interface] = normalized + return result + + +def _read_text(path: Path) -> str | None: + try: + return path.read_text().strip() + except (OSError, UnicodeError): + return None diff --git a/transfer_queue/storage/payload_transfer/ucx_runtime.py b/transfer_queue/storage/payload_transfer/ucx_runtime.py new file mode 100644 index 00000000..39184758 --- /dev/null +++ b/transfer_queue/storage/payload_transfer/ucx_runtime.py @@ -0,0 +1,597 @@ +# Copyright 2026 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. + +"""Low-level UCX Tagged transport implementation. + +This module owns UCX-specific discovery, descriptors and request lifecycle. The +SimpleStorage integration uses it only through :class:`UcxPayloadTransfer`. +""" + +from __future__ import annotations + +import hashlib +import os +from concurrent.futures import CancelledError, Future +from dataclasses import dataclass +from queue import Empty, Queue +from threading import Event, Thread +from time import perf_counter, sleep +from typing import Any, Callable, Literal + +from transfer_queue.storage.payload_transfer.base import PayloadTransferError +from transfer_queue.storage.payload_transfer.ucx_discovery import discover_ucx_device +from transfer_queue.utils.logging_utils import get_logger + +logger = get_logger(__name__) + +DEFAULT_TRANSFER_TIMEOUT_SECONDS = float(os.environ.get("TQ_UCX_TRANSFER_TIMEOUT_SECONDS", "200")) +DEFAULT_ENDPOINT_TIMEOUT_SECONDS = 30.0 +DEFAULT_CANCEL_TIMEOUT_SECONDS = 5.0 +_RDMA_UCX_TRANSPORTS = {"rc", "rc_mlx5", "rc_verbs", "rc_v", "rc_x"} + + +class UcxError(PayloadTransferError): + """A data-plane operation could not be completed safely.""" + + +class _UcxRequestFuture(Future[Any]): + """Future whose cancellation remains effective after a UCP request is posted.""" + + def __init__(self) -> None: + super().__init__() + self._cancel_requested = Event() + + def cancel(self) -> bool: + self._cancel_requested.set() + if super().cancel(): + return True + return not self.done() + + @property + def cancel_requested(self) -> bool: + return self._cancel_requested.is_set() + + +@dataclass +class _OwnerTask: + """One operation submitted to the thread-affine UCP worker.""" + + kind: Literal["call", "request", "cancel", "stop"] + operation: Callable[[], Any] + future: Future[Any] + poll: Callable[[Any], Any] | None = None + complete: Callable[[Any], Any] | None = None + timeout_seconds: float | None = None + + +@dataclass +class _ActiveRequest: + """A posted UCP request that the owner thread must keep progressing.""" + + request: Any + future: Future[Any] + poll: Callable[[Any], Any] + complete: Callable[[Any], Any] + deadline: float | None + + +class _UcxOwnerThread: + """Own one UCP worker and progress all active requests on one thread.""" + + def __init__(self, initialize: Callable[[], None], progress: Callable[[], None]): + self._tasks: Queue[_OwnerTask] = Queue() + self._active: list[_ActiveRequest] = [] + self._canceling: list[Any] = [] + self._cancel_waiters: list[Future[None]] = [] + self._progress_callback = progress + self._progress_sleep_seconds = 50 / 1_000_000 + self._ready = Event() + self._init_error: BaseException | None = None + self._thread = Thread(target=self._run, args=(initialize,), name="tq-ucx-owner", daemon=True) + self._thread.start() + self._ready.wait() + if self._init_error is not None: + raise self._init_error + + def _run(self, initialize: Callable[[], None]) -> None: + try: + initialize() + except BaseException as exc: + self._init_error = exc + self._ready.set() + return + self._ready.set() + while True: + try: + # Never delay UCP progress behind the control-task queue. A + # blocking get is useful only while the worker is idle; with + # active or canceling requests it added up to 1 ms to every + # progress iteration and dominated the configured 50 us yield. + task = self._tasks.get_nowait() if self._active or self._canceling else self._tasks.get(timeout=0.001) + except Empty: + task = None + if task is not None: + if task.kind == "stop": + self._cancel_active_requests(UcxError("UCX owner thread is stopping")) + task.future.set_result(None) + break + if task.kind == "cancel": + self._cancel_active_requests(UcxError("UCX request canceled during shutdown")) + if self._canceling: + self._cancel_waiters.append(task.future) + else: + task.future.set_result(None) + continue + if task.future.set_running_or_notify_cancel(): + try: + if task.kind == "call": + task.future.set_result(task.operation()) + else: + request = task.operation() + assert task.poll is not None + assert task.complete is not None + self._active.append( + _ActiveRequest( + request=request, + future=task.future, + poll=task.poll, + complete=task.complete, + deadline=( + None if task.timeout_seconds is None else perf_counter() + task.timeout_seconds + ), + ) + ) + except BaseException as exc: + task.future.set_exception(exc) + + self._progress() + if self._active: + remaining = [] + for active in self._active: + if active.future.cancelled() or ( + isinstance(active.future, _UcxRequestFuture) and active.future.cancel_requested + ): + self._start_cancel(active.request) + if not active.future.done(): + active.future.set_exception(CancelledError()) + continue + if active.deadline is not None and perf_counter() >= active.deadline: + self._start_cancel(active.request) + active.future.set_exception(TimeoutError("UCX request timed out")) + continue + try: + result = active.poll(active.request) + if result is None: + remaining.append(active) + else: + active.future.set_result(active.complete(result)) + except BaseException as exc: + # Retire the native request before dropping the logical + # request from _active. + self._start_cancel(active.request) + active.future.set_exception(exc) + self._active = remaining + if self._active and self._progress_sleep_seconds: + sleep(self._progress_sleep_seconds) + if self._canceling: + canceling = [] + for request in self._canceling: + try: + if request.test_cancel() is None: + canceling.append(request) + except Exception: + # The request is terminal and test_cancel() has already + # released its native handle before reporting failure. + pass + self._canceling = canceling + if not canceling: + waiters, self._cancel_waiters = self._cancel_waiters, [] + for waiter in waiters: + if not waiter.done(): + waiter.set_result(None) + elif not self._active and self._progress_sleep_seconds: + sleep(self._progress_sleep_seconds) + + def _start_cancel(self, request: Any) -> None: + """Begin native cancellation and retain the request until terminal.""" + try: + request.start_cancel() + self._canceling.append(request) + except Exception: + # A failed start means there is no safe action left for this + # best-effort cleanup path. Normal request failures are delivered + # through the caller-facing Future before reaching here. + pass + + def _cancel_active_requests(self, error: BaseException) -> None: + """Cancel native requests before their worker/context can be destroyed.""" + active, self._active = self._active, [] + for request in active: + self._start_cancel(request.request) + if not request.future.done(): + request.future.set_exception(error) + + def _progress(self) -> None: + try: + self._progress_callback() + except BaseException as exc: + self._cancel_active_requests(UcxError(f"UCX progress failed: {type(exc).__name__}: {exc}")) + + def call(self, operation: Callable[[], Any]) -> Any: + return self.submit(operation).result() + + def submit(self, operation: Callable[[], Any]) -> Future[Any]: + """Queue an owner-thread operation without blocking the caller.""" + future: Future[Any] = Future() + self._tasks.put(_OwnerTask(kind="call", operation=operation, future=future)) + return future + + def submit_request( + self, + operation: Callable[[], Any], + poll: Callable[[Any], Any], + complete: Callable[[Any], Any] | None, + timeout_seconds: float | None, + ) -> Future[Any]: + """Post a native request and poll it alongside other in-flight requests.""" + future: Future[Any] = _UcxRequestFuture() + # Keeping both callbacks in the owner thread avoids calling + # thread-affine UCP request methods from executor threads. + self._tasks.put( + _OwnerTask( + kind="request", + operation=operation, + future=future, + poll=poll, + complete=complete or (lambda value: value), + timeout_seconds=timeout_seconds, + ) + ) + return future + + def cancel_active_requests(self) -> None: + """Synchronously cancel all active native requests on the owner thread.""" + future: Future[None] = Future() + self._tasks.put(_OwnerTask(kind="cancel", operation=lambda: None, future=future)) + future.result() + + def stop(self) -> None: + future: Future[None] = Future() + self._tasks.put(_OwnerTask(kind="stop", operation=lambda: None, future=future)) + future.result() + self._thread.join() + + +@dataclass(frozen=True) +class UcxTransfer: + """Control-plane description of exactly one encoded payload.""" + + transfer_id: str + tag: int + payload_bytes: int + + def validate_identity(self) -> None: + """Validate fields that identify a transfer before its size is known.""" + if not self.transfer_id or len(self.transfer_id) > 128: + raise UcxError("UCX transfer_id must contain 1 to 128 characters") + expected_tag = transfer_tag(self.transfer_id) + if self.tag != expected_tag: + raise UcxError(f"UCX tag mismatch for {self.transfer_id}: expected {expected_tag}, got {self.tag}") + + def validate(self) -> None: + """Reject malformed control metadata before native allocation or I/O.""" + self.validate_identity() + if self.payload_bytes < 0: + raise UcxError(f"negative payload length for {self.transfer_id}") + + +def transfer_tag(transfer_id: str) -> int: + """Derive a stable positive 63-bit UCX tag from a UUID-like transfer id.""" + digest = hashlib.blake2b(transfer_id.encode(), digest_size=8).digest() + return int.from_bytes(digest, "big") & ((1 << 63) - 1) + + +def address_digest(address: bytes) -> str: + """Return the digest used to bind a UCX address to bootstrap metadata.""" + if not isinstance(address, bytes): + raise UcxError("UCX worker address must be bytes") + return hashlib.sha256(address).hexdigest() + + +class UcxRuntime: + """Thin synchronous/future wrapper over the native UCP Tagged binding. + + Synchronous calls are available for setup and small standalone tools. Large + transfer paths use the Future-returning methods, which keep UCP progress on + the owner thread without blocking the caller's asyncio loop. + """ + + def __init__( + self, + timeout_seconds: float | None, + require_address_digest: bool = True, + ucx_config: dict[str, str] | None = None, + ): + self._timeout_seconds = timeout_seconds + self._require_address_digest = require_address_digest + self._ucx_config = ucx_config or {} + # Keep the complete native object graph on one owner thread; callers + # may use this facade from asyncio or Ray/ZMQ worker threads. + self._owner = _UcxOwnerThread(self._init_worker, lambda: self._worker.progress()) + self._endpoints: dict[bytes, Any] = {} + self._receives: dict[str, Any] = {} + self._closed = False + + def _init_worker(self) -> None: + try: + from transfer_queue import _ucx + except ImportError as exc: # pragma: no cover - depends on optional native build + raise UcxError("the installed TransferQueue package does not include UCX support") from exc + self._worker = _ucx.Worker(self._ucx_config) + # A worker address is immutable for the lifetime of its worker. Cache + # it once on the owner thread instead of enqueueing a native call for + # every GET request and bootstrap read. + self._address = self._worker.address() + + def _call(self, operation: Callable[[], Any]) -> Any: + if self._closed: + raise UcxError("UCX data plane is closed") + return self._owner.call(operation) + + @property + def address(self) -> bytes: + if self._closed: + raise UcxError("UCX data plane is closed") + return self._address + + def prepare_receive(self, descriptor: UcxTransfer) -> None: + self._validate_descriptor(descriptor) + + def operation() -> None: + if descriptor.transfer_id in self._receives: + raise UcxError(f"duplicate receive preparation: {descriptor.transfer_id}") + request = self._worker.post_receive(descriptor.tag, descriptor.payload_bytes) + self._receives[descriptor.transfer_id] = request + + self._call(operation) + + def send( + self, + peer_address: bytes, + descriptor: UcxTransfer, + payload: bytes | bytearray | memoryview, + peer_address_digest: str | None = None, + ) -> None: + self._validate_send(peer_address, descriptor, len(payload), peer_address_digest) + + def operation() -> None: + request = self._post_send(peer_address, peer_address_digest, descriptor, payload) + request.wait(self._timeout_seconds) + + self._call(operation) + + def send_async( + self, + peer_address: bytes, + descriptor: UcxTransfer, + payload: bytes | bytearray | memoryview, + peer_address_digest: str | None = None, + ) -> Future[None]: + """Start a send without blocking the caller's control-plane handler. + + The future owns the payload through the submitted operation. UCX + posting, progress, completion polling, and endpoint access all remain + on the single owner thread; multiple requests can be in flight without + serializing on one blocking ``Request.wait()``. + """ + if self._closed: + raise UcxError("UCX data plane is closed") + self._validate_send(peer_address, descriptor, len(payload), peer_address_digest) + + def operation() -> Any: + return self._post_send(peer_address, peer_address_digest, descriptor, payload) + + future = self._owner.submit_request( + operation, + lambda request: request.test(), + lambda _: None, + self._timeout_seconds, + ) + + def discard_failed_endpoint(completed: Future[None]) -> None: + try: + completed.result() + except BaseException: + self._owner.submit(lambda: self._endpoints.pop(peer_address, None)) + + future.add_done_callback(discard_failed_endpoint) + return future + + def warmup(self, peer_address: bytes, peer_address_digest: str | None = None) -> None: + """Create and flush a cached endpoint before the first payload send.""" + if self._require_address_digest and peer_address_digest is None: + raise UcxError("UCX worker address digest is required") + + def operation() -> None: + endpoint = self._endpoint(peer_address, peer_address_digest) + endpoint.flush(DEFAULT_ENDPOINT_TIMEOUT_SECONDS) + + self._call(operation) + + def _finish_receive_operation(self, descriptor: UcxTransfer) -> memoryview: + try: + request = self._receives.pop(descriptor.transfer_id) + except KeyError as exc: + raise UcxError(f"receive was not prepared: {descriptor.transfer_id}") from exc + return self._received_payload(descriptor, request.wait(self._timeout_seconds)) + + def finish_receive(self, descriptor: UcxTransfer) -> memoryview: + return self._call(lambda: self._finish_receive_operation(descriptor)) + + def finish_receive_future(self, descriptor: UcxTransfer) -> Future[memoryview]: + """Complete a receive on the UCX owner thread without a Python worker thread.""" + if self._closed: + raise UcxError("UCX data plane is closed") + + def operation() -> Any: + try: + return self._receives.pop(descriptor.transfer_id) + except KeyError as exc: + raise UcxError(f"receive was not prepared: {descriptor.transfer_id}") from exc + + def complete(payload: Any) -> memoryview: + return self._received_payload(descriptor, payload) + + return self._owner.submit_request(operation, lambda request: request.test(), complete, self._timeout_seconds) + + def cancel_receive(self, transfer_id: str) -> Future[None]: + """Start canceling a posted receive without blocking the control plane.""" + if self._closed: + raise UcxError("UCX data plane is closed") + + def operation() -> Any: + request = self._receives.pop(transfer_id, None) + if request is not None: + request.start_cancel() + return request + + future = self._owner.submit_request( + operation, + lambda request: True if request is None else request.test_cancel(), + lambda _: None, + DEFAULT_CANCEL_TIMEOUT_SECONDS, + ) + + def report_cancel_failure(completed: Future[None]) -> None: + if completed.cancelled(): + return + try: + completed.result() + except Exception: + # Cancellation is best effort. The owner thread still applies + # its request timeout and keeps the native request alive until + # UCX reports a terminal state. + pass + + future.add_done_callback(report_cancel_failure) + return future + + @property + def pending_receive_count(self) -> int: + """Return the number of posted, not-yet-finished receives.""" + return self._call(lambda: len(self._receives)) + + def close(self) -> None: + if self._closed: + return + # Refuse new calls before enqueueing shutdown work. Existing requests + # are canceled below on the owner thread before the native worker dies. + self._closed = True + + def operation() -> None: + for transfer_id in list(self._receives): + request = self._receives.pop(transfer_id) + request.cancel() + for endpoint in self._endpoints.values(): + endpoint.close(DEFAULT_ENDPOINT_TIMEOUT_SECONDS) + self._endpoints.clear() + worker = self._worker + worker.close() + # Keep destruction on the UCX owner thread as well. The binding + # owns UCP objects whose cleanup is thread-sensitive on HNS. + self._worker = None + + try: + self._owner.cancel_active_requests() + self._owner.call(operation) + finally: + self._owner.stop() + + def _validate_send( + self, + peer_address: bytes, + descriptor: UcxTransfer, + payload_bytes: int, + peer_address_digest: str | None, + ) -> None: + self._validate_descriptor(descriptor) + if payload_bytes != descriptor.payload_bytes: + raise UcxError( + f"payload length mismatch for {descriptor.transfer_id}: " + f"expected {descriptor.payload_bytes}, got {payload_bytes}" + ) + if self._require_address_digest and peer_address_digest is None: + raise UcxError("UCX worker address digest is required") + self._validate_peer_address(peer_address, peer_address_digest) + + def _validate_descriptor(self, descriptor: UcxTransfer) -> None: + descriptor.validate() + + def _post_send( + self, + peer_address: bytes, + peer_address_digest: str | None, + descriptor: UcxTransfer, + payload: bytes | bytearray | memoryview, + ) -> Any: + endpoint = self._endpoint(peer_address, peer_address_digest) + return endpoint.post_send(descriptor.tag, memoryview(payload)) + + @staticmethod + def _received_payload(descriptor: UcxTransfer, received: Any) -> memoryview: + """Validate the single logical buffer filled by one or more receives.""" + payload = memoryview(received) + if payload.nbytes != descriptor.payload_bytes: + raise UcxError( + f"received length mismatch for {descriptor.transfer_id}: " + f"expected {descriptor.payload_bytes}, got {payload.nbytes}" + ) + return payload + + @staticmethod + def _validate_peer_address(peer_address: bytes, expected_digest: str | None = None) -> None: + """Validate Python control metadata before queueing native endpoint work.""" + if not isinstance(peer_address, bytes): + raise UcxError("UCX worker address must be bytes") + if expected_digest is not None and address_digest(peer_address) != expected_digest: + raise UcxError("UCX worker address digest does not match bootstrap metadata") + # This is a corruption guard, not native address parsing. + if not 16 <= len(peer_address) <= 1024 * 1024: + raise UcxError(f"invalid UCX worker address length: {len(peer_address)}") + + def _endpoint(self, peer_address: bytes, expected_digest: str | None = None): + self._validate_peer_address(peer_address, expected_digest) + endpoint = self._endpoints.get(peer_address) + if endpoint is None: + endpoint = self._worker.connect(peer_address) + self._endpoints[peer_address] = endpoint + return endpoint + + +def create_ucx_runtime(local_ip: str | None = None) -> UcxRuntime: + """Create the strict UCX/RDMA data plane for one local endpoint.""" + selection = discover_ucx_device(local_ip) + if selection.rdma_device is None: + raise UcxError(f"no RoCE-v2 device and GID match local IP {local_ip or 'unknown'}") + try: + ucx_config = selection.ucx_config + tls = {item.strip().split(":", 1)[0] for item in ucx_config.get("TLS", "").split(",") if item.strip()} + if not tls.intersection(_RDMA_UCX_TRANSPORTS): + raise UcxError("UCX selection has no reliable-connection RDMA transport") + transport = UcxRuntime( + timeout_seconds=DEFAULT_TRANSFER_TIMEOUT_SECONDS, + require_address_digest=True, + ucx_config=ucx_config, + ) + logger.info( + "SimpleStorage payload transfer selected: ucx local_ip=%s device=%s gid_index=%s tls=%s", + local_ip or "auto", + selection.net_devices, + selection.gid_index, + ucx_config.get("TLS", "ucx-default"), + ) + return transport + except Exception as exc: + raise UcxError(f"SimpleStorage UCX payload transfer is unavailable: {type(exc).__name__}: {exc}") from exc diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index aa23855f..61dfd191 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 dataclasses import dataclass from threading import Event, Thread from typing import TYPE_CHECKING, Any from uuid import uuid4 @@ -25,10 +26,18 @@ import ray import zmq +from transfer_queue.storage.payload_transfer import ( + DEFAULT_INLINE_PAYLOAD_BYTES, + PayloadDescriptor, + ReceiveToken, + TransferEndpoint, + 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 from transfer_queue.utils.perf_utils import IntervalPerfMonitor +from transfer_queue.utils.serial_utils import decode, encode, pack_frames, unpack_frames from transfer_queue.utils.zmq_utils import ( ZMQMessage, ZMQRequestType, @@ -46,6 +55,26 @@ TQ_STORAGE_POLLER_TIMEOUT = int(os.environ.get("TQ_STORAGE_POLLER_TIMEOUT", 5)) # in seconds TQ_NUM_THREADS = int(os.environ.get("TQ_NUM_THREADS", 8)) +_MAX_PENDING_PAYLOAD_TRANSFERS = 64 +_MAX_PENDING_PAYLOAD_BYTES = 2**32 - 1 +_PENDING_PAYLOAD_TTL_SECONDS = 300 + + +@dataclass +class _PendingPut: + descriptor: PayloadDescriptor + sender_id: str + global_indexes: tuple[int, ...] + data_parser: Any + created_at: float + + +@dataclass +class _PendingGet: + descriptor: PayloadDescriptor + sender_id: str + payload: bytearray + created_at: float class StorageUnitData: @@ -154,17 +183,31 @@ 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: str = "zmq", + ): """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: ``zmq`` keeps the existing path; ``ucx`` enables + the optional payload data plane. """ self.storage_unit_id = f"TQ_STORAGE_UNIT_{uuid4().hex[:8]}" self.storage_unit_size = storage_unit_size + self._node_ip = get_node_ip_address() self.storage_data = StorageUnitData(self.storage_unit_size) + self.payload_transfer = create_payload_transfer(payload_transfer, local_ip=self._node_ip) + self.inline_threshold_bytes = DEFAULT_INLINE_PAYLOAD_BYTES + self._pending_puts: dict[str, _PendingPut] = {} + # The submitted transfer owns the payload/frame references until + # completion. Keep only the protocol state needed by GET_COMMIT here, + # avoiding a second large-payload reference while the request is live. + self._pending_gets: dict[str, _PendingGet] = {} # Internal communication address for proxy and workers self._inproc_addr = f"inproc://simple_storage_workers_{self.storage_unit_id}" @@ -192,6 +235,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: @@ -201,8 +245,6 @@ def _init_zmq_socket(self) -> None: - worker_socket (DEALER): Backend socket for worker communication. """ self.zmq_context = zmq.Context() - self._node_ip = get_node_ip_address() - # Frontend: ROUTER for receiving client requests self.put_get_socket = create_zmq_socket(self.zmq_context, zmq.ROUTER, self._node_ip) @@ -303,9 +345,21 @@ def _worker_routine(self) -> None: 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.PUT_DATA_PREPARE: + response_msg = self._handle_put_prepare(request_msg) + elif operation == ZMQRequestType.PUT_DATA_COMMIT: + response_msg = self._handle_put_commit(request_msg) + elif operation == ZMQRequestType.PUT_DATA_CANCEL: + response_msg = self._handle_put_cancel(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.GET_DATA_PREPARE: + response_msg = self._handle_get_prepare(request_msg) + elif operation == ZMQRequestType.GET_DATA_COMMIT: + response_msg = self._handle_get_commit(request_msg) + elif operation == ZMQRequestType.GET_DATA_CANCEL: + response_msg = self._handle_get_cancel(request_msg) elif operation == ZMQRequestType.CLEAR_DATA: # type: ignore[arg-type] with monitor.measure(op_type="CLEAR_DATA"): response_msg = self._handle_clear(request_msg) @@ -363,49 +417,7 @@ def _handle_put(self, data_parts: ZMQMessage) -> ZMQMessage: 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) + self._put_decoded_data(global_indexes, field_data, data_parser) # After put operation finish, send a message to the client response_msg = ZMQMessage.create( @@ -466,6 +478,235 @@ def _handle_get(self, data_parts: ZMQMessage) -> ZMQMessage: ) return response_msg + def _handle_put_prepare(self, data_parts: ZMQMessage) -> ZMQMessage: + """Prepare a receive before acknowledging a large PUT descriptor.""" + if self.payload_transfer is None: + return self._payload_transfer_error("PUT prepare received while payload transfer is unavailable") + descriptor = None + receive_prepared = False + try: + descriptor = PayloadDescriptor.from_dict(data_parts.body["descriptor"]) + self._expire_pending_payloads() + self._check_pending_payload_capacity(descriptor.payload_bytes) + if descriptor.transfer_id in self._pending_puts: + raise RuntimeError(f"duplicate PUT transfer_id: {descriptor.transfer_id}") + global_indexes = tuple(data_parts.body["global_indexes"]) + token = self.payload_transfer.prepare_receive(descriptor) + receive_prepared = True + self._pending_puts[descriptor.transfer_id] = _PendingPut( + descriptor=descriptor, + sender_id=data_parts.sender_id, + global_indexes=global_indexes, + data_parser=data_parts.body.get("data_parser"), + created_at=time.monotonic(), + ) + return ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_READY, + sender_id=self.storage_unit_id, + body={"descriptor": descriptor.to_dict(), "receive_token": token.to_dict()}, + ) + except Exception as e: + if descriptor is not None and receive_prepared: + self._pending_puts.pop(descriptor.transfer_id, None) + self.payload_transfer.cancel_receive(descriptor.transfer_id) + return self._payload_transfer_error(f"PUT prepare failed: {e}") + + def _handle_put_commit(self, data_parts: ZMQMessage) -> ZMQMessage: + """Consume a completed PUT transfer and commit it through the existing store path.""" + transfer_id = data_parts.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 != data_parts.sender_id: + raise RuntimeError(f"PUT transfer {transfer_id} belongs to another sender") + self._pending_puts.pop(transfer_id) + owns_receive = True + payload = self.payload_transfer.receive(pending.descriptor).result() if self.payload_transfer else None + if payload is None: + raise RuntimeError("payload transfer is unavailable") + frames = unpack_frames(payload) + field_data = decode(frames) + self._put_decoded_data(list(pending.global_indexes), field_data, pending.data_parser) + return ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_RESPONSE, + sender_id=self.storage_unit_id, + body={"transfer_id": transfer_id}, + ) + except Exception as e: + if owns_receive and self.payload_transfer is not None: + self.payload_transfer.cancel_receive(transfer_id) + return self._payload_transfer_error(f"PUT commit failed: {e}") + + def _handle_put_cancel(self, data_parts: ZMQMessage) -> ZMQMessage: + """Release a prepared PUT that will not be committed.""" + transfer_id = data_parts.body["transfer_id"] + pending = self._pending_puts.get(transfer_id) + if pending is not None and pending.sender_id != data_parts.sender_id: + return self._payload_transfer_error(f"PUT transfer {transfer_id} belongs to another sender") + pending = self._pending_puts.pop(transfer_id, None) + if pending is not None and self.payload_transfer is not None: + self.payload_transfer.cancel_receive(transfer_id) + return ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_RESPONSE, + sender_id=self.storage_unit_id, + body={"transfer_id": transfer_id}, + ) + + def get_payload_transfer_pending_counts(self) -> dict[str, int]: + """Return pending transfer counts for lifecycle diagnostics.""" + return { + "pending_puts": len(self._pending_puts), + "pending_gets": len(self._pending_gets), + "pending_receives": self.payload_transfer.pending_receive_count if self.payload_transfer is not None else 0, + } + + def _expire_pending_payloads(self) -> None: + """Discard abandoned protocol state before accepting more payloads.""" + deadline = time.monotonic() - _PENDING_PAYLOAD_TTL_SECONDS + for transfer_id, pending in list(self._pending_puts.items()): + if pending.created_at < deadline: + self._pending_puts.pop(transfer_id) + if self.payload_transfer is not None: + self.payload_transfer.cancel_receive(transfer_id) + for transfer_id, pending in list(self._pending_gets.items()): + if pending.created_at < deadline: + self._pending_gets.pop(transfer_id) + + def _check_pending_payload_capacity(self, payload_bytes: int) -> None: + pending = [*self._pending_puts.values(), *self._pending_gets.values()] + if len(pending) >= _MAX_PENDING_PAYLOAD_TRANSFERS: + raise RuntimeError("too many pending payload transfers") + if sum(item.descriptor.payload_bytes for item in pending) + payload_bytes > _MAX_PENDING_PAYLOAD_BYTES: + raise RuntimeError("pending payload byte limit exceeded") + + def _handle_get_prepare(self, data_parts: ZMQMessage) -> ZMQMessage: + """Encode GET data and retain it until the requester posts its receive.""" + if self.payload_transfer is None: + return self._payload_transfer_error("GET prepare received while payload transfer is unavailable") + try: + transfer_id = str(data_parts.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}") + fields = data_parts.body["fields"] + global_indexes = data_parts.body["global_indexes"] + # StorageUnitData.get_data() gathers existing Python references; it + # does not perform tensor aggregation or mutate PyTorch settings. + result_data = self.storage_data.get_data(fields, global_indexes) + encoded_frames = encode(result_data) + payload = pack_frames(encoded_frames) + if len(payload) < self.inline_threshold_bytes: + return ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_READY, + sender_id=self.storage_unit_id, + body={"route": "zmq_inline", "data": result_data}, + ) + descriptor = PayloadDescriptor( + transfer_id=transfer_id, + payload_bytes=len(payload), + ) + descriptor.validate() + self._expire_pending_payloads() + self._check_pending_payload_capacity(descriptor.payload_bytes) + self._pending_gets[descriptor.transfer_id] = _PendingGet( + descriptor=descriptor, + sender_id=data_parts.sender_id, + payload=payload, + created_at=time.monotonic(), + ) + return ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_READY, + sender_id=self.storage_unit_id, + body={"descriptor": descriptor.to_dict()}, + ) + except Exception as e: + return self._payload_transfer_error(f"GET prepare failed: {e}") + + def _handle_get_commit(self, data_parts: ZMQMessage) -> ZMQMessage: + """Start a GET send after the requester has posted its receive.""" + transfer_id = data_parts.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 != data_parts.sender_id: + raise RuntimeError(f"GET transfer {transfer_id} belongs to another sender") + self._pending_gets.pop(transfer_id) + if self.payload_transfer is None: + raise RuntimeError("payload transfer is unavailable") + + # Do not acknowledge GET before the payload send completes. If + # the send failed after an early response, the requester could + # wait forever on its receive future with no control-plane error. + endpoint = TransferEndpoint.from_dict(data_parts.body["receiver_endpoint"]) + token = ReceiveToken.from_dict(data_parts.body["receive_token"]) + self.payload_transfer.send(endpoint, token, pending.descriptor, pending.payload).result() + return ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_RESPONSE, + sender_id=self.storage_unit_id, + body={"transfer_id": transfer_id}, + ) + except Exception as e: + return self._payload_transfer_error(f"GET commit failed: {e}") + + def _handle_get_cancel(self, data_parts: ZMQMessage) -> ZMQMessage: + """Discard a prepared GET before its send has started.""" + transfer_id = data_parts.body["transfer_id"] + pending = self._pending_gets.get(transfer_id) + if pending is not None and pending.sender_id != data_parts.sender_id: + return self._payload_transfer_error(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=self.storage_unit_id, + body={"transfer_id": transfer_id}, + ) + + 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) + + @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: + return len(value) + except Exception: + return None + + def _payload_transfer_error(self, message: str) -> ZMQMessage: + return ZMQMessage.create( + request_type=ZMQRequestType.PUT_GET_ERROR, + sender_id=self.storage_unit_id, + body={"message": message}, + ) + def _handle_clear(self, data_parts: ZMQMessage) -> ZMQMessage: """ Handle clear request, clear data in storage unit according to given global_indexes. @@ -676,6 +917,7 @@ def _shutdown_resources( proxy_thread: Thread | None, zmq_context: zmq.Context | None, put_get_socket: zmq.Socket | None, + payload_transfer: Any | None, ) -> None: """Clean up resources on garbage collection.""" logger.info("Shutting down SimpleStorageUnit resources...") @@ -697,6 +939,9 @@ def _shutdown_resources( if proxy_thread and proxy_thread.is_alive(): proxy_thread.join(timeout=5) + if payload_transfer is not None: + payload_transfer.close() + logger.info("SimpleStorageUnit resources shutdown complete.") def start_metrics(self, port: int = 0) -> str: @@ -727,3 +972,12 @@ 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.""" + if self.payload_transfer is None: + return None + return { + "id": self.storage_unit_id, + "endpoint": self.payload_transfer.endpoint().to_dict(), + } diff --git a/transfer_queue/utils/serial_utils.py b/transfer_queue/utils/serial_utils.py index e8ab267a..eebf64e3 100644 --- a/transfer_queue/utils/serial_utils.py +++ b/transfer_queue/utils/serial_utils.py @@ -47,6 +47,7 @@ bytestr: TypeAlias = bytes | bytearray | memoryview | zmq.Frame + logger = get_logger(__name__) # Ignore warnings about non-writable buffers from torch.frombuffer. Upper codes will ensure @@ -447,14 +448,11 @@ def pack_into(target_buffer: bytestr, items: Sequence[bytestr]) -> None: entry_offset = _PACK_HEADER_SIZE payload_offset = _PACK_HEADER_SIZE + len(items) * _PACK_ENTRY_SIZE - target_tensor = torch.frombuffer(target_mv, dtype=torch.uint8) - for item in items: item_mv = memoryview(item) nbytes = item_mv.nbytes struct.pack_into(_PACK_ENTRY_FMT, target_mv, entry_offset, payload_offset, nbytes) - src_tensor = torch.frombuffer(item_mv, dtype=torch.uint8) - target_tensor[payload_offset : payload_offset + nbytes].copy_(src_tensor) + target_mv[payload_offset : payload_offset + nbytes] = item_mv entry_offset += _PACK_ENTRY_SIZE payload_offset += nbytes @@ -462,14 +460,38 @@ 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 pack_frames(items: Sequence[bytestr]) -> bytearray: + """Pack codec frames into one contiguous mutable payload. + + Returning the existing bytearray avoids an additional bytearray-to-bytes + copy. A payload transfer retains this buffer until the send completes. + """ + payload = bytearray(calc_packed_size(items)) + pack_into(payload, items) + return payload + + +def unpack_frames(payload: bytes | bytearray | memoryview) -> list[memoryview]: + """Restore the exact codec frame layout from :func:`pack_frames`.""" + return unpack_from(payload) + + 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 49e8e674..a031bd02 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"