Skip to content
40 changes: 40 additions & 0 deletions nats-core/src/nats/client/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -1544,6 +1544,16 @@ def add_disconnected_callback(self, callback: Callable[[], None]) -> None:
"""
self._disconnected_callbacks.append(callback)

def remove_disconnected_callback(self, callback: Callable[[], None]) -> None:
"""Remove a previously registered disconnected callback.

Raises ``ValueError`` if ``callback`` was not registered.

Args:
callback: Function previously passed to :meth:`add_disconnected_callback`.
"""
self._disconnected_callbacks.remove(callback)

def add_reconnected_callback(self, callback: Callable[[], None]) -> None:
"""Add a callback to be invoked when the client is reconnected.

Expand All @@ -1552,6 +1562,16 @@ def add_reconnected_callback(self, callback: Callable[[], None]) -> None:
"""
self._reconnected_callbacks.append(callback)

def remove_reconnected_callback(self, callback: Callable[[], None]) -> None:
"""Remove a previously registered reconnected callback.

Raises ``ValueError`` if ``callback`` was not registered.

Args:
callback: Function previously passed to :meth:`add_reconnected_callback`.
"""
self._reconnected_callbacks.remove(callback)

def add_error_callback(self, callback: Callable[[Exception | str], None]) -> None:
"""Add a callback to be invoked when the client encounters an error.

Expand All @@ -1560,6 +1580,16 @@ def add_error_callback(self, callback: Callable[[Exception | str], None]) -> Non
"""
self._error_callbacks.append(callback)

def remove_error_callback(self, callback: Callable[[Exception | str], None]) -> None:
"""Remove a previously registered error callback.

Raises ``ValueError`` if ``callback`` was not registered.

Args:
callback: Function previously passed to :meth:`add_error_callback`.
"""
self._error_callbacks.remove(callback)

def add_lame_duck_mode_callback(self, callback: Callable[[], None]) -> None:
"""Add a callback to be invoked when the server enters lame duck mode.

Expand All @@ -1579,6 +1609,16 @@ def add_lame_duck_mode_callback(self, callback: Callable[[], None]) -> None:
"""
self._lame_duck_mode_callbacks.append(callback)

def remove_lame_duck_mode_callback(self, callback: Callable[[], None]) -> None:
"""Remove a previously registered lame duck mode callback.

Raises ``ValueError`` if ``callback`` was not registered.

