Skip to content
Open
Show file tree
Hide file tree
Changes from 1 commit
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
12 changes: 11 additions & 1 deletion fastembed/parallel_processor.py
Original file line number Diff line number Diff line change
Expand Up @@ -153,6 +153,7 @@ def ordered_map(self, stream: Iterable[Any], *args: Any, **kwargs: Any) -> Itera
def semi_ordered_map(
self, stream: Iterable[Any], *args: Any, **kwargs: Any
) -> Iterable[tuple[int, Any]]:
completed = False
Comment thread
coderabbitai[bot] marked this conversation as resolved.
try:
self.start(**kwargs)

Expand Down Expand Up @@ -196,10 +197,19 @@ def semi_ordered_map(
raise RuntimeError("Thread unexpectedly terminated")
yield out_item
read += 1
completed = True
finally:
assert self.input_queue is not None, "Input queue is None"
assert self.output_queue is not None, "Output queue is None"
self.join()
if completed:
self.join()
else:
# The generator may be closed before the input stream is exhausted (for
# example, when a caller only consumes a prefix of the embeddings). In that
# case workers have not necessarily received their stop signals and a normal
# join would wait forever for processes blocked on the input queue.
self.emergency_shutdown = True
self.join_or_terminate()
Comment thread
coderabbitai[bot] marked this conversation as resolved.
self.input_queue.close()
self.output_queue.close()
if self.emergency_shutdown:
Expand Down
73 changes: 73 additions & 0 deletions tests/test_parallel_processor.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
from queue import Empty
from unittest.mock import MagicMock

import pytest

from fastembed.parallel_processor import ParallelWorkerPool, Worker


class EchoWorker(Worker):
@classmethod
def start(cls, **kwargs):
return cls()

def process(self, items):
yield from items


def make_stubbed_pool(monkeypatch):
pool = ParallelWorkerPool(1, EchoWorker)
input_queue = MagicMock()
output_queue = MagicMock()
output_queue.get_nowait.side_effect = Empty
output_queue.get.return_value = (0, "result")
process = MagicMock()
process.is_alive.return_value = True

def start(**kwargs):
pool.input_queue = input_queue
pool.output_queue = output_queue
pool.processes = [process]

monkeypatch.setattr(pool, "start", start)
return pool, input_queue, output_queue, process


def test_closing_partial_parallel_iterator_terminates_workers(monkeypatch):
pool, input_queue, output_queue, process = make_stubbed_pool(monkeypatch)
results = pool.ordered_map(["input"])

assert next(results) == "result"
results.close()

process.join.assert_called_once_with(timeout=1)
process.terminate.assert_called_once_with()
input_queue.cancel_join_thread.assert_called_once_with()
output_queue.cancel_join_thread.assert_called_once_with()


def test_parallel_iterator_terminates_workers_when_input_raises(monkeypatch):
pool, input_queue, output_queue, process = make_stubbed_pool(monkeypatch)

def failing_stream():
yield "input"
raise RuntimeError("stream failed")

with pytest.raises(RuntimeError, match="stream failed"):
list(pool.semi_ordered_map(failing_stream()))

process.join.assert_called_once_with(timeout=1)
process.terminate.assert_called_once_with()
input_queue.cancel_join_thread.assert_called_once_with()
output_queue.cancel_join_thread.assert_called_once_with()


def test_exhausted_parallel_iterator_joins_workers_gracefully(monkeypatch):
pool, input_queue, output_queue, process = make_stubbed_pool(monkeypatch)

assert list(pool.semi_ordered_map(["input"])) == [(0, "result")]

process.join.assert_called_once_with()
process.terminate.assert_not_called()
input_queue.join_thread.assert_called_once_with()
output_queue.join_thread.assert_called_once_with()