diff --git a/README.md b/README.md index f2eef59d..eccf25e6 100644 --- a/README.md +++ b/README.md @@ -69,7 +69,7 @@ Optional integrations for different cloud providers can be installed using `plug Support for parallelisation and hyperparameter optimisation can be installed using `plugboard[ray]`. -Additional optional extras: `plugboard[llm]` for LLM components, `plugboard[redis]` for Redis-based connectors, and `plugboard[websockets]` for WebSocket I/O. +Additional optional extras: `plugboard[llm]` for LLM components, `plugboard[redis]` for Redis-based connectors, `plugboard[omq]` for the pyomq backend for ZMQ connectors, and `plugboard[websockets]` for WebSocket I/O. ## ⚡ Quickstart with AI diff --git a/docs/examples/tutorials/running-in-parallel.md b/docs/examples/tutorials/running-in-parallel.md index c8203abb..48b680f6 100644 --- a/docs/examples/tutorials/running-in-parallel.md +++ b/docs/examples/tutorials/running-in-parallel.md @@ -51,7 +51,7 @@ With some small changes we can make the same model run in parallel on Ray. First !!! info [`Channel`][plugboard.connector.Channel] objects are used by Plugboard to handle the communication between components. So far we have used [`AsyncioChannel`][plugboard.connector.AsyncioChannel], which is the best option for simple models that don't require parallelisation. - Plugboard provides different channel classes for use in parallel environments: [`RayChannel`][plugboard.connector.RayChannel] is suitable for single and multi-host Ray environments. [`ZMQChannel`][plugboard.connector.ZMQChannel] is faster, but currently only works on a single host. + Plugboard provides different channel classes for use in parallel environments: [`RayChannel`][plugboard.connector.RayChannel] is suitable for single and multi-host Ray environments. [`ZMQChannel`][plugboard.connector.ZMQChannel] is faster, but currently only works on a single host. Set `PLUGBOARD_ZMQ_BACKEND=pyomq` to use the optional pyomq backend instead of PyZMQ. ```python --8<-- "examples/tutorials/004_using_ray/hello_ray.py:ray" diff --git a/plugboard/_zmq/backend.py b/plugboard/_zmq/backend.py new file mode 100644 index 00000000..929bb44b --- /dev/null +++ b/plugboard/_zmq/backend.py @@ -0,0 +1,47 @@ +"""Selects the ZeroMQ Python backend.""" + +from __future__ import annotations + +import os +import typing as _t + + +ZMQ_BACKEND_ENV = "PLUGBOARD_ZMQ_BACKEND" +ZMQ_BACKEND_PYZMQ = "pyzmq" +ZMQ_BACKEND_PYOMQ = "pyomq" +ZMQ_BACKENDS = frozenset({ZMQ_BACKEND_PYZMQ, ZMQ_BACKEND_PYOMQ}) + + +class ZMQBackendImportError(ImportError): + """Raised when the selected ZeroMQ backend cannot be imported.""" + + +def _backend_name() -> str: + backend = os.environ.get(ZMQ_BACKEND_ENV, ZMQ_BACKEND_PYZMQ).strip().lower() + if not backend: + return ZMQ_BACKEND_PYZMQ + if backend not in ZMQ_BACKENDS: + choices = ", ".join(sorted(ZMQ_BACKENDS)) + raise ValueError( + f"Unsupported ZMQ backend {backend!r}. Set {ZMQ_BACKEND_ENV} to one of: {choices}." + ) + return backend + + +def _load_backend() -> tuple[str, _t.Any, _t.Any]: + backend = _backend_name() + try: + if backend == ZMQ_BACKEND_PYOMQ: + import pyomq as zmq + import pyomq.asyncio as zmq_asyncio + else: + import zmq + import zmq.asyncio as zmq_asyncio + except ImportError as e: + raise ZMQBackendImportError( + f"Failed to import {backend!r} ZMQ backend selected by {ZMQ_BACKEND_ENV}." + ) from e + return backend, zmq, zmq_asyncio + + +zmq_backend, zmq, zmq_asyncio = _load_backend() diff --git a/plugboard/_zmq/zmq_proxy.py b/plugboard/_zmq/zmq_proxy.py index 8134706c..aff202ba 100644 --- a/plugboard/_zmq/zmq_proxy.py +++ b/plugboard/_zmq/zmq_proxy.py @@ -7,8 +7,8 @@ import typing as _t from pydantic import BaseModel, Field, ValidationError -import zmq -import zmq.asyncio + +from plugboard._zmq.backend import zmq, zmq_asyncio try: @@ -23,8 +23,8 @@ def create_socket( socket_type: int, socket_opts: zmq_sockopts_t, - ctx: _t.Optional[zmq.asyncio.Context] = None, -) -> zmq.asyncio.Socket: + ctx: _t.Optional[zmq_asyncio.Context] = None, +) -> zmq_asyncio.Socket: """Creates a ZeroMQ socket with the given type and options. Args: @@ -35,7 +35,7 @@ def create_socket( Returns: The created ZMQ socket. """ - _ctx = ctx or zmq.asyncio.Context.instance() + _ctx = ctx or zmq_asyncio.Context.instance() socket = _ctx.socket(socket_type) for opt, value in socket_opts: socket.setsockopt(opt, value) @@ -184,7 +184,7 @@ def _connect_socket_req_socket(self) -> None: """Connects the REQ socket to the REP socket in the subprocess.""" if self._socket_rep_port is None: raise RuntimeError("ZMQ proxy socket REP port not set.") - self._socket_req_socket: zmq.asyncio.Socket = create_socket(zmq.REQ, []) + self._socket_req_socket: zmq_asyncio.Socket = create_socket(zmq.REQ, []) socket_rep_socket_address: str = f"{self._zmq_address}:{self._socket_rep_port}" self._socket_req_socket.connect(socket_rep_socket_address) self._socket_req_lock: asyncio.Lock = asyncio.Lock() @@ -205,8 +205,8 @@ async def add_push_socket(self, topic: str, maxsize: int = 2000) -> str: async def _run(self) -> None: """Async multiprocessing entrypoint to run ZMQ proxy.""" - self._push_poller: zmq.asyncio.Poller = zmq.asyncio.Poller() - self._push_sockets: dict[str, tuple[str, zmq.asyncio.Socket]] = {} + self._push_poller: zmq_asyncio.Poller = zmq_asyncio.Poller() + self._push_sockets: dict[str, tuple[str, zmq_asyncio.Socket]] = {} self._create_proxy_sockets() @@ -291,7 +291,7 @@ async def _poll_push_sockets(self) -> None: for socket in events: tg.create_task(self._handle_push_socket(socket)) - async def _handle_push_socket(self, socket: zmq.asyncio.Socket) -> None: + async def _handle_push_socket(self, socket: zmq_asyncio.Socket) -> None: msg = await socket.recv_multipart() topic = msg[0].decode("utf8") _, push_socket = self._push_sockets[topic] diff --git a/plugboard/connector/zmq_channel.py b/plugboard/connector/zmq_channel.py index d74ad849..4e773214 100644 --- a/plugboard/connector/zmq_channel.py +++ b/plugboard/connector/zmq_channel.py @@ -7,9 +7,8 @@ import typing as _t from that_depends import Provide, inject -import zmq -import zmq.asyncio +from plugboard._zmq.backend import ZMQ_BACKEND_PYOMQ, zmq, zmq_asyncio, zmq_backend from plugboard._zmq.zmq_proxy import ZMQ_ADDR, ZMQProxy, create_socket, zmq_sockopts_t from plugboard.connector.connector import Connector from plugboard.connector.serde_channel import SerdeChannel @@ -19,6 +18,7 @@ ZMQ_CONFIRM_MSG: str = "__PLUGBOARD_CHAN_CONFIRM_MSG__" +PYOMQ_CLOSE_DRAIN_SECONDS: float = 0.1 # Collection of poll tasks for ZMQ channels required to create strong refs to polling tasks # to avoid destroying tasks before they are done on garbage collection. Is there a better way? @@ -32,8 +32,8 @@ class ZMQChannel(SerdeChannel): def __init__( # noqa: D417 self, *args: _t.Any, - send_socket: _t.Optional[zmq.asyncio.Socket] = None, - recv_socket: _t.Optional[zmq.asyncio.Socket] = None, + send_socket: _t.Optional[zmq_asyncio.Socket] = None, + recv_socket: _t.Optional[zmq_asyncio.Socket] = None, topic: str = "", maxsize: int = 2000, **kwargs: _t.Any, @@ -54,8 +54,8 @@ def __init__( # noqa: D417 maxsize: Optional; Queue maximum item capacity, defaults to 2000. """ super().__init__(*args, **kwargs) - self._send_socket: _t.Optional[zmq.asyncio.Socket] = send_socket - self._recv_socket: _t.Optional[zmq.asyncio.Socket] = recv_socket + self._send_socket: _t.Optional[zmq_asyncio.Socket] = send_socket + self._recv_socket: _t.Optional[zmq_asyncio.Socket] = recv_socket self._is_send_closed = send_socket is None self._is_recv_closed = recv_socket is None self._send_hwm = max(maxsize // 2, 1) @@ -83,6 +83,10 @@ async def close(self) -> None: """Closes the `ZMQChannel`.""" if self._send_socket is not None: await super().close() + if zmq_backend == ZMQ_BACKEND_PYOMQ: + # pyomq does not expose an awaitable socket drain; give queued PUB frames, + # including the close sentinel, a short window to reach the proxy. + await asyncio.sleep(PYOMQ_CLOSE_DRAIN_SECONDS) self._send_socket.close() self._send_socket = None if self._recv_socket is not None: @@ -232,7 +236,7 @@ def __init__(self, *args: _t.Any, **kwargs: _t.Any) -> None: self._xsub_port = self._xsub_socket.bind_to_random_port("tcp://*") self._xpub_socket = create_socket(zmq.XPUB, [(zmq.SNDHWM, self._maxsize)]) self._xpub_port = self._xpub_socket.bind_to_random_port("tcp://*") - self._poller = zmq.asyncio.Poller() + self._poller = zmq_asyncio.Poller() self._poller.register(self._xsub_socket, zmq.POLLIN) self._poller.register(self._xpub_socket, zmq.POLLIN) self._poll_task = asyncio.create_task(self._poll()) @@ -251,7 +255,7 @@ async def _poll(self) -> None: poll_fn, xps, xss = self._poller.poll, self._xpub_socket, self._xsub_socket try: while True: - events = dict(await poll_fn()) + events = dict(await poll_fn(timeout=1000)) if xps in events: await xss.send_multipart(await xps.recv_multipart()) if xss in events: diff --git a/pyproject.toml b/pyproject.toml index d64d18f0..e81e3201 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -49,6 +49,7 @@ llm = [ "llama-index-core>=0.12.30,<1", "llama-index-llms-openai>=0.3.33,<1", ] +omq = ["pyomq>=0.20.1,<1"] # Pinning jsonschema due to performance issues with Lark and rfc3987-syntax parser # https://github.com/python-jsonschema/jsonschema/issues/1392 ray = ["ray[tune]>=2.47.1,<3", "jsonschema<4.25.0", "optuna>=3.0,<5"] diff --git a/tests/unit/test_zmq_backend.py b/tests/unit/test_zmq_backend.py new file mode 100644 index 00000000..824aca8f --- /dev/null +++ b/tests/unit/test_zmq_backend.py @@ -0,0 +1,131 @@ +"""Tests for ZMQ backend selection.""" + +from __future__ import annotations + +import importlib.util +import os +import subprocess +import sys +import textwrap + +import pytest + +from plugboard._zmq.backend import ZMQ_BACKEND_ENV + + +def _run_backend_probe( + code: str, + backend: str | None = None, + extra_env: dict[str, str] | None = None, +) -> subprocess.CompletedProcess[str]: + env = os.environ.copy() + if backend is not None: + env[ZMQ_BACKEND_ENV] = backend + if extra_env is not None: + env.update(extra_env) + return subprocess.run( # noqa: S603 + [sys.executable, "-c", textwrap.dedent(code)], + check=False, + capture_output=True, + env=env, + text=True, + ) + + +def test_default_zmq_backend_is_pyzmq() -> None: + """The default ZMQ backend remains PyZMQ.""" + result = _run_backend_probe( + """ + from plugboard._zmq.backend import zmq_backend, zmq + print(zmq_backend) + print(zmq.__name__) + """, + extra_env={ZMQ_BACKEND_ENV: ""}, + ) + + assert result.returncode == 0, result.stderr + assert result.stdout.splitlines() == ["pyzmq", "zmq"] + + +def test_invalid_zmq_backend_fails_with_clear_error() -> None: + """Unsupported backend names fail during import with a clear error.""" + result = _run_backend_probe( + """ + import plugboard._zmq.backend + """, + backend="not-a-backend", + ) + + assert result.returncode != 0 + assert "Unsupported ZMQ backend" in result.stderr + assert ZMQ_BACKEND_ENV in result.stderr + + +@pytest.mark.skipif(importlib.util.find_spec("pyomq") is None, reason="pyomq not installed") +def test_pyomq_backend_supports_create_socket() -> None: + """The pyomq backend can run the ZMQ socket helper.""" + result = _run_backend_probe( + """ + import asyncio + + from plugboard._zmq.backend import zmq, zmq_backend + from plugboard._zmq.zmq_proxy import create_socket + + async def main() -> None: + pull = create_socket(zmq.PULL, [(zmq.RCVHWM, 100)]) + port = pull.bind_to_random_port("tcp://127.0.0.1") + push = create_socket(zmq.PUSH, [(zmq.SNDHWM, 100)]) + push.connect(f"tcp://127.0.0.1:{port}") + await asyncio.sleep(0.2) + await push.send_multipart([b"", b"payload"]) + got = await asyncio.wait_for(pull.recv_multipart(), timeout=1.0) + assert got == [b"", b"payload"] + push.close(linger=0) + pull.close(linger=0) + print(zmq_backend) + + asyncio.run(main()) + """, + backend="pyomq", + ) + + assert result.returncode == 0, result.stderr + assert result.stdout.strip() == "pyomq" + + +@pytest.mark.skipif(importlib.util.find_spec("pyomq") is None, reason="pyomq not installed") +def test_pyomq_backend_supports_zmq_proxy() -> None: + """The pyomq backend can run the ZMQ proxy process.""" + result = _run_backend_probe( + """ + import asyncio + + from plugboard._zmq.backend import zmq + from plugboard._zmq.zmq_proxy import ZMQProxy, create_socket + + async def main() -> None: + proxy = ZMQProxy(maxsize=100) + try: + topic = b"topic" + sub = create_socket( + zmq.SUB, + [(zmq.RCVHWM, 100), (zmq.SUBSCRIBE, topic)], + ) + sub.connect(proxy.xpub_addr) + pub = create_socket(zmq.PUB, [(zmq.SNDHWM, 100)]) + pub.connect(proxy.xsub_addr) + await asyncio.sleep(0.3) + await pub.send_multipart([topic, b"payload"]) + got = await asyncio.wait_for(sub.recv_multipart(), timeout=1.0) + assert got == [topic, b"payload"] + pub.close(linger=0) + sub.close(linger=0) + finally: + proxy.terminate(timeout=5.0) + + asyncio.run(main()) + """, + backend="pyomq", + ) + + assert result.returncode == 0, result.stderr diff --git a/uv.lock b/uv.lock index ba5c716e..42453ee4 100644 --- a/uv.lock +++ b/uv.lock @@ -3964,6 +3964,9 @@ llm = [ { name = "llama-index-core" }, { name = "llama-index-llms-openai" }, ] +omq = [ + { name = "pyomq" }, +] ray = [ { name = "jsonschema" }, { name = "optuna" }, @@ -4081,6 +4084,7 @@ requires-dist = [ { name = "pydantic", marker = "python_full_version < '3.14'", specifier = ">=2.8.0,<3" }, { name = "pydantic", marker = "python_full_version >= '3.14'", specifier = ">=2.13.1,<3" }, { name = "pydantic-settings", specifier = ">=2.7.1,<3" }, + { name = "pyomq", marker = "extra == 'omq'", specifier = ">=0.20.1,<1" }, { name = "pyzmq", specifier = ">=26.2,<28" }, { name = "ray", extras = ["tune"], marker = "extra == 'ray'", specifier = ">=2.47.1,<3" }, { name = "redis", marker = "extra == 'redis'", specifier = ">=7.1,<9" }, @@ -4093,7 +4097,7 @@ requires-dist = [ { name = "uvloop", marker = "sys_platform != 'win32'", specifier = ">=0.21.0,<1" }, { name = "websockets", marker = "extra == 'websockets'", specifier = ">=14.2,<17" }, ] -provides-extras = ["aws", "azure", "gcp", "llm", "ray", "redis", "websockets"] +provides-extras = ["aws", "azure", "gcp", "llm", "omq", "ray", "redis", "websockets"] [package.metadata.requires-dev] all = [ @@ -4636,6 +4640,22 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/f7/27/a2fc51a4a122dfd1015e921ae9d22fee3d20b0b8080d9a704578bf9deece/pymdown_extensions-10.21.2-py3-none-any.whl", hash = "sha256:5c0fd2a2bea14eb39af8ff284f1066d898ab2187d81b889b75d46d4348c01638", size = 268901, upload-time = "2026-03-29T15:01:53.244Z" }, ] +[[package]] +name = "pyomq" +version = "0.20.1" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/4b/f8/80ab2fc7b7ecab59bcdd4ddcae72dbd819e708ae12109a7005b473510486/pyomq-0.20.1.tar.gz", hash = "sha256:2f42c4c48666fdc4b621e9e7c9d5c94082ef3c9147f35ceb5c3eb42e6ed23819", size = 649938, upload-time = "2026-08-23T21:03:55.159Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/4c/b5/fd6a4a81754a5b5736f15f1a6215752f38e61cca41c7ff420de63a7448ae/pyomq-0.20.1-cp311-abi3-macosx_10_12_x86_64.whl", hash = "sha256:0b62acc91d0882889cba23aca60530062f6109bebc5b1c68a8aea29334f75b82", size = 1779867, upload-time = "2026-08-23T21:03:44.028Z" }, + { url = "https://files.pythonhosted.org/packages/ba/d7/a14d7176942a9bca468584c7147cf6e6eb86574102c7e1bc4dc7fadaffdf/pyomq-0.20.1-cp311-abi3-macosx_11_0_arm64.whl", hash = "sha256:9f737cabb6f3f6300c55236f2f6f4ded09a6a6d31b0e8f381965f9f4f16c436c", size = 1725287, upload-time = "2026-08-23T21:03:45.442Z" }, + { url = "https://files.pythonhosted.org/packages/02/67/5b4a4c648f945ad3b8dc5ac1dac1495a89e22d62795a922bf1e5759e9210/pyomq-0.20.1-cp311-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:7362231d5e222f30ad43ee6e66572c98b665dde282fbcb94c0952282d873d0af", size = 1806696, upload-time = "2026-08-23T21:03:46.781Z" }, + { url = "https://files.pythonhosted.org/packages/c9/d1/7d0baccf1eaa34e25efcffd4e750e198af1dd3a703a8abda1f0cf1ba1dc3/pyomq-0.20.1-cp311-abi3-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:ebe3ea509a7de86be592763ce90af4c53065701bd4b9f58c99887df2d389a269", size = 1829920, upload-time = "2026-08-23T21:03:48.24Z" }, + { url = "https://files.pythonhosted.org/packages/d9/dc/c9e5c16b2925761dee05f0c430d68e4772b1347620a9b84111c51a8633e1/pyomq-0.20.1-cp311-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:85f47fa8c6a66a9ee0cddc8c7ec940582e9bb80dcc555c1c1fae182c15764d42", size = 1983314, upload-time = "2026-08-23T21:03:49.594Z" }, + { url = "https://files.pythonhosted.org/packages/c5/c8/aca1a3a033f4ca1f8359eae1cb5e458a9798181d43bfbf474b4ce6a357a1/pyomq-0.20.1-cp311-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:1c57a649272d11f162b4d8ca126b0096278a38fb27a51e3319f8e4e4dd744555", size = 2041966, upload-time = "2026-08-23T21:03:51.117Z" }, + { url = "https://files.pythonhosted.org/packages/e8/32/217f446433516bcf168d4c262e31c30006fda30eabbe8cf3e259b7b0caed/pyomq-0.20.1-cp311-abi3-win_amd64.whl", hash = "sha256:4aa0ad5b91adeb264739944c4ba505e73885db1a04ac2a7bf6d6956db14e6a9a", size = 1874489, upload-time = "2026-08-23T21:03:52.48Z" }, + { url = "https://files.pythonhosted.org/packages/b1/61/6c5469e45eaa62325eaf1ca9c6258fe28534f5d763498dd5c58c5aa2e807/pyomq-0.20.1-cp311-abi3-win_arm64.whl", hash = "sha256:24361da804107a7431cfaba3c6bb49d89746441827f9532875d7c3b84725bd64", size = 1750463, upload-time = "2026-08-23T21:03:53.747Z" }, +] + [[package]] name = "pyparsing" version = "3.3.2"