Args:
callback: Function previously passed to :meth:`add_lame_duck_mode_callback`.
"""
self._lame_duck_mode_callbacks.remove(callback)


def _setup_nkey_auth(
nkey: str | Path | tuple[Callable[[], str], Callable[[str], bytes]],
Expand Down
148 changes: 148 additions & 0 deletions nats-core/tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -3595,3 +3595,151 @@ async def test_force_reconnect_raises_when_drained(server):

with pytest.raises(ConnectionError):
await client.force_reconnect()


@pytest.mark.asyncio
async def test_remove_disconnected_callback_skips_invocation():
"""A removed disconnected callback is not invoked when the client disconnects."""
server = await run(port=0)

client = await connect(
server.client_url,
timeout=1.0,
allow_reconnect=True,
reconnect_time_wait=0.1,
)

removed_called = False
kept_called = asyncio.Event()

def removed_cb():
nonlocal removed_called
removed_called = True

def kept_cb():
kept_called.set()

client.add_disconnected_callback(removed_cb)
client.add_disconnected_callback(kept_cb)
client.remove_disconnected_callback(removed_cb)

try:
await server.shutdown()
await asyncio.wait_for(kept_called.wait(), timeout=2.0)
assert removed_called is False
finally:
await client.close()


@pytest.mark.asyncio
async def test_remove_reconnected_callback_skips_invocation():
"""A removed reconnected callback is not invoked when the client reconnects."""
server = await run(port=0)
server_port = server.port

client = await connect(
server.client_url,
timeout=1.0,
allow_reconnect=True,
reconnect_time_wait=0.1,
)

removed_called = False
kept_called = asyncio.Event()

def removed_cb():
nonlocal removed_called
removed_called = True

def kept_cb():
kept_called.set()

client.add_reconnected_callback(removed_cb)
client.add_reconnected_callback(kept_cb)
client.remove_reconnected_callback(removed_cb)

try:
await server.shutdown()
new_server = await run(port=server_port)
try:
await asyncio.wait_for(kept_called.wait(), timeout=5.0)
assert removed_called is False
finally:
await new_server.shutdown()
finally:
await client.close()


@pytest.mark.asyncio
async def test_remove_error_callback_skips_invocation(client):
"""A removed error callback is not invoked when the client surfaces an error."""
test_subject = f"test.remove_error_callback.{uuid.uuid4()}"

removed_called = False
kept_errors: list[Exception | str] = []

def removed_cb(_error):
nonlocal removed_called
removed_called = True

def kept_cb(error):
if isinstance(error, SlowConsumerError):
kept_errors.append(error)

client.add_error_callback(removed_cb)
client.add_error_callback(kept_cb)
client.remove_error_callback(removed_cb)

await client.subscribe(test_subject, max_pending_messages=5)
await client.flush()

for i in range(20):
await client.publish(test_subject, f"message-{i}".encode())
await client.flush()
await asyncio.sleep(0.2)

assert len(kept_errors) == 1
assert removed_called is False


@pytest.mark.skipif(sys.platform == "win32", reason="SIGUSR2 is POSIX only")
@pytest.mark.asyncio
async def test_remove_lame_duck_mode_callback_skips_invocation(client, server):
"""A removed lame duck mode callback is not invoked when LDM is signalled."""
removed_called = False
kept_called = asyncio.Event()

def removed_cb():
nonlocal removed_called
removed_called = True

def kept_cb():
kept_called.set()

client.add_lame_duck_mode_callback(removed_cb)
client.add_lame_duck_mode_callback(kept_cb)
client.remove_lame_duck_mode_callback(removed_cb)

server.lame_duck_mode()
await asyncio.wait_for(kept_called.wait(), timeout=5.0)
assert removed_called is False


@pytest.mark.asyncio
async def test_remove_callback_raises_when_not_registered(client):
"""remove_*_callback raises ValueError when the callback was never registered."""

def never_registered():
pass

def never_registered_error(_error):
pass

with pytest.raises(ValueError):
client.remove_disconnected_callback(never_registered)
with pytest.raises(ValueError):
client.remove_reconnected_callback(never_registered)
with pytest.raises(ValueError):
client.remove_error_callback(never_registered_error)
with pytest.raises(ValueError):
client.remove_lame_duck_mode_callback(never_registered)
35 changes: 28 additions & 7 deletions nats-jetstream/src/nats/jetstream/consumer/pull.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@
from nats.jetstream.util import new_inbox

if TYPE_CHECKING:
from nats.client import Subscription
from nats.client import Client, Subscription
from nats.client.message import Message as ClientMessage
from nats.jetstream.stream import Stream

Expand Down Expand Up @@ -59,6 +59,7 @@ class PullMessageBatch(MessageBatch):
_heartbeat_deadline: float | None
_heartbeat_paused: bool
_heartbeat_remaining: float | None
_client: Client | None

def __init__(
self,
Expand All @@ -80,10 +81,11 @@ def __init__(
self._heartbeat_remaining = None

# Register disconnect/reconnect callbacks for heartbeat timer (ADR-37)
self._client = None
if heartbeat is not None:
client = jetstream._client
client.add_disconnected_callback(self._pause_heartbeat_timer)
client.add_reconnected_callback(self._resume_heartbeat_timer)
self._client = jetstream._client
self._client.add_disconnected_callback(self._pause_heartbeat_timer)
self._client.add_reconnected_callback(self._resume_heartbeat_timer)

def _pause_heartbeat_timer(self) -> None:
"""Pause the heartbeat timer on disconnect (ADR-37)."""
Expand All @@ -98,6 +100,13 @@ def _resume_heartbeat_timer(self) -> None:
self._heartbeat_paused = False
self._heartbeat_remaining = None

def _deregister_callbacks(self) -> None:
"""Remove the heartbeat callbacks registered on the client (ADR-37)."""
if self._client is not None:
self._client.remove_disconnected_callback(self._pause_heartbeat_timer)
self._client.remove_reconnected_callback(self._resume_heartbeat_timer)
self._client = None

@property
def error(self) -> Exception | None:
return self._error
Expand All @@ -109,6 +118,7 @@ async def __anext__(self) -> Message:
if self._terminated or self._pending_messages <= 0:
if not self._terminated:
await self._subscription.unsubscribe()
self._deregister_callbacks()
self._terminated = True
raise StopAsyncIteration

Expand Down Expand Up @@ -207,6 +217,7 @@ async def __anext__(self) -> Message:
except (StopAsyncIteration, asyncio.TimeoutError):
if not self._terminated:
await self._subscription.unsubscribe()
self._deregister_callbacks()
self._terminated = True
raise StopAsyncIteration

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The except (StopAsyncIteration, asyncio.TimeoutError) clause doesn't catch asyncio.CancelledError (a BaseException since Python 3.8). If the task iterating the batch is cancelled mid-flight, _deregister_callbacks() is skipped and the heartbeat callbacks leak — the same bug this PR fixes for the normal termination path.

This is technically a pre-existing gap (the subscription also isn't unsubscribed on cancellation), but it means the fix is incomplete for cancellation. Consider a try/finally wrapping both the unsubscribe and deregister calls.


Expand All @@ -233,6 +244,7 @@ class PullMessageStream(MessageStream):
_heartbeat_deadline: float | None
_heartbeat_paused: bool
_heartbeat_remaining: float | None
_client: Client | None

def __init__(
self,
Expand Down Expand Up @@ -277,10 +289,11 @@ def __init__(
self._heartbeat_deadline = time.time() + (heartbeat * 2) if heartbeat is not None else None

# Register disconnect/reconnect callbacks for heartbeat timer (ADR-37)
self._client = None
if heartbeat is not None:
client = consumer._stream._jetstream._client
client.add_disconnected_callback(self._pause_heartbeat_timer)
client.add_reconnected_callback(self._resume_heartbeat_timer)
self._client = consumer._stream._jetstream._client
self._client.add_disconnected_callback(self._pause_heartbeat_timer)
self._client.add_reconnected_callback(self._resume_heartbeat_timer)

def _pause_heartbeat_timer(self) -> None:
"""Pause the heartbeat timer on disconnect (ADR-37)."""
Expand All @@ -295,6 +308,13 @@ def _resume_heartbeat_timer(self) -> None:
self._heartbeat_paused = False
self._heartbeat_remaining = None

def _deregister_callbacks(self) -> None:
"""Remove the heartbeat callbacks registered on the client (ADR-37)."""
if self._client is not None:
self._client.remove_disconnected_callback(self._pause_heartbeat_timer)
self._client.remove_reconnected_callback(self._resume_heartbeat_timer)
self._client = None

@property
def is_active(self) -> bool:
"""Check if the message stream is still active."""
Expand Down Expand Up @@ -495,6 +515,7 @@ async def _cleanup(self):
pass
self._heartbeat_task = None

self._deregister_callbacks()
await self._subscription.unsubscribe()


Expand Down
51 changes: 51 additions & 0 deletions nats-jetstream/tests/test_consumer.py
Original file line number Diff line number Diff line change
Expand Up @@ -377,6 +377,57 @@ async def collect_messages():
await message_stream.stop()


@pytest.mark.asyncio
async def test_messages_deregisters_heartbeat_callbacks_on_stop(jetstream: JetStream):
"""Regression for #962: stopping a heartbeat message stream removes the
disconnect/reconnect callbacks it registered, so repeatedly creating and
stopping streams over one connection does not leak callbacks."""
client = jetstream._client
stream = await jetstream.create_stream(name="hb_leak_stream", subjects=["HBLEAK.*"])
consumer = await stream.create_consumer(name="hb_leak_consumer")

disconnected = len(client._disconnected_callbacks)
reconnected = len(client._reconnected_callbacks)

for _ in range(5):
message_stream = await consumer.messages(max_messages=10, heartbeat=5.0)
# Registration is observable while the stream is active.
assert len(client._disconnected_callbacks) == disconnected + 1
assert len(client._reconnected_callbacks) == reconnected + 1
await message_stream.stop()
# Stopping deregisters exactly what it registered, leaving no residue.
assert len(client._disconnected_callbacks) == disconnected
assert len(client._reconnected_callbacks) == reconnected

# A second stop() is a no-op and does not raise.
await message_stream.stop()
assert len(client._disconnected_callbacks) == disconnected
assert len(client._reconnected_callbacks) == reconnected


@pytest.mark.asyncio
async def test_fetch_deregisters_heartbeat_callbacks_on_exhaustion(jetstream: JetStream):
"""Regression for #962: a heartbeat fetch batch deregisters its callbacks
once the batch is exhausted (StopAsyncIteration), not just on stream stop()."""
client = jetstream._client
stream = await jetstream.create_stream(name="hb_fetch_stream", subjects=["HBFETCH.*"])
consumer = await stream.create_consumer(name="hb_fetch_consumer")

disconnected = len(client._disconnected_callbacks)
reconnected = len(client._reconnected_callbacks)

# No messages published — the batch ends via timeout, exhausting the iterator.
batch = await consumer.fetch(max_messages=5, max_wait=0.5, heartbeat=1.0)
assert len(client._disconnected_callbacks) == disconnected + 1
assert len(client._reconnected_callbacks) == reconnected + 1

async for _ in batch:
pass

assert len(client._disconnected_callbacks) == disconnected
assert len(client._reconnected_callbacks) == reconnected


@pytest.mark.asyncio
async def test_messages_rejects_both_max_messages_and_max_bytes(jetstream: JetStream):
"""Test ADR-37: messages() cannot accept both max_messages and max_bytes simultaneously."""
Expand Down
Loading