From d832fa9eaea41cb9f834b25010d378c1fd360003 Mon Sep 17 00:00:00 2001 From: overloadedHenry <115620759+overloadedHenry@users.noreply.github.com> Date: Mon, 24 Aug 2026 03:24:51 +0800 Subject: [PATCH 1/2] fix(storage): enforce Mooncake correctness contracts --- tests/test_mooncake_utils.py | 60 ++++++++- tests/test_storage_manager_notifications.py | 115 ++++++++++++++++++ .../storage/clients/mooncake_client.py | 31 +++-- transfer_queue/storage/managers/base.py | 82 +++++++------ 4 files changed, 240 insertions(+), 48 deletions(-) create mode 100644 tests/test_storage_manager_notifications.py diff --git a/tests/test_mooncake_utils.py b/tests/test_mooncake_utils.py index 42cbcd68..1cb4f251 100644 --- a/tests/test_mooncake_utils.py +++ b/tests/test_mooncake_utils.py @@ -376,12 +376,18 @@ def test_no_gdr_meta_none_no_warning(self, caplog): client.clear(["k0"], custom_backend_meta=None) assert "custom_backend_meta" not in caplog.text - def test_error_code_triggers_log(self, caplog): + def test_non_idempotent_failure_is_raised(self): client = _make_clear_client(use_gdr=False) client._store.batch_remove.side_effect = lambda keys, force: [-1] * len(keys) - with caplog.at_level(logging.ERROR, logger="transfer_queue.storage.clients.mooncake_client"): + with pytest.raises(RuntimeError, match=r"batch_remove failed: k0=-1"): client.clear(["k0"]) - assert "remove failed" in caplog.text + + def test_short_result_is_raised(self): + client = _make_clear_client(use_gdr=False) + client._store.batch_remove.return_value = [0] + client._store.batch_remove.side_effect = None + with pytest.raises(RuntimeError, match="returned 1 results, expected 2"): + client.clear(["k0", "k1"]) def test_already_removed_code_704_is_silent(self, caplog): client = _make_clear_client(use_gdr=False) @@ -395,3 +401,51 @@ def test_success_code_zero_is_silent(self, caplog): with caplog.at_level(logging.ERROR): client.clear(["k0"]) assert "remove failed" not in caplog.text + + +class _SequenceStore: + """Return configured results from each low-level Mooncake batch call.""" + + def __init__(self, results): + self.results = iter(results) + + def batch_upsert_from(self, keys, ptrs, sizes, config=None): + return next(self.results) + + def batch_get_into(self, keys, ptrs, sizes): + return next(self.results) + + +def _make_retry_client(store): + from transfer_queue.storage.clients.mooncake_client import MooncakeStoreClient + + client = object.__new__(MooncakeStoreClient) + client._store = store + client.replica_config = None + return client + + +class TestBatchResultValidation: + def test_upsert_retry_short_result_is_raised(self, monkeypatch): + monkeypatch.setattr("transfer_queue.storage.clients.mooncake_client.RETRY_DELAY_SECONDS", 0) + client = _make_retry_client(_SequenceStore([[-1, -1], [0]])) + + with pytest.raises(RuntimeError, match="batch_upsert_from returned 1 results, expected 2"): + client._batch_upsert_with_retry(["k0", "k1"], [1, 2], [8, 8]) + + def test_get_retry_short_result_is_raised(self, monkeypatch): + monkeypatch.setattr("transfer_queue.storage.clients.mooncake_client.RETRY_DELAY_SECONDS", 0) + client = _make_retry_client(_SequenceStore([[-1, -1], [0]])) + + with pytest.raises(RuntimeError, match="batch_get_into returned 1 results, expected 2"): + client._batch_get_into_with_retry(["k0", "k1"], [1, 2], [8, 8]) + + @pytest.mark.parametrize("operation", ["upsert", "get"]) + def test_non_sized_result_is_raised(self, operation): + client = _make_retry_client(_SequenceStore([None])) + + with pytest.raises(RuntimeError, match="returned a non-sized result, expected 2 codes"): + if operation == "upsert": + client._batch_upsert_with_retry(["k0", "k1"], [1, 2], [8, 8]) + else: + client._batch_get_into_with_retry(["k0", "k1"], [1, 2], [8, 8]) diff --git a/tests/test_storage_manager_notifications.py b/tests/test_storage_manager_notifications.py new file mode 100644 index 00000000..24884a34 --- /dev/null +++ b/tests/test_storage_manager_notifications.py @@ -0,0 +1,115 @@ +# 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. + +from types import SimpleNamespace + +import pytest + +from transfer_queue.storage.managers.base import StorageManager +from transfer_queue.utils.zmq_utils import ZMQRequestType + + +class _FakeNotifySocket: + def __init__(self, connect_error: Exception | None = None) -> None: + self.closed = False + self.connect_error = connect_error + + def setsockopt(self, *args, **kwargs) -> None: + pass + + def connect(self, *args, **kwargs) -> None: + if self.connect_error is not None: + raise self.connect_error + + async def send_multipart(self, request) -> None: + pass + + async def recv_multipart(self, copy=False): + return [b"ack"] + + def close(self, linger=0) -> None: + self.closed = True + + +def _manager(controller_info=None): + if controller_info is None: + controller_info = SimpleNamespace( + id="controller", + ip="127.0.0.1", + to_addr=lambda name: "inproc://controller", + ) + return SimpleNamespace( + storage_manager_id="notification-test", + zmq_context=object(), + controller_info=controller_info, + ) + + +def _ack(success: bool, partition_id: str = "p0"): + return SimpleNamespace( + request_type=ZMQRequestType.NOTIFY_DATA_UPDATE_ACK, + sender_id="controller", + body={"success": success, "partition_id": partition_id}, + ) + + +@pytest.mark.asyncio +async def test_notify_data_update_rejects_missing_controller(): + manager = _manager(controller_info=False) + + with pytest.raises(RuntimeError, match="has no controller"): + await StorageManager.notify_data_update(manager, "p0", [], {}, {}) + + +@pytest.mark.asyncio +async def test_notify_and_wait_requires_positive_ack(monkeypatch): + socket = _FakeNotifySocket() + monkeypatch.setattr("transfer_queue.storage.managers.base.create_zmq_socket", lambda **kwargs: socket) + monkeypatch.setattr("transfer_queue.storage.managers.base.ZMQMessage.deserialize", lambda messages: _ack(False)) + + with pytest.raises(RuntimeError, match="rejected the production-status update"): + await StorageManager._notify_and_wait(_manager(), [b"request"]) + assert socket.closed is True + + +@pytest.mark.asyncio +async def test_notify_and_wait_accepts_positive_ack(monkeypatch): + socket = _FakeNotifySocket() + monkeypatch.setattr("transfer_queue.storage.managers.base.create_zmq_socket", lambda **kwargs: socket) + monkeypatch.setattr("transfer_queue.storage.managers.base.ZMQMessage.deserialize", lambda messages: _ack(True)) + + await StorageManager._notify_and_wait(_manager(), [b"request"]) + assert socket.closed is True + + +@pytest.mark.asyncio +async def test_notify_and_wait_times_out_without_ack(monkeypatch): + socket = _FakeNotifySocket() + monkeypatch.setattr("transfer_queue.storage.managers.base.create_zmq_socket", lambda **kwargs: socket) + monkeypatch.setattr("transfer_queue.storage.managers.base.TQ_DATA_UPDATE_RESPONSE_TIMEOUT", 0) + + with pytest.raises(TimeoutError, match="production-status ACK"): + await StorageManager._notify_and_wait(_manager(), [b"request"]) + assert socket.closed is True + + +@pytest.mark.asyncio +async def test_notify_and_wait_closes_socket_when_connect_fails(monkeypatch): + socket = _FakeNotifySocket(connect_error=ConnectionError("controller unavailable")) + monkeypatch.setattr("transfer_queue.storage.managers.base.create_zmq_socket", lambda **kwargs: socket) + + with pytest.raises(ConnectionError, match="controller unavailable"): + await StorageManager._notify_and_wait(_manager(), [b"request"]) + assert socket.closed is True diff --git a/transfer_queue/storage/clients/mooncake_client.py b/transfer_queue/storage/clients/mooncake_client.py index d6914902..8c30326c 100644 --- a/transfer_queue/storage/clients/mooncake_client.py +++ b/transfer_queue/storage/clients/mooncake_client.py @@ -45,6 +45,17 @@ MAX_SERIAL_WORKER_THREADS = 4 MAX_RETRIES = 3 RETRY_DELAY_SECONDS = 1.0 +_MOONCAKE_OBJECT_NOT_FOUND = -704 + + +def _validate_batch_result_count(operation: str, keys: list[str], results: Any) -> None: + """Require one Mooncake result code for every requested key.""" + try: + actual = len(results) + except TypeError as error: + raise RuntimeError(f"{operation} returned a non-sized result, expected {len(keys)} codes") from error + if actual != len(keys): + raise RuntimeError(f"{operation} returned {actual} results, expected {len(keys)}") @StorageClientFactory.register("MooncakeStoreClient") @@ -529,9 +540,15 @@ def clear(self, keys: list[str], custom_backend_meta: list[Any] | None = None) - actual_keys = keys ret_codes = self._store.batch_remove(actual_keys, force=True) - for i, ret in enumerate(ret_codes): - if not (ret == 0 or ret == -704): - logger.error(f"remove failed for key `{actual_keys[i]}` with error code: {ret}") + _validate_batch_result_count("batch_remove", actual_keys, ret_codes) + failures = [ + (key, code) + for key, code in zip(actual_keys, ret_codes, strict=True) + if code not in (0, _MOONCAKE_OBJECT_NOT_FOUND) + ] + if failures: + detail = ", ".join(f"{key}={code}" for key, code in failures) + raise RuntimeError(f"batch_remove failed: {detail}") def close(self): """Closes MooncakeStore.""" @@ -549,8 +566,7 @@ def _batch_upsert_with_retry(self, batch_keys: list[str], batch_ptrs: list[int], backing tensors/buffers). """ results = self._store.batch_upsert_from(batch_keys, batch_ptrs, batch_sizes, config=self.replica_config) - if len(results) != len(batch_keys): - raise RuntimeError(f"batch_upsert_from returned {len(results)} results, expected {len(batch_keys)}") + _validate_batch_result_count("batch_upsert_from", batch_keys, results) failed_indices = [j for j, r in enumerate(results) if r != 0] if not failed_indices: @@ -572,6 +588,7 @@ def _batch_upsert_with_retry(self, batch_keys: list[str], batch_ptrs: list[int], retry_results = self._store.batch_upsert_from( current_failed_keys, retry_ptrs, retry_sizes, config=self.replica_config ) + _validate_batch_result_count("batch_upsert_from", current_failed_keys, retry_results) next_failed_indices = [] next_failed_keys = [] @@ -612,8 +629,7 @@ def _batch_get_into_with_retry( Caller owns the receive buffers (allocate/register/unregister). """ ret_codes = self._store.batch_get_into(batch_keys, batch_buffer_ptrs, batch_nbytes) - if len(ret_codes) != len(batch_keys): - raise RuntimeError(f"batch_get_into returned {len(ret_codes)} results, expected {len(batch_keys)}") + _validate_batch_result_count("batch_get_into", batch_keys, ret_codes) failed_indices = [i for i, ret in enumerate(ret_codes) if ret < 0] if not failed_indices: @@ -634,6 +650,7 @@ def _batch_get_into_with_retry( retry_nbytes = [batch_nbytes[i] for i in current_failed_indices] retry_codes = self._store.batch_get_into(current_failed_keys, retry_ptrs, retry_nbytes) + _validate_batch_result_count("batch_get_into", current_failed_keys, retry_codes) next_failed_indices = [] next_failed_keys = [] diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index 5deb5ac8..640c8e61 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -21,7 +21,7 @@ import weakref from abc import ABC, abstractmethod from concurrent.futures import ThreadPoolExecutor -from typing import Any, Callable, Optional +from typing import Any, Callable from uuid import uuid4 import ray @@ -206,8 +206,8 @@ async def notify_data_update( partition_id: str, global_indexes: list[int], field_schema: dict[str, dict[str, Any]], - custom_backend_meta: Optional[dict[int, dict[str, Any]]] = None, - user_custom_meta: Optional[dict[int, dict[str, Any]]] = None, + custom_backend_meta: dict[int, dict[str, Any]] | None = None, + user_custom_meta: dict[int, dict[str, Any]] | None = None, ) -> None: """ Notify controller that new data is ready. @@ -222,8 +222,9 @@ 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"Storage manager {self.storage_manager_id} has no controller for production-status notification" + ) normalized_field_schema = {} for field_name, field in field_schema.items(): @@ -261,55 +262,60 @@ async def notify_data_update( async def _notify_and_wait(self, request_msg: list) -> None: """Send a data status notification to the controller and block until ACK is received.""" identity = f"{self.storage_manager_id}-notify-{uuid4().hex[:8]}".encode() - sock = create_zmq_socket( - ctx=self.zmq_context, socket_type=zmq.DEALER, ip=self.controller_info.ip, identity=identity - ) - sock.setsockopt(zmq.LINGER, 0) - sock.connect(self.controller_info.to_addr("request_handle_socket")) - + sock = None try: + sock = create_zmq_socket( + ctx=self.zmq_context, socket_type=zmq.DEALER, ip=self.controller_info.ip, identity=identity + ) + sock.setsockopt(zmq.LINGER, 0) + sock.connect(self.controller_info.to_addr("request_handle_socket")) + await sock.send_multipart(request_msg) logger.debug( f"[{self.storage_manager_id}]: Sent data status update request " f"to controller id #{self.controller_info.id} successfully." ) - response_received = False - timeout = TQ_DATA_UPDATE_RESPONSE_TIMEOUT - - while not response_received and timeout > 0: + loop = asyncio.get_running_loop() + deadline = loop.time() + TQ_DATA_UPDATE_RESPONSE_TIMEOUT + while True: + remaining = deadline - loop.time() + if remaining <= 0: + raise TimeoutError( + f"Timed out waiting for production-status ACK after {TQ_DATA_UPDATE_RESPONSE_TIMEOUT}s" + ) try: - poll_interval = min(TQ_STORAGE_POLLER_TIMEOUT, timeout) messages = await asyncio.wait_for( sock.recv_multipart(copy=False), - timeout=poll_interval, + timeout=min(TQ_STORAGE_POLLER_TIMEOUT, remaining), ) - response_msg = ZMQMessage.deserialize(messages) - - if response_msg.request_type == ZMQRequestType.NOTIFY_DATA_UPDATE_ACK: # type: ignore[arg-type] - response_received = True - logger.debug( - f"[{self.storage_manager_id}]: Get data status update ACK response " - f"from controller id #{response_msg.sender_id} successfully." - ) - break except asyncio.TimeoutError: - timeout -= poll_interval - except Exception as e: - logger.warning(f"[{self.storage_manager_id}]: Error receiving response: {e}") - break + continue + except Exception as error: + raise RuntimeError("Failed while waiting for production-status ACK") from error + + response_msg = ZMQMessage.deserialize(messages) + if response_msg.request_type != ZMQRequestType.NOTIFY_DATA_UPDATE_ACK: # type: ignore[arg-type] + continue + + response_body = response_msg.body if isinstance(response_msg.body, dict) else {} + if response_body.get("success") is not True: + raise RuntimeError( + "Controller rejected the production-status update " + f"for partition={response_body.get('partition_id', 'unknown')}" + ) - if not response_received: - logger.error( - f"[{self.storage_manager_id}]: Timeout waiting for data status update ACK " - f"from controller after {TQ_DATA_UPDATE_RESPONSE_TIMEOUT}s." + logger.debug( + f"[{self.storage_manager_id}]: Get data status update ACK response " + f"from controller id #{response_msg.sender_id} successfully." ) + return finally: try: - if not sock.closed: + if sock is not None and not sock.closed: sock.close(linger=0) - except Exception: - pass + except Exception as error: + logger.debug(f"Failed to close production-status notification socket: {error}") @abstractmethod async def put_data( @@ -694,7 +700,7 @@ async def put_data( # atomically with the readiness notification (avoids the put/set_custom_meta # race for streaming consumers). Only sent when at least one sample has it. user_custom_meta_list = metadata.get_all_custom_meta() - user_custom_meta: Optional[dict[int, dict[str, Any]]] = None + user_custom_meta: dict[int, dict[str, Any]] | None = None if any(user_custom_meta_list): user_custom_meta = { metadata.global_indexes[i]: user_custom_meta_list[i] From 6c7a587292910af0827f027de99e005e1900310e Mon Sep 17 00:00:00 2001 From: overloadedHenry <115620759+overloadedHenry@users.noreply.github.com> Date: Thu, 27 Aug 2026 01:19:05 +0800 Subject: [PATCH 2/2] feat(storage): add Mooncake correctness contract version --- tests/test_mooncake_utils.py | 7 +++++++ transfer_queue/__init__.py | 5 +++++ 2 files changed, 12 insertions(+) diff --git a/tests/test_mooncake_utils.py b/tests/test_mooncake_utils.py index 1cb4f251..e20da77f 100644 --- a/tests/test_mooncake_utils.py +++ b/tests/test_mooncake_utils.py @@ -34,6 +34,13 @@ _has_cuda_python = importlib.util.find_spec("cuda") is not None +def test_mooncake_correctness_contract_version_is_public(): + import transfer_queue as tq + + assert tq.MOONCAKE_CORRECTNESS_CONTRACT_VERSION == 1 + assert "MOONCAKE_CORRECTNESS_CONTRACT_VERSION" in tq.__all__ + + def _aligned(n: int) -> int: return (n + _DEFAULT_ALIGN - 1) // _DEFAULT_ALIGN * _DEFAULT_ALIGN diff --git a/transfer_queue/__init__.py b/transfer_queue/__init__.py index 0ed98b2f..2c3b354e 100644 --- a/transfer_queue/__init__.py +++ b/transfer_queue/__init__.py @@ -45,6 +45,10 @@ from .sampler.sequential_sampler import SequentialSampler from .sampler.streaming_token_budget_sampler import StreamingTokenBudgetSampler +# Version 1 guarantees Mooncake batch/retry result-count validation, +# batch_remove failure propagation, and fail-closed production-status ACKs. +MOONCAKE_CORRECTNESS_CONTRACT_VERSION = 1 + __all__ = ( [ # High-Level KV Interface @@ -80,6 +84,7 @@ "get_client", "BatchMeta", "TransferQueueClient", + "MOONCAKE_CORRECTNESS_CONTRACT_VERSION", ] + [ # Sampler