Skip to content
Merged
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
58 changes: 39 additions & 19 deletions src/blazingmq/_ext.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -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)._session.stop(True)
cdef shared_ptr[NativeSession] session
ext_session = weakref_ext_session()
if ext_session is not None:
session = (<Session?>ext_session)._session
if session.get() is not NULL:
session.get().stop(True)


def ensure_stop_session(weakref_ext_session):
Expand All @@ -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

Expand Down Expand Up @@ -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,
Expand All @@ -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.
Expand Down Expand Up @@ -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')
Expand All @@ -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,
Expand All @@ -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')
Expand All @@ -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,
Expand All @@ -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)
62 changes: 62 additions & 0 deletions tests/unit/test_ext_session.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down
Loading