From bca3a31351162fa77c2582644d732a731707a4f5 Mon Sep 17 00:00:00 2001 From: Robert Washbourne Date: Sat, 3 Oct 2026 13:09:14 +0000 Subject: [PATCH] Propagate failed data notifications to producers --- tests/test_notification_failures.py | 173 ++++++++++++++++++++++++ transfer_queue/storage/managers/base.py | 56 ++++---- 2 files changed, 205 insertions(+), 24 deletions(-) create mode 100644 tests/test_notification_failures.py diff --git a/tests/test_notification_failures.py b/tests/test_notification_failures.py new file mode 100644 index 00000000..8da3a170 --- /dev/null +++ b/tests/test_notification_failures.py @@ -0,0 +1,173 @@ +# Copyright 2025 Huawei Technologies Co., Ltd. All Rights Reserved. +# Copyright 2025 The TransferQueue Team +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Public producer controls with real controller and storage services.""" + +import asyncio +import os +import time +from uuid import uuid4 + +import psutil +import pytest +import ray +import torch +from tensordict import TensorDict + +from transfer_queue.client import TransferQueueClient +from transfer_queue.controller import TransferQueueController +from transfer_queue.storage.managers import base +from transfer_queue.storage.simple_storage import SimpleStorageUnit +from transfer_queue.utils.zmq_utils import ZMQRequestType + + +@ray.remote(num_cpus=1) +class NotificationController(TransferQueueController.__ray_metadata__.modified_class): + response_mode = "normal" + + def set_response_mode(self, mode): + self.response_mode = mode + + def process_id(self): + return os.getpid() + + def _handle_notify_data_update_request(self, request): + if self.response_mode == "normal": + return super()._handle_notify_data_update_request(request) + response = self._make_response(request, ZMQRequestType.NOTIFY_DATA_UPDATE_ACK, {"success": True}) + if self.response_mode == "wrong_sender": + response.sender_id = "unexpected_controller" + elif self.response_mode == "missing_success": + response.body = {} + elif self.response_mode == "error": + response.request_type = ZMQRequestType.REQUEST_ERROR + response.body = {"message": "notification rejected by control"} + elif self.response_mode == "unexpected_type": + response.request_type = ZMQRequestType.GET_META_RESPONSE + return response + + +@pytest.fixture(scope="module") +def ray_services(): + assert not ray.is_initialized(), "requires a test-owned local Ray instance" + ray.init( + address="local", + num_cpus=4, + num_gpus=0, + include_dashboard=False, + object_store_memory=256 * 1024 * 1024, + namespace="notify-" + uuid4().hex, + ) + yield + ray.shutdown() + + +@pytest.fixture +def queue(ray_services, monkeypatch): + monkeypatch.setattr(base, "TQ_DATA_UPDATE_RESPONSE_TIMEOUT", 2) + controller = NotificationController.remote(polling_mode=True) + storage = SimpleStorageUnit.remote(config={"num_data_storage_units": 1}) + info, storage_info = ray.get([controller.get_zmq_server_info.remote(), storage.get_zmq_server_info.remote()]) + client = TransferQueueClient(client_id="producer-" + uuid4().hex, controller_info=info) + client.initialize_storage_manager("SimpleStorage", {"zmq_info": {storage_info.id: storage_info}}) + try: + yield client, controller + finally: + client.close() + try: + ray.get(storage.shutdown.remote(), timeout=10) + finally: + ray.kill(storage, no_restart=True) + ray.kill(controller, no_restart=True) + + +def _data(): + return TensorDict({"value": torch.tensor([[22]])}, batch_size=1) + + +def test_successful_ack_completes_public_put(queue): + client, controller = queue + meta = client.put(_data(), partition_id="valid") + assert [sample.item() for sample in client.get_data(meta)["value"].unbind()] == [22] + partition = ray.get(controller.get_partition_snapshot.remote("valid")) + assert partition.global_indexes == set(meta.global_indexes) + assert partition.production_status[meta.global_indexes, : partition.total_fields_num].eq(1).all() + + +def test_real_negative_ack_reaches_public_put(queue): + client, controller = queue + meta = client.put(_data(), partition_id="removed") + client.clear_partition("removed") + # The retained handle exercises the real controller's negative ACK without reallocating. + with pytest.raises(RuntimeError, match="rejected data status update"): + client.put(_data(), metadata=meta) + assert "removed" not in ray.get(controller.list_partitions.remote()) + asyncio.run(client.storage_manager.clear_data(meta)) + + +@pytest.mark.parametrize( + "mode, error", + [ + ("wrong_sender", "Unexpected data status update sender"), + ("missing_success", "rejected data status update"), + ("error", "notification rejected by control"), + ("unexpected_type", "Controller data status update failed"), + ], +) +def test_invalid_service_response_reaches_public_put(queue, mode, error): + client, controller = queue + ray.get(controller.set_response_mode.remote(mode)) + with pytest.raises(RuntimeError, match=error): + client.put(_data(), partition_id="invalid") + partition = ray.get(controller.get_partition_snapshot.remote("invalid")) + assert partition.production_status.eq(0).all() + # A failed response must not poison the pooled socket used by the next valid put. + ray.get(controller.set_response_mode.remote("normal")) + meta = client.put(_data(), partition_id="next") + assert [sample.item() for sample in client.get_data(meta)["value"].unbind()] == [22] + + +def test_missing_controller_ack_fails_public_put_before_deadline(queue): + client, controller = queue + meta = client.put(_data(), partition_id="timeout") + process = psutil.Process(ray.get(controller.process_id.remote())) + ray.kill(controller, no_restart=True) + # Ray marks the actor dead before its process necessarily stops serving ZMQ. + process.wait(timeout=10) + assert not process.is_running() + print({"controller_pid": process.pid, "process_terminated": True}) + with pytest.raises(ray.exceptions.RayActorError): + ray.get(controller.list_partitions.remote(), timeout=1) + + async def put(): + await asyncio.wait_for(client.async_put(_data(), metadata=meta), timeout=8) + + started = time.monotonic() + with pytest.raises(TimeoutError, match="no ACK from controller"): + asyncio.run(put()) + assert time.monotonic() - started < 8 + asyncio.run(client.storage_manager.clear_data(meta)) + + +def test_missing_controller_configuration_fails_public_put(queue): + client, _ = queue + meta = client.put(_data(), partition_id="missing") + info = client.storage_manager.controller_info + client.storage_manager.controller_info = None + try: + with pytest.raises(RuntimeError, match="No controller connected"): + client.put(_data(), metadata=meta) + finally: + client.storage_manager.controller_info = info diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index 369f800e..eaa771c3 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -238,8 +238,7 @@ async def notify_data_update( """ if not self.controller_info: - logger.warning(f"No controller connected for storage manager {self.storage_manager_id}") - return + raise RuntimeError(f"No controller connected for storage manager {self.storage_manager_id}") normalized_field_schema = {} for field_name, field in field_schema.items(): @@ -285,29 +284,38 @@ async def _notify_and_wait(self, request_msg: list) -> None: f"to controller id #{self.controller_info.id} successfully." ) - # One deadline for the whole wait, so unrelated traffic on this socket cannot - # extend it and a quiet controller still gets the full budget. - deadline = time.monotonic() + TQ_DATA_UPDATE_RESPONSE_TIMEOUT - while True: - remaining = deadline - time.monotonic() - if remaining <= 0: - raise TimeoutError( - f"no ACK from controller {self.controller_info.id} after {TQ_DATA_UPDATE_RESPONSE_TIMEOUT}s" - ) - messages = await asyncio.wait_for(sock.recv_multipart(copy=False), timeout=remaining) - response_msg = ZMQMessage.deserialize(messages) - - if response_msg.request_type == ZMQRequestType.NOTIFY_DATA_UPDATE_ACK: # type: ignore[arg-type] - logger.debug( - f"[{self.storage_manager_id}]: Get data status update ACK response " - f"from controller id #{response_msg.sender_id} successfully." - ) - return - except Exception as e: - # Logged rather than raised, so a slow controller does not fail the put. Close - # the socket: a late ACK would otherwise be read as the next lessee's reply. - logger.error(f"[{self.storage_manager_id}]: Data status update failed: {type(e).__name__}: {e}") + try: + messages = await asyncio.wait_for( + sock.recv_multipart(copy=False), timeout=TQ_DATA_UPDATE_RESPONSE_TIMEOUT + ) + except asyncio.TimeoutError as e: + raise TimeoutError( + f"no ACK from controller {self.controller_info.id} after {TQ_DATA_UPDATE_RESPONSE_TIMEOUT}s" + ) from e + response_msg = ZMQMessage.deserialize(messages) + if response_msg.sender_id != self.controller_info.id: + raise RuntimeError( + f"Unexpected data status update sender {response_msg.sender_id}; " + f"expected {self.controller_info.id}" + ) + if response_msg.request_type != ZMQRequestType.NOTIFY_DATA_UPDATE_ACK: + raise RuntimeError( + f"Controller data status update failed: {response_msg.request_type}: " + f"{response_msg.body.get('message', 'unexpected response')}" + ) + if response_msg.body.get("success") is not True: + raise RuntimeError( + f"Controller rejected data status update: " + f"{response_msg.body.get('message', 'success is not true')}" + ) + logger.debug( + f"[{self.storage_manager_id}]: Get data status update ACK response " + f"from controller id #{response_msg.sender_id} successfully." + ) + except BaseException: + # A failed or cancelled wait must not leave a late ACK for the next lessee. sock.close(linger=0) + raise @abstractmethod async def put_data(