diff --git a/tests/test_nixl_payload_lifecycle.py b/tests/test_nixl_payload_lifecycle.py new file mode 100644 index 00000000..07e489d1 --- /dev/null +++ b/tests/test_nixl_payload_lifecycle.py @@ -0,0 +1,687 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team + +"""Focused tests for the frame-native NIXL payload lifecycle.""" + +from __future__ import annotations + +import asyncio +import gc +import socket +import sys +import threading +import time +import types +from concurrent.futures import Future +from queue import SimpleQueue +from typing import Any, Callable + +import numpy as np +import pytest +import torch +import zmq + +from transfer_queue.storage.payload_transfer import DeferredResponse +from transfer_queue.storage.payload_transfer.nixl import ( + NixlPayloadTransfer, + PayloadDescriptor, + ReceiveToken, + TransferEndpoint, + _PendingGet, +) +from transfer_queue.storage.payload_transfer.nixl_ucx_runtime import ( + NixlError, + NixlRuntime, + _frame_regions, + _FrameSource, + _RegisteredReceiveBuffer, +) +from transfer_queue.storage.simple_storage import _drain_deferred_responses, _queue_deferred_response +from transfer_queue.utils.serial_utils import decode, encode +from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType + + +class _FakeAgent: + def __init__(self) -> None: + self.backends = {"UCX": object()} + self.register_calls: list[Any] = [] + self.deregister_calls: list[Any] = [] + self.add_remote_calls: list[bytes] = [] + self.remove_remote_calls: list[str] = [] + self.transfer_calls: list[Any] = [] + self.metadata_names: dict[bytes, str] = {} + self.status_by_remote: dict[str, str] = {} + self.release_error = False + self.deregister_error = False + self._registration_id = 0 + self._handle_id = 0 + + def get_agent_metadata(self) -> bytes: + return f"metadata-{len(self.register_calls)}".encode() + + def register_memory(self, regions: Any, **kwargs: Any) -> tuple[str, int]: + self.register_calls.append(regions) + self._registration_id += 1 + return ("registration", self._registration_id) + + def deregister_memory(self, registration: Any, **kwargs: Any) -> None: + if self.deregister_error: + raise RuntimeError("deregister failed") + self.deregister_calls.append(registration) + + def get_xfer_descs(self, regions: Any, **kwargs: Any) -> Any: + return list(regions) + + def get_serialized_descs(self, descs: Any) -> bytes: + return repr(descs).encode() + + def deserialize_descs(self, descs: bytes) -> bytes: + return descs + + def add_remote_agent(self, metadata: bytes) -> str: + self.add_remote_calls.append(metadata) + return self.metadata_names[metadata] + + def remove_remote_agent(self, remote_name: str) -> None: + self.remove_remote_calls.append(remote_name) + + def initialize_xfer(self, operation: str, local: Any, remote: Any, remote_name: str, **kwargs: Any) -> Any: + self._handle_id += 1 + return (self._handle_id, remote_name, local, remote) + + def transfer(self, handle: Any) -> str: + self.transfer_calls.append(handle) + return self.status_by_remote.get(handle[1], "DONE") + + def check_xfer_state(self, handle: Any) -> str: + return self.status_by_remote.get(handle[1], "DONE") + + def release_xfer_handle(self, handle: Any) -> None: + if self.release_error: + raise RuntimeError("release failed") + + +@pytest.fixture +def runtime(monkeypatch: pytest.MonkeyPatch): + agents: list[_FakeAgent] = [] + module = types.ModuleType("nixl") + + def make_agent(*args: Any, **kwargs: Any) -> _FakeAgent: + agent = _FakeAgent() + agents.append(agent) + return agent + + module.nixl_agent = make_agent + module.nixl_agent_config = lambda **kwargs: kwargs + monkeypatch.setitem(sys.modules, "nixl", module) + instance = NixlRuntime() + yield instance, agents[0] + instance.close() + + +def _descriptor(transfer_id: str, *sizes: int) -> PayloadDescriptor: + return PayloadDescriptor(transfer_id, sum(sizes), tuple(sizes)) + + +def _token(remote_name: str, metadata: bytes, descriptor: PayloadDescriptor) -> dict[str, Any]: + return { + "agent_name": remote_name, + "agent_metadata": metadata, + "frame_remote_descs": b"remote-descs", + "payload_bytes": descriptor.payload_bytes, + } + + +def test_frame_native_layout_preserves_empty_frames_and_large_offsets(): + sizes = (3 * 1024**3, 0, 3 * 1024**3, 2 * 1024**3) + descriptor = _descriptor("large", *sizes) + descriptor.validate() + + address = 0x1000 + regions = _frame_regions(address, descriptor.frame_sizes) + + assert descriptor.payload_bytes == 8 * 1024**3 + assert regions == [ + (address, 3 * 1024**3, 0), + (address + 3 * 1024**3, 3 * 1024**3, 0), + (address + 6 * 1024**3, 2 * 1024**3, 0), + ] + + +def test_all_empty_frames_skip_registration_and_write(runtime): + instance, agent = runtime + descriptor = _descriptor("empty", 0, 0) + + token = instance.prepare_receive(descriptor) + frames = instance.receive(descriptor).result() + sent = instance.send({"agent_name": "peer"}, token, descriptor, (b"", b"")) + + assert [bytes(frame) for frame in frames] == [b"", b""] + assert sent.result() is None + assert agent.register_calls == [] + assert agent.transfer_calls == [] + + +def test_receive_frames_decode_without_copy_and_lease_until_last_tensor(runtime): + instance, _ = runtime + source = {"value": torch.arange(8, dtype=torch.int32)} + encoded = tuple(encode(source)) + descriptor = _descriptor("receive", *(memoryview(frame).nbytes for frame in encoded)) + + instance.prepare_receive(descriptor) + scratch = instance._receives[descriptor.transfer_id] + assert scratch is not None + offset = 0 + for frame in encoded: + view = memoryview(frame) + scratch.buffer[offset : offset + view.nbytes] = view + offset += view.nbytes + + frames = instance.receive(descriptor).result() + decoded = decode(list(frames)) + detached = decoded["value"].detach() + del frames, decoded + gc.collect() + + assert instance._idle_receive_buffers == [] + assert torch.equal(detached, source["value"]) + + del detached + gc.collect() + assert instance._idle_receive_buffers == [scratch] + + +def test_receive_pool_uses_first_usable_buffer_without_deregistering(runtime): + instance, agent = runtime + first = _descriptor("first", 16) + second = _descriptor("second", 8) + + instance.prepare_receive(first) + first_buffer = instance._receives["first"] + instance.cancel_receive("first") + instance.prepare_receive(second) + + assert instance._receives["second"] is first_buffer + assert len(agent.register_calls) == 1 + assert agent.deregister_calls == [] + + +def test_sources_reuse_leased_mr_and_external_registration_is_transfer_scoped(runtime): + instance, agent = runtime + descriptor = _descriptor("lease", 8) + instance.prepare_receive(descriptor) + leased_frame = instance.receive(descriptor).result()[0] + + reused = instance._acquire_frame_source(leased_frame) + writable = bytearray(b"writable") + direct = instance._acquire_frame_source(writable) + readonly = instance._acquire_frame_source(b"readonly") + non_contiguous = instance._acquire_frame_source(memoryview(np.arange(8, dtype=np.uint8))[::2]) + direct_again = instance._acquire_frame_source(writable) + + assert reused.registration is None + assert isinstance(direct.owner, memoryview) + assert isinstance(readonly.owner, bytearray) + assert isinstance(non_contiguous.owner, bytearray) + assert isinstance(direct_again.owner, bytearray) + assert direct_again.address != direct.address + assert direct.registration != direct_again.registration + assert len(agent.register_calls) == 5 # one receive MR plus four transfer-local registrations + + +def test_full_metadata_reused_then_replaced_in_peer_executor(runtime): + instance, agent = runtime + descriptor = _descriptor("send", 1) + agent.metadata_names.update({b"m1": "peer", b"m2": "peer"}) + + instance.send({}, _token("peer", b"m1", descriptor), descriptor, (bytearray(b"a"),)).result() + instance.send({}, _token("peer", b"m1", descriptor), descriptor, (bytearray(b"b"),)).result() + instance.send({}, _token("peer", b"m2", descriptor), descriptor, (bytearray(b"c"),)).result() + + assert agent.add_remote_calls == [b"m1", b"m2"] + assert agent.remove_remote_calls == ["peer"] + + +def test_same_peer_serializes_while_different_peer_runs_concurrently(runtime, monkeypatch: pytest.MonkeyPatch): + instance, _ = runtime + descriptor = _descriptor("concurrency", 1) + first_started = threading.Event() + other_started = threading.Event() + release = threading.Event() + starts: list[str] = [] + + def blocking_send(remote_name: str, *args: Any) -> None: + starts.append(remote_name) + (first_started if remote_name == "peer-a" else other_started).set() + release.wait(timeout=2) + + monkeypatch.setattr(instance, "_send_scatter", blocking_send) + first = instance.send({}, _token("peer-a", b"a", descriptor), descriptor, (bytearray(b"a"),)) + assert first_started.wait(timeout=1) + queued = instance.send({}, _token("peer-a", b"a", descriptor), descriptor, (bytearray(b"b"),)) + other = instance.send({}, _token("peer-b", b"b", descriptor), descriptor, (bytearray(b"c"),)) + + assert other_started.wait(timeout=1) + assert starts.count("peer-a") == 1 + release.set() + first.result(timeout=1) + queued.result(timeout=1) + other.result(timeout=1) + assert starts.count("peer-a") == 2 + + +def test_completed_write_cleanup_failure_keeps_success_and_fails_peer(runtime): + instance, agent = runtime + descriptor = _descriptor("cleanup", 1) + agent.metadata_names[b"metadata"] = "peer" + agent.release_error = True + + completed = instance.send({}, _token("peer", b"metadata", descriptor), descriptor, (bytearray(b"x"),)) + + assert completed.result(timeout=1) is None + assert instance._retained_resources + with pytest.raises(NixlError, match="has failed"): + instance.send({}, _token("peer", b"metadata", descriptor), descriptor, (bytearray(b"y"),)) + + +def test_failed_write_retains_resources_and_other_peer_still_works(runtime): + instance, agent = runtime + descriptor = _descriptor("failure", 1) + agent.metadata_names.update({b"bad": "bad-peer", b"good": "good-peer"}) + agent.status_by_remote["bad-peer"] = "ERR" + + with pytest.raises(NixlError, match="status"): + instance.send({}, _token("bad-peer", b"bad", descriptor), descriptor, (bytearray(b"x"),)).result() + + assert instance._retained_resources + assert instance.send({}, _token("good-peer", b"good", descriptor), descriptor, (bytearray(b"y"),)).result() is None + + +def test_failed_future_does_not_retain_handle_or_sources(runtime): + instance, agent = runtime + descriptor = _descriptor("failure-traceback", 1) + agent.metadata_names[b"metadata"] = "peer" + agent.status_by_remote["peer"] = "ERR" + + future = instance.send({}, _token("peer", b"metadata", descriptor), descriptor, (bytearray(b"x"),)) + error = future.exception(timeout=1) + + assert isinstance(error, NixlError) + traceback = error.__traceback__ + while traceback is not None and traceback.tb_frame.f_code.co_name != "_send_scatter": + traceback = traceback.tb_next + assert traceback is not None + assert traceback.tb_frame.f_locals["handle"] is None + assert traceback.tb_frame.f_locals["sources"] == [] + + +def test_close_interrupts_active_polling(runtime): + instance, agent = runtime + descriptor = _descriptor("closing", 1) + agent.metadata_names[b"metadata"] = "peer" + agent.status_by_remote["peer"] = "PROC" + future = instance.send({}, _token("peer", b"metadata", descriptor), descriptor, (bytearray(b"x"),)) + + deadline = time.monotonic() + 1 + while not agent.transfer_calls and time.monotonic() < deadline: + time.sleep(0.001) + started = time.monotonic() + instance.close() + + assert time.monotonic() - started < 1 + with pytest.raises(NixlError, match="closing"): + future.result() + + +def test_close_tears_down_agent_before_retained_owners(): + events: list[str] = [] + + class TrackingAgent: + def __del__(self) -> None: + events.append("agent") + + class AgentHandle: + def __init__(self, agent: TrackingAgent, release_lease: Callable[[], None]) -> None: + self.agent = agent + self.release_lease = release_lease + + def __del__(self) -> None: + events.append("handle") + self.release_lease() + + class TrackingBuffer(bytearray): + def __init__(self, name: str) -> None: + super().__init__(b"x") + self.name = name + + def __del__(self) -> None: + events.append(self.name) + + instance = object.__new__(NixlRuntime) + instance._lock = threading.RLock() + instance._closing = threading.Event() + instance._closed = False + instance._peer_executors = {} + instance._registered_sources = {} + instance._receives = {} + instance._idle_receive_buffers = [] + instance._leased_receive_buffers = {} + instance._remote_metadata = {} + instance._registered_bytes = 2 + instance._quarantined_bytes = 1 + + agent = TrackingAgent() + source = _FrameSource(TrackingBuffer("source"), None, 0, 1) + receiver = _RegisteredReceiveBuffer(TrackingBuffer("receiver"), object(), 0) + leased_receiver = _RegisteredReceiveBuffer(TrackingBuffer("leased_receiver"), object(), 1) + lease_key = id(leased_receiver) + handle = AgentHandle(agent, lambda: instance._release_lease(lease_key)) + instance._agent = agent + instance._retained_resources = [(handle, [source])] + instance._quarantined_receive_buffers = [receiver] + instance._leased_receive_buffers[lease_key] = leased_receiver + del agent, handle, source, receiver, leased_receiver + + instance.close() + gc.collect() + + assert events.index("handle") < events.index("agent") + assert events.index("agent") < events.index("source") + assert events.index("agent") < events.index("receiver") + assert events.index("agent") < events.index("leased_receiver") + + +def test_quarantine_never_returns_exposed_receiver_to_pool(runtime): + instance, _ = runtime + descriptor = _descriptor("unsafe", 32) + instance.prepare_receive(descriptor) + scratch = instance._receives["unsafe"] + + instance.quarantine_receive("unsafe") + + assert instance._idle_receive_buffers == [] + assert instance._quarantined_receive_buffers == [scratch] + assert instance.diagnostics["quarantined_bytes"] == 32 + + +def test_get_commit_returns_deferred_response_and_maps_completion(): + transfer = object.__new__(NixlPayloadTransfer) + descriptor = _descriptor("get", 1) + transfer._pending_gets = {"get": _PendingGet(descriptor, "manager", (bytearray(b"x"),))} + send_future: Future[None] = Future() + transfer.send = lambda *args, **kwargs: send_future + request = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_COMMIT, + sender_id="manager", + body={ + "transfer_id": "get", + "receiver_endpoint": {"transport": "nixl-ucx", "data": {}}, + "receive_token": {"data": {}}, + }, + ) + + response = transfer._handle_get_commit(request, "storage") + + assert isinstance(response, DeferredResponse) + assert not response.future.done() + send_future.set_result(None) + assert response.future.result().request_type == ZMQRequestType.GET_DATA_RESPONSE + + +def test_get_commit_maps_transfer_failure_to_normal_error_response(): + transfer = object.__new__(NixlPayloadTransfer) + descriptor = _descriptor("failed-get", 1) + transfer._pending_gets = {"failed-get": _PendingGet(descriptor, "manager", (bytearray(b"x"),))} + send_future: Future[None] = Future() + transfer.send = lambda *args, **kwargs: send_future + request = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_COMMIT, + sender_id="manager", + body={ + "transfer_id": "failed-get", + "receiver_endpoint": {"transport": "nixl-ucx", "data": {}}, + "receive_token": {"data": {}}, + }, + ) + + response = transfer._handle_get_commit(request, "storage") + send_future.set_exception(NixlError("write failed")) + + assert isinstance(response, DeferredResponse) + assert response.future.result().request_type == ZMQRequestType.PUT_GET_ERROR + + +class _ProtocolSocket: + def __init__(self, *, fail_commit: bool = False, fail_ready: bool = False) -> None: + self.fail_commit = fail_commit + self.fail_ready = fail_ready + self.last_request: ZMQMessage | None = None + + async def send_multipart(self, frames: Any, **kwargs: Any) -> None: + request = ZMQMessage.deserialize(frames) + if self.fail_commit and request.request_type == ZMQRequestType.GET_DATA_COMMIT: + raise RuntimeError("commit send failed") + self.last_request = request + + async def recv_multipart(self, **kwargs: Any) -> Any: + assert self.last_request is not None + request = self.last_request + if request.request_type == ZMQRequestType.PUT_DATA_PREPARE: + if self.fail_ready: + raise RuntimeError("ready receive failed") + response = ZMQMessage.create( + request_type=ZMQRequestType.PUT_DATA_READY, + sender_id="storage", + body={ + "descriptor": request.body["descriptor"], + "receive_token": {"data": {}}, + }, + ) + elif request.request_type == ZMQRequestType.GET_DATA_PREPARE: + descriptor = _descriptor(str(request.body["transfer_id"]), 1) + response = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_READY, + sender_id="storage", + body={"descriptor": descriptor.to_dict()}, + ) + else: + raise AssertionError(f"unexpected receive after {request.request_type}") + return response.serialize() + + +def test_put_does_not_cancel_after_send_is_attempted(): + transfer = object.__new__(NixlPayloadTransfer) + failed_send: Future[None] = Future() + failed_send.set_exception(NixlError("write failed")) + transfer._peer_endpoint = lambda target_id: TransferEndpoint("nixl-ucx", {}) + transfer.send = lambda *args, **kwargs: failed_send + cancellations: list[Any] = [] + + async def cancel(*args: Any) -> None: + cancellations.append(args) + + transfer._cancel = cancel + with pytest.raises(NixlError, match="write failed"): + asyncio.run( + transfer.put( + control_socket=_ProtocolSocket(), + sender_id="manager", + target_id="storage", + global_indexes=[1], + data={"value": [1]}, + data_parser=None, + ) + ) + + assert cancellations == [] + + +def test_put_cancels_prepared_receiver_before_send_is_attempted(): + transfer = object.__new__(NixlPayloadTransfer) + + def missing_endpoint(target_id: str) -> TransferEndpoint: + raise RuntimeError("endpoint unavailable") + + transfer._peer_endpoint = missing_endpoint + cancellations: list[Any] = [] + + async def cancel(*args: Any) -> None: + cancellations.append(args) + + transfer._cancel = cancel + with pytest.raises(RuntimeError, match="endpoint unavailable"): + asyncio.run( + transfer.put( + control_socket=_ProtocolSocket(), + sender_id="manager", + target_id="storage", + global_indexes=[1], + data={"value": [1]}, + data_parser=None, + ) + ) + + assert cancellations[0][2] == ZMQRequestType.PUT_DATA_CANCEL + + +def test_put_cancels_when_send_fails_before_submission(): + transfer = object.__new__(NixlPayloadTransfer) + transfer._peer_endpoint = lambda target_id: TransferEndpoint("nixl-ucx", {}) + + def failed_send(*args: Any, **kwargs: Any) -> Future[None]: + raise NixlError("metadata invalid") + + transfer.send = failed_send + cancellations: list[Any] = [] + + async def cancel(*args: Any) -> None: + cancellations.append(args) + + transfer._cancel = cancel + with pytest.raises(NixlError, match="metadata invalid"): + asyncio.run( + transfer.put( + control_socket=_ProtocolSocket(), + sender_id="manager", + target_id="storage", + global_indexes=[1], + data={"value": [1]}, + data_parser=None, + ) + ) + + assert cancellations[0][2] == ZMQRequestType.PUT_DATA_CANCEL + + +def test_put_attempts_cancel_when_ready_is_lost(): + transfer = object.__new__(NixlPayloadTransfer) + cancellations: list[Any] = [] + + async def cancel(*args: Any) -> None: + cancellations.append(args) + + transfer._cancel = cancel + with pytest.raises(RuntimeError, match="ready receive failed"): + asyncio.run( + transfer.put( + control_socket=_ProtocolSocket(fail_ready=True), + sender_id="manager", + target_id="storage", + global_indexes=[1], + data={"value": [1]}, + data_parser=None, + ) + ) + + assert cancellations[0][2] == ZMQRequestType.PUT_DATA_CANCEL + + +def test_get_cancels_remote_source_when_local_receive_prepare_fails(): + transfer = object.__new__(NixlPayloadTransfer) + + def failed_prepare(descriptor: PayloadDescriptor) -> ReceiveToken: + raise NixlError("registration failed") + + transfer.prepare_receive = failed_prepare + cancellations: list[Any] = [] + + async def cancel(*args: Any) -> None: + cancellations.append(args) + + transfer._cancel = cancel + with pytest.raises(NixlError, match="registration failed"): + asyncio.run( + transfer.get( + control_socket=_ProtocolSocket(), + sender_id="manager", + target_id="storage", + global_indexes=[1], + fields=["value"], + ) + ) + + assert cancellations[0][2] == ZMQRequestType.GET_DATA_CANCEL + + +def test_get_commit_attempt_failure_quarantines_without_cancel(): + transfer = object.__new__(NixlPayloadTransfer) + transfer.prepare_receive = lambda descriptor: ReceiveToken({"agent_name": "manager"}) + quarantined: list[str] = [] + cancellations: list[Any] = [] + transfer.quarantine_receive = quarantined.append + transfer.cancel_receive = lambda transfer_id: pytest.fail("exposed receiver was cancelled") + + async def cancel(*args: Any) -> None: + cancellations.append(args) + + transfer._cancel = cancel + with pytest.raises(RuntimeError, match="commit send failed"): + asyncio.run( + transfer.get( + control_socket=_ProtocolSocket(fail_commit=True), + sender_id="manager", + target_id="storage", + global_indexes=[1], + fields=["value"], + ) + ) + + assert len(quarantined) == 1 + assert cancellations == [] + + +class _CapturingSocket: + def __init__(self) -> None: + self.messages: list[Any] = [] + + def send_multipart(self, message: Any, **kwargs: Any) -> None: + self.messages.append(message) + + +def test_deferred_response_copies_identity_and_worker_sends_once(): + future: Future[ZMQMessage] = Future() + response = DeferredResponse(future) + completions: SimpleQueue[tuple[bytes, Future[ZMQMessage]]] = SimpleQueue() + reader, writer = socket.socketpair() + reader.settimeout(1) + poller = zmq.Poller() + poller.register(reader.fileno(), zmq.POLLIN) + shutdown = threading.Event() + identity = bytearray(b"client-a") + worker = _CapturingSocket() + try: + _queue_deferred_response(identity, response, completions, writer, shutdown) + identity[:] = b"client-b" + future.set_result( + ZMQMessage.create(request_type=ZMQRequestType.GET_DATA_RESPONSE, sender_id="storage", body={}) + ) + assert dict(poller.poll(1000))[reader.fileno()] == zmq.POLLIN + assert reader.recv(1) == b"\0" + _drain_deferred_responses(completions, worker, "storage") + _drain_deferred_responses(completions, worker, "storage") + finally: + reader.close() + writer.close() + + assert len(worker.messages) == 1 + assert worker.messages[0][0] == b"client-a" diff --git a/tests/test_payload_transfer.py b/tests/test_payload_transfer.py index 8495b56f..83a5cab2 100644 --- a/tests/test_payload_transfer.py +++ b/tests/test_payload_transfer.py @@ -35,12 +35,12 @@ def test_payload_descriptor_preserves_frame_layout(): - descriptor = PayloadDescriptor("framed", 4 + 8 * 2 + 5, (2, 3)) + descriptor = PayloadDescriptor("framed", 5, (2, 3)) descriptor.validate() assert PayloadDescriptor.from_dict(descriptor.to_dict()) == descriptor - with pytest.raises(PayloadTransferError, match="packed payload length"): - PayloadDescriptor("framed", 5, (2, 3)).validate() + with pytest.raises(PayloadTransferError, match="payload length mismatch"): + PayloadDescriptor("framed", 6, (2, 3)).validate() def test_payload_descriptor_requires_frame_layout_and_rejects_negative_lengths(): diff --git a/transfer_queue/storage/payload_transfer/__init__.py b/transfer_queue/storage/payload_transfer/__init__.py index 8df2ac9c..28af2773 100644 --- a/transfer_queue/storage/payload_transfer/__init__.py +++ b/transfer_queue/storage/payload_transfer/__init__.py @@ -16,6 +16,7 @@ """Payload transfer strategies for SimpleStorage.""" from transfer_queue.storage.payload_transfer.base import ( + DeferredResponse, PayloadTransfer, PayloadTransferError, ) @@ -25,6 +26,7 @@ ) __all__ = [ + "DeferredResponse", "PayloadTransfer", "PayloadTransferError", "create_payload_transfer", diff --git a/transfer_queue/storage/payload_transfer/base.py b/transfer_queue/storage/payload_transfer/base.py index fa2003e0..fd67e830 100644 --- a/transfer_queue/storage/payload_transfer/base.py +++ b/transfer_queue/storage/payload_transfer/base.py @@ -18,6 +18,8 @@ from __future__ import annotations from abc import ABC, abstractmethod +from concurrent.futures import Future +from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Callable if TYPE_CHECKING: @@ -28,6 +30,13 @@ class PayloadTransferError(RuntimeError): """A payload transfer could not be completed safely.""" +@dataclass(frozen=True) +class DeferredResponse: + """A response that the SimpleStorage worker must send after completion.""" + + future: Future[ZMQMessage] + + class PayloadTransfer(ABC): """Complete SimpleStorage payload strategy, including its wire protocol.""" @@ -64,7 +73,7 @@ def handle_request( storage_id: str, load_data: Callable[..., dict[str, Any]], store_data: Callable[..., None], - ) -> ZMQMessage | None: + ) -> ZMQMessage | DeferredResponse | None: """Handle a strategy-owned request on a SimpleStorageUnit.""" def bootstrap_info(self) -> dict[str, Any] | None: diff --git a/transfer_queue/storage/payload_transfer/nixl.py b/transfer_queue/storage/payload_transfer/nixl.py index 4350905a..f1243499 100644 --- a/transfer_queue/storage/payload_transfer/nixl.py +++ b/transfer_queue/storage/payload_transfer/nixl.py @@ -27,11 +27,11 @@ import zmq.asyncio -from transfer_queue.storage.payload_transfer.base import PayloadTransfer, PayloadTransferError +from transfer_queue.storage.payload_transfer.base import DeferredResponse, PayloadTransfer, PayloadTransferError from transfer_queue.storage.payload_transfer.nixl_ucx_runtime import NixlError, NixlRuntime from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads from transfer_queue.utils.logging_utils import get_logger -from transfer_queue.utils.serial_utils import calc_packed_size, decode, encode, unpack_from +from transfer_queue.utils.serial_utils import decode, encode from transfer_queue.utils.zmq_utils import ( ZMQMessage, ZMQRequestType, @@ -64,11 +64,10 @@ def validate(self) -> None: raise PayloadTransferError(f"negative payload length for {self.transfer_id}") if any(size < 0 for size in self.frame_sizes): raise PayloadTransferError(f"negative frame length for {self.transfer_id}") - packed_size = 4 + 8 * len(self.frame_sizes) + sum(self.frame_sizes) - if packed_size != self.payload_bytes: + frame_bytes = sum(self.frame_sizes) + if frame_bytes != self.payload_bytes: raise PayloadTransferError( - f"packed payload length mismatch for {self.transfer_id}: " - f"expected {packed_size}, got {self.payload_bytes}" + f"payload length mismatch for {self.transfer_id}: expected {frame_bytes}, got {self.payload_bytes}" ) def to_dict(self) -> dict[str, int | str | list[int]]: @@ -182,11 +181,12 @@ async def put( frames = tuple(encode(data)) descriptor = PayloadDescriptor( transfer_id=uuid4().hex, - payload_bytes=calc_packed_size(frames), + payload_bytes=sum(memoryview(frame).nbytes for frame in frames), frame_sizes=tuple(memoryview(frame).nbytes for frame in frames), ) descriptor.validate() - remote_may_be_prepared = True + remote_may_be_prepared = False + send_attempted = False try: prepare = ZMQMessage.create( request_type=ZMQRequestType.PUT_DATA_PREPARE, @@ -199,13 +199,16 @@ async def put( }, ) await control_socket.send_multipart(prepare.serialize(), copy=False) + remote_may_be_prepared = True ready = ZMQMessage.deserialize(await control_socket.recv_multipart(copy=False)) self._expect(ready, ZMQRequestType.PUT_DATA_READY, target_id) if PayloadDescriptor.from_dict(ready.body["descriptor"]) != descriptor: raise RuntimeError(f"PUT descriptor changed by storage unit {target_id}") token = ReceiveToken.from_dict(ready.body["receive_token"]) endpoint = self._peer_endpoint(target_id) - await asyncio.wrap_future(self.send(endpoint, token, descriptor, frames)) + send_future = self.send(endpoint, token, descriptor, frames) + send_attempted = True + await asyncio.wrap_future(send_future) commit = ZMQMessage.create( request_type=ZMQRequestType.PUT_DATA_COMMIT, @@ -216,9 +219,8 @@ async def put( await control_socket.send_multipart(commit.serialize(), copy=False) response = ZMQMessage.deserialize(await control_socket.recv_multipart(copy=False)) self._expect(response, ZMQRequestType.PUT_DATA_RESPONSE, target_id) - remote_may_be_prepared = False except BaseException: - if remote_may_be_prepared: + if remote_may_be_prepared and not send_attempted: await self._cancel(sender_id, target_id, ZMQRequestType.PUT_DATA_CANCEL, descriptor.transfer_id) raise @@ -234,6 +236,7 @@ async def get( transfer_id = uuid4().hex remote_prepared = False receive_prepared = False + commit_attempted = False descriptor = None try: prepare = ZMQMessage.create( @@ -257,20 +260,26 @@ async def get( receiver_id=target_id, body={ "transfer_id": descriptor.transfer_id, - "receiver_endpoint": self.endpoint().to_dict(), + "receiver_endpoint": TransferEndpoint( + self.transport, {"agent_name": token.data["agent_name"]} + ).to_dict(), "receive_token": token.to_dict(), }, ) + commit_attempted = True await control_socket.send_multipart(commit.serialize(), copy=False) response = ZMQMessage.deserialize(await control_socket.recv_multipart(copy=False)) self._expect(response, ZMQRequestType.GET_DATA_RESPONSE, target_id) remote_prepared = False - payload = await asyncio.wrap_future(self.receive(descriptor)) - return decode(unpack_from(payload)) + frames = await asyncio.wrap_future(self.receive(descriptor)) + return decode(list(frames)) except BaseException: if receive_prepared and descriptor is not None: - self.cancel_receive(descriptor.transfer_id) - if remote_prepared: + if commit_attempted: + self.quarantine_receive(descriptor.transfer_id) + else: + self.cancel_receive(descriptor.transfer_id) + if remote_prepared and not commit_attempted: await self._cancel(sender_id, target_id, ZMQRequestType.GET_DATA_CANCEL, transfer_id) raise @@ -281,7 +290,7 @@ def handle_request( storage_id: str, load_data: Callable[..., dict[str, Any]], store_data: Callable[..., None], - ) -> ZMQMessage | None: + ) -> ZMQMessage | DeferredResponse | None: if request.request_type == ZMQRequestType.PUT_DATA_PREPARE: return self._handle_put_prepare(request, storage_id) if request.request_type == ZMQRequestType.PUT_DATA_COMMIT: @@ -324,7 +333,6 @@ def _handle_put_prepare(self, request: ZMQMessage, storage_id: str) -> ZMQMessag def _handle_put_commit(self, request: ZMQMessage, storage_id: str, store_data: Callable[..., None]) -> ZMQMessage: transfer_id = request.body["transfer_id"] - owns_receive = False try: pending = self._pending_puts.get(transfer_id) if pending is None: @@ -332,20 +340,17 @@ def _handle_put_commit(self, request: ZMQMessage, storage_id: str, store_data: C if pending.sender_id != request.sender_id: raise RuntimeError(f"PUT transfer {transfer_id} belongs to another sender") self._pending_puts.pop(transfer_id) - owns_receive = True - payload = self.receive(pending.descriptor).result() + frames = self.receive(pending.descriptor).result() with limit_pytorch_auto_parallel_threads( target_num_threads=TQ_NUM_THREADS, info=f"[{storage_id}] PUT commit" ): - store_data(list(pending.global_indexes), decode(unpack_from(payload)), pending.data_parser) + store_data(list(pending.global_indexes), decode(list(frames)), pending.data_parser) return ZMQMessage.create( request_type=ZMQRequestType.PUT_DATA_RESPONSE, sender_id=storage_id, body={"transfer_id": transfer_id}, ) except Exception as exc: - if owns_receive: - self.cancel_receive(transfer_id) return self._error(storage_id, "PUT commit", exc) def _handle_put_cancel(self, request: ZMQMessage, storage_id: str) -> ZMQMessage: @@ -382,7 +387,7 @@ def _handle_get_prepare( frames = tuple(encode(load_data(request.body["fields"], request.body["global_indexes"]))) descriptor = PayloadDescriptor( transfer_id=transfer_id, - payload_bytes=calc_packed_size(frames), + payload_bytes=sum(memoryview(frame).nbytes for frame in frames), frame_sizes=tuple(memoryview(frame).nbytes for frame in frames), ) descriptor.validate() @@ -395,7 +400,7 @@ def _handle_get_prepare( except Exception as exc: return self._error(storage_id, "GET prepare", exc) - def _handle_get_commit(self, request: ZMQMessage, storage_id: str) -> ZMQMessage: + def _handle_get_commit(self, request: ZMQMessage, storage_id: str) -> DeferredResponse | ZMQMessage: transfer_id = request.body["transfer_id"] try: pending = self._pending_gets.get(transfer_id) @@ -406,12 +411,23 @@ def _handle_get_commit(self, request: ZMQMessage, storage_id: str) -> ZMQMessage self._pending_gets.pop(transfer_id) endpoint = TransferEndpoint.from_dict(request.body["receiver_endpoint"]) token = ReceiveToken.from_dict(request.body["receive_token"]) - self.send(endpoint, token, pending.descriptor, pending.frames).result() - return ZMQMessage.create( - request_type=ZMQRequestType.GET_DATA_RESPONSE, - sender_id=storage_id, - body={"transfer_id": transfer_id}, - ) + send_future = self.send(endpoint, token, pending.descriptor, pending.frames) + response_future: Future[ZMQMessage] = Future() + + def complete_response(completed: Future[None]) -> None: + try: + completed.result() + response = ZMQMessage.create( + request_type=ZMQRequestType.GET_DATA_RESPONSE, + sender_id=storage_id, + body={"transfer_id": transfer_id}, + ) + except Exception as exc: + response = self._error(storage_id, "GET commit", exc) + response_future.set_result(response) + + send_future.add_done_callback(complete_response) + return DeferredResponse(response_future) except Exception as exc: return self._error(storage_id, "GET commit", exc) @@ -444,13 +460,18 @@ def send( self._validate_metadata(endpoint, token, descriptor) return self._runtime.send(endpoint.data, token.data, descriptor, tuple(frames)) - def receive(self, descriptor: PayloadDescriptor) -> Future[memoryview]: + def receive(self, descriptor: PayloadDescriptor) -> Future[tuple[memoryview, ...]]: return self._runtime.receive(descriptor) def cancel_receive(self, transfer_id: str) -> None: self._runtime.cancel_receive(transfer_id) + def quarantine_receive(self, transfer_id: str) -> None: + self._runtime.quarantine_receive(transfer_id) + def close(self) -> None: + self._pending_puts.clear() + self._pending_gets.clear() self._runtime.close() def _peer_endpoint(self, target_id: str) -> TransferEndpoint: diff --git a/transfer_queue/storage/payload_transfer/nixl_ucx_runtime.py b/transfer_queue/storage/payload_transfer/nixl_ucx_runtime.py index 3bf19482..8181b89e 100644 --- a/transfer_queue/storage/payload_transfer/nixl_ucx_runtime.py +++ b/transfer_queue/storage/payload_transfer/nixl_ucx_runtime.py @@ -13,7 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""Small NIXL runtime used by the SimpleStorage H2H payload adapter.""" +"""NIXL runtime for the SimpleStorage host payload fast path.""" from __future__ import annotations @@ -22,18 +22,21 @@ import socket import threading import time +import weakref from concurrent.futures import Future, ThreadPoolExecutor from dataclasses import dataclass from typing import Any from uuid import uuid4 +import numpy as np + from transfer_queue.storage.payload_transfer.base import PayloadTransferError from transfer_queue.utils.logging_utils import get_logger -from transfer_queue.utils.serial_utils import initialize_packed_frame_table logger = get_logger(__name__) DEFAULT_NIXL_TRANSFER_TIMEOUT_SECONDS = 180 +_POLL_INTERVAL_SECONDS = 0.0005 class NixlError(PayloadTransferError): @@ -43,7 +46,9 @@ class NixlError(PayloadTransferError): @dataclass class _RegisteredReceiveBuffer: buffer: bytearray - registrations: Any + registration: Any + address: int + lease_finalizer: weakref.finalize | None = None @property def capacity(self) -> int: @@ -51,26 +56,43 @@ def capacity(self) -> int: @dataclass -class _ReceiveState: - scratch: _RegisteredReceiveBuffer - serialized_frame_descs: bytes - - -@dataclass -class _FrameRegistration: - """One registered source frame kept alive for an active transfer.""" - +class _FrameSource: owner: bytearray | memoryview - registrations: Any + registration: Any | None address: int size: int def _buffer_address(buffer: bytearray | memoryview) -> int: - """Return the address of a writable, contiguous host buffer.""" - if not buffer: - raise NixlError("NIXL does not support an empty payload buffer") - return ctypes.addressof(ctypes.c_ubyte.from_buffer(buffer)) + """Return the address of a non-empty writable contiguous host buffer.""" + view = memoryview(buffer) + if not view.nbytes: + raise NixlError("NIXL does not support an empty memory region") + if view.readonly or not view.c_contiguous: + raise NixlError("NIXL memory regions must be writable and C-contiguous") + return ctypes.addressof(ctypes.c_ubyte.from_buffer(view)) + + +def _frame_regions(address: int, frame_sizes: tuple[int, ...]) -> list[tuple[int, int, int]]: + """Build non-empty frame regions without narrowing cumulative offsets.""" + regions = [] + offset = 0 + for size in frame_sizes: + if size: + regions.append((address + offset, size, 0)) + offset += size + return regions + + +def _frame_views(owner: Any, frame_sizes: tuple[int, ...]) -> tuple[memoryview, ...]: + """Split a contiguous owner into frame views, preserving empty frames.""" + view = memoryview(owner).cast("B") + frames = [] + offset = 0 + for size in frame_sizes: + frames.append(view[offset : offset + size]) + offset += size + return tuple(frames) def _configure_ucx_environment(ucx_env_vars: dict[str, object] | None) -> dict[str, str]: @@ -94,12 +116,7 @@ def _warn_if_tcp_fallback_possible() -> None: class NixlRuntime: - """Own one NIXL agent and serialize metadata updates safely. - - SimpleStorage's control plane already orders prepare/send/commit. The - runtime therefore only keeps registered receive buffers alive and waits - for the sender-side NIXL request to finish. - """ + """Own the NIXL agent, registered memory, and transfer completion.""" def __init__(self, ucx_env_vars: dict[str, object] | None = None): _configure_ucx_environment(ucx_env_vars) @@ -131,14 +148,24 @@ def __init__(self, ucx_env_vars: dict[str, object] | None = None): os.environ.get("UCX_TLS", "ucx-auto"), ) - self._receives: dict[str, _ReceiveState] = {} - self._reusable_receive_buffer: _RegisteredReceiveBuffer | None = None + self._receives: dict[str, _RegisteredReceiveBuffer | None] = {} + self._idle_receive_buffers: list[_RegisteredReceiveBuffer] = [] + self._leased_receive_buffers: dict[int, _RegisteredReceiveBuffer] = {} + self._quarantined_receive_buffers: list[_RegisteredReceiveBuffer] = [] + self._retained_resources: list[tuple[Any | None, list[_FrameSource]]] = [] + self._registered_sources: dict[int, _FrameSource] = {} self._remote_metadata: dict[str, bytes] = {} - self._deferred_sends: list[tuple[Any, list[_FrameRegistration]]] = [] + self._failed_peers: set[str] = set() + self._peer_executors: dict[str, ThreadPoolExecutor] = {} self._lock = threading.RLock() + self._closing = threading.Event() self._closed = False self._timeout_seconds = DEFAULT_NIXL_TRANSFER_TIMEOUT_SECONDS - self._send_executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="tq-nixl-send") + self._registered_bytes = 0 + self._quarantined_bytes = 0 + self._registration_seconds = 0.0 + self._data_transfer_seconds = 0.0 + self._total_seconds = 0.0 @staticmethod def _make_agent_name() -> str: @@ -146,48 +173,55 @@ def _make_agent_name() -> str: @property def agent_name(self) -> str: - """Return the stable name advertised by this runtime instance.""" return self._agent_name + @property + def diagnostics(self) -> dict[str, float | int]: + """Return the small runtime-only metric set from the design.""" + with self._lock: + return { + "registration_seconds": self._registration_seconds, + "data_transfer_seconds": self._data_transfer_seconds, + "total_seconds": self._total_seconds, + "registered_bytes": self._registered_bytes, + "quarantined_bytes": self._quarantined_bytes, + } + def endpoint_metadata(self) -> bytes: - """Return serialized metadata that peers need to address this agent.""" with self._lock: self._ensure_open() return self._agent.get_agent_metadata() def prepare_receive(self, descriptor: Any) -> dict[str, Any]: - """Allocate or reuse registered storage for a scatter receive.""" + """Prepare frame-native remote descriptors and publish full metadata.""" descriptor.validate() - if not descriptor.frame_sizes or not any(descriptor.frame_sizes): - raise NixlError("NIXL direct-frame receive requires a non-empty frame") with self._lock: self._ensure_open() if descriptor.transfer_id in self._receives: raise NixlError(f"duplicate NIXL receive: {descriptor.transfer_id}") - scratch = self._acquire_receive_buffer(descriptor.payload_bytes) + + scratch = None try: - initialize_packed_frame_table(scratch.buffer, descriptor.frame_sizes) - address = _buffer_address(scratch.buffer) - payload_offset = 4 + 8 * len(descriptor.frame_sizes) - regions = [] - for size in descriptor.frame_sizes: - if size: - regions.append((address + payload_offset, size, 0)) - payload_offset += size - frame_descs = self._agent.get_serialized_descs(self._agent.get_xfer_descs(regions, mem_type="DRAM")) - state = _ReceiveState(scratch, frame_descs) + if descriptor.payload_bytes: + scratch = self._acquire_receive_buffer(descriptor.payload_bytes) + regions = _frame_regions(scratch.address, descriptor.frame_sizes) + remote_descs = self._agent.get_xfer_descs(regions, mem_type="DRAM") + serialized = self._agent.get_serialized_descs(remote_descs) + else: + serialized = b"" metadata = self._agent.get_agent_metadata() - self._receives[descriptor.transfer_id] = state + self._receives[descriptor.transfer_id] = scratch except Exception as exc: - self._return_receive_buffer(scratch) - raise NixlError(f"failed to register NIXL receive buffer: {exc}") from exc - result = { + if scratch is not None: + self._idle_receive_buffers.append(scratch) + raise NixlError(f"failed to prepare NIXL receive buffer: {exc}") from exc + + return { "agent_name": self._agent_name, "agent_metadata": metadata, - "frame_remote_descs": state.serialized_frame_descs, + "frame_remote_descs": serialized, "payload_bytes": descriptor.payload_bytes, } - return result def send( self, @@ -196,21 +230,24 @@ def send( descriptor: Any, frames: tuple[bytes | bytearray | memoryview, ...], ) -> Future[None]: - """Run the NIXL transfer on the dedicated send thread.""" + """Submit one transfer to the remote peer's single-worker executor.""" descriptor.validate() - if not descriptor.frame_sizes or not any(descriptor.frame_sizes): - raise NixlError("NIXL direct-frame send requires a non-empty frame") if tuple(memoryview(frame).nbytes for frame in frames) != descriptor.frame_sizes: raise NixlError(f"frame lengths do not match descriptor for {descriptor.transfer_id}") - try: - remote_name, metadata = self._validate_send_metadata(endpoint, token, descriptor) - except Exception as exc: - future: Future[None] = Future() - future.set_exception(exc) - return future + remote_name, metadata = self._validate_send_metadata(endpoint, token, descriptor) with self._lock: self._ensure_open() - return self._send_executor.submit( + if remote_name in self._failed_peers: + raise NixlError(f"NIXL peer session {remote_name!r} has failed") + if not descriptor.payload_bytes: + future: Future[None] = Future() + future.set_result(None) + return future + executor = self._peer_executors.get(remote_name) + if executor is None: + executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix=f"tq-nixl-{remote_name}") + self._peer_executors[remote_name] = executor + return executor.submit( self._send_scatter, remote_name, metadata, @@ -218,68 +255,147 @@ def send( token["frame_remote_descs"], ) + def receive(self, descriptor: Any) -> Future[tuple[memoryview, ...]]: + """Return zero-copy frame views and lease the receive MR to their owner.""" + descriptor.validate() + future: Future[tuple[memoryview, ...]] = Future() + with self._lock: + try: + self._ensure_open() + scratch = self._receives.pop(descriptor.transfer_id) + if scratch is None: + future.set_result(tuple(memoryview(b"") for _ in descriptor.frame_sizes)) + return future + + owner = np.frombuffer(scratch.buffer, dtype=np.uint8, count=descriptor.payload_bytes) + key = id(scratch) + self._leased_receive_buffers[key] = scratch + scratch.lease_finalizer = weakref.finalize(owner, self._release_lease, key) + future.set_result(_frame_views(owner, descriptor.frame_sizes)) + except KeyError: + future.set_exception(NixlError(f"no prepared NIXL receive for {descriptor.transfer_id}")) + except Exception as exc: + future.set_exception(NixlError(f"failed to expose NIXL receive frames: {exc}")) + return future + + def cancel_receive(self, transfer_id: str) -> None: + """Release a receive that was never exposed to a possible WRITE.""" + with self._lock: + scratch = self._receives.pop(transfer_id, None) + if scratch is not None: + self._idle_receive_buffers.append(scratch) + + def quarantine_receive(self, transfer_id: str) -> None: + """Keep an exposed receiver out of the reuse pool until agent teardown.""" + with self._lock: + scratch = self._receives.pop(transfer_id, None) + if scratch is not None: + self._quarantined_receive_buffers.append(scratch) + self._quarantined_bytes += scratch.capacity + + def close(self) -> None: + """Stop submissions, stop polling, then tear down the agent before owners.""" + with self._lock: + if self._closed: + return + self._closed = True + self._closing.set() + executors = list(self._peer_executors.values()) + self._peer_executors.clear() + + for executor in executors: + executor.shutdown(wait=True, cancel_futures=True) + + with self._lock: + agent = self._agent + self._agent = None + retained = self._retained_resources + self._retained_resources = [] + retained_handles = [handle for handle, _ in retained if handle is not None] + retained_sources = [source for _, sources in retained for source in sources] + retained.clear() + + # NIXL handles keep their agent alive. Drop them while every registered + # owner is still retained, then tear down the agent before those owners. + retained_handles.clear() + del agent + + with self._lock: + self._receives.clear() + self._idle_receive_buffers.clear() + self._leased_receive_buffers.clear() + self._quarantined_receive_buffers.clear() + self._registered_sources.clear() + self._remote_metadata.clear() + self._registered_bytes = 0 + self._quarantined_bytes = 0 + retained_sources.clear() + def _acquire_receive_buffer(self, required_capacity: int) -> _RegisteredReceiveBuffer: - reusable = self._reusable_receive_buffer - if reusable is not None and reusable.capacity >= required_capacity: - self._reusable_receive_buffer = None - return reusable + for index, scratch in enumerate(self._idle_receive_buffers): + if scratch.capacity >= required_capacity: + return self._idle_receive_buffers.pop(index) buffer = bytearray(required_capacity) address = _buffer_address(buffer) - registrations = self._agent.register_memory( + started = time.monotonic() + registration = self._agent.register_memory( [(address, required_capacity, 0, "")], mem_type="DRAM", backends=["UCX"] ) - if registrations is None: - raise NixlError("failed to register NIXL receive scratch buffer") - allocated = _RegisteredReceiveBuffer(buffer, registrations) - if reusable is not None: - self._reusable_receive_buffer = None - self._deregister_registration(reusable.registrations) - return allocated - - def _return_receive_buffer(self, scratch: _RegisteredReceiveBuffer) -> None: - reusable = self._reusable_receive_buffer - if reusable is None: - self._reusable_receive_buffer = scratch - elif scratch.capacity > reusable.capacity: - self._deregister_registration(reusable.registrations) - self._reusable_receive_buffer = scratch - else: - self._deregister_registration(scratch.registrations) - - def _acquire_frame_registration(self, frame: bytes | bytearray | memoryview) -> _FrameRegistration: + self._registration_seconds += time.monotonic() - started + if registration is None: + raise NixlError("failed to register NIXL receive buffer") + self._registered_bytes += required_capacity + return _RegisteredReceiveBuffer(buffer, registration, address) + + def _register_frame_source(self, owner: bytearray | memoryview) -> _FrameSource: + """Register one writable contiguous source while the runtime lock is held.""" + address = _buffer_address(owner) + size = memoryview(owner).nbytes + started = time.monotonic() + registration = self._agent.register_memory([(address, size, 0, "")], mem_type="DRAM", backends=["UCX"]) + self._registration_seconds += time.monotonic() - started + if registration is None: + raise NixlError("failed to register NIXL source frame") + self._registered_bytes += size + source = _FrameSource(owner, registration, address, size) + self._registered_sources[address] = source + return source + + def _acquire_frame_source(self, frame: bytes | bytearray | memoryview) -> _FrameSource: view = memoryview(frame) + + # Staging can copy large payloads. Keep that work outside the runtime + # lock so an unrelated peer is not blocked on Python memory copies. if view.readonly or not view.c_contiguous: owner: bytearray | memoryview = bytearray(view) - else: - owner = view.cast("B") + with self._lock: + self._ensure_open() + return self._register_frame_source(owner) + + owner = view.cast("B") address = _buffer_address(owner) - registrations = self._agent.register_memory( - [(address, memoryview(owner).nbytes, 0, "")], mem_type="DRAM", backends=["UCX"] - ) - if registrations is None: - raise NixlError("failed to register NIXL source frame") - return _FrameRegistration(owner, registrations, address, memoryview(owner).nbytes) + size = memoryview(owner).nbytes + with self._lock: + self._ensure_open() + for scratch in self._leased_receive_buffers.values(): + if scratch.address <= address and address + size <= scratch.address + scratch.capacity: + return _FrameSource(owner, None, address, size) + + # NIXL resolves registrations by address. An overlapping external + # source needs an independent staging owner before it can be + # registered and cleaned up by this transfer. + overlaps = any( + address < source.address + source.size and source.address < address + size + for source in self._registered_sources.values() + ) + if not overlaps: + return self._register_frame_source(owner) - def _release_send_resources(self, handle: Any, registrations: list[_FrameRegistration]) -> bool: - """Release a send only after NIXL confirms its handle can be released.""" - try: - self._agent.release_xfer_handle(handle) - except Exception as exc: - logger.warning("failed to release NIXL transfer handle; keeping source memory registered: %s", exc) - return False - for registration in registrations: - self._deregister_registration(registration.registrations) - return True - - def _reap_deferred_sends(self) -> None: - if not self._deferred_sends: - return - pending = [] - for handle, registrations in self._deferred_sends: - if not self._release_send_resources(handle, registrations): - pending.append((handle, registrations)) - self._deferred_sends = pending + owner = bytearray(owner) + with self._lock: + self._ensure_open() + return self._register_frame_source(owner) def _send_scatter( self, @@ -288,126 +404,136 @@ def _send_scatter( frames: tuple[bytes | bytearray | memoryview, ...], serialized_remote_descs: bytes, ) -> None: - """Send frames directly to matching registered remote regions.""" - frame_registrations: list[_FrameRegistration] = [] + started_total = time.monotonic() + sources: list[_FrameSource] = [] handle = None + error_message = None try: with self._lock: self._ensure_open() - self._reap_deferred_sends() - previous = self._remote_metadata.get(remote_name) - if previous != metadata: - if previous is not None: - self._agent.remove_remote_agent(remote_name) - loaded_name = self._agent.add_remote_agent(metadata) - if isinstance(loaded_name, bytes): - loaded_name = loaded_name.decode() - if loaded_name != remote_name: - raise NixlError( - f"NIXL remote agent name mismatch: expected {remote_name!r}, got {loaded_name!r}" - ) - self._remote_metadata[remote_name] = metadata - - for frame in frames: - if memoryview(frame).nbytes: - frame_registrations.append(self._acquire_frame_registration(frame)) + if remote_name in self._failed_peers: + raise NixlError(f"NIXL peer session {remote_name!r} has failed") + + for frame in frames: + if memoryview(frame).nbytes: + sources.append(self._acquire_frame_source(frame)) + + with self._lock: + self._ensure_open() + if remote_name in self._failed_peers: + raise NixlError(f"NIXL peer session {remote_name!r} has failed") + self._load_remote_metadata(remote_name, metadata) local_descs = self._agent.get_xfer_descs( - [(registration.address, registration.size, 0) for registration in frame_registrations], - mem_type="DRAM", + [(source.address, source.size, 0) for source in sources], mem_type="DRAM" ) remote_descs = self._agent.deserialize_descs(serialized_remote_descs) handle = self._agent.initialize_xfer("WRITE", local_descs, remote_descs, remote_name, backends=["UCX"]) + started_transfer = time.monotonic() status = self._agent.transfer(handle) - deadline = time.monotonic() + self._timeout_seconds + deadline = started_transfer + self._timeout_seconds while status == "PROC": + if self._closing.is_set(): + raise NixlError("NIXL runtime is closing") if time.monotonic() >= deadline: raise NixlError(f"NIXL WRITE timed out after {self._timeout_seconds:g}s") with self._lock: - self._ensure_open() status = self._agent.check_xfer_state(handle) if status == "PROC": - time.sleep(0.0005) + time.sleep(_POLL_INTERVAL_SECONDS) + with self._lock: + self._data_transfer_seconds += time.monotonic() - started_transfer if status != "DONE": raise NixlError(f"NIXL WRITE failed with status {status!r}") - except NixlError: - raise except Exception as exc: - raise NixlError(f"NIXL H2H scatter WRITE failed: {exc}") from exc - finally: + error_message = f"NIXL H2H scatter WRITE failed: {exc}" with self._lock: + self._failed_peers.add(remote_name) if handle is None: - for registration in frame_registrations: - self._deregister_registration(registration.registrations) - elif not self._release_send_resources(handle, frame_registrations): - self._deferred_sends.append((handle, frame_registrations)) + self._retain_cleanup_failures(None, self._cleanup_sources(sources)) + else: + self._retained_resources.append((handle, sources)) + self._total_seconds += time.monotonic() - started_total + handle = None + sources = [] - @staticmethod - def _validate_send_metadata( - endpoint: dict[str, Any], - token: dict[str, Any], - descriptor: Any, - ) -> tuple[str, bytes]: - remote_name = str(token.get("agent_name") or endpoint.get("agent_name") or "") - metadata = token.get("agent_metadata") or endpoint.get("agent_metadata") - if not remote_name or not isinstance(metadata, bytes): - raise NixlError("NIXL endpoint is missing remote agent metadata") - if int(token.get("payload_bytes", -1)) != descriptor.payload_bytes: - raise NixlError("NIXL receive token length does not match descriptor") - if not isinstance(token.get("frame_remote_descs"), bytes): - raise NixlError("NIXL receive token is missing frame descriptors") - return remote_name, metadata + if error_message is not None: + # A retained Future must not keep the native handle and agent alive. + raise NixlError(error_message) - def receive(self, descriptor: Any) -> Future[memoryview]: - """Complete a prepared receive and return detached packed payload bytes.""" - descriptor.validate() - future: Future[memoryview] = Future() with self._lock: - self._ensure_open() - state = self._receives.pop(descriptor.transfer_id, None) - if state is None: - future.set_exception(NixlError(f"no prepared NIXL receive for {descriptor.transfer_id}")) - return future + cleanup_failed = False try: - payload = state.scratch.buffer[: descriptor.payload_bytes] - future.set_result(memoryview(payload)) + self._agent.release_xfer_handle(handle) + handle = None except Exception as exc: - future.set_exception(NixlError(f"failed to copy NIXL receive payload: {exc}")) - finally: - self._return_receive_buffer(state.scratch) - return future + logger.warning("failed to release completed NIXL transfer handle: %s", exc) + self._retained_resources.append((handle, sources)) + cleanup_failed = True + if not cleanup_failed: + failed_sources = self._cleanup_sources(sources) + self._retain_cleanup_failures(None, failed_sources) + cleanup_failed = bool(failed_sources) + if cleanup_failed: + self._failed_peers.add(remote_name) + self._total_seconds += time.monotonic() - started_total + + def _load_remote_metadata(self, remote_name: str, metadata: bytes) -> None: + previous = self._remote_metadata.get(remote_name) + if previous == metadata: + return + if previous is not None: + self._agent.remove_remote_agent(remote_name) + loaded_name = self._agent.add_remote_agent(metadata) + if isinstance(loaded_name, bytes): + loaded_name = loaded_name.decode() + if loaded_name != remote_name: + raise NixlError(f"NIXL remote agent name mismatch: expected {remote_name!r}, got {loaded_name!r}") + self._remote_metadata[remote_name] = metadata + + def _cleanup_sources(self, sources: list[_FrameSource]) -> list[_FrameSource]: + failed = [] + for source in sources: + if source.registration is None: + continue + try: + self._agent.deregister_memory(source.registration, backends=["UCX"]) + self._registered_bytes -= source.size + self._registered_sources.pop(source.address) + except Exception as exc: + logger.warning("failed to deregister NIXL source memory: %s", exc) + failed.append(source) + return failed - def cancel_receive(self, transfer_id: str) -> None: - """Cancel a prepared receive and deregister its scratch buffer.""" - with self._lock: - state = self._receives.pop(transfer_id, None) - if state is None: - return - self._deregister_registration(state.scratch.registrations) + def _retain_cleanup_failures(self, handle: Any | None, sources: list[_FrameSource]) -> None: + if handle is not None or sources: + self._retained_resources.append((handle, sources)) - def close(self) -> None: - """Stop the send executor and release all NIXL registrations.""" + def _release_lease(self, key: int) -> None: with self._lock: - if self._closed: + scratch = self._leased_receive_buffers.get(key) + if scratch is None or self._closed: return - self._closed = True - self._send_executor.shutdown(wait=True, cancel_futures=True) - with self._lock: - self._reap_deferred_sends() - for state in self._receives.values(): - self._deregister_registration(state.scratch.registrations) - self._receives.clear() - if self._reusable_receive_buffer is not None: - self._deregister_registration(self._reusable_receive_buffer.registrations) - self._reusable_receive_buffer = None - self._agent = None + self._leased_receive_buffers.pop(key) + scratch.lease_finalizer = None + self._idle_receive_buffers.append(scratch) def _ensure_open(self) -> None: if self._closed: raise NixlError("NIXL runtime is closed") - def _deregister_registration(self, registrations: Any) -> None: - try: - self._agent.deregister_memory(registrations, backends=["UCX"]) - except Exception as exc: - logger.warning("failed to deregister NIXL memory: %s", exc) + @staticmethod + def _validate_send_metadata( + endpoint: dict[str, Any], + token: dict[str, Any], + descriptor: Any, + ) -> tuple[str, bytes]: + remote_name = str(token.get("agent_name") or endpoint.get("agent_name") or "") + metadata = token.get("agent_metadata") + if not remote_name or not isinstance(metadata, bytes): + raise NixlError("NIXL receive token is missing current agent metadata") + if int(token.get("payload_bytes", -1)) != descriptor.payload_bytes: + raise NixlError("NIXL receive token length does not match descriptor") + if not isinstance(token.get("frame_remote_descs"), bytes): + raise NixlError("NIXL receive token is missing frame descriptors") + return remote_name, metadata diff --git a/transfer_queue/storage/simple_storage.py b/transfer_queue/storage/simple_storage.py index eb1f4384..f373c877 100644 --- a/transfer_queue/storage/simple_storage.py +++ b/transfer_queue/storage/simple_storage.py @@ -15,9 +15,12 @@ import os import pickle +import socket import time import weakref from collections.abc import Mapping +from concurrent.futures import Future +from queue import Empty, SimpleQueue from threading import Event, Thread from typing import TYPE_CHECKING, Any from uuid import uuid4 @@ -26,7 +29,7 @@ import ray import zmq -from transfer_queue.storage.payload_transfer import PayloadTransfer, create_payload_transfer +from transfer_queue.storage.payload_transfer import DeferredResponse, PayloadTransfer, create_payload_transfer from transfer_queue.utils.common import limit_pytorch_auto_parallel_threads from transfer_queue.utils.enum_utils import Role from transfer_queue.utils.logging_utils import get_logger @@ -53,6 +56,50 @@ KEY_NOT_FOUND_MARKER = "TQKeyNotFound" +def _queue_deferred_response( + identity: bytes, + response: DeferredResponse, + completions: SimpleQueue[tuple[bytes, Future[ZMQMessage]]], + wakeup: socket.socket, + shutdown_event: Event, +) -> None: + """Wake the owning worker when an asynchronous response becomes ready.""" + routing_identity = bytes(identity) + + def ready(future: Future[ZMQMessage]) -> None: + if shutdown_event.is_set(): + return + completions.put((routing_identity, future)) + try: + wakeup.send(b"\0") + except (BlockingIOError, OSError): + pass + + response.future.add_done_callback(ready) + + +def _drain_deferred_responses( + completions: SimpleQueue[tuple[bytes, Future[ZMQMessage]]], + worker_socket: zmq.Socket, + storage_id: str, +) -> None: + """Send completed responses from the only thread that owns the ZMQ socket.""" + while True: + try: + identity, future = completions.get_nowait() + except Empty: + return + try: + response = future.result() + except Exception as exc: + response = ZMQMessage.create( + request_type=ZMQRequestType.PUT_GET_ERROR, + sender_id=storage_id, + body={"message": f"{storage_id}, deferred response failed: {exc}."}, + ) + worker_socket.send_multipart([identity] + response.serialize(), copy=False) + + class StorageKeyNotFoundError(KeyError): """Raised when a requested global index is absent from a storage unit. @@ -289,6 +336,12 @@ def _worker_routine(self) -> None: poller = zmq.Poller() poller.register(worker_socket, zmq.POLLIN) + wakeup_reader, wakeup_writer = socket.socketpair() + wakeup_reader.setblocking(False) + wakeup_writer.setblocking(False) + wakeup_fd = wakeup_reader.fileno() + poller.register(wakeup_fd, zmq.POLLIN) + deferred_completions: SimpleQueue[tuple[bytes, Future[ZMQMessage]]] = SimpleQueue() logger.info(f"[{self.storage_unit_id}]: worker thread started...") perf_monitor = IntervalPerfMonitor(caller_name=f"{self.storage_unit_id}") @@ -308,6 +361,18 @@ def _worker_routine(self) -> None: if self._shutdown_event.is_set(): break + if wakeup_fd in socks: + try: + while wakeup_reader.recv(4096): + pass + except BlockingIOError: + pass + try: + _drain_deferred_responses(deferred_completions, worker_socket, self.storage_unit_id) + except zmq.ZMQError as exc: + logger.warning(f"[{self.storage_unit_id}]: deferred response send failed: {exc}") + break + if worker_socket in socks: # Messages received from proxy: [identity, serialized_msg_frame1, ...] messages = worker_socket.recv_multipart(copy=False) @@ -374,11 +439,22 @@ def _worker_routine(self) -> None: }, ) - # Send response back with identity for routing - worker_socket.send_multipart([identity] + response_msg.serialize(), copy=False) + if isinstance(response_msg, DeferredResponse): + _queue_deferred_response( + bytes(identity), + response_msg, + deferred_completions, + wakeup_writer, + self._shutdown_event, + ) + else: + worker_socket.send_multipart([identity] + response_msg.serialize(), copy=False) logger.info(f"[{self.storage_unit_id}]: worker stopped.") + poller.unregister(wakeup_fd) poller.unregister(worker_socket) + wakeup_reader.close() + wakeup_writer.close() worker_socket.close(linger=0) def _put_decoded_data(self, global_indexes: list[int], field_data: dict[str, Any], data_parser: Any) -> None: