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
4 changes: 3 additions & 1 deletion src/ezmsg/core/messagechannel.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
35 changes: 33 additions & 2 deletions src/ezmsg/core/pubclient.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down Expand Up @@ -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()
Expand All @@ -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):
Expand Down Expand Up @@ -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)

Expand All @@ -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:
Expand Down Expand Up @@ -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())

Expand Down
24 changes: 23 additions & 1 deletion src/ezmsg/core/shm.py
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -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:
Expand Down
77 changes: 77 additions & 0 deletions src/ezmsg/core/shm_grow_test_support.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@
from dataclasses import dataclass

import json
import pickle

import ezmsg.core as ez

Expand Down Expand Up @@ -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()
30 changes: 30 additions & 0 deletions tests/test_shm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
32 changes: 31 additions & 1 deletion tests/test_shm_grow.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down Expand Up @@ -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)