From 8aadcb348dd4c17ed4b2592f9b8eb50f5a831e77 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Thu, 23 Jul 2026 17:04:16 +0800 Subject: [PATCH 1/6] feat: reuse long-lived ZMQ context in with_zmq_socket instead of per-call The with_zmq_socket decorator created a brand-new zmq.asyncio.Context() per RPC call and context.term()'d it in the finally block. Under the high-concurrency agent-loop store path this churned libzmq's signaler file descriptors and crashed the worker (signaler.cpp Bad file descriptor -> SIGABRT), and could also hang on the blocking term() (aggravated by sock.close(linger=-1)). Fix: the decorator now reuses the owner's long-lived context via a required get_context callable, and only creates/closes the DEALER socket per call. The context is created once per owner and terminated once at close(). Contexts are thread-safe and event-loop-agnostic, so a single shared context is safe across loops/threads; each socket stays per-call on one loop. - zmq_utils.with_zmq_socket: add required get_context; drop per-call Context()/term(); change sock.close(linger=-1) -> linger=0. - client.AsyncTransferQueueClient: own a shared self.zmq_context; destroy(linger=0) in close(). - simple_storage_manager: feed the base StorageManager's self.zmq_context via get_context. - base.StorageManager.close(): term() -> destroy(linger=0) so a leaked socket cannot hang shutdown. - tests: add test_zmq_shared_context.py asserting concurrent RPCs reuse one context and it is closed exactly once. Microbenchmark: ~7.6x faster socket setup/teardown (~210us saved per call). Co-Authored-By: Claude Opus 4.8 (1M context) --- tests/test_zmq_shared_context.py | 168 ++++++++++++++++++ transfer_queue/client.py | 13 ++ transfer_queue/storage/managers/base.py | 4 +- .../managers/simple_storage_manager.py | 4 + transfer_queue/utils/zmq_utils.py | 33 ++-- 5 files changed, 211 insertions(+), 11 deletions(-) create mode 100644 tests/test_zmq_shared_context.py diff --git a/tests/test_zmq_shared_context.py b/tests/test_zmq_shared_context.py new file mode 100644 index 00000000..37700189 --- /dev/null +++ b/tests/test_zmq_shared_context.py @@ -0,0 +1,168 @@ +# 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. + +"""Regression tests for the shared long-lived ZMQ context in with_zmq_socket. + +Background: with_zmq_socket used to create a brand-new ``zmq.asyncio.Context()`` per RPC +call and ``context.term()`` it in the finally block. Under concurrency this churned +libzmq signaler file descriptors and crashed the process (``signaler.cpp`` Bad file +descriptor -> SIGABRT). The fix makes the decorator reuse the owner's long-lived context +(``get_context``) and only create/close the DEALER socket per call. + +These tests assert that concurrent decorated calls all reuse the SAME context object and +that the context is never terminated between calls, only when the client is closed. +""" + +import asyncio +from threading import Thread + +import pytest +import zmq + +import transfer_queue.utils.zmq_utils as zmq_utils +from transfer_queue.client import AsyncTransferQueueClient +from transfer_queue.metadata import BatchMeta +from transfer_queue.utils.enum_utils import Role +from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType, ZMQServerInfo + + +class _EchoController: + """Minimal in-process ROUTER controller that answers GET_META requests.""" + + def __init__(self, controller_id="controller_0"): + self.controller_id = controller_id + self.context = zmq.Context() + self.request_socket = self.context.socket(zmq.ROUTER) + self.request_port = self.request_socket.bind_to_random_port("tcp://127.0.0.1") + self.zmq_server_info = ZMQServerInfo( + role=Role.CONTROLLER, + id=controller_id, + ip="127.0.0.1", + ports={"request_handle_socket": self.request_port}, + ) + self.running = True + self.request_thread = Thread(target=self._handle_requests, daemon=True) + self.request_thread.start() + + def _handle_requests(self): + poller = zmq.Poller() + poller.register(self.request_socket, zmq.POLLIN) + while self.running: + try: + socks = dict(poller.poll(100)) + if self.request_socket not in socks: + continue + messages = self.request_socket.recv_multipart(copy=False) + identity = messages.pop(0) + request_msg = ZMQMessage.deserialize(messages) + + batch_size = request_msg.body.get("batch_size", 1) + data_fields = request_msg.body.get("data_fields", []) + field_schema = { + name: {"dtype": None, "shape": None, "is_nested": False, "is_non_tensor": False} + for name in data_fields + } + metadata = BatchMeta( + global_indexes=list(range(batch_size)), + partition_ids=["0"] * batch_size, + field_schema=field_schema, + ) + response_msg = ZMQMessage.create( + request_type=ZMQRequestType.GET_META_RESPONSE, + sender_id=self.controller_id, + receiver_id=request_msg.sender_id, + body={"metadata": metadata}, + ) + self.request_socket.send_multipart([identity, *response_msg.serialize()]) + except zmq.Again: + continue + except Exception as e: # pragma: no cover - surfaced via test failure + print(f"_EchoController ERROR: {e}") + + def stop(self): + self.running = False + self.request_thread.join(timeout=2.0) + self.request_socket.close(linger=0) + self.context.term() + + +@pytest.fixture +def echo_controller(): + controller = _EchoController() + yield controller + controller.stop() + + +@pytest.mark.asyncio +async def test_shared_context_reused_across_concurrent_calls(echo_controller, monkeypatch): + """Many concurrent decorated RPCs must all reuse the client's single context. + + This is the core regression guard: pre-fix, each call created and term()ed its own + context, which is what corrupted libzmq's signaler FDs under concurrency. + """ + client = AsyncTransferQueueClient( + client_id="client_shared_ctx", + controller_info=echo_controller.zmq_server_info, + ) + + # Record the context object handed to every socket creation inside the decorator. + seen_contexts = [] + original_create = zmq_utils.create_zmq_socket + + def _spy_create(ctx, *args, **kwargs): + seen_contexts.append(ctx) + return original_create(ctx, *args, **kwargs) + + # The decorator resolves create_zmq_socket from the zmq_utils module globals. + monkeypatch.setattr(zmq_utils, "create_zmq_socket", _spy_create) + + num_calls = 200 + coros = [ + client.async_get_meta(data_fields=["tokens", "labels"], batch_size=2, partition_id="0") + for _ in range(num_calls) + ] + # wait_for guards against the hang failure mode. + results = await asyncio.wait_for(asyncio.gather(*coros), timeout=60) + + assert len(results) == num_calls + assert all(isinstance(meta, BatchMeta) for meta in results) + + # Every call must have used the SAME context, and it must be the client's context. + assert len(seen_contexts) == num_calls + assert all(ctx is client.zmq_context for ctx in seen_contexts) + # The shared context must NOT have been terminated by any call. + assert not client.zmq_context.closed + + client.close() + + +@pytest.mark.asyncio +async def test_close_destroys_context(echo_controller): + """close() must terminate the shared context exactly once (no leak, no hang).""" + client = AsyncTransferQueueClient( + client_id="client_close_ctx", + controller_info=echo_controller.zmq_server_info, + ) + assert not client.zmq_context.closed + + # A normal call before shutdown leaves the context alive. + await asyncio.wait_for( + client.async_get_meta(data_fields=["tokens"], batch_size=1, partition_id="0"), + timeout=30, + ) + assert not client.zmq_context.closed + + client.close() + assert client.zmq_context.closed diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 4c0db125..6b088785 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -43,6 +43,7 @@ "request_handle_socket", get_identity=lambda self: self.client_id, get_peer=lambda self, target: self._controller, + get_context=lambda self: self.zmq_context, ) @@ -70,6 +71,10 @@ def __init__( raise TypeError(f"controller_info must be ZMQServerInfo, got {type(controller_info)}") self.client_id = client_id self._controller: ZMQServerInfo = controller_info + # Long-lived ZMQ context shared by all controller RPCs on this client. Contexts are + # thread-safe and event-loop-agnostic; each RPC creates and closes its own DEALER + # socket from this context (see with_controller_socket). Terminated once in close(). + self.zmq_context = zmq.asyncio.Context() logger.info(f"[{self.client_id}]: Registered Controller server {controller_info.id} at {controller_info.ip}") def initialize_storage_manager( @@ -1095,6 +1100,14 @@ def close(self) -> None: except Exception as e: logger.warning(f"Error closing storage manager: {e}") + # Tear down the shared context last. destroy(linger=0) force-closes any socket that + # leaked from an interrupted RPC then terminates, so shutdown cannot hang. + try: + if hasattr(self, "zmq_context") and self.zmq_context is not None: + self.zmq_context.destroy(linger=0) + except Exception as e: + logger.warning(f"[{self.client_id}]: Error terminating zmq_context: {e}") + # ==================== Checkpoint API ==================== @with_controller_socket async def async_save_controller_checkpoint( diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index e6b0faf4..c6c959c4 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -391,7 +391,9 @@ def close(self) -> None: else: logger.debug(f"[{self.storage_manager_id}]: Notify ZMQ thread shut down.") - self.zmq_context.term() + # destroy(linger=0) force-closes any socket still open (e.g. from an interrupted + # request or the notify path) then terminates, so shutdown cannot hang on term(). + self.zmq_context.destroy(linger=0) def __del__(self): """Destructor to ensure resources are cleaned up.""" diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index 0b6777fb..c6b20eff 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -49,6 +49,10 @@ "put_get_socket", get_identity=lambda self: self.storage_manager_id, get_peer=lambda self, target: self.storage_unit_infos[target], + # Long-lived context from the base StorageManager (base.py). Now shared by both the + # notify path (_notify_and_wait) and the per-call storage-unit request sockets below; + # this is safe because the context is loop-agnostic and each socket stays per-call. + get_context=lambda self: self.zmq_context, resolve_target=lambda args, kwargs: kwargs.get("target_storage_unit"), timeout=TQ_SIMPLE_STORAGE_SEND_RECV_TIMEOUT, ) diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index 93656542..147340e4 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -316,6 +316,7 @@ def with_zmq_socket( *, get_identity: Callable[[Any], str], get_peer: Callable[[Any, str | None], ZMQServerInfo], + get_context: Callable[[Any], "zmq.asyncio.Context"], resolve_target: Callable[[tuple, dict], str | None] | None = None, timeout: int | None = None, ): @@ -323,7 +324,16 @@ def with_zmq_socket( This decorator encapsulates the common socket lifecycle used by both client-side and storage-manager-side request paths: - create context/socket -> connect -> inject socket -> close/term. + get owner's shared context -> create/connect socket -> inject socket -> close socket. + + The ZMQ context is owned by ``self`` (via ``get_context``) and is long-lived: it is + created once per owner and reused across all calls, then terminated once when the owner + is closed. Only the DEALER *socket* is created and closed per call. Do NOT create or + terminate a context here -- per-call context churn corrupts libzmq's signaler file + descriptors under concurrency (Bad file descriptor / SIGABRT) and can hang on term(). + Contexts are thread-safe and event-loop-agnostic, so a single shared context is safe + even when decorated methods run on different loops/threads; each socket is created and + fully used within one awaited call on one loop. Args: socket_name: Socket port key in ``ZMQServerInfo.ports``. @@ -333,6 +343,8 @@ def with_zmq_socket( For single-target scenarios, ignore the target parameter. Example: ``lambda self, target: self.server_info`` Example: ``lambda self, target: self.storage_unit_infos[target]`` + get_context: Callable that returns the owner's long-lived ``zmq.asyncio.Context``. + Example: ``lambda self: self.zmq_context`` resolve_target: Optional callable that extracts target identifier from function arguments. Receives (args, kwargs) and returns target name. Example: ``lambda args, kwargs: kwargs.get("target_storage_unit")`` @@ -358,7 +370,11 @@ async def wrapper(self, *args, **kwargs): if port is None: raise RuntimeError(f"Socket '{socket_name}' not configured for server '{server_info.id}'") - context = zmq.asyncio.Context() + # Reuse the owner's long-lived context; only the socket is per-call. + context = get_context(self) + if context is None: + raise RuntimeError("get_context returned None") + sock = None try: address = format_zmq_address(server_info.ip, port) @@ -371,14 +387,11 @@ async def wrapper(self, *args, **kwargs): kwargs["socket"] = sock return await func(self, *args, **kwargs) finally: - if sock is not None: - try: - if not sock.closed: - sock.close(linger=-1) - finally: - context.term() - else: - context.term() + # Close the per-call socket only; the context outlives the call. + # linger=0 drops any unsent frames immediately (the reply is already + # received on the happy path) so close never blocks the event loop. + if sock is not None and not sock.closed: + sock.close(linger=0) return wrapper From f27ff24998b4b4b19cabb4f0ff15bf7a5c556feb Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 29 Jul 2026 22:10:36 +0800 Subject: [PATCH 2/6] feat: share fixed ZMQ context pool across client requests --- tests/test_zmq_shared_context.py | 49 +++++++++++++++++- transfer_queue/client.py | 22 ++++++-- transfer_queue/storage/managers/base.py | 51 +++++++++++++++---- .../storage/managers/mooncake_manager.py | 11 +++- .../storage/managers/ray_storage_manager.py | 15 +++++- .../managers/simple_storage_manager.py | 9 +++- .../storage/managers/yuanrong_manager.py | 11 +++- transfer_queue/utils/zmq_utils.py | 22 +++++++- 8 files changed, 163 insertions(+), 27 deletions(-) diff --git a/tests/test_zmq_shared_context.py b/tests/test_zmq_shared_context.py index 37700189..808df2a3 100644 --- a/tests/test_zmq_shared_context.py +++ b/tests/test_zmq_shared_context.py @@ -21,12 +21,13 @@ descriptor -> SIGABRT). The fix makes the decorator reuse the owner's long-lived context (``get_context``) and only create/close the DEALER socket per call. -These tests assert that concurrent decorated calls all reuse the SAME context object and -that the context is never terminated between calls, only when the client is closed. +These tests assert that concurrent decorated calls all reuse the SAME fixed-size context +pool and that the context is never terminated between calls, only when the client is closed. """ import asyncio from threading import Thread +from unittest.mock import patch import pytest import zmq @@ -148,6 +149,50 @@ def _spy_create(ctx, *args, **kwargs): client.close() +def test_context_uses_configured_fixed_io_thread_pool(echo_controller): + """All sockets from a client share the configured native ZMQ I/O-thread pool.""" + client = AsyncTransferQueueClient( + client_id="client_fixed_ctx_pool", + controller_info=echo_controller.zmq_server_info, + zmq_io_threads=4, + ) + + assert client.zmq_context.get(zmq.IO_THREADS) == 4 + + client.close() + + +def test_context_rejects_invalid_io_thread_pool_size(echo_controller): + with pytest.raises(ValueError, match="at least 1"): + AsyncTransferQueueClient( + client_id="client_invalid_ctx_pool", + controller_info=echo_controller.zmq_server_info, + zmq_io_threads=0, + ) + + +def test_client_shares_context_pool_with_storage_manager(echo_controller): + client = AsyncTransferQueueClient( + client_id="client_storage_ctx_pool", + controller_info=echo_controller.zmq_server_info, + zmq_io_threads=4, + ) + config = {"client_name": "unused"} + + with patch("transfer_queue.client.StorageManagerFactory.create") as create_manager: + client.initialize_storage_manager("unused", config) + + create_manager.assert_called_once_with( + "unused", + controller_info=echo_controller.zmq_server_info, + config=config, + zmq_context=client.zmq_context, + ) + assert config == {"client_name": "unused"} + + client.close() + + @pytest.mark.asyncio async def test_close_destroys_context(echo_controller): """close() must terminate the shared context exactly once (no leak, no hang).""" diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 6b088785..fda39a0b 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -31,6 +31,7 @@ ZMQMessage, ZMQRequestType, ZMQServerInfo, + create_zmq_context, with_zmq_socket, ) @@ -58,12 +59,15 @@ def __init__( self, client_id: str, controller_info: ZMQServerInfo, + zmq_io_threads: int | None = None, ): """Initialize the asynchronous TransferQueue client. Args: client_id: Unique identifier for this client instance controller_info: Single controller ZMQ server information + zmq_io_threads: Size of the long-lived ZMQ context's native I/O-thread + pool. Defaults to ``TQ_ZMQ_IO_THREADS`` (8). """ if controller_info is None: raise ValueError("controller_info cannot be None") @@ -71,10 +75,11 @@ def __init__( raise TypeError(f"controller_info must be ZMQServerInfo, got {type(controller_info)}") self.client_id = client_id self._controller: ZMQServerInfo = controller_info - # Long-lived ZMQ context shared by all controller RPCs on this client. Contexts are - # thread-safe and event-loop-agnostic; each RPC creates and closes its own DEALER - # socket from this context (see with_controller_socket). Terminated once in close(). - self.zmq_context = zmq.asyncio.Context() + self._zmq_io_threads = zmq_io_threads + # One long-lived ZMQ context per client. Its fixed native I/O-thread pool is shared + # by all concurrent RPC sockets; sockets remain per-request because ZMQ sockets are + # not thread-safe. The context is terminated once in close(). + self.zmq_context = create_zmq_context(zmq_io_threads) logger.info(f"[{self.client_id}]: Registered Controller server {controller_info.id} at {controller_info.ip}") def initialize_storage_manager( @@ -93,7 +98,10 @@ def initialize_storage_manager( """ self.storage_manager = StorageManagerFactory.create( - manager_type, controller_info=self._controller, config=config + manager_type, + controller_info=self._controller, + config=config, + zmq_context=self.zmq_context, ) # ==================== Basic API ==================== @@ -1243,16 +1251,20 @@ def __init__( self, client_id: str, controller_info: ZMQServerInfo, + zmq_io_threads: int | None = None, ): """Initialize the synchronous TransferQueue client. Args: client_id: Unique identifier for this client instance controller_info: Single controller ZMQ server information + zmq_io_threads: Size of the long-lived ZMQ context's native I/O-thread + pool. Defaults to ``TQ_ZMQ_IO_THREADS`` (8). """ super().__init__( client_id, controller_info, + zmq_io_threads, ) # create new event loop in a separate thread diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index c6c959c4..a0effa6f 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -14,6 +14,7 @@ # limitations under the License. import asyncio +import inspect import itertools import os import threading @@ -35,7 +36,13 @@ from transfer_queue.metadata import BatchMeta, extract_field_schema from transfer_queue.storage.clients.base import StorageClientFactory from transfer_queue.utils.logging_utils import get_logger -from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType, ZMQServerInfo, create_zmq_socket +from transfer_queue.utils.zmq_utils import ( + ZMQMessage, + ZMQRequestType, + ZMQServerInfo, + create_zmq_context, + create_zmq_socket, +) logger = get_logger(__name__) @@ -63,7 +70,12 @@ class StorageManager(ABC): """Base class for storage layer. It defines the interface for data operations and generally provides handshake & notification capabilities.""" - def __init__(self, controller_info: ZMQServerInfo, config: DictConfig): + def __init__( + self, + controller_info: ZMQServerInfo, + config: DictConfig, + zmq_context: zmq.asyncio.Context | None = None, + ): self.storage_manager_id = f"TQ_STORAGE_{uuid4().hex[:8]}" self.config = config self.controller_info = controller_info @@ -71,7 +83,11 @@ def __init__(self, controller_info: ZMQServerInfo, config: DictConfig): # Handshake socket is sync (used only during initialization) self.controller_handshake_socket: zmq.Socket | None = None - self.zmq_context = zmq.asyncio.Context() + # A manager created by TransferQueueClient borrows the client's context so + # controller and storage requests share one fixed native I/O-thread pool. + # Standalone managers create and own an equivalent long-lived context. + self._owns_zmq_context = zmq_context is None + self.zmq_context = zmq_context or create_zmq_context(config.get("zmq_io_threads", None)) self._connect_to_controller() # Dedicated asyncio loop for ZMQ notify traffic, isolated from the caller's loop @@ -391,9 +407,10 @@ def close(self) -> None: else: logger.debug(f"[{self.storage_manager_id}]: Notify ZMQ thread shut down.") - # destroy(linger=0) force-closes any socket still open (e.g. from an interrupted - # request or the notify path) then terminates, so shutdown cannot hang on term(). - self.zmq_context.destroy(linger=0) + if self._owns_zmq_context: + # destroy(linger=0) force-closes any socket still open (e.g. from an interrupted + # request or the notify path) then terminates, so shutdown cannot hang on term(). + self.zmq_context.destroy(linger=0) def __del__(self): """Destructor to ensure resources are cleaned up.""" @@ -424,12 +441,21 @@ def decorator(manager_cls: type[StorageManager]): return decorator @classmethod - def create(cls, manager_type: str, controller_info: ZMQServerInfo, config: dict[str, Any]) -> StorageManager: + def create( + cls, + manager_type: str, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ) -> StorageManager: """Create and return a StorageManager instance.""" assert manager_type in cls._registry, ( f"Unknown manager_type: {manager_type}. Supported managers include: {list(cls._registry.keys())}" ) - return cls._registry[manager_type](controller_info, config) + manager_cls = cls._registry[manager_type] + if zmq_context is not None and "zmq_context" in inspect.signature(manager_cls).parameters: + return manager_cls(controller_info, config, zmq_context=zmq_context) + return manager_cls(controller_info, config) class KVStorageManager(StorageManager): @@ -438,14 +464,19 @@ class KVStorageManager(StorageManager): It maps structured metadata (BatchMeta) to flat lists of keys and values for efficient KV operations. """ - def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): + def __init__( + self, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ): """ Initialize the KVStorageManager with configuration. """ client_name = config.get("client_name", None) if client_name is None: raise ValueError("Missing client_name in config") - super().__init__(controller_info, config) + super().__init__(controller_info, config, zmq_context=zmq_context) self.storage_client = StorageClientFactory.create(client_name, config) self._multi_threads_executor: ThreadPoolExecutor | None = None self._executor_finalizer = weakref.finalize(self, self._shutdown_executor, self._multi_threads_executor) diff --git a/transfer_queue/storage/managers/mooncake_manager.py b/transfer_queue/storage/managers/mooncake_manager.py index c3e8f5ce..a929d6b7 100644 --- a/transfer_queue/storage/managers/mooncake_manager.py +++ b/transfer_queue/storage/managers/mooncake_manager.py @@ -15,6 +15,8 @@ from typing import Any +import zmq.asyncio + from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -29,6 +31,11 @@ class MooncakeStorageManager(KVStorageManager): pybind bindings. """ - def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): + def __init__( + self, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ): config["client_name"] = "MooncakeStoreClient" - super().__init__(controller_info, config) + super().__init__(controller_info, config, zmq_context=zmq_context) diff --git a/transfer_queue/storage/managers/ray_storage_manager.py b/transfer_queue/storage/managers/ray_storage_manager.py index 0cc2a09c..f91f008a 100644 --- a/transfer_queue/storage/managers/ray_storage_manager.py +++ b/transfer_queue/storage/managers/ray_storage_manager.py @@ -15,6 +15,8 @@ from typing import Any +import zmq.asyncio + from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -23,8 +25,17 @@ class RayStorageManager(KVStorageManager): """Storage manager for Ray-RDT backend.""" - def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): + def __init__( + self, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ): config = (config or {}).copy() if config.get("client_name") not in (None, "RayStorageClient"): raise ValueError(f"RayStorageManager only supports 'RayStorageClient', got: {config.get('client_name')}") - super().__init__(controller_info, {**config, "client_name": "RayStorageClient"}) + super().__init__( + controller_info, + {**config, "client_name": "RayStorageClient"}, + zmq_context=zmq_context, + ) diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index c6b20eff..aec260b6 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -73,8 +73,13 @@ class AsyncSimpleStorageManager(StorageManager): instances using ZMQ communication and dynamic socket management. """ - def __init__(self, controller_info: ZMQServerInfo, config: DictConfig): - super().__init__(controller_info, config) + def __init__( + self, + controller_info: ZMQServerInfo, + config: DictConfig, + zmq_context: zmq.asyncio.Context | None = None, + ): + super().__init__(controller_info, config, zmq_context=zmq_context) self.config = config server_infos: ZMQServerInfo | dict[str, ZMQServerInfo] | None = config.get("zmq_info", None) diff --git a/transfer_queue/storage/managers/yuanrong_manager.py b/transfer_queue/storage/managers/yuanrong_manager.py index f76b47b2..a409cb66 100644 --- a/transfer_queue/storage/managers/yuanrong_manager.py +++ b/transfer_queue/storage/managers/yuanrong_manager.py @@ -15,6 +15,8 @@ from typing import Any +import zmq.asyncio + from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -26,7 +28,12 @@ class YuanrongStorageManager(KVStorageManager): """Storage manager for Yuanrong backend.""" - def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): + def __init__( + self, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ): worker_port = config.get("worker_port", None) client_name = config.get("client_name", None) @@ -38,4 +45,4 @@ def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): config["client_name"] = "YuanrongStorageClient" elif client_name != "YuanrongStorageClient": raise ValueError(f"Invalid 'client_name': {client_name} in config. Expecting 'YuanrongStorageClient'") - super().__init__(controller_info, config) + super().__init__(controller_info, config, zmq_context=zmq_context) diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index 147340e4..a43aa7b5 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -13,6 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import os import socket import time from dataclasses import dataclass @@ -265,6 +266,22 @@ def get_free_port(ip: str) -> int: return sock.getsockname()[1] +TQ_ZMQ_IO_THREADS = int(os.environ.get("TQ_ZMQ_IO_THREADS", 8)) + + +def create_zmq_context(io_threads: int | None = None) -> "zmq.asyncio.Context": + """Create a long-lived async ZMQ context with a fixed I/O-thread pool. + + A ZMQ context owns the native I/O-thread pool used by all sockets created from + that context. Keeping one context per owner lets concurrent request sockets share + the whole pool without creating or terminating contexts per request. + """ + pool_size = TQ_ZMQ_IO_THREADS if io_threads is None else io_threads + if pool_size < 1: + raise ValueError(f"ZMQ I/O thread pool size must be at least 1, got {pool_size}") + return zmq.asyncio.Context(io_threads=pool_size) + + def create_zmq_socket( ctx: zmq.Context, socket_type: Any, @@ -332,8 +349,9 @@ def with_zmq_socket( terminate a context here -- per-call context churn corrupts libzmq's signaler file descriptors under concurrency (Bad file descriptor / SIGABRT) and can hang on term(). Contexts are thread-safe and event-loop-agnostic, so a single shared context is safe - even when decorated methods run on different loops/threads; each socket is created and - fully used within one awaited call on one loop. + even when decorated methods run on different loops/threads. The context's fixed native + I/O-thread pool is shared by all request sockets; each socket is created and fully used + within one awaited call on one loop. Args: socket_name: Socket port key in ``ZMQServerInfo.ports``. From c2a83106359589fa9f11c7ffab0800572f4d9b0a Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 29 Jul 2026 22:16:43 +0800 Subject: [PATCH 3/6] fix: limit shared ZMQ context to SimpleStorage --- tests/test_zmq_shared_context.py | 26 ++++++++++++++++--- transfer_queue/client.py | 5 +++- transfer_queue/storage/managers/base.py | 20 +++++--------- .../storage/managers/mooncake_manager.py | 11 ++------ .../storage/managers/ray_storage_manager.py | 15 ++--------- .../storage/managers/yuanrong_manager.py | 11 ++------ 6 files changed, 39 insertions(+), 49 deletions(-) diff --git a/tests/test_zmq_shared_context.py b/tests/test_zmq_shared_context.py index 808df2a3..ce0f5b62 100644 --- a/tests/test_zmq_shared_context.py +++ b/tests/test_zmq_shared_context.py @@ -171,7 +171,7 @@ def test_context_rejects_invalid_io_thread_pool_size(echo_controller): ) -def test_client_shares_context_pool_with_storage_manager(echo_controller): +def test_client_shares_context_pool_with_simple_storage_manager(echo_controller): client = AsyncTransferQueueClient( client_id="client_storage_ctx_pool", controller_info=echo_controller.zmq_server_info, @@ -180,10 +180,10 @@ def test_client_shares_context_pool_with_storage_manager(echo_controller): config = {"client_name": "unused"} with patch("transfer_queue.client.StorageManagerFactory.create") as create_manager: - client.initialize_storage_manager("unused", config) + client.initialize_storage_manager("SimpleStorage", config) create_manager.assert_called_once_with( - "unused", + "SimpleStorage", controller_info=echo_controller.zmq_server_info, config=config, zmq_context=client.zmq_context, @@ -193,6 +193,26 @@ def test_client_shares_context_pool_with_storage_manager(echo_controller): client.close() +def test_client_does_not_pass_context_to_other_storage_backends(echo_controller): + client = AsyncTransferQueueClient( + client_id="client_other_storage_ctx", + controller_info=echo_controller.zmq_server_info, + zmq_io_threads=4, + ) + config = {"client_name": "unused"} + + with patch("transfer_queue.client.StorageManagerFactory.create") as create_manager: + client.initialize_storage_manager("OtherStorage", config) + + create_manager.assert_called_once_with( + "OtherStorage", + controller_info=echo_controller.zmq_server_info, + config=config, + ) + + client.close() + + @pytest.mark.asyncio async def test_close_destroys_context(echo_controller): """close() must terminate the shared context exactly once (no leak, no hang).""" diff --git a/transfer_queue/client.py b/transfer_queue/client.py index fda39a0b..26ccd0b5 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -97,11 +97,14 @@ def initialize_storage_manager( - zmq_info: ZMQ server information about the storage units """ + create_kwargs = {} + if manager_type == "SimpleStorage": + create_kwargs["zmq_context"] = self.zmq_context self.storage_manager = StorageManagerFactory.create( manager_type, controller_info=self._controller, config=config, - zmq_context=self.zmq_context, + **create_kwargs, ) # ==================== Basic API ==================== diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index a0effa6f..6c683670 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -14,7 +14,6 @@ # limitations under the License. import asyncio -import inspect import itertools import os import threading @@ -40,7 +39,6 @@ ZMQMessage, ZMQRequestType, ZMQServerInfo, - create_zmq_context, create_zmq_socket, ) @@ -83,11 +81,10 @@ def __init__( # Handshake socket is sync (used only during initialization) self.controller_handshake_socket: zmq.Socket | None = None - # A manager created by TransferQueueClient borrows the client's context so - # controller and storage requests share one fixed native I/O-thread pool. - # Standalone managers create and own an equivalent long-lived context. + # SimpleStorage can borrow the client's long-lived context. Other storage + # backends retain the original behavior and own their default ZMQ context. self._owns_zmq_context = zmq_context is None - self.zmq_context = zmq_context or create_zmq_context(config.get("zmq_io_threads", None)) + self.zmq_context = zmq_context or zmq.asyncio.Context() self._connect_to_controller() # Dedicated asyncio loop for ZMQ notify traffic, isolated from the caller's loop @@ -453,7 +450,7 @@ def create( f"Unknown manager_type: {manager_type}. Supported managers include: {list(cls._registry.keys())}" ) manager_cls = cls._registry[manager_type] - if zmq_context is not None and "zmq_context" in inspect.signature(manager_cls).parameters: + if manager_type == "SimpleStorage" and zmq_context is not None: return manager_cls(controller_info, config, zmq_context=zmq_context) return manager_cls(controller_info, config) @@ -464,19 +461,14 @@ class KVStorageManager(StorageManager): It maps structured metadata (BatchMeta) to flat lists of keys and values for efficient KV operations. """ - def __init__( - self, - controller_info: ZMQServerInfo, - config: dict[str, Any], - zmq_context: zmq.asyncio.Context | None = None, - ): + def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): """ Initialize the KVStorageManager with configuration. """ client_name = config.get("client_name", None) if client_name is None: raise ValueError("Missing client_name in config") - super().__init__(controller_info, config, zmq_context=zmq_context) + super().__init__(controller_info, config) self.storage_client = StorageClientFactory.create(client_name, config) self._multi_threads_executor: ThreadPoolExecutor | None = None self._executor_finalizer = weakref.finalize(self, self._shutdown_executor, self._multi_threads_executor) diff --git a/transfer_queue/storage/managers/mooncake_manager.py b/transfer_queue/storage/managers/mooncake_manager.py index a929d6b7..c3e8f5ce 100644 --- a/transfer_queue/storage/managers/mooncake_manager.py +++ b/transfer_queue/storage/managers/mooncake_manager.py @@ -15,8 +15,6 @@ from typing import Any -import zmq.asyncio - from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -31,11 +29,6 @@ class MooncakeStorageManager(KVStorageManager): pybind bindings. """ - def __init__( - self, - controller_info: ZMQServerInfo, - config: dict[str, Any], - zmq_context: zmq.asyncio.Context | None = None, - ): + def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): config["client_name"] = "MooncakeStoreClient" - super().__init__(controller_info, config, zmq_context=zmq_context) + super().__init__(controller_info, config) diff --git a/transfer_queue/storage/managers/ray_storage_manager.py b/transfer_queue/storage/managers/ray_storage_manager.py index f91f008a..0cc2a09c 100644 --- a/transfer_queue/storage/managers/ray_storage_manager.py +++ b/transfer_queue/storage/managers/ray_storage_manager.py @@ -15,8 +15,6 @@ from typing import Any -import zmq.asyncio - from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -25,17 +23,8 @@ class RayStorageManager(KVStorageManager): """Storage manager for Ray-RDT backend.""" - def __init__( - self, - controller_info: ZMQServerInfo, - config: dict[str, Any], - zmq_context: zmq.asyncio.Context | None = None, - ): + def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): config = (config or {}).copy() if config.get("client_name") not in (None, "RayStorageClient"): raise ValueError(f"RayStorageManager only supports 'RayStorageClient', got: {config.get('client_name')}") - super().__init__( - controller_info, - {**config, "client_name": "RayStorageClient"}, - zmq_context=zmq_context, - ) + super().__init__(controller_info, {**config, "client_name": "RayStorageClient"}) diff --git a/transfer_queue/storage/managers/yuanrong_manager.py b/transfer_queue/storage/managers/yuanrong_manager.py index a409cb66..f76b47b2 100644 --- a/transfer_queue/storage/managers/yuanrong_manager.py +++ b/transfer_queue/storage/managers/yuanrong_manager.py @@ -15,8 +15,6 @@ from typing import Any -import zmq.asyncio - from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -28,12 +26,7 @@ class YuanrongStorageManager(KVStorageManager): """Storage manager for Yuanrong backend.""" - def __init__( - self, - controller_info: ZMQServerInfo, - config: dict[str, Any], - zmq_context: zmq.asyncio.Context | None = None, - ): + def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): worker_port = config.get("worker_port", None) client_name = config.get("client_name", None) @@ -45,4 +38,4 @@ def __init__( config["client_name"] = "YuanrongStorageClient" elif client_name != "YuanrongStorageClient": raise ValueError(f"Invalid 'client_name': {client_name} in config. Expecting 'YuanrongStorageClient'") - super().__init__(controller_info, config, zmq_context=zmq_context) + super().__init__(controller_info, config) From 5798f36dca55c4f4ac671ec86f92c47f868480b0 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 29 Jul 2026 22:28:45 +0800 Subject: [PATCH 4/6] Revert "fix: limit shared ZMQ context to SimpleStorage" This reverts commit b54c4093d71348b3c82572ff18c3a35a0bb9dd2b. --- tests/test_zmq_shared_context.py | 26 +++---------------- transfer_queue/client.py | 5 +--- transfer_queue/storage/managers/base.py | 20 +++++++++----- .../storage/managers/mooncake_manager.py | 11 ++++++-- .../storage/managers/ray_storage_manager.py | 15 +++++++++-- .../storage/managers/yuanrong_manager.py | 11 ++++++-- 6 files changed, 49 insertions(+), 39 deletions(-) diff --git a/tests/test_zmq_shared_context.py b/tests/test_zmq_shared_context.py index ce0f5b62..808df2a3 100644 --- a/tests/test_zmq_shared_context.py +++ b/tests/test_zmq_shared_context.py @@ -171,7 +171,7 @@ def test_context_rejects_invalid_io_thread_pool_size(echo_controller): ) -def test_client_shares_context_pool_with_simple_storage_manager(echo_controller): +def test_client_shares_context_pool_with_storage_manager(echo_controller): client = AsyncTransferQueueClient( client_id="client_storage_ctx_pool", controller_info=echo_controller.zmq_server_info, @@ -180,10 +180,10 @@ def test_client_shares_context_pool_with_simple_storage_manager(echo_controller) config = {"client_name": "unused"} with patch("transfer_queue.client.StorageManagerFactory.create") as create_manager: - client.initialize_storage_manager("SimpleStorage", config) + client.initialize_storage_manager("unused", config) create_manager.assert_called_once_with( - "SimpleStorage", + "unused", controller_info=echo_controller.zmq_server_info, config=config, zmq_context=client.zmq_context, @@ -193,26 +193,6 @@ def test_client_shares_context_pool_with_simple_storage_manager(echo_controller) client.close() -def test_client_does_not_pass_context_to_other_storage_backends(echo_controller): - client = AsyncTransferQueueClient( - client_id="client_other_storage_ctx", - controller_info=echo_controller.zmq_server_info, - zmq_io_threads=4, - ) - config = {"client_name": "unused"} - - with patch("transfer_queue.client.StorageManagerFactory.create") as create_manager: - client.initialize_storage_manager("OtherStorage", config) - - create_manager.assert_called_once_with( - "OtherStorage", - controller_info=echo_controller.zmq_server_info, - config=config, - ) - - client.close() - - @pytest.mark.asyncio async def test_close_destroys_context(echo_controller): """close() must terminate the shared context exactly once (no leak, no hang).""" diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 26ccd0b5..fda39a0b 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -97,14 +97,11 @@ def initialize_storage_manager( - zmq_info: ZMQ server information about the storage units """ - create_kwargs = {} - if manager_type == "SimpleStorage": - create_kwargs["zmq_context"] = self.zmq_context self.storage_manager = StorageManagerFactory.create( manager_type, controller_info=self._controller, config=config, - **create_kwargs, + zmq_context=self.zmq_context, ) # ==================== Basic API ==================== diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index 6c683670..a0effa6f 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -14,6 +14,7 @@ # limitations under the License. import asyncio +import inspect import itertools import os import threading @@ -39,6 +40,7 @@ ZMQMessage, ZMQRequestType, ZMQServerInfo, + create_zmq_context, create_zmq_socket, ) @@ -81,10 +83,11 @@ def __init__( # Handshake socket is sync (used only during initialization) self.controller_handshake_socket: zmq.Socket | None = None - # SimpleStorage can borrow the client's long-lived context. Other storage - # backends retain the original behavior and own their default ZMQ context. + # A manager created by TransferQueueClient borrows the client's context so + # controller and storage requests share one fixed native I/O-thread pool. + # Standalone managers create and own an equivalent long-lived context. self._owns_zmq_context = zmq_context is None - self.zmq_context = zmq_context or zmq.asyncio.Context() + self.zmq_context = zmq_context or create_zmq_context(config.get("zmq_io_threads", None)) self._connect_to_controller() # Dedicated asyncio loop for ZMQ notify traffic, isolated from the caller's loop @@ -450,7 +453,7 @@ def create( f"Unknown manager_type: {manager_type}. Supported managers include: {list(cls._registry.keys())}" ) manager_cls = cls._registry[manager_type] - if manager_type == "SimpleStorage" and zmq_context is not None: + if zmq_context is not None and "zmq_context" in inspect.signature(manager_cls).parameters: return manager_cls(controller_info, config, zmq_context=zmq_context) return manager_cls(controller_info, config) @@ -461,14 +464,19 @@ class KVStorageManager(StorageManager): It maps structured metadata (BatchMeta) to flat lists of keys and values for efficient KV operations. """ - def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): + def __init__( + self, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ): """ Initialize the KVStorageManager with configuration. """ client_name = config.get("client_name", None) if client_name is None: raise ValueError("Missing client_name in config") - super().__init__(controller_info, config) + super().__init__(controller_info, config, zmq_context=zmq_context) self.storage_client = StorageClientFactory.create(client_name, config) self._multi_threads_executor: ThreadPoolExecutor | None = None self._executor_finalizer = weakref.finalize(self, self._shutdown_executor, self._multi_threads_executor) diff --git a/transfer_queue/storage/managers/mooncake_manager.py b/transfer_queue/storage/managers/mooncake_manager.py index c3e8f5ce..a929d6b7 100644 --- a/transfer_queue/storage/managers/mooncake_manager.py +++ b/transfer_queue/storage/managers/mooncake_manager.py @@ -15,6 +15,8 @@ from typing import Any +import zmq.asyncio + from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -29,6 +31,11 @@ class MooncakeStorageManager(KVStorageManager): pybind bindings. """ - def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): + def __init__( + self, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ): config["client_name"] = "MooncakeStoreClient" - super().__init__(controller_info, config) + super().__init__(controller_info, config, zmq_context=zmq_context) diff --git a/transfer_queue/storage/managers/ray_storage_manager.py b/transfer_queue/storage/managers/ray_storage_manager.py index 0cc2a09c..f91f008a 100644 --- a/transfer_queue/storage/managers/ray_storage_manager.py +++ b/transfer_queue/storage/managers/ray_storage_manager.py @@ -15,6 +15,8 @@ from typing import Any +import zmq.asyncio + from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -23,8 +25,17 @@ class RayStorageManager(KVStorageManager): """Storage manager for Ray-RDT backend.""" - def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): + def __init__( + self, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ): config = (config or {}).copy() if config.get("client_name") not in (None, "RayStorageClient"): raise ValueError(f"RayStorageManager only supports 'RayStorageClient', got: {config.get('client_name')}") - super().__init__(controller_info, {**config, "client_name": "RayStorageClient"}) + super().__init__( + controller_info, + {**config, "client_name": "RayStorageClient"}, + zmq_context=zmq_context, + ) diff --git a/transfer_queue/storage/managers/yuanrong_manager.py b/transfer_queue/storage/managers/yuanrong_manager.py index f76b47b2..a409cb66 100644 --- a/transfer_queue/storage/managers/yuanrong_manager.py +++ b/transfer_queue/storage/managers/yuanrong_manager.py @@ -15,6 +15,8 @@ from typing import Any +import zmq.asyncio + from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -26,7 +28,12 @@ class YuanrongStorageManager(KVStorageManager): """Storage manager for Yuanrong backend.""" - def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): + def __init__( + self, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ): worker_port = config.get("worker_port", None) client_name = config.get("client_name", None) @@ -38,4 +45,4 @@ def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): config["client_name"] = "YuanrongStorageClient" elif client_name != "YuanrongStorageClient": raise ValueError(f"Invalid 'client_name': {client_name} in config. Expecting 'YuanrongStorageClient'") - super().__init__(controller_info, config) + super().__init__(controller_info, config, zmq_context=zmq_context) From 0c58d1293aa2622f4b241a9a399b37517bcab769 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 29 Jul 2026 22:31:51 +0800 Subject: [PATCH 5/6] Revert "feat: share fixed ZMQ context pool across client requests" This reverts commit a8bfbd81c68226f0c679ce3673c467866d470881. --- tests/test_zmq_shared_context.py | 49 +----------------- transfer_queue/client.py | 22 ++------ transfer_queue/storage/managers/base.py | 51 ++++--------------- .../storage/managers/mooncake_manager.py | 11 +--- .../storage/managers/ray_storage_manager.py | 15 +----- .../managers/simple_storage_manager.py | 9 +--- .../storage/managers/yuanrong_manager.py | 11 +--- transfer_queue/utils/zmq_utils.py | 22 +------- 8 files changed, 27 insertions(+), 163 deletions(-) diff --git a/tests/test_zmq_shared_context.py b/tests/test_zmq_shared_context.py index 808df2a3..37700189 100644 --- a/tests/test_zmq_shared_context.py +++ b/tests/test_zmq_shared_context.py @@ -21,13 +21,12 @@ descriptor -> SIGABRT). The fix makes the decorator reuse the owner's long-lived context (``get_context``) and only create/close the DEALER socket per call. -These tests assert that concurrent decorated calls all reuse the SAME fixed-size context -pool and that the context is never terminated between calls, only when the client is closed. +These tests assert that concurrent decorated calls all reuse the SAME context object and +that the context is never terminated between calls, only when the client is closed. """ import asyncio from threading import Thread -from unittest.mock import patch import pytest import zmq @@ -149,50 +148,6 @@ def _spy_create(ctx, *args, **kwargs): client.close() -def test_context_uses_configured_fixed_io_thread_pool(echo_controller): - """All sockets from a client share the configured native ZMQ I/O-thread pool.""" - client = AsyncTransferQueueClient( - client_id="client_fixed_ctx_pool", - controller_info=echo_controller.zmq_server_info, - zmq_io_threads=4, - ) - - assert client.zmq_context.get(zmq.IO_THREADS) == 4 - - client.close() - - -def test_context_rejects_invalid_io_thread_pool_size(echo_controller): - with pytest.raises(ValueError, match="at least 1"): - AsyncTransferQueueClient( - client_id="client_invalid_ctx_pool", - controller_info=echo_controller.zmq_server_info, - zmq_io_threads=0, - ) - - -def test_client_shares_context_pool_with_storage_manager(echo_controller): - client = AsyncTransferQueueClient( - client_id="client_storage_ctx_pool", - controller_info=echo_controller.zmq_server_info, - zmq_io_threads=4, - ) - config = {"client_name": "unused"} - - with patch("transfer_queue.client.StorageManagerFactory.create") as create_manager: - client.initialize_storage_manager("unused", config) - - create_manager.assert_called_once_with( - "unused", - controller_info=echo_controller.zmq_server_info, - config=config, - zmq_context=client.zmq_context, - ) - assert config == {"client_name": "unused"} - - client.close() - - @pytest.mark.asyncio async def test_close_destroys_context(echo_controller): """close() must terminate the shared context exactly once (no leak, no hang).""" diff --git a/transfer_queue/client.py b/transfer_queue/client.py index fda39a0b..6b088785 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -31,7 +31,6 @@ ZMQMessage, ZMQRequestType, ZMQServerInfo, - create_zmq_context, with_zmq_socket, ) @@ -59,15 +58,12 @@ def __init__( self, client_id: str, controller_info: ZMQServerInfo, - zmq_io_threads: int | None = None, ): """Initialize the asynchronous TransferQueue client. Args: client_id: Unique identifier for this client instance controller_info: Single controller ZMQ server information - zmq_io_threads: Size of the long-lived ZMQ context's native I/O-thread - pool. Defaults to ``TQ_ZMQ_IO_THREADS`` (8). """ if controller_info is None: raise ValueError("controller_info cannot be None") @@ -75,11 +71,10 @@ def __init__( raise TypeError(f"controller_info must be ZMQServerInfo, got {type(controller_info)}") self.client_id = client_id self._controller: ZMQServerInfo = controller_info - self._zmq_io_threads = zmq_io_threads - # One long-lived ZMQ context per client. Its fixed native I/O-thread pool is shared - # by all concurrent RPC sockets; sockets remain per-request because ZMQ sockets are - # not thread-safe. The context is terminated once in close(). - self.zmq_context = create_zmq_context(zmq_io_threads) + # Long-lived ZMQ context shared by all controller RPCs on this client. Contexts are + # thread-safe and event-loop-agnostic; each RPC creates and closes its own DEALER + # socket from this context (see with_controller_socket). Terminated once in close(). + self.zmq_context = zmq.asyncio.Context() logger.info(f"[{self.client_id}]: Registered Controller server {controller_info.id} at {controller_info.ip}") def initialize_storage_manager( @@ -98,10 +93,7 @@ def initialize_storage_manager( """ self.storage_manager = StorageManagerFactory.create( - manager_type, - controller_info=self._controller, - config=config, - zmq_context=self.zmq_context, + manager_type, controller_info=self._controller, config=config ) # ==================== Basic API ==================== @@ -1251,20 +1243,16 @@ def __init__( self, client_id: str, controller_info: ZMQServerInfo, - zmq_io_threads: int | None = None, ): """Initialize the synchronous TransferQueue client. Args: client_id: Unique identifier for this client instance controller_info: Single controller ZMQ server information - zmq_io_threads: Size of the long-lived ZMQ context's native I/O-thread - pool. Defaults to ``TQ_ZMQ_IO_THREADS`` (8). """ super().__init__( client_id, controller_info, - zmq_io_threads, ) # create new event loop in a separate thread diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index a0effa6f..c6c959c4 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -14,7 +14,6 @@ # limitations under the License. import asyncio -import inspect import itertools import os import threading @@ -36,13 +35,7 @@ from transfer_queue.metadata import BatchMeta, extract_field_schema from transfer_queue.storage.clients.base import StorageClientFactory from transfer_queue.utils.logging_utils import get_logger -from transfer_queue.utils.zmq_utils import ( - ZMQMessage, - ZMQRequestType, - ZMQServerInfo, - create_zmq_context, - create_zmq_socket, -) +from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType, ZMQServerInfo, create_zmq_socket logger = get_logger(__name__) @@ -70,12 +63,7 @@ class StorageManager(ABC): """Base class for storage layer. It defines the interface for data operations and generally provides handshake & notification capabilities.""" - def __init__( - self, - controller_info: ZMQServerInfo, - config: DictConfig, - zmq_context: zmq.asyncio.Context | None = None, - ): + def __init__(self, controller_info: ZMQServerInfo, config: DictConfig): self.storage_manager_id = f"TQ_STORAGE_{uuid4().hex[:8]}" self.config = config self.controller_info = controller_info @@ -83,11 +71,7 @@ def __init__( # Handshake socket is sync (used only during initialization) self.controller_handshake_socket: zmq.Socket | None = None - # A manager created by TransferQueueClient borrows the client's context so - # controller and storage requests share one fixed native I/O-thread pool. - # Standalone managers create and own an equivalent long-lived context. - self._owns_zmq_context = zmq_context is None - self.zmq_context = zmq_context or create_zmq_context(config.get("zmq_io_threads", None)) + self.zmq_context = zmq.asyncio.Context() self._connect_to_controller() # Dedicated asyncio loop for ZMQ notify traffic, isolated from the caller's loop @@ -407,10 +391,9 @@ def close(self) -> None: else: logger.debug(f"[{self.storage_manager_id}]: Notify ZMQ thread shut down.") - if self._owns_zmq_context: - # destroy(linger=0) force-closes any socket still open (e.g. from an interrupted - # request or the notify path) then terminates, so shutdown cannot hang on term(). - self.zmq_context.destroy(linger=0) + # destroy(linger=0) force-closes any socket still open (e.g. from an interrupted + # request or the notify path) then terminates, so shutdown cannot hang on term(). + self.zmq_context.destroy(linger=0) def __del__(self): """Destructor to ensure resources are cleaned up.""" @@ -441,21 +424,12 @@ def decorator(manager_cls: type[StorageManager]): return decorator @classmethod - def create( - cls, - manager_type: str, - controller_info: ZMQServerInfo, - config: dict[str, Any], - zmq_context: zmq.asyncio.Context | None = None, - ) -> StorageManager: + def create(cls, manager_type: str, controller_info: ZMQServerInfo, config: dict[str, Any]) -> StorageManager: """Create and return a StorageManager instance.""" assert manager_type in cls._registry, ( f"Unknown manager_type: {manager_type}. Supported managers include: {list(cls._registry.keys())}" ) - manager_cls = cls._registry[manager_type] - if zmq_context is not None and "zmq_context" in inspect.signature(manager_cls).parameters: - return manager_cls(controller_info, config, zmq_context=zmq_context) - return manager_cls(controller_info, config) + return cls._registry[manager_type](controller_info, config) class KVStorageManager(StorageManager): @@ -464,19 +438,14 @@ class KVStorageManager(StorageManager): It maps structured metadata (BatchMeta) to flat lists of keys and values for efficient KV operations. """ - def __init__( - self, - controller_info: ZMQServerInfo, - config: dict[str, Any], - zmq_context: zmq.asyncio.Context | None = None, - ): + def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): """ Initialize the KVStorageManager with configuration. """ client_name = config.get("client_name", None) if client_name is None: raise ValueError("Missing client_name in config") - super().__init__(controller_info, config, zmq_context=zmq_context) + super().__init__(controller_info, config) self.storage_client = StorageClientFactory.create(client_name, config) self._multi_threads_executor: ThreadPoolExecutor | None = None self._executor_finalizer = weakref.finalize(self, self._shutdown_executor, self._multi_threads_executor) diff --git a/transfer_queue/storage/managers/mooncake_manager.py b/transfer_queue/storage/managers/mooncake_manager.py index a929d6b7..c3e8f5ce 100644 --- a/transfer_queue/storage/managers/mooncake_manager.py +++ b/transfer_queue/storage/managers/mooncake_manager.py @@ -15,8 +15,6 @@ from typing import Any -import zmq.asyncio - from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -31,11 +29,6 @@ class MooncakeStorageManager(KVStorageManager): pybind bindings. """ - def __init__( - self, - controller_info: ZMQServerInfo, - config: dict[str, Any], - zmq_context: zmq.asyncio.Context | None = None, - ): + def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): config["client_name"] = "MooncakeStoreClient" - super().__init__(controller_info, config, zmq_context=zmq_context) + super().__init__(controller_info, config) diff --git a/transfer_queue/storage/managers/ray_storage_manager.py b/transfer_queue/storage/managers/ray_storage_manager.py index f91f008a..0cc2a09c 100644 --- a/transfer_queue/storage/managers/ray_storage_manager.py +++ b/transfer_queue/storage/managers/ray_storage_manager.py @@ -15,8 +15,6 @@ from typing import Any -import zmq.asyncio - from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -25,17 +23,8 @@ class RayStorageManager(KVStorageManager): """Storage manager for Ray-RDT backend.""" - def __init__( - self, - controller_info: ZMQServerInfo, - config: dict[str, Any], - zmq_context: zmq.asyncio.Context | None = None, - ): + def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): config = (config or {}).copy() if config.get("client_name") not in (None, "RayStorageClient"): raise ValueError(f"RayStorageManager only supports 'RayStorageClient', got: {config.get('client_name')}") - super().__init__( - controller_info, - {**config, "client_name": "RayStorageClient"}, - zmq_context=zmq_context, - ) + super().__init__(controller_info, {**config, "client_name": "RayStorageClient"}) diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index aec260b6..c6b20eff 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -73,13 +73,8 @@ class AsyncSimpleStorageManager(StorageManager): instances using ZMQ communication and dynamic socket management. """ - def __init__( - self, - controller_info: ZMQServerInfo, - config: DictConfig, - zmq_context: zmq.asyncio.Context | None = None, - ): - super().__init__(controller_info, config, zmq_context=zmq_context) + def __init__(self, controller_info: ZMQServerInfo, config: DictConfig): + super().__init__(controller_info, config) self.config = config server_infos: ZMQServerInfo | dict[str, ZMQServerInfo] | None = config.get("zmq_info", None) diff --git a/transfer_queue/storage/managers/yuanrong_manager.py b/transfer_queue/storage/managers/yuanrong_manager.py index a409cb66..f76b47b2 100644 --- a/transfer_queue/storage/managers/yuanrong_manager.py +++ b/transfer_queue/storage/managers/yuanrong_manager.py @@ -15,8 +15,6 @@ from typing import Any -import zmq.asyncio - from transfer_queue.storage.managers.base import KVStorageManager, StorageManagerFactory from transfer_queue.utils.logging_utils import get_logger from transfer_queue.utils.zmq_utils import ZMQServerInfo @@ -28,12 +26,7 @@ class YuanrongStorageManager(KVStorageManager): """Storage manager for Yuanrong backend.""" - def __init__( - self, - controller_info: ZMQServerInfo, - config: dict[str, Any], - zmq_context: zmq.asyncio.Context | None = None, - ): + def __init__(self, controller_info: ZMQServerInfo, config: dict[str, Any]): worker_port = config.get("worker_port", None) client_name = config.get("client_name", None) @@ -45,4 +38,4 @@ def __init__( config["client_name"] = "YuanrongStorageClient" elif client_name != "YuanrongStorageClient": raise ValueError(f"Invalid 'client_name': {client_name} in config. Expecting 'YuanrongStorageClient'") - super().__init__(controller_info, config, zmq_context=zmq_context) + super().__init__(controller_info, config) diff --git a/transfer_queue/utils/zmq_utils.py b/transfer_queue/utils/zmq_utils.py index a43aa7b5..147340e4 100644 --- a/transfer_queue/utils/zmq_utils.py +++ b/transfer_queue/utils/zmq_utils.py @@ -13,7 +13,6 @@ # See the License for the specific language governing permissions and # limitations under the License. -import os import socket import time from dataclasses import dataclass @@ -266,22 +265,6 @@ def get_free_port(ip: str) -> int: return sock.getsockname()[1] -TQ_ZMQ_IO_THREADS = int(os.environ.get("TQ_ZMQ_IO_THREADS", 8)) - - -def create_zmq_context(io_threads: int | None = None) -> "zmq.asyncio.Context": - """Create a long-lived async ZMQ context with a fixed I/O-thread pool. - - A ZMQ context owns the native I/O-thread pool used by all sockets created from - that context. Keeping one context per owner lets concurrent request sockets share - the whole pool without creating or terminating contexts per request. - """ - pool_size = TQ_ZMQ_IO_THREADS if io_threads is None else io_threads - if pool_size < 1: - raise ValueError(f"ZMQ I/O thread pool size must be at least 1, got {pool_size}") - return zmq.asyncio.Context(io_threads=pool_size) - - def create_zmq_socket( ctx: zmq.Context, socket_type: Any, @@ -349,9 +332,8 @@ def with_zmq_socket( terminate a context here -- per-call context churn corrupts libzmq's signaler file descriptors under concurrency (Bad file descriptor / SIGABRT) and can hang on term(). Contexts are thread-safe and event-loop-agnostic, so a single shared context is safe - even when decorated methods run on different loops/threads. The context's fixed native - I/O-thread pool is shared by all request sockets; each socket is created and fully used - within one awaited call on one loop. + even when decorated methods run on different loops/threads; each socket is created and + fully used within one awaited call on one loop. Args: socket_name: Socket port key in ``ZMQServerInfo.ports``. From 114a2590f0370e3c0a261c3b026f449afb2aa097 Mon Sep 17 00:00:00 2001 From: OutstanderWang Date: Wed, 29 Jul 2026 22:38:38 +0800 Subject: [PATCH 6/6] feat: share fixed ZMQ context pool with SimpleStorage --- tests/test_zmq_shared_context.py | 85 +++++++++++++++++++ transfer_queue/client.py | 36 ++++++-- transfer_queue/storage/managers/base.py | 33 +++++-- .../managers/simple_storage_manager.py | 9 +- 4 files changed, 149 insertions(+), 14 deletions(-) diff --git a/tests/test_zmq_shared_context.py b/tests/test_zmq_shared_context.py index 37700189..a4dde94f 100644 --- a/tests/test_zmq_shared_context.py +++ b/tests/test_zmq_shared_context.py @@ -27,6 +27,7 @@ import asyncio from threading import Thread +from unittest.mock import patch import pytest import zmq @@ -34,6 +35,7 @@ import transfer_queue.utils.zmq_utils as zmq_utils from transfer_queue.client import AsyncTransferQueueClient from transfer_queue.metadata import BatchMeta +from transfer_queue.storage.managers.simple_storage_manager import AsyncSimpleStorageManager from transfer_queue.utils.enum_utils import Role from transfer_queue.utils.zmq_utils import ZMQMessage, ZMQRequestType, ZMQServerInfo @@ -148,6 +150,89 @@ def _spy_create(ctx, *args, **kwargs): client.close() +def test_client_context_has_fixed_io_thread_pool(echo_controller): + client = AsyncTransferQueueClient( + client_id="client_fixed_context_pool", + controller_info=echo_controller.zmq_server_info, + simple_storage_zmq_io_threads=4, + ) + + assert client.zmq_context.get(zmq.IO_THREADS) == 4 + + client.close() + + +def test_client_rejects_invalid_context_pool_size(echo_controller): + with pytest.raises(ValueError, match="at least 1"): + AsyncTransferQueueClient( + client_id="client_invalid_context_pool", + controller_info=echo_controller.zmq_server_info, + simple_storage_zmq_io_threads=0, + ) + + +def test_simple_storage_borrows_client_context(echo_controller): + client = AsyncTransferQueueClient( + client_id="client_simple_storage_context", + controller_info=echo_controller.zmq_server_info, + ) + config = {"zmq_info": {}} + + with patch("transfer_queue.client.StorageManagerFactory.create") as create_manager: + client.initialize_storage_manager("SimpleStorage", config) + + create_manager.assert_called_once_with( + "SimpleStorage", + controller_info=echo_controller.zmq_server_info, + config=config, + zmq_context=client.zmq_context, + ) + + client.close() + + +def test_simple_storage_does_not_destroy_borrowed_context(echo_controller): + client = AsyncTransferQueueClient( + client_id="client_borrowed_context_lifecycle", + controller_info=echo_controller.zmq_server_info, + ) + + with patch("transfer_queue.storage.managers.base.StorageManager._connect_to_controller"): + manager = AsyncSimpleStorageManager( + echo_controller.zmq_server_info, + {"zmq_info": {"storage_0": echo_controller.zmq_server_info}}, + zmq_context=client.zmq_context, + ) + + assert manager.zmq_context is client.zmq_context + assert not manager._owns_zmq_context + + manager.close() + assert not client.zmq_context.closed + + client.close() + assert client.zmq_context.closed + + +def test_other_backends_do_not_borrow_client_context(echo_controller): + client = AsyncTransferQueueClient( + client_id="client_other_storage_context", + controller_info=echo_controller.zmq_server_info, + ) + config = {"client_name": "unused"} + + with patch("transfer_queue.client.StorageManagerFactory.create") as create_manager: + client.initialize_storage_manager("OtherStorage", config) + + create_manager.assert_called_once_with( + "OtherStorage", + controller_info=echo_controller.zmq_server_info, + config=config, + ) + + client.close() + + @pytest.mark.asyncio async def test_close_destroys_context(echo_controller): """close() must terminate the shared context exactly once (no leak, no hang).""" diff --git a/transfer_queue/client.py b/transfer_queue/client.py index 6b088785..656331cc 100644 --- a/transfer_queue/client.py +++ b/transfer_queue/client.py @@ -37,6 +37,7 @@ logger = get_logger(__name__) TQ_NUM_THREADS = int(os.environ.get("TQ_NUM_THREADS", 8)) +TQ_SIMPLE_STORAGE_ZMQ_IO_THREADS = int(os.environ.get("TQ_SIMPLE_STORAGE_ZMQ_IO_THREADS", 8)) # Pre-bound decorator for controller socket operations. with_controller_socket = with_zmq_socket( @@ -58,12 +59,16 @@ def __init__( self, client_id: str, controller_info: ZMQServerInfo, + simple_storage_zmq_io_threads: int | None = None, ): """Initialize the asynchronous TransferQueue client. Args: client_id: Unique identifier for this client instance controller_info: Single controller ZMQ server information + simple_storage_zmq_io_threads: Fixed size of the client context's + native I/O-thread pool. Defaults to + ``TQ_SIMPLE_STORAGE_ZMQ_IO_THREADS`` (8). """ if controller_info is None: raise ValueError("controller_info cannot be None") @@ -71,10 +76,20 @@ def __init__( raise TypeError(f"controller_info must be ZMQServerInfo, got {type(controller_info)}") self.client_id = client_id self._controller: ZMQServerInfo = controller_info - # Long-lived ZMQ context shared by all controller RPCs on this client. Contexts are - # thread-safe and event-loop-agnostic; each RPC creates and closes its own DEALER - # socket from this context (see with_controller_socket). Terminated once in close(). - self.zmq_context = zmq.asyncio.Context() + # One long-lived context per client. Its fixed native I/O-thread pool is shared + # by all controller RPCs and, for SimpleStorage only, storage-unit requests. + # Sockets remain per-request because ZMQ sockets are not thread-safe. + io_threads = ( + TQ_SIMPLE_STORAGE_ZMQ_IO_THREADS + if simple_storage_zmq_io_threads is None + else simple_storage_zmq_io_threads + ) + if io_threads < 1: + raise ValueError( + "SimpleStorage ZMQ I/O thread pool size must be at least 1, " + f"got {io_threads}" + ) + self.zmq_context = zmq.asyncio.Context(io_threads=io_threads) logger.info(f"[{self.client_id}]: Registered Controller server {controller_info.id} at {controller_info.ip}") def initialize_storage_manager( @@ -92,8 +107,14 @@ def initialize_storage_manager( - zmq_info: ZMQ server information about the storage units """ + create_kwargs = {} + if manager_type == "SimpleStorage": + create_kwargs["zmq_context"] = self.zmq_context self.storage_manager = StorageManagerFactory.create( - manager_type, controller_info=self._controller, config=config + manager_type, + controller_info=self._controller, + config=config, + **create_kwargs, ) # ==================== Basic API ==================== @@ -1243,16 +1264,21 @@ def __init__( self, client_id: str, controller_info: ZMQServerInfo, + simple_storage_zmq_io_threads: int | None = None, ): """Initialize the synchronous TransferQueue client. Args: client_id: Unique identifier for this client instance controller_info: Single controller ZMQ server information + simple_storage_zmq_io_threads: Fixed size of the client context's + native I/O-thread pool. Defaults to + ``TQ_SIMPLE_STORAGE_ZMQ_IO_THREADS`` (8). """ super().__init__( client_id, controller_info, + simple_storage_zmq_io_threads, ) # create new event loop in a separate thread diff --git a/transfer_queue/storage/managers/base.py b/transfer_queue/storage/managers/base.py index c6c959c4..f435fd1e 100644 --- a/transfer_queue/storage/managers/base.py +++ b/transfer_queue/storage/managers/base.py @@ -63,7 +63,12 @@ class StorageManager(ABC): """Base class for storage layer. It defines the interface for data operations and generally provides handshake & notification capabilities.""" - def __init__(self, controller_info: ZMQServerInfo, config: DictConfig): + def __init__( + self, + controller_info: ZMQServerInfo, + config: DictConfig, + zmq_context: zmq.asyncio.Context | None = None, + ): self.storage_manager_id = f"TQ_STORAGE_{uuid4().hex[:8]}" self.config = config self.controller_info = controller_info @@ -71,7 +76,11 @@ def __init__(self, controller_info: ZMQServerInfo, config: DictConfig): # Handshake socket is sync (used only during initialization) self.controller_handshake_socket: zmq.Socket | None = None - self.zmq_context = zmq.asyncio.Context() + # SimpleStorage may borrow a client-owned context whose fixed native I/O + # thread pool is shared by controller and storage-unit request sockets. + # Other backends and standalone managers retain their own context. + self._owns_zmq_context = zmq_context is None + self.zmq_context = zmq_context or zmq.asyncio.Context() self._connect_to_controller() # Dedicated asyncio loop for ZMQ notify traffic, isolated from the caller's loop @@ -391,9 +400,10 @@ def close(self) -> None: else: logger.debug(f"[{self.storage_manager_id}]: Notify ZMQ thread shut down.") - # destroy(linger=0) force-closes any socket still open (e.g. from an interrupted - # request or the notify path) then terminates, so shutdown cannot hang on term(). - self.zmq_context.destroy(linger=0) + if self._owns_zmq_context: + # destroy(linger=0) force-closes any socket still open (e.g. from an interrupted + # request or the notify path) then terminates, so shutdown cannot hang on term(). + self.zmq_context.destroy(linger=0) def __del__(self): """Destructor to ensure resources are cleaned up.""" @@ -424,12 +434,21 @@ def decorator(manager_cls: type[StorageManager]): return decorator @classmethod - def create(cls, manager_type: str, controller_info: ZMQServerInfo, config: dict[str, Any]) -> StorageManager: + def create( + cls, + manager_type: str, + controller_info: ZMQServerInfo, + config: dict[str, Any], + zmq_context: zmq.asyncio.Context | None = None, + ) -> StorageManager: """Create and return a StorageManager instance.""" assert manager_type in cls._registry, ( f"Unknown manager_type: {manager_type}. Supported managers include: {list(cls._registry.keys())}" ) - return cls._registry[manager_type](controller_info, config) + manager_cls = cls._registry[manager_type] + if manager_type == "SimpleStorage" and zmq_context is not None: + return manager_cls(controller_info, config, zmq_context=zmq_context) + return manager_cls(controller_info, config) class KVStorageManager(StorageManager): diff --git a/transfer_queue/storage/managers/simple_storage_manager.py b/transfer_queue/storage/managers/simple_storage_manager.py index c6b20eff..aec260b6 100644 --- a/transfer_queue/storage/managers/simple_storage_manager.py +++ b/transfer_queue/storage/managers/simple_storage_manager.py @@ -73,8 +73,13 @@ class AsyncSimpleStorageManager(StorageManager): instances using ZMQ communication and dynamic socket management. """ - def __init__(self, controller_info: ZMQServerInfo, config: DictConfig): - super().__init__(controller_info, config) + def __init__( + self, + controller_info: ZMQServerInfo, + config: DictConfig, + zmq_context: zmq.asyncio.Context | None = None, + ): + super().__init__(controller_info, config, zmq_context=zmq_context) self.config = config server_infos: ZMQServerInfo | dict[str, ZMQServerInfo] | None = config.get("zmq_info", None)