From ad9d6fb3c4c35f8c062bb046d7af2e7c87ff5f49 Mon Sep 17 00:00:00 2001 From: "Chris A. Evans" Date: Mon, 28 Sep 2026 11:27:51 -0400 Subject: [PATCH] Fix: Hold `pybmq::Session` by `shared_ptr` in `_ext.Session` `Session.stop()` stops the underlying `pybmq::Session` but does not destroy it. Destruction is left to `_ext.Session.__dealloc__`, which runs whenever the garbage collector gets to it, on whichever thread triggers the collection. Destroying a `pybmq::Session` blocks until the libbmq FSM thread finishes stopping it, and the FSM thread logs through the Python `logging` module before it can finish. If the collection runs on a thread that holds a `logging.Handler` lock, the destructor waits for the FSM thread and the FSM thread waits for the lock, and the process deadlocks. Any reference cycle hands a stopped `Session` to the collector; a traceback retained after an exception is a common one. This patch holds the `pybmq::Session` by `shared_ptr` and releases it in `stop()`, so a `Session` used as a context manager is destroyed when the `with` block exits. Each method takes its own copy of the pointer before calling into C++, so a concurrent `stop()` cannot destroy the session while it is in use. Calling a method after `stop()` now raises `Error`. Signed-off-by: Chris A. Evans --- src/blazingmq/_ext.pyx | 58 ++++++++++++++++++++----------- tests/unit/test_ext_session.py | 62 ++++++++++++++++++++++++++++++++++ 2 files changed, 101 insertions(+), 19 deletions(-) diff --git a/src/blazingmq/_ext.pyx b/src/blazingmq/_ext.pyx index 4d07e74..9373ef1 100644 --- a/src/blazingmq/_ext.pyx +++ b/src/blazingmq/_ext.pyx @@ -127,9 +127,12 @@ cdef TimeInterval create_time_interval(timeout: Optional[int|float]=None): cdef ensure_stop_session_impl(weakref_ext_session): - session = weakref_ext_session() - if session is not None: - (session)._session.stop(True) + cdef shared_ptr[NativeSession] session + ext_session = weakref_ext_session() + if ext_session is not None: + session = (ext_session)._session + if session.get() is not NULL: + session.get().stop(True) def ensure_stop_session(weakref_ext_session): @@ -156,7 +159,7 @@ cdef class FakeHostHealthMonitor: cdef class Session: cdef object __weakref__ - cdef NativeSession* _session + cdef shared_ptr[NativeSession] _session cdef readonly object monitor_host_health cdef readonly bint owned_by_session @@ -245,7 +248,7 @@ cdef class Session: config.close_queue_timeout = c_close_queue_timeout config.monitor_host_health = monitor_host_health - self._session = new NativeSession( + self._session = shared_ptr[NativeSession](new NativeSession( session_cb, message_cb, ack_cb, @@ -254,12 +257,24 @@ cdef class Session: Error, BrokerTimeoutError, _mock, - c_user_agent_prefix) - self._session.start(c_connect_timeout) + c_user_agent_prefix)) + self._session.get().start(c_connect_timeout) atexit.register(ensure_stop_session_impl, weakref.ref(self)) + cdef shared_ptr[NativeSession] _get_session(self) except *: + # Return a copy, so the session outlives a concurrent stop(). + cdef shared_ptr[NativeSession] session = self._session + if session.get() is NULL: + raise Error("Method called after session was stopped") + return session + def stop(self) -> None: - self._session.stop(False) + cdef shared_ptr[NativeSession] session = self._session + if session.get() is not NULL: + try: + session.get().stop(False) + finally: + self._session.reset() def set_owned_by_session(self): """Mark that a Session holds a strong reference to this object. @@ -306,6 +321,7 @@ cdef class Session: cdef optional[int] c_max_unconfirmed_bytes cdef optional[cppbool] c_suspends_on_bad_host_health cdef TimeInterval c_timeout = create_time_interval(timeout) + cdef shared_ptr[NativeSession] session = self._get_session() if b'\x00' in queue_uri: raise ValueError('queue_uri must not contain an embedded NUL byte') @@ -322,7 +338,7 @@ cdef class Session: if suspends_on_bad_host_health is not None: c_suspends_on_bad_host_health = optional[cppbool](suspends_on_bad_host_health) - self._session.open_queue_sync(queue_uri, + session.get().open_queue_sync(queue_uri, read, write, c_consumer_priority, @@ -344,6 +360,7 @@ cdef class Session: cdef optional[int] c_max_unconfirmed_bytes cdef optional[cppbool] c_suspends_on_bad_host_health cdef TimeInterval c_timeout = create_time_interval(timeout) + cdef shared_ptr[NativeSession] session = self._get_session() if b'\x00' in queue_uri: raise ValueError('queue_uri must not contain an embedded NUL byte') @@ -360,7 +377,7 @@ cdef class Session: if suspends_on_bad_host_health is not None: c_suspends_on_bad_host_health = optional[cppbool](suspends_on_bad_host_health) - self._session.configure_queue_sync(queue_uri, + session.get().configure_queue_sync(queue_uri, c_consumer_priority, c_max_unconfirmed_messages, c_max_unconfirmed_bytes, @@ -371,35 +388,38 @@ cdef class Session: queue_uri not None: bytes, timeout: Optional[int|float] = None) -> None: cdef TimeInterval c_timeout = create_time_interval(timeout) + cdef shared_ptr[NativeSession] session = self._get_session() if b'\x00' in queue_uri: raise ValueError('queue_uri must not contain an embedded NUL byte') - self._session.close_queue_sync(queue_uri, c_timeout) + session.get().close_queue_sync(queue_uri, c_timeout) def get_queue_options(self, queue_uri not None: bytes) -> object: + cdef shared_ptr[NativeSession] session = self._get_session() + if b'\x00' in queue_uri: raise ValueError('queue_uri must not contain an embedded NUL byte') - return self._session.get_queue_options(queue_uri) + return session.get().get_queue_options(queue_uri) def post(self, queue_uri not None: bytes, payload not None: bytes, properties=None, on_ack=None) -> None: + cdef shared_ptr[NativeSession] session = self._get_session() + if b'\x00' in queue_uri: raise ValueError('queue_uri must not contain an embedded NUL byte') - self._session.post(queue_uri, payload, len(payload), properties, on_ack) + session.get().post(queue_uri, payload, len(payload), properties, on_ack) def confirm(self, message not None) -> None: - self._session.confirm(message.queue_uri, message.guid, len(message.guid)) + cdef shared_ptr[NativeSession] session = self._get_session() + session.get().confirm(message.queue_uri, message.guid, len(message.guid)) def __dealloc__(self) -> None: - if self._session: - try: - self._session.stop(True) - finally: - del self._session + if self._session.get() is not NULL: + self._session.get().stop(True) diff --git a/tests/unit/test_ext_session.py b/tests/unit/test_ext_session.py index eb13bcd..91f1573 100644 --- a/tests/unit/test_ext_session.py +++ b/tests/unit/test_ext_session.py @@ -13,6 +13,7 @@ # See the License for the specific language governing permissions and # limitations under the License. +import gc import os import queue import sys @@ -132,6 +133,67 @@ def test_stopped_session_not_stopped_on_dealloc(): mock.stop.assert_called_once_with() +def test_stop_twice(): + # GIVEN + mock = sdk_mock(start=0, stop=None) + session = Session(dummy_callback, _mock=mock) + + # WHEN + session.stop() + session.stop() + + # THEN + mock.stop.assert_called_once_with() + + +@pytest.mark.parametrize( + "call_method", + [ + lambda session: session.post(QUEUE_NAME, b"payload"), + lambda session: session.get_queue_options(QUEUE_NAME), + lambda session: session.close_queue_sync(QUEUE_NAME), + lambda session: session.open_queue_sync(QUEUE_NAME, read=True, write=False), + lambda session: session.configure_queue_sync(QUEUE_NAME), + ], + ids=[ + "post", + "get_queue_options", + "close_queue_sync", + "open_queue_sync", + "configure_queue_sync", + ], +) +def test_method_after_stop_raises(call_method): + # GIVEN + mock = sdk_mock(start=0, stop=None) + session = Session(dummy_callback, _mock=mock) + session.stop() + + # WHEN + with pytest.raises(Exception) as exc: + call_method(session) + + # THEN + assert exc.type is exceptions.Error + assert exc.match("stopped") + + +def test_stop_releases_native_session(): + # GIVEN + mock = sdk_mock(start=0, stop=None) + mock_ref = weakref.ref(mock) + session = Session(dummy_callback, _mock=mock) + del mock + + # WHEN + session.stop() + gc.collect() + + # THEN + assert mock_ref() is None + del session + + def test_start_connect_timeout(): # GIVEN mock = sdk_mock(start=0, stop=None)