diff --git a/ml-worker/tests/test_worker_audit.py b/ml-worker/tests/test_worker_audit.py index a936cbd75..b148296aa 100644 --- a/ml-worker/tests/test_worker_audit.py +++ b/ml-worker/tests/test_worker_audit.py @@ -24,7 +24,7 @@ import worker # noqa: E402 from worker import contributing_detectors # noqa: E402 -from models.isolation_forest import IsoForestModel # noqa: E402 +from models.isolation_forest import IsoForestModel, RetrainResult # noqa: E402 from models.lstm_autoencoder import LSTMAEModel, SEQ_LEN # noqa: E402 @@ -367,5 +367,146 @@ def test_malformed_event_metric_write_failure_does_not_raise(self): worker.write_malformed_event_metric(es, {"_id": "x", "_index": "y"}, ValueError("boom")) # must not raise +class TestRetrainBatchIsolation: + """#2984: run_worker()'s retrain call chain (compute_batch_session_features() + -> extract_features()/featurise_temporal()) had no per-event guard, unlike + score_and_write_events()'s live path (#171 above). One unreadable row in a + retrain batch (#2978's non-numeric port shape) raised out of retrain() and + killed the whole worker process. quarantine_retrain_events() gives the + retrain path the same drop-and-count contract, in the caller, without + touching extract_features().""" + + def _valid_events(self, n: int) -> list: + return [ + {"_id": f"valid-{i}", "_index": "honeypot-v2-2026.07.31", "_source": dict(REAL_SHAPED_DOCUMENT["_source"])} + for i in range(n) + ] + + def test_quarantine_drops_malformed_row_and_keeps_the_rest(self): + iso_model = IsoForestModel(model_dir=_placeholder_model_dir("does-not-matter")) + es = MagicMock() + + events = self._valid_events(120) + malformed_event = { + "_id": "malformed-1", "_index": "honeypot-v2-2026.07.31", + "_source": {"@timestamp": "2026-07-31T00:00:00Z", "destination": {"port": "not-a-port"}}, + } + events.insert(60, malformed_event) + + sources, dropped = worker.quarantine_retrain_events(es, iso_model, events) + + assert dropped == 1 + assert len(sources) == 120 + + malformed_calls = [ + call for call in es.index.call_args_list + if call.kwargs.get("document", {}).get("kind") == "malformed_event" + ] + assert len(malformed_calls) == 1 + assert malformed_calls[0].kwargs["document"]["source_event_id"] == "malformed-1" + + def test_retrain_succeeds_on_the_quarantined_survivors(self): + iso_model = IsoForestModel(model_dir=_placeholder_model_dir("does-not-matter")) + es = MagicMock() + + events = self._valid_events(120) + malformed_event = { + "_id": "malformed-2", "_index": "honeypot-v2-2026.07.31", + "_source": {"@timestamp": "2026-07-31T00:00:00Z", "destination": {"port": "not-a-port"}}, + } + events.insert(30, malformed_event) + + sources, dropped = worker.quarantine_retrain_events(es, iso_model, events) + assert dropped == 1 + + # Must not raise -- this is exactly the call run_worker() makes with + # quarantine_retrain_events()'s output, and exactly what raised + # unguarded before #2984 when a malformed row reached retrain() + # directly. + result = iso_model.retrain(sources) + assert result.train_samples + result.holdout_samples == 120 + + def test_retrain_without_quarantine_still_raises_on_the_malformed_row(self): + """Proves the pre-#2984 failure mode still exists at the retrain() + layer itself -- retrain() (like extract_features()) is intentionally + left strict/unguarded per #171's "readers stay strict" rule. The fix + is quarantine_retrain_events() in the caller, not a change here.""" + iso_model = IsoForestModel(model_dir=_placeholder_model_dir("does-not-matter")) + sources = [e["_source"] for e in self._valid_events(120)] + sources.insert(60, {"@timestamp": "2026-07-31T00:00:00Z", "destination": {"port": "not-a-port"}}) + + with pytest.raises(ValueError): + iso_model.retrain(sources) + + def test_quarantine_drops_row_that_only_compute_batch_session_features_rejects(self): + """extract_features() never reads honeypot.session -- only + compute_batch_session_features() does, as a dict key + (session_features.py's session_counts.get(session, 0)). A list + there survives extract_features()'s probe and then raises + TypeError: unhashable type: 'list' inside retrain(). The probe must + call compute_batch_session_features() too, not just + extract_features().""" + iso_model = IsoForestModel(model_dir=_placeholder_model_dir("does-not-matter")) + es = MagicMock() + + events = self._valid_events(120) + unhashable_session_event = { + "_id": "unhashable-session-1", "_index": "honeypot-v2-2026.07.31", + "_source": dict(REAL_SHAPED_DOCUMENT["_source"], honeypot={ + **REAL_SHAPED_DOCUMENT["_source"].get("honeypot", {}), "session": ["s1", "s2"], + }), + } + events.insert(60, unhashable_session_event) + + sources, dropped = worker.quarantine_retrain_events(es, iso_model, events) + + assert dropped == 1 + assert len(sources) == 120 + + malformed_calls = [ + call for call in es.index.call_args_list + if call.kwargs.get("document", {}).get("kind") == "malformed_event" + ] + assert len(malformed_calls) == 1 + assert malformed_calls[0].kwargs["document"]["source_event_id"] == "unhashable-session-1" + + # Must not raise -- same contract as the malformed-port test above. + result = iso_model.retrain(sources) + assert result.train_samples + result.holdout_samples == 120 + + def test_write_retrain_metric_reports_dropped_count(self): + es = MagicMock() + result = RetrainResult( + accepted=True, reason="accepted", train_samples=100, holdout_samples=20, + anomaly_rate_new=0.05, anomaly_rate_previous=0.05, + ) + + worker.write_retrain_metric(es, "isolation_forest_hbos", result, dropped_count=1) + + retrain_calls = [ + call for call in es.index.call_args_list + if call.kwargs.get("document", {}).get("kind") == "retrain" + ] + assert len(retrain_calls) == 1 + assert retrain_calls[0].kwargs["document"]["dropped_count"] == 1 + + def test_write_retrain_metric_defaults_dropped_count_to_zero(self): + """Existing 3-positional-arg callers (test_model_lifecycle.py) must + keep working unchanged.""" + es = MagicMock() + result = RetrainResult( + accepted=True, reason="accepted", train_samples=100, holdout_samples=20, + anomaly_rate_new=0.05, anomaly_rate_previous=0.05, + ) + + worker.write_retrain_metric(es, "isolation_forest_hbos", result) + + retrain_calls = [ + call for call in es.index.call_args_list + if call.kwargs.get("document", {}).get("kind") == "retrain" + ] + assert retrain_calls[0].kwargs["document"]["dropped_count"] == 0 + + if __name__ == "__main__": sys.exit(pytest.main([__file__, "-v"])) diff --git a/ml-worker/worker.py b/ml-worker/worker.py index 491253333..d87d5ecdb 100644 --- a/ml-worker/worker.py +++ b/ml-worker/worker.py @@ -24,7 +24,7 @@ _get_dst_ip, _get_src_port, _get_transport_proto, is_our_own_address) from models.lstm_autoencoder import LSTMAEModel -from models.session_features import SessionFeatureTracker +from models.session_features import SessionFeatureTracker, compute_batch_session_features # #1971: the poll/checkpoint mechanics below (#168 boundary semantics, # #188 failure-vs-empty shape, #190 batch caps) are the reference @@ -403,11 +403,17 @@ def _warn(message): ) -def write_retrain_metric(es: Elasticsearch, model_name: str, result) -> None: +def write_retrain_metric(es: Elasticsearch, model_name: str, result, dropped_count: int = 0) -> None: """Evidence for one retrain() call, accepted or not (#65, docs/ml-worker-plan.md §11.1/§11.4). Best-effort like write_anomaly()'s Redis publish: a metrics-write failure must never take down the retrain - cycle that already succeeded or failed on its own terms.""" + cycle that already succeeded or failed on its own terms. + + dropped_count (#2984): how many rows quarantine_retrain_events() dropped + from this cycle's fetched batch before either model ever saw it -- same + reviewable-metric contract as write_malformed_event_metric() (#171), just + surfaced as a count here since a retrain batch is scored as one unit, not + per-row.""" doc = { "@timestamp": datetime.now(timezone.utc).isoformat(), "kind": "retrain", @@ -418,6 +424,7 @@ def write_retrain_metric(es: Elasticsearch, model_name: str, result) -> None: "holdout_samples": result.holdout_samples, "anomaly_rate_new": round(result.anomaly_rate_new, 4), "anomaly_rate_previous": round(result.anomaly_rate_previous, 4) if result.anomaly_rate_previous is not None else None, + "dropped_count": dropped_count, } try: es.index(index=METRICS_INDEX, document=doc) @@ -507,6 +514,37 @@ def write_malformed_event_metric(es: Elasticsearch, event: dict, exc: Exception) logger.warning(f"Skipping malformed event {doc['source_event_id']} in {doc['source_index']}: {exc}") +def quarantine_retrain_events(es: Elasticsearch, iso_model, events: list) -> tuple: + """Per-event guard for the retrain batch (#2984), same quarantine + contract score_and_write_events() already has for live scoring (#171): + a source that can't survive extract_features() is dropped and recorded + as a reviewable metric instead of raising out of retrain(). + + Probes with both iso_model.extract_features() and + compute_batch_session_features() -- the two per-row entry points both + models' retrain() calls go through (isolation_forest.retrain() and + lstm_autoencoder.retrain() each call compute_batch_session_features() + themselves; extract_features() covers the rest of each model's own + field reads). + + Returns (sources, dropped_count) -- sources is `events` filtered down to + just the ones that survived, in original order, ready to pass straight + to IsoForestModel.retrain()/LSTMAEModel.retrain().""" + sources = [] + dropped = 0 + for event in events: + src = event.get("_source", {}) + try: + iso_model.extract_features(src) + compute_batch_session_features([src]) + except Exception as exc: + write_malformed_event_metric(es, event, exc) + dropped += 1 + continue + sources.append(src) + return sources, dropped + + def score_and_write_events(es: Elasticsearch, rdb, iso_model, lstm_model, events: list, recent_flags, session_tracker=None) -> None: """Score every event in one fetched batch, writing anomalies over @@ -1005,6 +1043,7 @@ def run_worker() -> None: "holdout_samples": {"type": "integer"}, "anomaly_rate_new": {"type": "float"}, "anomaly_rate_previous": {"type": "float"}, + "dropped_count": {"type": "integer"}, # kind="retrain" (#2984) "drift_window": {"type": "integer"}, "drift_rate": {"type": "float"}, "source_event_id": {"type": "keyword"}, # kind="malformed_event" (#171) @@ -1170,12 +1209,12 @@ def run_worker() -> None: all_events.extend(idx_events) if len(all_events) > 100: - sources = [e["_source"] for e in all_events] + sources, dropped_count = quarantine_retrain_events(es, iso_model, all_events) iso_result = iso_model.retrain(sources, source_index_counts=index_counts) lstm_result = lstm_model.retrain(sources) - write_retrain_metric(es, "isolation_forest_hbos", iso_result) - write_retrain_metric(es, "lstm_ae", lstm_result) - logger.info(f"Retrain cycle on {len(all_events)} events complete") + write_retrain_metric(es, "isolation_forest_hbos", iso_result, dropped_count=dropped_count) + write_retrain_metric(es, "lstm_ae", lstm_result, dropped_count=dropped_count) + logger.info(f"Retrain cycle on {len(all_events)} events complete ({dropped_count} dropped)") if scheduled_due: # Mark this slot handled regardless of whether len(all_events)