From ef07cd0353fb0dedb9ad6ed7efd4bfdc76325232 Mon Sep 17 00:00:00 2001 From: Enoch Tang Date: Thu, 24 Sep 2026 15:28:01 -0400 Subject: [PATCH 1/2] Signal draining to push taskworkers from preStop --- .../python/src/taskbroker_client/constants.py | 8 ++ .../src/taskbroker_client/worker/worker.py | 90 ++++++++++-- clients/python/tests/worker/test_worker.py | 129 +++++++++++++++++- 3 files changed, 218 insertions(+), 9 deletions(-) diff --git a/clients/python/src/taskbroker_client/constants.py b/clients/python/src/taskbroker_client/constants.py index 4e700d54..dbbd2128 100644 --- a/clients/python/src/taskbroker_client/constants.py +++ b/clients/python/src/taskbroker_client/constants.py @@ -68,6 +68,14 @@ flipping gRPC health to SERVING anyway. """ +DEFAULT_WORKER_DRAIN_FILE_PATH = "/tmp/taskworker-draining" +""" +File whose presence tells a push worker that its pod is terminating. The pod's +preStop hook creates it before sleeping; SIGTERM only arrives once preStop +finishes, so this is the earliest the worker can learn brokers have stopped +routing to it. The worker stops publishing occupancy when it sees the file. +""" + ALWAYS_EAGER = False """ diff --git a/clients/python/src/taskbroker_client/worker/worker.py b/clients/python/src/taskbroker_client/worker/worker.py index c4e402a7..6ea2b2d1 100644 --- a/clients/python/src/taskbroker_client/worker/worker.py +++ b/clients/python/src/taskbroker_client/worker/worker.py @@ -31,6 +31,7 @@ from taskbroker_client.constants import ( DEFAULT_GRPC_MAX_MESSAGE_SIZE, DEFAULT_REBALANCE_AFTER, + DEFAULT_WORKER_DRAIN_FILE_PATH, DEFAULT_WORKER_HEALTH_CHECK_SEC_PER_TOUCH, DEFAULT_WORKER_QUEUE_SIZE, DEFAULT_WORKER_WARMUP_TIMEOUT_SEC, @@ -258,6 +259,7 @@ def __init__( warmup_timeout: float = DEFAULT_WORKER_WARMUP_TIMEOUT_SEC, prometheus_port: int | None = None, future_checking_frequency: float = 0.1, + drain_file_path: str | None = DEFAULT_WORKER_DRAIN_FILE_PATH, ) -> None: app = import_app(app_module) @@ -286,6 +288,7 @@ def __init__( skip_awaiting_futures=skip_awaiting_futures, prometheus_port=prometheus_port, future_checking_frequency=future_checking_frequency, + drain_file_path=drain_file_path, ) logger.info("Running in PUSH mode") @@ -871,6 +874,7 @@ def __init__( skip_awaiting_futures: bool = True, prometheus_port: int | None = None, future_checking_frequency: float = 0.1, + drain_file_path: str | None = None, ) -> None: self._concurrency = concurrency @@ -911,6 +915,53 @@ def __init__( self._metrics_thread: threading.Thread | None = None self._spawn_children_thread: threading.Thread | None = None + self._received_task = False + self._occupancy_stopped = False + self._drain_file_path = drain_file_path + if drain_file_path is not None: + self._clear_stale_drain_file(drain_file_path) + + def _clear_stale_drain_file(self, path: str) -> None: + """ + preStop also runs before a liveness restart, and the file survives that + restart on an emptyDir /tmp. Left in place it would silence this process's + occupancy for its whole life. + """ + try: + Path(path).unlink() + except FileNotFoundError: + return + except OSError as e: + logger.warning( + "taskworker.worker.drain_file.clear_failed", + extra={"path": path, "error": e, "processing_pool": self._processing_pool_name}, + ) + return + logger.info( + "taskworker.worker.drain_file.cleared", + extra={"path": path, "processing_pool": self._processing_pool_name}, + ) + + def stop_occupancy_reporting(self) -> None: + """ + Withdraw occupancy for the rest of this process's life. + + Only flips a flag: the metrics thread is the sole writer of the gauge, and + on its next flush it removes the Prometheus series instead of leaving the + last value up, so the scraper marks it stale rather than averaging a + draining pod's idle slots into the pool. Idempotent. + """ + if self._occupancy_stopped: + return + self._occupancy_stopped = True + logger.info( + "taskworker.worker.occupancy.stopped", + extra={"processing_pool": self._processing_pool_name, "pod_name": self._pod_name}, + ) + + def _should_publish_occupancy(self) -> bool: + return self._received_task and not self._occupancy_stopped + @property def ready_count(self) -> int: """Number of children that have finished warming up and are consuming.""" @@ -934,6 +985,26 @@ def _emit_periodic_metrics(self) -> None: "pod_name": self._pod_name, } + if ( + self._drain_file_path is not None + and not self._occupancy_stopped + and os.path.exists(self._drain_file_path) + ): + logger.info( + "taskworker.worker.drain_file.detected", + extra={ + "path": self._drain_file_path, + "processing_pool": self._processing_pool_name, + }, + ) + self.stop_occupancy_reporting() + + # Every flush, not just the first: this thread may have set the gauge + # between the flag flipping and now. Outside the occupancy block below, + # which is skipped once no children are running, as during shutdown. + if self._occupancy_stopped and self._prom is not None: + self._prom.occupancy.remove(self._processing_pool_name) + # Emit queue size metrics try: # Method 'qsize' not implemented on all platforms, such as macOS @@ -1029,15 +1100,16 @@ def _emit_periodic_metrics(self) -> None: ) occupancy = min(busy_time / ceiling, 1.0) - self._metrics.gauge( - "taskworker.worker.occupancy", - occupancy, - tags=tags, - ) - if self._prom is not None: - self._prom.occupancy.labels(processing_pool=self._processing_pool_name).set( - occupancy + if self._should_publish_occupancy(): + self._metrics.gauge( + "taskworker.worker.occupancy", + occupancy, + tags=tags, ) + if self._prom is not None: + self._prom.occupancy.labels(processing_pool=self._processing_pool_name).set( + occupancy + ) self._metrics.gauge( "taskworker.worker.concurrency", @@ -1380,6 +1452,7 @@ def push_task(self, inflight: InflightTaskActivation, timeout: float | None = No ) return False + self._received_task = True self._metrics.distribution( "taskworker.worker.child_task.put.duration", time.monotonic() - start_time, @@ -1441,6 +1514,7 @@ def shutdown(self) -> None: """ logger.info("taskworker.worker.shutdown.start") shutdown_start = time.monotonic() + self.stop_occupancy_reporting() self._shutdown_event.set() logger.info("taskworker.worker.shutdown.spawn_children") diff --git a/clients/python/tests/worker/test_worker.py b/clients/python/tests/worker/test_worker.py index dec277a3..5f1ed41e 100644 --- a/clients/python/tests/worker/test_worker.py +++ b/clients/python/tests/worker/test_worker.py @@ -19,6 +19,7 @@ import grpc import msgpack +import prometheus_client import pytest import zstandard as zstd from arroyo.backends.kafka import KafkaPayload @@ -360,8 +361,10 @@ def _make_result_thread_pool( concurrency: int = 3, result_queue_maxsize: int = 3, update_in_batches: bool = False, + received_task: bool = True, + drain_file_path: str | None = None, ) -> TaskWorkerProcessingPool: - return TaskWorkerProcessingPool( + pool = TaskWorkerProcessingPool( app_module="examples.app:app", send_result_fn=capture, mp_context=get_context("fork"), @@ -371,7 +374,10 @@ def _make_result_thread_pool( processing_pool_name="test", update_in_batches=update_in_batches, process_type="fork", + drain_file_path=drain_file_path, ) + pool._received_task = received_task + return pool class _FakeProcess: @@ -983,6 +989,7 @@ def test_push_worker_health_check_touches_while_idle(tmp_path: Path) -> None: def _make_push_worker(**kwargs: Any) -> PushTaskWorker: + kwargs.setdefault("drain_file_path", None) return PushTaskWorker( app_module="examples.app:app", broker_service="127.0.0.1:50051", @@ -3611,3 +3618,123 @@ def test_shutdown_reports_whether_the_result_thread_joined() -> None: assert calls[0].kwargs["tags"]["outcome"] == "timeout" finally: stuck.set() + + +def _prom_with_registry() -> tuple[Any, prometheus_client.CollectorRegistry]: + """WorkerPrometheusMetrics' series on a private registry, without its HTTP server.""" + registry = prometheus_client.CollectorRegistry() + prom = mock.Mock() + prom.occupancy = prometheus_client.Gauge( + "taskworker_worker_occupancy", "", ["processing_pool"], registry=registry + ) + prom.child_busy_seconds = prometheus_client.Counter( + "taskworker_worker_child_busy_seconds", "", ["processing_pool"], registry=registry + ) + prom.child_wait_seconds = prometheus_client.Counter( + "taskworker_worker_child_wait_seconds", "", ["processing_pool"], registry=registry + ) + return prom, registry + + +def _prom_occupancy(registry: prometheus_client.CollectorRegistry) -> float | None: + return registry.get_sample_value("taskworker_worker_occupancy", {"processing_pool": "test"}) + + +def test_occupancy_withheld_until_first_task() -> None: + # A warm pod brokers are not yet routing to: running children, no task yet. + pool = _make_result_thread_pool(_SendResultCapture(), concurrency=2, received_task=False) + pool._metrics = mock.Mock() + pool._prom, registry = _prom_with_registry() + with pool._children_lock: + pool._children[uuid4()] = _make_tracked_child("running", wait_since=10.0) + + with mock.patch("taskbroker_client.worker.worker.time.monotonic", return_value=11.0): + pool._emit_periodic_metrics() + + assert _gauge_calls(pool._metrics, "taskworker.worker.occupancy") == [] + # Absent, not zero, so the scaler averages only pods doing real work. + assert _prom_occupancy(registry) is None + # The counters still report the idle time. + assert _distribution_calls(pool._metrics, "taskworker.worker.child_wait_seconds") + + +def test_push_task_opens_the_occupancy_gate() -> None: + pool = _make_result_thread_pool(_SendResultCapture(), concurrency=2, received_task=False) + assert pool.push_task(SIMPLE_TASK, timeout=1) + assert pool._should_publish_occupancy() + + +def test_stop_occupancy_reporting_removes_the_series() -> None: + pool = _make_result_thread_pool(_SendResultCapture(), concurrency=2) + pool._metrics = mock.Mock() + pool._prom, registry = _prom_with_registry() + with pool._children_lock: + pool._children[uuid4()] = _make_tracked_child("running", busy_since=10.0) + + with mock.patch("taskbroker_client.worker.worker.time.monotonic", return_value=11.0): + pool._emit_periodic_metrics() + assert _prom_occupancy(registry) == pytest.approx(1.0) + + pool.stop_occupancy_reporting() + # Only the metrics thread writes the gauge; it removes the series on its next flush. + assert _prom_occupancy(registry) == pytest.approx(1.0) + + # Children already gone, as during shutdown, so the occupancy block is skipped. + with pool._children_lock: + pool._children.clear() + pool._emit_periodic_metrics() + # Gone rather than frozen at its last value, so the scraper marks it stale. + assert _prom_occupancy(registry) is None + + pool._metrics.reset_mock() + with pool._children_lock: + pool._children[uuid4()] = _make_tracked_child("running", busy_since=11.0) + with mock.patch("taskbroker_client.worker.worker.time.monotonic", return_value=12.0): + pool._emit_periodic_metrics() + # Later flushes must not re-create the series. + assert _prom_occupancy(registry) is None + assert _gauge_calls(pool._metrics, "taskworker.worker.occupancy") == [] + + +def test_pool_shutdown_stops_occupancy() -> None: + # The single SIGTERM hook: push and pull workers both end in pool.shutdown(). + pool = _make_result_thread_pool(_SendResultCapture(), concurrency=2) + pool.shutdown() + assert pool._occupancy_stopped + + +def test_drain_file_stops_occupancy(tmp_path: Path) -> None: + drain_file = tmp_path / "draining" + pool = _make_result_thread_pool( + _SendResultCapture(), concurrency=2, drain_file_path=str(drain_file) + ) + pool._metrics = mock.Mock() + pool._prom, registry = _prom_with_registry() + with pool._children_lock: + pool._children[uuid4()] = _make_tracked_child("running", busy_since=10.0) + + with mock.patch("taskbroker_client.worker.worker.time.monotonic", return_value=11.0): + pool._emit_periodic_metrics() + assert _prom_occupancy(registry) == pytest.approx(1.0) + + # preStop creates the file; SIGTERM is still preStop-sleep seconds away. + drain_file.touch() + pool._metrics.reset_mock() + with mock.patch("taskbroker_client.worker.worker.time.monotonic", return_value=12.0): + pool._emit_periodic_metrics() + + assert _prom_occupancy(registry) is None + assert _gauge_calls(pool._metrics, "taskworker.worker.occupancy") == [] + + +def test_stale_drain_file_cleared_on_startup(tmp_path: Path) -> None: + # preStop ran before a liveness restart and the file survived on /tmp. + drain_file = tmp_path / "draining" + drain_file.touch() + + pool = _make_result_thread_pool( + _SendResultCapture(), concurrency=2, drain_file_path=str(drain_file) + ) + + assert not drain_file.exists() + assert pool._should_publish_occupancy() From d8cea3a0560f5f1b10227633b4e71e9030aca8cc Mon Sep 17 00:00:00 2001 From: Enoch Tang Date: Thu, 24 Sep 2026 16:01:18 -0400 Subject: [PATCH 2/2] Tolerate missing occupancy series on removal --- clients/python/pyproject.toml | 2 +- .../src/taskbroker_client/worker/worker.py | 7 +++---- clients/python/tests/worker/test_worker.py | 16 ++++++++++++++++ uv.lock | 2 +- 4 files changed, 21 insertions(+), 6 deletions(-) diff --git a/clients/python/pyproject.toml b/clients/python/pyproject.toml index 7b0b0637..147aa77f 100644 --- a/clients/python/pyproject.toml +++ b/clients/python/pyproject.toml @@ -12,7 +12,7 @@ dependencies = [ "grpcio>=1.67.1", "grpcio-health-checking>=1.67.1", "msgpack>=1.2.1", - "prometheus_client>=0.20", + "prometheus_client>=0.22", "protobuf>=5.28.3", "redis>=3.4.1", "zstandard>=0.18.0", diff --git a/clients/python/src/taskbroker_client/worker/worker.py b/clients/python/src/taskbroker_client/worker/worker.py index 6ea2b2d1..cdd71f71 100644 --- a/clients/python/src/taskbroker_client/worker/worker.py +++ b/clients/python/src/taskbroker_client/worker/worker.py @@ -1,5 +1,6 @@ from __future__ import annotations +import contextlib import logging import multiprocessing import os @@ -999,11 +1000,9 @@ def _emit_periodic_metrics(self) -> None: ) self.stop_occupancy_reporting() - # Every flush, not just the first: this thread may have set the gauge - # between the flag flipping and now. Outside the occupancy block below, - # which is skipped once no children are running, as during shutdown. if self._occupancy_stopped and self._prom is not None: - self._prom.occupancy.remove(self._processing_pool_name) + with contextlib.suppress(KeyError): + self._prom.occupancy.remove(self._processing_pool_name) # Emit queue size metrics try: diff --git a/clients/python/tests/worker/test_worker.py b/clients/python/tests/worker/test_worker.py index 5f1ed41e..4a1febe6 100644 --- a/clients/python/tests/worker/test_worker.py +++ b/clients/python/tests/worker/test_worker.py @@ -3696,6 +3696,22 @@ def test_stop_occupancy_reporting_removes_the_series() -> None: assert _gauge_calls(pool._metrics, "taskworker.worker.occupancy") == [] +def test_occupancy_removal_error_does_not_abort_the_flush() -> None: + # prometheus_client < 0.22 raises KeyError when removing a missing series. + pool = _make_result_thread_pool(_SendResultCapture(), concurrency=2) + pool._metrics = mock.Mock() + pool._prom = mock.Mock() + pool._prom.occupancy.remove.side_effect = KeyError(("test",)) + pool.stop_occupancy_reporting() + + pool._emit_periodic_metrics() + + pool._prom.occupancy.remove.assert_called_once_with("test") + # The rest of the flush still ran. + assert _gauge_calls(pool._metrics, "taskworker.worker.concurrency") + assert _distribution_calls(pool._metrics, "taskworker.worker.child_busy_seconds") + + def test_pool_shutdown_stops_occupancy() -> None: # The single SIGTERM hook: push and pull workers both end in pool.shutdown(). pool = _make_result_thread_pool(_SendResultCapture(), concurrency=2) diff --git a/uv.lock b/uv.lock index 00811086..92d97bb2 100644 --- a/uv.lock +++ b/uv.lock @@ -850,7 +850,7 @@ requires-dist = [ { name = "grpcio", specifier = ">=1.67.1" }, { name = "grpcio-health-checking", specifier = ">=1.67.1" }, { name = "msgpack", specifier = ">=1.2.1" }, - { name = "prometheus-client", specifier = ">=0.20" }, + { name = "prometheus-client", specifier = ">=0.22" }, { name = "protobuf", specifier = ">=5.28.3" }, { name = "redis", specifier = ">=3.4.1" }, { name = "redis-py-cluster", marker = "extra == 'cluster'", specifier = ">=2.1.0" },