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
7 changes: 5 additions & 2 deletions streamz/core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
60 changes: 60 additions & 0 deletions streamz/tests/test_core.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()
Expand Down
Loading