Skip to content
52 changes: 41 additions & 11 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 @@ -108,10 +117,12 @@ def __aiter__(self) -> AsyncIterator[Message]:
async def __anext__(self) -> Message:
if self._terminated or self._pending_messages <= 0:
if not self._terminated:
await self._subscription.unsubscribe()
self._terminated = True
self._deregister_callbacks()
await self._subscription.unsubscribe()
raise StopAsyncIteration

delivering = False
try:
while True:
# Check heartbeat timeout (ADR-37: warn at 2x idle_heartbeat)
Expand Down Expand Up @@ -203,12 +214,21 @@ async def __anext__(self) -> Message:
)

self._pending_messages -= 1
delivering = True
return js_msg
except (StopAsyncIteration, asyncio.TimeoutError):
if not self._terminated:
await self._subscription.unsubscribe()
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.

finally:
# Any exit other than delivering a message terminates the batch:
# exhaustion, timeout, cancellation, or an unexpected error. All
# of them must release the subscription and the heartbeat
# callbacks, or the callbacks leak on the client. Deregister
# before the await: unsubscribing can itself be interrupted by a
# (second) cancellation.
if not delivering and not self._terminated:
self._terminated = True
self._deregister_callbacks()
await self._subscription.unsubscribe()


class PullMessageStream(MessageStream):
Expand All @@ -233,6 +253,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 +298,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 +317,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 +524,7 @@ async def _cleanup(self):
pass
self._heartbeat_task = None

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


Expand Down
81 changes: 81 additions & 0 deletions nats-jetstream/tests/test_consumer.py
Original file line number Diff line number Diff line change
Expand Up @@ -377,6 +377,87 @@ 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_fetch_deregisters_heartbeat_callbacks_on_cancellation(jetstream: JetStream):
"""Regression for #962: cancelling a heartbeat fetch mid-iteration releases
the callbacks too, not just normal exhaustion."""
client = jetstream._client
stream = await jetstream.create_stream(name="hb_cancel_stream", subjects=["HBCANCEL.*"])
consumer = await stream.create_consumer(name="hb_cancel_consumer")

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

# No messages published — iteration blocks until cancelled.
batch = await consumer.fetch(max_messages=5, max_wait=5.0, heartbeat=1.0)
assert len(client._disconnected_callbacks) == disconnected + 1
assert len(client._reconnected_callbacks) == reconnected + 1

async def consume() -> None:
async for _ in batch:
pass

task = asyncio.create_task(consume())
await asyncio.sleep(0.2)
task.cancel()
with pytest.raises(asyncio.CancelledError):
await task

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