From faba351c07078eaceeb3962f60c174872c975db7 Mon Sep 17 00:00:00 2001 From: Xore Date: Sat, 12 Sep 2026 22:55:12 +0200 Subject: [PATCH] fix(ml-worker): cooldown gate for drift-triggered retrains so a storm cannot re-fire every poll cycle (#3168) Count-only drift trigger had no time floor: window cleared on fire, then refired next poll, accepting 50+ retrains in one hour on 2026-09-08. Adds a persisted drift_retrain_allowed() cooldown gate that survives worker restarts. Gate only delays retrains, never mutes drift alerts. Window still clears only on the accepted retrain path. Closes #3168 --- ml-worker/tests/test_model_lifecycle.py | 60 ++++++++++++++++++++- ml-worker/worker.py | 69 ++++++++++++++++++++++--- 2 files changed, 121 insertions(+), 8 deletions(-) diff --git a/ml-worker/tests/test_model_lifecycle.py b/ml-worker/tests/test_model_lifecycle.py index 7f36fea91..5e38d6c8a 100644 --- a/ml-worker/tests/test_model_lifecycle.py +++ b/ml-worker/tests/test_model_lifecycle.py @@ -9,7 +9,7 @@ import os import random import sys -from datetime import datetime, timezone +from datetime import datetime, timezone, timedelta from pathlib import Path import numpy as np @@ -929,6 +929,64 @@ def test_drift_metric_write_failure_does_not_raise(self): worker.write_drift_metric(es, window=500, rate=0.5) # must not raise +class TestDriftRetrainCooldown: + """#3168: 2026-09-08 03:00-04:00 CEST, 50+ retrain-accepted events in + under an hour at 15-26% drift against a 15% trigger -- genuine sustained + churn, not noise. drift_rate_if_triggered() has no time dimension, so a + DRIFT_WINDOW-sized recent_flags window that refills within a single + POLL_INTERVAL re-triggered on the very next cycle with zero spacing. + drift_retrain_allowed() adds that spacing as a separate, directly + testable gate.""" + def test_first_ever_drift_is_allowed(self): + now = datetime(2026, 9, 8, 3, 0, tzinfo=timezone.utc) + assert worker.drift_retrain_allowed(None, now, cooldown_minutes=60) is True + + def test_immediate_retrigger_within_cooldown_is_blocked(self): + last = datetime(2026, 9, 8, 3, 0, tzinfo=timezone.utc) + one_minute_later = datetime(2026, 9, 8, 3, 1, tzinfo=timezone.utc) + assert worker.drift_retrain_allowed(last, one_minute_later, cooldown_minutes=60) is False + + def test_sustained_drift_still_retrains_once_cooldown_elapses(self): + """The throttle delays, it never mutes: elevated drift outlasting + the cooldown must retrain again, not go silent forever.""" + last = datetime(2026, 9, 8, 3, 0, tzinfo=timezone.utc) + after_cooldown = datetime(2026, 9, 8, 4, 0, tzinfo=timezone.utc) + assert worker.drift_retrain_allowed(last, after_cooldown, cooldown_minutes=60) is True + + def test_storm_shape_collapses_to_one_retrain_per_cooldown(self): + """Simulates the incident: a poll cycle every 30s, each one finding + the recent_flags window already over threshold (sustained churn). + Without the gate this fires every cycle (the reported storm); with + it, only the cycles at least cooldown_minutes apart accept.""" + start = datetime(2026, 9, 8, 3, 0, tzinfo=timezone.utc) + cycles = [start + timedelta(seconds=30 * i) for i in range(120)] # 1 hour, storm cadence + last_triggered_at = None + accepted = 0 + for now in cycles: + if worker.drift_retrain_allowed(last_triggered_at, now, cooldown_minutes=60): + accepted += 1 + last_triggered_at = now + assert accepted == 1 # not the 50+ the storm produced + + def test_last_drift_retrain_persists_and_loads_back(self): + es = MagicMock() + when = datetime(2026, 9, 8, 3, 0, tzinfo=timezone.utc) + worker.save_last_drift_retrain(es, when) + doc = es.index.call_args.kwargs["document"] + assert doc["last_triggered_at"] == when.isoformat() + + def test_no_persisted_drift_retrain_yet_loads_as_none(self): + es = MagicMock() + es.get.side_effect = Exception("not found") + assert worker.load_last_drift_retrain(es) is None + + def test_persisted_drift_retrain_loads_back_correctly(self): + es = MagicMock() + when = datetime(2026, 9, 8, 3, 0, tzinfo=timezone.utc) + es.get.return_value = {"_source": {"last_triggered_at": when.isoformat()}} + assert worker.load_last_drift_retrain(es) == when + + class _StageWarnCapture: """Collect loguru records at WARNING+ while active (#2230 tests).""" diff --git a/ml-worker/worker.py b/ml-worker/worker.py index d87d5ecdb..0dba8016b 100644 --- a/ml-worker/worker.py +++ b/ml-worker/worker.py @@ -91,6 +91,7 @@ # to the traffic itself. DRIFT_WINDOW = int(os.getenv("DRIFT_WINDOW", "500")) DRIFT_ANOMALY_RATE = float(os.getenv("DRIFT_ANOMALY_RATE", "0.15")) +DRIFT_RETRAIN_COOLDOWN_MINUTES = int(os.getenv("DRIFT_RETRAIN_COOLDOWN_MINUTES", "60")) # #190: the regular poll path had no cap, unlike the retrain path # (MAX_TRAIN_SAMPLES) -- a backlog large enough to need more than one @@ -724,6 +725,30 @@ def load_last_fired_slot(es: Elasticsearch) -> "str | None": def save_last_fired_slot(es: Elasticsearch, slot_id: str) -> None: es.index(index=STATE_INDEX, id=RETRAIN_SCHEDULE_STATE_ID, document={"last_fired_slot_id": slot_id}) +DRIFT_RETRAIN_STATE_ID = "drift-retrain" +def load_last_drift_retrain(es: Elasticsearch) -> "datetime | None": + """Persisted for the same reason as load_last_fired_slot: a restart + mid-storm must not reset the cooldown and let the storm resume (#3168).""" + try: + doc = es.get(index=STATE_INDEX, id=DRIFT_RETRAIN_STATE_ID) + return datetime.fromisoformat(doc["_source"]["last_triggered_at"]) + except Exception: + return None +def save_last_drift_retrain(es: Elasticsearch, when: datetime) -> None: + es.index(index=STATE_INDEX, id=DRIFT_RETRAIN_STATE_ID, + document={"last_triggered_at": when.isoformat()}) +def drift_retrain_allowed(last_triggered_at: "datetime | None", now: datetime, + cooldown_minutes: int) -> bool: + """Pure gate (#3168): a drift-triggered retrain is allowed only once + cooldown_minutes have passed since the last one. Without this, a + sustained drift condition (genuine model churn, not noise) refills the + DRIFT_WINDOW-sized recent_flags window and re-fires on the very next + poll cycle -- observed as 50+ retrains in under an hour on 2026-09-08. + Does not mute drift: once cooldown elapses, a still-elevated rate + triggers again on the next check.""" + if last_triggered_at is None: + return True + return now - last_triggered_at >= timedelta(minutes=cooldown_minutes) def _persist_best_effort(label: str, fn, *args) -> bool: @@ -1063,6 +1088,7 @@ def run_worker() -> None: retrain_slots = parse_retrain_slots(RETRAIN_SLOTS_UTC) last_fired_slot_id = load_last_fired_slot(es) # #172: persisted, not restart-relative + last_drift_retrain_at = load_last_drift_retrain(es) # #3168: persisted, not restart-relative recent_flags = deque(maxlen=DRIFT_WINDOW) # composite >= THRESHOLD, drift detection (#65) consecutive_es_failures = {idx: 0 for idx in SOURCE_INDICES} # #188 @@ -1151,15 +1177,34 @@ def run_worker() -> None: # triggering so a persistent drift condition retrains once and then # re-accumulates a fresh window, rather than firing again every # single poll cycle on the same stale evidence. + # + # #3168: that assumption breaks under real throughput -- 500 events + # can refill inside a single POLL_INTERVAL, so sustained (genuine, + # not noisy) churn re-triggered on the very next cycle with zero + # spacing (50+ accepted retrains in under an hour on 2026-09-08). + # drift_retrain_allowed() adds a cooldown floor between accepted + # drift retrains, same persisted-state shape as last_fired_slot_id. + # It only delays, never mutes: once cooldown elapses, still-elevated + # drift triggers again on the next check. drift_rate = drift_rate_if_triggered(recent_flags, DRIFT_WINDOW, DRIFT_ANOMALY_RATE) + drift_due = False if drift_rate is not None: - logger.warning( - f"Drift detected: {drift_rate:.1%} anomaly rate over the last " - f"{DRIFT_WINDOW} events (threshold {DRIFT_ANOMALY_RATE:.0%}) -- " - "triggering an early retrain" - ) write_drift_metric(es, DRIFT_WINDOW, drift_rate) - recent_flags.clear() + if drift_retrain_allowed(last_drift_retrain_at, datetime.now(timezone.utc), + DRIFT_RETRAIN_COOLDOWN_MINUTES): + logger.warning( + f"Drift detected: {drift_rate:.1%} anomaly rate over the last " + f"{DRIFT_WINDOW} events (threshold {DRIFT_ANOMALY_RATE:.0%}) -- " + "triggering an early retrain" + ) + drift_due = True + recent_flags.clear() + else: + logger.info( + f"Drift detected: {drift_rate:.1%} anomaly rate but still within " + f"the {DRIFT_RETRAIN_COOLDOWN_MINUTES}min cooldown since the last " + "drift-triggered retrain -- skipping this cycle" + ) # #172: scheduled retraining now fires at explicit UTC slot # boundaries (persisted last_fired_slot_id), not a restart-relative @@ -1183,7 +1228,7 @@ def run_worker() -> None: # spawning a thread. Revisit if Milestone I's shared-GPU scheduling # (#84) needs retrain and polling decoupled; MAX_POLL_BATCH above at # least bounds how large a backlog blocking this causes. - if drift_rate is not None or scheduled_due: + if drift_due or scheduled_due: logger.info("Starting model retraining...") # #3097: fetching each index's most-recent MAX_TRAIN_SAMPLES # ascending-from-24h-ago systematically starves whichever index @@ -1228,6 +1273,16 @@ def run_worker() -> None: f"save_last_fired_slot({due_slot_id})", save_last_fired_slot, es, due_slot_id) + if drift_due: + # Same reasoning as scheduled_due above: start the cooldown + # now regardless of the 100-event floor, so a quiet cycle + # can't immediately re-arm (recent_flags was already + # cleared above the moment drift was accepted). + last_drift_retrain_at = datetime.now(timezone.utc) + _persist_best_effort( + "save_last_drift_retrain", save_last_drift_retrain, + es, last_drift_retrain_at) + elapsed = time.time() - cycle_start sleep_for = max(0, POLL_INTERVAL - elapsed) time.sleep(sleep_for)