diff --git a/src/conductor/client/automator/task_handler.py b/src/conductor/client/automator/task_handler.py index 7256d59f..df9aafe3 100644 --- a/src/conductor/client/automator/task_handler.py +++ b/src/conductor/client/automator/task_handler.py @@ -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): @@ -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() @@ -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: diff --git a/tests/unit/automator/test_task_handler.py b/tests/unit/automator/test_task_handler.py index 51d50a7a..b109052b 100644 --- a/tests/unit/automator/test_task_handler.py +++ b/tests/unit/automator/test_task_handler.py @@ -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(): diff --git a/tests/unit/automator/test_task_handler_coverage.py b/tests/unit/automator/test_task_handler_coverage.py index b97e0083..dd85af22 100644 --- a/tests/unit/automator/test_task_handler_coverage.py +++ b/tests/unit/automator/test_task_handler_coverage.py @@ -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], @@ -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