From 6d6bc01a1834cb05526615b414c551a0c4eb50da Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 1 Oct 2026 02:16:57 -0400 Subject: [PATCH 1/2] Harden SHM against retained zero-copy views (#272) Split from #273: the parts that stand on their own, without the ChannelFailed policy that is still under discussion. - SHMContext: if closing the mapping raises BufferError because views into it are still alive (a unit kept a zero-copy array), drop the handle and let the mmap be unmapped when the last view is collected, instead of killing the channel when the publisher grows its segment. - SHMContext.wait_closed() also closes the segment: a monitor task cancelled before it first ran never executed its finally, leaking the mapping (the "Exception ignored in SharedMemory.__del__" noise). - Channel: SHM-delivered messages are read-only, matching the TCP path, so a subscriber cannot write into memory shared with the publisher and every other subscriber. - Tests: closing with a live view, and a cross-process grow with a receiver that retains every message. --- src/ezmsg/core/messagechannel.py | 4 +- src/ezmsg/core/shm.py | 24 +++++++- src/ezmsg/core/shm_grow_test_support.py | 77 +++++++++++++++++++++++++ tests/test_shm.py | 30 ++++++++++ tests/test_shm_grow.py | 32 +++++++++- 5 files changed, 164 insertions(+), 3 deletions(-) diff --git a/src/ezmsg/core/messagechannel.py b/src/ezmsg/core/messagechannel.py index 6fab584b..033ace55 100644 --- a/src/ezmsg/core/messagechannel.py +++ b/src/ezmsg/core/messagechannel.py @@ -383,7 +383,9 @@ def _deliver_from_shm(self, msg_id: int) -> None: publisher named. """ assert self.shm is not None - shm_buf = self.shm[msg_id % self.num_buffers] + # Read-only: this memory is shared with the publisher and every other + # subscriber (matches the TCP path below). + shm_buf = self.shm[msg_id % self.num_buffers].toreadonly() # The slot for this msg_id may be uninitialized after a mid-stream # resize; msg_id() raises UninitializedMemory in that case. Treat it as # a mismatch (drop + release) rather than letting it kill the channel. diff --git a/src/ezmsg/core/shm.py b/src/ezmsg/core/shm.py index 279aeb9c..09ea2f41 100644 --- a/src/ezmsg/core/shm.py +++ b/src/ezmsg/core/shm.py @@ -1,6 +1,7 @@ import asyncio from collections.abc import Generator import logging +import os from dataclasses import dataclass, field from contextlib import contextmanager, suppress @@ -124,9 +125,28 @@ async def _graph_connection( except (ConnectionResetError, BrokenPipeError) as e: logger.debug(f"SHMContext {self.name} GraphServer {type(e)}") finally: - self._shm.close() + self._close_shm() await close_stream_writer(writer) + def _close_shm(self) -> None: + try: + self._shm.close() + except BufferError: + # Something still holds a view into this segment (e.g. a subscriber + # retained a zero-copy numpy array past its callback). The mmap + # can't be closed while exported, so drop our handle to it instead; + # it is unmapped when the last view is garbage collected. The + # segment itself is unlinked by the GraphServer once leases end. + # SharedMemory.close() has already released and cleared ``_buf``. + logger.debug( + f"SHMContext {self.name} has live views; deferring unmap to GC" + ) + self._shm._mmap = None # type: ignore[attr-defined] + fd = getattr(self._shm, "_fd", -1) + if fd >= 0: + os.close(fd) + self._shm._fd = -1 # type: ignore[attr-defined] + def __getitem__(self, idx: int) -> memoryview: """ Get a memory view of a specific buffer in the shared memory segment. @@ -180,6 +200,8 @@ async def wait_closed(self) -> None: """ with suppress(asyncio.CancelledError): await self._graph_task + # A task cancelled before it first ran never executes its finally. + self._close_shm() @property def name(self) -> str: diff --git a/src/ezmsg/core/shm_grow_test_support.py b/src/ezmsg/core/shm_grow_test_support.py index 25c0b5f6..07c2933c 100644 --- a/src/ezmsg/core/shm_grow_test_support.py +++ b/src/ezmsg/core/shm_grow_test_support.py @@ -2,6 +2,7 @@ from dataclasses import dataclass import json +import pickle import ezmsg.core as ez @@ -98,3 +99,79 @@ def network(self) -> ez.NetworkDefinition: def process_components(self): return (self.PUB, self.SUB) + + +class ViewPayload: + """ + Pickles its data out-of-band, so on receipt ``data`` is a memoryview into + the channel's SHM (like a zero-copy numpy array) rather than a copy. + """ + + def __init__(self, data) -> None: + # memoryview() registers a fresh export on the underlying buffer, so a + # retained ViewPayload pins the SHM mapping just like an ndarray does. + self.data = memoryview(data) + + def __reduce_ex__(self, protocol): + return (ViewPayload, (pickle.PickleBuffer(self.data),)) + + +class ViewGenerator(ez.Unit): + SETTINGS = BlobGeneratorSettings + + OUTPUT = ez.OutputStream( + ViewPayload, + num_buffers=4, + buf_size=4096, + allow_local=False, + ) + + async def initialize(self) -> None: + self.OUTPUT.buf_size = self.SETTINGS.buf_size + self.OUTPUT.num_buffers = self.SETTINGS.num_buffers + + @ez.publisher(OUTPUT) + async def spawn(self) -> AsyncGenerator: + for seq, size in enumerate(self.SETTINGS.sizes): + yield self.OUTPUT, ViewPayload(bytes([seq]) * size) + raise ez.Complete + + +class RetainingReceiverState(ez.State): + num_received: int = 0 + retained: list | None = None + + +class RetainingReceiver(ez.Unit): + """Keeps every received message past its callback (issue #272).""" + + STATE = RetainingReceiverState + SETTINGS = BlobReceiverSettings + + INPUT = ez.InputStream(ViewPayload) + + async def initialize(self) -> None: + self.STATE.retained = [] + + @ez.subscriber(INPUT) + async def on_message(self, msg: ViewPayload) -> None: + self.STATE.num_received += 1 + self.STATE.retained.append(msg) + with open(self.SETTINGS.output_fn, "a") as output_file: + output_file.write( + json.dumps( + { + "seq": msg.data[0], + "len": len(msg.data), + "readonly": msg.data.readonly, + } + ) + + "\n" + ) + if self.STATE.num_received == self.SETTINGS.num_msgs: + raise ez.Complete + + +class RetainingGrowSystem(GrowSystem): + PUB = ViewGenerator() + SUB = RetainingReceiver() diff --git a/tests/test_shm.py b/tests/test_shm.py index 32e2e0f5..ca93e954 100644 --- a/tests/test_shm.py +++ b/tests/test_shm.py @@ -210,3 +210,33 @@ def start(self, address): asyncio.run(test_rw()) asyncio.run(test_shm_detach_order()) asyncio.run(test_shutdown()) + + +@pytest.mark.asyncio +async def test_close_with_live_view() -> None: + """Closing an SHMContext while a view into it is alive must not raise; + the mapping stays valid for the view until it is released (issue #272).""" + service = GraphService() + server = service.create_server() + + shm = await service.create_shm(4, 2**16) + attach_shm = await service.attach_shm(shm.name) + + content = b"HELLO" + with shm.buffer(0) as mem: + mem[0 : len(content)] = content[:] + + # Like a retained zero-copy ndarray: a fresh export on the mapping. + view = memoryview(attach_shm[0]) + + attach_shm.close() + await attach_shm.wait_closed() + + assert bytes(view[0 : len(content)]) == content + with pytest.raises(BufferError): + attach_shm[0] + view.release() + + shm.close() + await shm.wait_closed() + server.stop() diff --git a/tests/test_shm_grow.py b/tests/test_shm_grow.py index f4d76e8a..d6cf7870 100644 --- a/tests/test_shm_grow.py +++ b/tests/test_shm_grow.py @@ -18,7 +18,11 @@ import ezmsg.core as ez from ez_test_utils import get_test_fn -from ezmsg.core.shm_grow_test_support import GrowSystem, GrowSystemSettings +from ezmsg.core.shm_grow_test_support import ( + GrowSystem, + GrowSystemSettings, + RetainingGrowSystem, +) @pytest.mark.parametrize( @@ -48,3 +52,29 @@ def test_cross_process_grow_delivers_all(sizes): assert [r["seq"] for r in results] == list(range(len(sizes))) assert [r["len"] for r in results] == list(sizes) + + +def test_cross_process_grow_with_retained_views(): + """Issue #272: a subscriber retaining zero-copy views into the old SHM + segment must not kill the channel when the publisher grows it.""" + sizes = (8, 8, 16384, 8, 65536, 8) + with get_test_fn() as test_filename: + system = RetainingGrowSystem( + GrowSystemSettings( + sizes=sizes, + buf_size=4096, + num_buffers=4, + output_fn=str(test_filename), + ) + ) + ez.run(SYSTEM=system) + + results = [] + with open(test_filename, "r") as file: + for line in file: + results.append(json.loads(line)) + + assert [r["seq"] for r in results] == list(range(len(sizes))) + assert [r["len"] for r in results] == list(sizes) + # Zero-copy SHM views must not let subscribers write shared memory. + assert all(r["readonly"] for r in results) From d71704feff65d8867c06800525a674fa08d00801 Mon Sep 17 00:00:00 2001 From: Chadwick Boulay Date: Thu, 1 Oct 2026 02:39:59 -0400 Subject: [PATCH 2/2] Keep a grown-out SHM segment open until every channel has attached it On a grow the publisher closed the old segment at once. A message already sent names that segment, and a channel that has not processed it yet still has to attach it; if the publisher's lease was the last one, the GraphServer unlinked it first and the channel's attach failed with FileNotFoundError, killing the channel. Two quick grows (num_buffers > 1, no wait for acks between them) lost that race whenever timing shifted. The publisher now retires the old segment with the last msg_id sent under its name, and closes it once the backpressure wait for msg_id L + num_buffers is over: by then every channel has released message L, and a channel attaches the segment a message names before it can release it. --- src/ezmsg/core/pubclient.py | 35 +++++++++++++++++++++++++++++++++-- 1 file changed, 33 insertions(+), 2 deletions(-) diff --git a/src/ezmsg/core/pubclient.py b/src/ezmsg/core/pubclient.py index 11abbe7c..30b3c882 100644 --- a/src/ezmsg/core/pubclient.py +++ b/src/ezmsg/core/pubclient.py @@ -272,6 +272,11 @@ def __init__( self.pid = os.getpid() self.topic = topic self._shm = shm + # Segments replaced by a grow, each with the last msg_id sent under its + # name. Kept open until every channel must have attached it (see + # _close_retired_shms): closing at once let the GraphServer unlink a + # segment a lagging channel had yet to attach. + self._retired_shms: list[tuple[int, SHMContext]] = [] self._msg_id = 0 self._channels = dict() self._channel_tasks = dict() @@ -300,6 +305,8 @@ def close(self) -> None: self._graph_task.cancel() PROFILES.unregister_publisher(self.id) self._shm.close() + for _, retired in self._retired_shms: + retired.close() self._connection_task.cancel() for task in self._channel_tasks.values(): task.cancel() @@ -312,6 +319,9 @@ async def wait_closed(self) -> None: connection server shutdown, and all subscriber tasks to complete. """ await self._shm.wait_closed() + for _, retired in self._retired_shms: + await retired.wait_closed() + self._retired_shms.clear() with suppress(asyncio.CancelledError): await self._graph_task with suppress(asyncio.CancelledError): @@ -500,6 +510,9 @@ async def broadcast(self, obj: Any) -> None: PROFILE_TIME() - wait_start_ns, msg_seq=self._msg_id ) + if self._retired_shms: + await self._close_retired_shms() + if self._should_use_local_fast_path(): self._local_channel.put_local(self._msg_id, obj) @@ -525,8 +538,10 @@ async def broadcast(self, obj: Any) -> None: except UninitializedMemory: pass - self._shm.close() - await self._shm.wait_closed() + # Not closed yet: messages already sent name the old + # segment, and a channel that has not processed them + # still has to attach it. + self._retired_shms.append((self._msg_id - 1, self._shm)) self._shm = new_shm with self._shm.buffer(buf_idx) as mem: @@ -596,6 +611,22 @@ async def broadcast(self, obj: Any) -> None: ) self._msg_id += 1 + async def _close_retired_shms(self) -> None: + """Close replaced segments no channel can still need. + + Called once the backpressure wait for this message's slot is over, so + message ``msg_id - num_buffers`` has been released by every channel, + and with it every earlier one. A channel attaches the segment a message + names before it can release that message, so a segment whose last + message is at or below that floor has been attached by every channel + that will ever ask for it. + """ + floor = self._msg_id - self._num_buffers + while self._retired_shms and self._retired_shms[0][0] <= floor: + _, retired = self._retired_shms.pop(0) + retired.close() + await retired.wait_closed() + def _should_use_local_fast_path(self) -> bool: return any(self._can_deliver_locally(ch) for ch in self._channels.values())