Skip to content
12 changes: 12 additions & 0 deletions src/conductor/client/automator/task_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -331,9 +331,17 @@ def __init__(
self._next_restart_at: List[float] = [0.0 for _ in self.workers]
# Lock to protect process list during concurrent access (monitor thread vs main thread)
self._process_lock = threading.Lock()
self._processes_started = False
logger.info("TaskHandler initialized")

def __enter__(self):
try:
self.start_processes()
except BaseException:
# __exit__ is not called if __enter__ raises, so clean up any
# partially-spawned workers here before propagating.
self.stop_processes()
raise
return self

def __exit__(self, exc_type, exc_value, traceback):
Expand All @@ -347,11 +355,14 @@ def stop_processes(self) -> None:
with self._process_lock:
self.__stop_task_runner_processes()
self.__stop_metrics_provider_process()
self._processes_started = False
logger.info("Stopped worker processes...")
self.queue.put(None)
self.logger_process.terminate()

def start_processes(self) -> None:
if self._processes_started:
return
logger.info("Starting worker processes...")
freeze_support()
self._monitor_stop_event.clear()
Expand All @@ -376,6 +387,7 @@ def start_processes(self) -> None:
self.stop_processes()
raise
self.__start_monitor_thread()
self._processes_started = True
logger.info("Started all processes")

def join_processes(self) -> None:
Expand Down
7 changes: 6 additions & 1 deletion tests/unit/automator/test_task_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,8 +57,13 @@ def test_metrics_directory_cleaned_once_in_parent_init(self):
workers=[ClassWorker('task')],
metrics_settings=metrics_settings,
)
with task_handler:
try:
# Assert before any processes are started — the mock on
# metrics_settings is not picklable, so we must not call
# start_processes() (which __enter__ now does automatically).
metrics_settings.clean_metrics_directory.assert_called_once_with()
finally:
task_handler.stop_processes()


def _get_valid_task_handler():
Expand Down
15 changes: 14 additions & 1 deletion tests/unit/automator/test_task_handler_coverage.py
Original file line number Diff line number Diff line change
Expand Up @@ -663,12 +663,19 @@ def test_context_manager_enter(self, mock_process_class, mock_import, mock_loggi

@patch('conductor.client.automator.task_handler._setup_logging_queue')
@patch('importlib.import_module')
def test_context_manager_exit(self, mock_import, mock_logging):
@patch('conductor.client.automator.task_handler.Process')
def test_context_manager_exit(self, mock_process_class, mock_import, mock_logging):
"""Test context manager __exit__ method."""
mock_queue = Mock()
mock_logger_process = Mock()
mock_logging.return_value = (mock_logger_process, mock_queue)

mock_process = Mock()
mock_process.terminate = Mock()
mock_process.kill = Mock()
mock_process.is_alive = Mock(return_value=False)
mock_process_class.return_value = mock_process

worker = ClassWorker('test_task')
handler = TaskHandler(
workers=[worker],
Expand All @@ -679,10 +686,16 @@ def test_context_manager_exit(self, mock_import, mock_logging):
# Override the queue and logger_process with fresh mocks
handler.queue = Mock()
handler.logger_process = Mock()
handler.logger_process.is_alive = Mock(return_value=False)
handler.metrics_provider_process = Mock()
handler.metrics_provider_process.terminate = Mock()
handler.metrics_provider_process.is_alive = Mock(return_value=False)

# Mock terminate on all processes
for process in handler.task_runner_processes:
process.terminate = Mock()
process.kill = Mock()
process.is_alive = Mock(return_value=False)

with handler:
pass
Expand Down
Loading