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
143 changes: 142 additions & 1 deletion ml-worker/tests/test_worker_audit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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"]))
53 changes: 46 additions & 7 deletions ml-worker/worker.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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",
Expand All @@ -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)
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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)
Expand Down
Loading