diff --git a/clients/python/src/taskbroker_client/worker/worker.py b/clients/python/src/taskbroker_client/worker/worker.py index 0cce2164..a46a1310 100644 --- a/clients/python/src/taskbroker_client/worker/worker.py +++ b/clients/python/src/taskbroker_client/worker/worker.py @@ -415,16 +415,17 @@ def start(self) -> int: self.worker_pool.start_result_thread() self.worker_pool.start_spawn_children_thread() - # Convert signals into KeyboardInterrupt. - # Running shutdown() within the signal handler can lead to deadlocks + # Signal shutdown without raising an exception while multiprocessing + # synchronization primitives may be held. server: grpc.Server | None = None + server_started = False health_servicer: health.HealthServicer | None = None def signal_handler(*args: Any) -> None: - if server: + self._grpc_sync_event.set() + if server is not None and server_started: server.stop(grace=5) - raise KeyboardInterrupt() signal.signal(signal.SIGINT, signal_handler) signal.signal(signal.SIGTERM, signal_handler) @@ -459,7 +460,13 @@ def signal_handler(*args: Any) -> None: health_servicer.set(WORKER_SERVICE_NAME, health_pb2.HealthCheckResponse.NOT_SERVING) server.add_insecure_port(f"[::]:{self._grpc_port}") + if self._grpc_sync_event.is_set(): + return 0 + server.start() + server_started = True + if self._grpc_sync_event.is_set(): + return 0 # Hold NOT_SERVING until children are warm so the pod stays out of # the NEG/readiness set while its child processes are still loading. @@ -477,18 +484,14 @@ def signal_handler(*args: Any) -> None: logger.info("taskworker.grpc_server.started", extra={"port": self._grpc_port}) self._start_health_check_thread() - try: - server.wait_for_termination() - except KeyboardInterrupt: - # Signals are converted to KeyboardInterrupt, swallow for exit code 0 - pass + server.wait_for_termination() finally: if health_servicer is not None: health_servicer.set("", health_pb2.HealthCheckResponse.NOT_SERVING) health_servicer.set(WORKER_SERVICE_NAME, health_pb2.HealthCheckResponse.NOT_SERVING) - if server is not None: + if server is not None and server_started: server.stop(grace=5) self._stop_health_check_thread() @@ -595,20 +598,21 @@ def start(self) -> int: self.worker_pool.start_result_thread() self.worker_pool.start_spawn_children_thread() - # Convert signals into KeyboardInterrupt. - # Running shutdown() within the signal handler can lead to deadlocks + # Signal shutdown without raising an exception while multiprocessing + # synchronization primitives may be held. def signal_handler(*args: Any) -> None: - raise KeyboardInterrupt() + self._grpc_sync_event.set() signal.signal(signal.SIGINT, signal_handler) signal.signal(signal.SIGTERM, signal_handler) try: - while True: + while not self._grpc_sync_event.is_set(): self.run_once() - except KeyboardInterrupt: + finally: self.shutdown() - raise + + return 0 def run_once(self) -> None: """Access point for tests to run a single worker loop""" @@ -667,6 +671,8 @@ def _send_update_task( ) self._grpc_sync_event.wait(self._setstatus_backoff_seconds) + if self._grpc_sync_event.is_set(): + fetch_next = None try: next_task = self.client.update_task(result, fetch_next) @@ -691,6 +697,9 @@ def _send_update_task( def fetch_task(self) -> InflightTaskActivation | None: self._grpc_sync_event.wait(self._gettask_backoff_seconds) + if self._grpc_sync_event.is_set(): + return None + try: activation = self.client.get_task(self._namespace) except grpc.RpcError as e: diff --git a/clients/python/tests/worker/test_worker.py b/clients/python/tests/worker/test_worker.py index 7253a74e..640dcc1f 100644 --- a/clients/python/tests/worker/test_worker.py +++ b/clients/python/tests/worker/test_worker.py @@ -30,6 +30,7 @@ TASK_ACTIVATION_STATUS_COMPLETE, TASK_ACTIVATION_STATUS_FAILURE, TASK_ACTIVATION_STATUS_RETRY, + FetchNextTask, PushTaskRequest, PushTaskResponse, RetryState, @@ -453,6 +454,36 @@ def test_fetch_task(self) -> None: assert task assert task.activation.id == SIMPLE_TASK.activation.id + def test_fetch_task_skips_request_during_shutdown(self) -> None: + taskworker = TaskWorker( + app_module="examples.app:app", + broker_hosts=["127.0.0.1:50051"], + max_child_task_count=100, + process_type="fork", + ) + taskworker._grpc_sync_event.set() + + with mock.patch.object(taskworker.client, "get_task") as mock_get: + task = taskworker.fetch_task() + + assert task is None + mock_get.assert_not_called() + + def test_send_update_does_not_fetch_next_during_shutdown(self) -> None: + taskworker = TaskWorker( + app_module="examples.app:app", + broker_hosts=["127.0.0.1:50051"], + max_child_task_count=100, + process_type="fork", + ) + taskworker._grpc_sync_event.set() + result = _make_processing_result("completed") + + with mock.patch.object(taskworker.client, "update_task", return_value=None) as update: + taskworker._send_update_task(result, FetchNextTask(namespace="examples")) + + update.assert_called_once_with(result, None) + def test_fetch_no_task(self) -> None: taskworker = TaskWorker( app_module="examples.app:app", @@ -710,6 +741,37 @@ def update_task_response(*args: Any, **kwargs: Any) -> None: assert redis.get("no-retries-remaining"), "key should exist if except block was hit" redis.delete("no-retries-remaining") + def test_start_handles_sigterm_without_raising(self) -> None: + taskworker = TaskWorker( + app_module="examples.app:app", + broker_hosts=["127.0.0.1:50051"], + max_child_task_count=100, + process_type="fork", + ) + handlers: dict[signal.Signals, Callable[..., None]] = {} + + def install_handler(signum: signal.Signals, handler: Callable[..., None]) -> None: + handlers[signum] = handler + + def request_shutdown() -> None: + handlers[signal.SIGTERM]() + + with ( + mock.patch( + "taskbroker_client.worker.worker.signal.signal", side_effect=install_handler + ), + mock.patch.object(taskworker.worker_pool, "start_metrics_thread"), + mock.patch.object(taskworker.worker_pool, "start_result_thread"), + mock.patch.object(taskworker.worker_pool, "start_spawn_children_thread"), + mock.patch.object(taskworker.worker_pool, "shutdown") as shutdown, + mock.patch.object(taskworker, "run_once", side_effect=request_shutdown), + ): + exitcode = taskworker.start() + + assert exitcode == 0 + assert taskworker._grpc_sync_event.is_set() + shutdown.assert_called_once_with() + def test_constructor_push_mode(self) -> None: taskworker = PushTaskWorker( app_module="examples.app:app", @@ -831,6 +893,82 @@ def warm_up() -> None: assert timeout_calls == [] +def test_push_worker_handles_sigterm_without_raising() -> None: + taskworker = _make_push_worker(concurrency=2, warmup_timeout=5) + handlers: dict[signal.Signals, Callable[..., None]] = {} + + def install_handler(signum: signal.Signals, handler: Callable[..., None]) -> None: + handlers[signum] = handler + + fake_health = mock.MagicMock() + fake_server = mock.MagicMock() + fake_server.wait_for_termination.side_effect = lambda: handlers[signal.SIGTERM]() + + with ( + mock.patch("taskbroker_client.worker.worker.signal.signal", side_effect=install_handler), + mock.patch.object(taskworker.worker_pool, "start_metrics_thread"), + mock.patch.object(taskworker.worker_pool, "start_result_thread"), + mock.patch.object(taskworker.worker_pool, "start_spawn_children_thread"), + mock.patch.object(taskworker.worker_pool, "shutdown") as shutdown, + mock.patch.object(taskworker, "_start_health_check_thread"), + mock.patch.object(taskworker, "_stop_health_check_thread"), + mock.patch.object( + TaskWorkerProcessingPool, "ready_count", new_callable=mock.PropertyMock, return_value=2 + ), + mock.patch("taskbroker_client.worker.worker.grpc.server", return_value=fake_server), + mock.patch( + "taskbroker_client.worker.worker.health.HealthServicer", return_value=fake_health + ), + mock.patch("taskbroker_client.worker.worker.health_pb2_grpc.add_HealthServicer_to_server"), + mock.patch( + "taskbroker_client.worker.worker.taskbroker_pb2_grpc" + ".add_WorkerServiceServicer_to_server" + ), + ): + exitcode = taskworker.start() + + assert exitcode == 0 + assert taskworker._grpc_sync_event.is_set() + fake_server.stop.assert_called_with(grace=5) + shutdown.assert_called_once_with() + + +def test_push_worker_does_not_start_server_when_signal_arrives_during_setup() -> None: + taskworker = _make_push_worker(concurrency=2, warmup_timeout=5) + handlers: dict[signal.Signals, Callable[..., None]] = {} + + def install_handler(signum: signal.Signals, handler: Callable[..., None]) -> None: + handlers[signum] = handler + + fake_health = mock.MagicMock() + fake_server = mock.MagicMock() + fake_server.add_insecure_port.side_effect = lambda *args: handlers[signal.SIGTERM]() + + with ( + mock.patch("taskbroker_client.worker.worker.signal.signal", side_effect=install_handler), + mock.patch.object(taskworker.worker_pool, "start_metrics_thread"), + mock.patch.object(taskworker.worker_pool, "start_result_thread"), + mock.patch.object(taskworker.worker_pool, "start_spawn_children_thread"), + mock.patch.object(taskworker.worker_pool, "shutdown") as shutdown, + mock.patch.object(taskworker, "_stop_health_check_thread"), + mock.patch("taskbroker_client.worker.worker.grpc.server", return_value=fake_server), + mock.patch( + "taskbroker_client.worker.worker.health.HealthServicer", return_value=fake_health + ), + mock.patch("taskbroker_client.worker.worker.health_pb2_grpc.add_HealthServicer_to_server"), + mock.patch( + "taskbroker_client.worker.worker.taskbroker_pb2_grpc" + ".add_WorkerServiceServicer_to_server" + ), + ): + exitcode = taskworker.start() + + assert exitcode == 0 + fake_server.start.assert_not_called() + fake_server.stop.assert_not_called() + shutdown.assert_called_once_with() + + def test_start_does_not_serve_when_shutdown_during_warmup() -> None: from grpc_health.v1 import health_pb2 @@ -866,7 +1004,8 @@ def test_start_does_not_serve_when_shutdown_during_warmup() -> None: if c.args[1] == health_pb2.HealthCheckResponse.SERVING ] assert serving_calls == [] - # We never reached server.wait_for_termination() (returned before it). + # We never started the server or reached wait_for_termination(). + fake_server.start.assert_not_called() fake_server.wait_for_termination.assert_not_called()