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
21 changes: 16 additions & 5 deletions dapr/clients/grpc/subscription.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,8 @@

logger = logging.getLogger(__name__)

_SEND_QUEUE_POLL_SECONDS = 1.0


class Subscription:
def __init__(self, stub, pubsub_name, topic, metadata=None, dead_letter_topic=None):
Expand All @@ -31,6 +33,10 @@ def __init__(self, stub, pubsub_name, topic, metadata=None, dead_letter_topic=No
self._stream_lock = threading.Lock() # Protects _stream_active

def start(self):
# Each stream gets its own send queue so a request iterator left over from a
# previous stream cannot steal acknowledgements meant for the current one.
send_queue: queue.Queue = queue.Queue()

def outgoing_request_iterator():
"""
Generator function to create the request iterator for the stream.
Expand All @@ -49,22 +55,27 @@ def outgoing_request_iterator():
yield initial_request

# Start sending back acknowledgement messages from the send queue
while self._is_stream_active():
while self._is_stream_active() and self._send_queue is send_queue:
try:
# Wait for responses/acknowledgements to send from the send queue.
response = self._send_queue.get()
response = send_queue.get(timeout=_SEND_QUEUE_POLL_SECONDS)
yield response
except queue.Empty:
continue
except Exception as e:
raise Exception(f'Error while writing to stream: {e}')

# Create the bidirectional stream
self._stream = self._stub.SubscribeTopicEventsAlpha1(outgoing_request_iterator())
self._send_queue = send_queue
# gotcha: gRPC starts consuming the request iterator on its own thread before
# SubscribeTopicEventsAlpha1 returns. The stream must be marked active first,
# otherwise the iterator can end right after the initial request,
# half-closing the stream and making the sidecar drop the subscription with EOF.
self._set_stream_active()
try:
# Create the bidirectional stream
self._stream = self._stub.SubscribeTopicEventsAlpha1(outgoing_request_iterator())
next(self._stream) # type: ignore[arg-type] # discard the initial message
except Exception as e:
self._set_stream_inactive()
raise Exception(f'Error while initializing stream: {e}')

def reconnect_stream(self):
Expand Down
85 changes: 84 additions & 1 deletion tests/clients/test_subscription.py
Original file line number Diff line number Diff line change
@@ -1,8 +1,12 @@
import queue
import threading
import unittest
from collections.abc import Iterator

from google.protobuf.struct_pb2 import Struct

from dapr.clients.grpc.subscription import SubscriptionMessage
from dapr.clients.grpc.subscription import Subscription, SubscriptionMessage
from dapr.proto import api_v1
from dapr.proto.runtime.v1.appcallback_pb2 import TopicEventRequest


Expand Down Expand Up @@ -108,3 +112,82 @@ def test_subscription_message_init_unknown_content_type(self):

self.assertEqual(b'{"a": 1}', subscription_message.raw_data())
self.assertIsNone(subscription_message.data())


class _EagerStreamCall:
"""Fake bidi call that yields the initial response and records cancellation."""

def __init__(self) -> None:
self.cancelled = False
self._responses = iter([api_v1.SubscribeTopicEventsResponseAlpha1()])

def __iter__(self) -> '_EagerStreamCall':
return self

def __next__(self) -> api_v1.SubscribeTopicEventsResponseAlpha1:
return next(self._responses)

def cancel(self) -> bool:
self.cancelled = True
return True


class _EagerStub:
"""Fake stub that drains the request iterator on a background thread before returning.

This mimics gRPC's request-consumer thread pulling from the iterator before
``SubscribeTopicEventsAlpha1`` hands the call back to the caller.
"""

def __init__(self) -> None:
self.requests: queue.Queue = queue.Queue()
self.iterator_exhausted = threading.Event()

def SubscribeTopicEventsAlpha1(self, request_iterator: Iterator) -> _EagerStreamCall:
first_request_pulled = threading.Event()

def consume() -> None:
for request in request_iterator:
self.requests.put(request)
first_request_pulled.set()
self.iterator_exhausted.set()

threading.Thread(target=consume, daemon=True).start()
first_request_pulled.wait(timeout=_TEST_TIMEOUT_SECONDS)
# Give the consumer a chance to re-check the stream state before we return.
self.iterator_exhausted.wait(timeout=0.2)
return _EagerStreamCall()


_TEST_TIMEOUT_SECONDS = 5


class SubscriptionStreamTests(unittest.TestCase):
def test_request_stream_stays_open_when_consumer_runs_before_start_returns(self):
stub = _EagerStub()
subscription = Subscription(stub, 'pubsub', 'topic')
subscription.start()
try:
self.assertFalse(
stub.iterator_exhausted.is_set(),
'request iterator ended right after the initial request, half-closing the stream',
)

initial_request = stub.requests.get(timeout=_TEST_TIMEOUT_SECONDS)
self.assertTrue(initial_request.HasField('initial_request'))

event_message = SubscriptionMessage(TopicEventRequest(id='msg-1'))
subscription.respond_success(event_message)

ack_request = stub.requests.get(timeout=_TEST_TIMEOUT_SECONDS)
self.assertEqual('msg-1', ack_request.event_processed.id)
finally:
subscription.close()

def test_request_iterator_exits_after_close(self):
stub = _EagerStub()
subscription = Subscription(stub, 'pubsub', 'topic')
subscription.start()
subscription.close()

self.assertTrue(stub.iterator_exhausted.wait(timeout=_TEST_TIMEOUT_SECONDS))
6 changes: 5 additions & 1 deletion tests/integration/test_pubsub.py
Original file line number Diff line number Diff line change
Expand Up @@ -101,7 +101,11 @@ def test_streaming_subscribe_receives_published_message(client):

def read_next_message() -> None:
try:
next_message_future.set_result(subscription.next_message())
# next_message() returns None after a transient reconnect; keep reading.
message = subscription.next_message()
while message is None:
message = subscription.next_message()
next_message_future.set_result(message)
except Exception as exc:
next_message_future.set_exception(exc)

Expand Down
Loading