diff --git a/streamz/core.py b/streamz/core.py index 81d89832..6d49d5a8 100644 --- a/streamz/core.py +++ b/streamz/core.py @@ -1939,10 +1939,13 @@ def update(self, x, who=None, metadata=None): def flush(self, _=None): out = tuple(self.cache) metadata = list(self.metadata_cache) - self._emit(out, metadata) - self._release_refs(metadata) + # Downstream callbacks may collect more values while emit waits. self.cache.clear() self.metadata_cache.clear() + try: + return self.emit(out, metadata=metadata) + finally: + self._release_refs(metadata) @Stream.register_api() diff --git a/streamz/tests/test_core.py b/streamz/tests/test_core.py index 6a8ca3ab..67261255 100644 --- a/streamz/tests/test_core.py +++ b/streamz/tests/test_core.py @@ -1061,6 +1061,66 @@ def test_collect(): assert L == [(1, 2), (), (3,)] +@pytest.mark.parametrize("triggered", [False, True]) +def test_collect_async_sink(triggered): + source = Stream(asynchronous=False) + trigger = Stream(asynchronous=False) + collector = source.collect() + received = [] + + async def write(value): + await asyncio.sleep(0) + received.append(value) + + collector.sink(write) + trigger.sink(collector.flush) + source.emit(1) + source.emit(2) + if triggered: + trigger.emit(None) + else: + collector.flush() + assert received == [(1, 2)] + + +def test_collect_async_sink_can_collect_more(): + source = Stream(asynchronous=False) + collector = source.collect() + received = [] + + async def write(value): + await asyncio.sleep(0) + received.append(value) + if value == (1,): + await source.emit(2) + + collector.sink(write) + source.emit(1) + collector.flush() + collector.flush() + assert received == [(1,), (2,)] + + +def test_collect_await_flush(): + async def run(): + source = Stream(asynchronous=True) + collector = source.collect() + received = [] + + async def write(value): + await asyncio.sleep(0) + received.append(value) + + collector.sink(write) + await source.emit(1) + await collector.flush() + assert received == [(1,)] + await collector.flush() + assert received == [(1,), ()] + + asyncio.run(run()) + + def test_collect_ref_counts(): source = Stream() collector = source.collect()