Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
173 changes: 173 additions & 0 deletions tests/test_notification_failures.py
Original file line number Diff line number Diff line change
@@ -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
56 changes: 32 additions & 24 deletions transfer_queue/storage/managers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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():
Expand Down Expand Up @@ -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(
Expand Down