diff --git a/src/ezmsg/core/backendprocess.py b/src/ezmsg/core/backendprocess.py index 4fc829e4..59b24c91 100644 --- a/src/ezmsg/core/backendprocess.py +++ b/src/ezmsg/core/backendprocess.py @@ -48,6 +48,7 @@ from .processclient import ProcessControlClient from .pubclient import Publisher from .subclient import Subscriber +from .messagechannel import ChannelFailed from .netprotocol import AddressType from .settingsmeta import ( coerce_settings_field_value, @@ -867,28 +868,33 @@ async def next_message() -> AsyncGenerator[Any, None]: await sub.wait_closed() break - async with next_message() as msg: - try: - if on_message is not None: - try: - await on_message(msg) - except Exception as exc: - logger.warning( - f"Failed to report subscriber message metadata: {exc}" - ) - for callable in list(callables): - try: - span_start_ns = sub.begin_profile() + try: + async with next_message() as msg: + try: + if on_message is not None: try: - await callable(msg) - finally: - sub.end_profile( - span_start_ns, getattr(callable, "__name__", None) + await on_message(msg) + except Exception as exc: + logger.warning( + f"Failed to report subscriber message metadata: {exc}" ) - except (Complete, NormalTermination): - callables.remove(callable) - finally: - del msg + for callable in list(callables): + try: + span_start_ns = sub.begin_profile() + try: + await callable(msg) + finally: + sub.end_profile( + span_start_ns, getattr(callable, "__name__", None) + ) + except (Complete, NormalTermination): + callables.remove(callable) + finally: + del msg + except ChannelFailed as exc: + # Keep serving this stream's other publishers. + logger.error(f"{exc}: {exc.__cause__!r}") + continue if len(callables) > 1: await asyncio.sleep(0) diff --git a/src/ezmsg/core/graphserver.py b/src/ezmsg/core/graphserver.py index 46c41a3e..4ea4f7fe 100644 --- a/src/ezmsg/core/graphserver.py +++ b/src/ezmsg/core/graphserver.py @@ -273,6 +273,10 @@ async def api( num_buffers = await read_int(reader) buf_size = await read_int(reader) + # Forget segments whose last lease has ended. + for name in [n for n, i in self.shms.items() if i.unlinked]: + del self.shms[name] + # Create segment shm_info = SHMInfo.create(num_buffers, buf_size) self.shms[shm_info.shm.name] = shm_info @@ -281,6 +285,11 @@ async def api( elif req == Command.SHM_ATTACH.value: shm_name = await read_str(reader) shm_info = self.shms.get(shm_name, None) + if shm_info is not None and shm_info.unlinked: + # Its last lease ended and it is gone; attaching would + # hand the client a name it cannot open. + del self.shms[shm_name] + shm_info = None if shm_info is None: await close_stream_writer(writer) diff --git a/src/ezmsg/core/messagechannel.py b/src/ezmsg/core/messagechannel.py index 08b02ebe..f5ff1be0 100644 --- a/src/ezmsg/core/messagechannel.py +++ b/src/ezmsg/core/messagechannel.py @@ -32,6 +32,17 @@ logger = logging.getLogger("ezmsg") +# Notification msg_id a Channel sends its clients when its publisher connection +# dies unexpectedly; real msg_ids are unsigned. +CHANNEL_FAILED = -1 + + +class ChannelFailed(RuntimeError): + """ + Raised to a subscriber when the Channel delivering messages from one of its + publishers crashes; no further messages will arrive from that publisher. + """ + class LeakyQueue(asyncio.Queue[typing.Tuple[UUID, int]]): """ @@ -94,9 +105,19 @@ def bind(self, channel: "Channel") -> None: self._channel = channel def connection_lost(self, exc: BaseException | None) -> None: + # A peer that simply went away is a normal disconnect; anything else is + # a crash. An exception out of frames_available is caught by + # _drain_frames, which keeps it as _close_exc and aborts (so exc is None + # here) -- read it before the base class overwrites it with exc. + failure = exc if exc is not None else self._close_exc super().connection_lost(exc) - if self._channel is not None: - self._channel._on_disconnected() + chan = self._channel + if chan is not None: + if failure is not None and chan.failure is None and not isinstance( + failure, (ConnectionResetError, BrokenPipeError) + ): + chan.failure = failure + chan._on_disconnected() def frames_available(self) -> None: chan = self._channel @@ -161,7 +182,10 @@ async def _reattach_shm(self, shm_name: str, frame_end: int, msg_id: int) -> Non try: chan.shm = await GraphService(chan._graph_address).attach_shm(shm_name) - except ValueError: + except (ValueError, FileNotFoundError): + # ValueError: the GraphServer does not know the name. + # FileNotFoundError: an older GraphServer still knew a segment + # it had already unlinked. logger.warning( "Channel %s received stale SHM %s for publisher %s; waiting for next valid SHM", chan.id, @@ -181,8 +205,9 @@ async def _reattach_shm(self, shm_name: str, frame_end: int, msg_id: int) -> Non ) del self._buffer[:frame_end] chan._release_backpressure(msg_id, chan.id) - except Exception: + except Exception as exc: logger.exception("Channel %s failed to reattach SHM", chan.id) + chan.failure = exc self.abort() return @@ -250,6 +275,7 @@ def __init__( # Axes the publisher sent in full, for resolving its elided references. self._axis_table = AxisTable() self._axes_requested = False + self.failure: BaseException | None = None @classmethod async def create( @@ -468,6 +494,8 @@ def _on_disconnected(self) -> None: self.cache.clear() if self.shm is not None: self.shm.close() + if self.failure is not None: + self._notify_failure() logger.debug(f"disconnected: channel:{self.id} -> pub:{self.pub_id}") def _set_channel_kind(self, kind: ProfileChannelType) -> None: @@ -496,6 +524,12 @@ def _notify_clients(self, msg_id: int) -> bool: queue.put_nowait((self.pub_id, msg_id)) return not self.backpressure.available(buf_idx) + def _notify_failure(self) -> None: + """wake every client so it can surface this channel's failure""" + for queue in self.clients.values(): + if queue is not None: + queue.put_nowait((self.pub_id, CHANNEL_FAILED)) + def put_local(self, msg_id: int, msg: typing.Any) -> None: """ Put a message DIRECTLY into cache and notify all clients. diff --git a/src/ezmsg/core/shm.py b/src/ezmsg/core/shm.py index 09ea2f41..c43bb9ea 100644 --- a/src/ezmsg/core/shm.py +++ b/src/ezmsg/core/shm.py @@ -235,6 +235,7 @@ class SHMInfo: shm: SharedMemory leases: set["asyncio.Task[None]"] = field(default_factory=set) + unlinked: bool = False @classmethod def create(cls, num_buffers: int, buf_size: int) -> "SHMInfo": @@ -278,7 +279,8 @@ async def _wait_for_eof() -> None: def _release(self, task: "asyncio.Task[None]"): self.leases.discard(task) logger.debug(f"discarded lease from {self.shm.name}; {len(self.leases)} left") - if len(self.leases) == 0: + if len(self.leases) == 0 and not self.unlinked: logger.debug(f"unlinking {self.shm.name}") + self.unlinked = True self.shm.close() self.shm.unlink() diff --git a/src/ezmsg/core/subclient.py b/src/ezmsg/core/subclient.py index 4ae67f0f..e277d33c 100644 --- a/src/ezmsg/core/subclient.py +++ b/src/ezmsg/core/subclient.py @@ -8,7 +8,13 @@ from .graphserver import GraphService from .channelmanager import CHANNELS -from .messagechannel import NotificationQueue, LeakyQueue, Channel +from .messagechannel import ( + CHANNEL_FAILED, + ChannelFailed, + NotificationQueue, + LeakyQueue, + Channel, +) from .profiling import PROFILES, PROFILE_TIME from .netprotocol import ( @@ -154,6 +160,8 @@ def _handle_dropped_notification( :type notification: tuple[UUID, int] """ pub_id, msg_id = notification + if msg_id == CHANNEL_FAILED: + return if pub_id in self._channels: self._channels[pub_id].release_without_get(msg_id, self.id) @@ -302,6 +310,11 @@ async def recv_zero_copy(self) -> typing.AsyncGenerator[typing.Any, None]: # Stale notification from an unregistered publisher — skip. channel = self._channels[pub_id] + if msg_id < 0: # CHANNEL_FAILED; literal compare keeps the hot path cheap + raise ChannelFailed( + f"Subscriber {self.topic}({self.id}) will receive no more messages " + f"from publisher {channel.topic}({pub_id}): its channel crashed" + ) from channel.failure channel_kind = channel.channel_kind self._active_msg_seq = msg_id sampled, trace_lease, _trace_user_span = self._profile.trace_receive_state( diff --git a/tests/test_shm.py b/tests/test_shm.py index ca93e954..02d814b0 100644 --- a/tests/test_shm.py +++ b/tests/test_shm.py @@ -240,3 +240,23 @@ async def test_close_with_live_view() -> None: shm.close() await shm.wait_closed() server.stop() + + +@pytest.mark.asyncio +async def test_attaching_an_unlinked_segment_is_refused() -> None: + """Once a segment's last lease ends it is unlinked; the GraphServer must + then refuse the name (ValueError, which channels treat as stale) rather + than hand back one the client cannot open (FileNotFoundError).""" + service = GraphService() + server = service.create_server() + + shm = await service.create_shm(2, 2**12) + name = shm.name + shm.close() + await shm.wait_closed() + await asyncio.sleep(0.05) # let the server see the lease end + + with pytest.raises(ValueError): + await service.attach_shm(name) + + server.stop() diff --git a/tests/test_subclient.py b/tests/test_subclient.py index 20c5ea27..6552cee5 100644 --- a/tests/test_subclient.py +++ b/tests/test_subclient.py @@ -5,6 +5,7 @@ import pytest from ezmsg.core.subclient import Subscriber +from ezmsg.core.messagechannel import CHANNEL_FAILED, ChannelFailed from ezmsg.core.graphmeta import ProfileChannelType from ezmsg.core.netprotocol import Command, encode_str from ezmsg.core import channelmanager as channelmanager_module @@ -148,3 +149,58 @@ async def test_recv_zero_copy_skips_stale_notification(): # Before the fix this would raise KeyError for stale_pub. async with sub.recv_zero_copy() as msg: assert msg == "msg-1" + + +@pytest.mark.asyncio +async def test_recv_zero_copy_raises_channel_failed(): + """A crashed channel must surface to the subscriber instead of leaving it + waiting forever (issue #272).""" + sub = Subscriber( + id=uuid4(), + topic="test/topic", + graph_address=None, + _guard=Subscriber._SENTINEL, + ) + + pub_id = uuid4() + channel = DummyChannel() + channel.failure = BufferError("boom") + sub._channels[pub_id] = channel + + await sub._incoming.put((pub_id, CHANNEL_FAILED)) + await sub._incoming.put((pub_id, 1)) + + with pytest.raises(ChannelFailed) as excinfo: + async with sub.recv_zero_copy(): + pass + assert excinfo.value.__cause__ is channel.failure + + # Other notifications still flow afterwards. + async with sub.recv_zero_copy() as msg: + assert msg == "msg-1" + + +@pytest.mark.asyncio +async def test_a_channel_crash_reaches_the_subscriber(): + """An exception while the channel delivers a message ends its connection; + the subscriber must get ChannelFailed (chained from it), not wait forever.""" + import ezmsg.core as ez + + async with ez.GraphContext(auto_start=True) as ctx: + pub = await ctx.publisher("/CRASH", host="127.0.0.1", num_buffers=2, allow_local=False) + sub = await ctx.subscriber("/CRASH") + await pub.broadcast(b"ok") + async with sub.recv_zero_copy() as msg: + assert msg == b"ok" + + channel = sub._channels[pub.id] + boom = RuntimeError("delivery failed") + + def explode(*args, **kwargs): + raise boom + + channel._deliver_from_shm = explode + await pub.broadcast(b"never delivered") + with pytest.raises(ChannelFailed) as excinfo: + await asyncio.wait_for(sub.recv(), timeout=5.0) + assert excinfo.value.__cause__ is boom