Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
2 changes: 1 addition & 1 deletion clients/python/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
8 changes: 8 additions & 0 deletions clients/python/src/taskbroker_client/constants.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
"""
Expand Down
89 changes: 81 additions & 8 deletions clients/python/src/taskbroker_client/worker/worker.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
from __future__ import annotations

import contextlib
import logging
import multiprocessing
import os
Expand Down Expand Up @@ -31,6 +32,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,
Expand Down Expand Up @@ -258,6 +260,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)

Expand Down Expand Up @@ -286,6 +289,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")
Expand Down Expand Up @@ -871,6 +875,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

Expand Down Expand Up @@ -911,6 +916,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."""
Expand All @@ -934,6 +986,24 @@ 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()

if self._occupancy_stopped and self._prom is not None:
with contextlib.suppress(KeyError):
self._prom.occupancy.remove(self._processing_pool_name)

# Emit queue size metrics
try:
# Method 'qsize' not implemented on all platforms, such as macOS
Expand Down Expand Up @@ -1029,15 +1099,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",
Expand Down Expand Up @@ -1380,6 +1451,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,
Expand Down Expand Up @@ -1441,6 +1513,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")
Expand Down
145 changes: 144 additions & 1 deletion clients/python/tests/worker/test_worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@

import grpc
import msgpack
import prometheus_client
import pytest
import zstandard as zstd
from arroyo.backends.kafka import KafkaPayload
Expand Down Expand Up @@ -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"),
Expand All @@ -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:
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -3611,3 +3618,139 @@ 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_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)
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()
2 changes: 1 addition & 1 deletion uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading