Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 26 additions & 20 deletions src/ezmsg/core/backendprocess.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
9 changes: 9 additions & 0 deletions src/ezmsg/core/graphserver.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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)
Expand Down
42 changes: 38 additions & 4 deletions src/ezmsg/core/messagechannel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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]]):
"""
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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

Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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.
Expand Down
4 changes: 3 additions & 1 deletion src/ezmsg/core/shm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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":
Expand Down Expand Up @@ -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()
15 changes: 14 additions & 1 deletion src/ezmsg/core/subclient.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 (
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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(
Expand Down
20 changes: 20 additions & 0 deletions tests/test_shm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
56 changes: 56 additions & 0 deletions tests/test_subclient.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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