diff --git a/dapr/clients/grpc/subscription.py b/dapr/clients/grpc/subscription.py index cd5798321..c2ef30bb2 100644 --- a/dapr/clients/grpc/subscription.py +++ b/dapr/clients/grpc/subscription.py @@ -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): @@ -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. @@ -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): diff --git a/tests/clients/test_subscription.py b/tests/clients/test_subscription.py index 21018aaac..1d9bc70c7 100644 --- a/tests/clients/test_subscription.py +++ b/tests/clients/test_subscription.py @@ -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 @@ -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)) diff --git a/tests/integration/test_pubsub.py b/tests/integration/test_pubsub.py index 8c07d5ae8..acdc54dbd 100644 --- a/tests/integration/test_pubsub.py +++ b/tests/integration/test_pubsub.py @@ -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)