diff --git a/docs/ml-worker-evaluation.md b/docs/ml-worker-evaluation.md index 7bf11c5ab..bc01bd48e 100644 --- a/docs/ml-worker-evaluation.md +++ b/docs/ml-worker-evaluation.md @@ -102,6 +102,30 @@ Wiring Tier 2 — the BETH rail and disposition-corpus calibration — is tracke [#1974](https://github.com/Xore/APIARY/issues/1974). It is not tracked by this paragraph. +### 2026-09-25 — disposition export/census landed; Tier 2 calibration remains gated + +`ml-worker/benchmarks/disposition_corpus.py` now exports the closed +operator-disposition population and writes a hashed census outside the +repository. The report carries a canonical-content SHA-256 (the digest field +is excluded from its own hash), so the saved artifact can be independently +verified. It is read-only: open and legacy unlabelled alerts remain in the +full alert denominator but are excluded from the labelled snapshot, and the +command never updates Elasticsearch or synthesises a verdict. + +The census is a gate, not a calibration result. A zero-label or single-class +labelled subset is reported as non-calibratable. Even after labels accumulate, +the corpus is **precision-only**: `write_anomaly()` returns before persistence +below `ML_ALERT_THRESHOLD`, so it can describe precision within alerts and +within-alert calibration, but it can never measure deployment recall or +ordinary below-threshold calibration. Any Tier 2 report that consumes this +corpus must repeat that limitation and must not treat unlabelled or absent +events as negatives. + +The decision record should receive a result only after a concrete snapshot is +attached to the run and its class/model-state diversity is sufficient. The +live census is therefore not copied into this file as a durable result; rerun +the command against the deployment and retain the generated report and hash. + ## Findings carried in from the research phase Recorded so they are not rediscovered, each with the reason it matters here. diff --git a/ml-worker/benchmarks/README.md b/ml-worker/benchmarks/README.md index b25f03b1f..1822bdd67 100644 --- a/ml-worker/benchmarks/README.md +++ b/ml-worker/benchmarks/README.md @@ -8,6 +8,61 @@ governance model. The decision record is **This produces the ruler, not a detector.** +## Operator-disposition corpus + +`disposition_corpus.py` is a separate read-only export/census for the +`ml-anomalies` index. It does not train, calibrate, mutate, re-label, or +auto-disposition an alert. + +```bash +python3 ml-worker/benchmarks/disposition_corpus.py \ + --es-host "$ES_HOST" \ + --output "$HOME/ml-worker-qualification/dispositions.ndjson" \ + --report "$HOME/ml-worker-qualification/dispositions-census.json" +``` + +The endpoint may come from `--es-host`, `ES_HOST`, or `ELASTICSEARCH_URL`; no +deployment endpoint is hardcoded. Credentials may come from `ES_API_KEY` or +`ELASTICSEARCH_API_KEY`, or from the `ES_USERNAME`/`ES_PASSWORD` pair. The +tool issues search and PIT lifecycle requests only. It first opens one PIT, +uses it for the all-alert status/time census and the closed-label export, then +closes it. The exported rows and hashed report are written atomically outside +the repository. The report records a canonical-content SHA-256 (computed +without the digest field) and prints the same digest after the census. + +Open and legacy documents without a disposition are counted in the full alert +denominator but excluded from the labelled export. Every closed row contains +the production anomaly timestamp, score, detector scores/contributors, +threshold, model state, sensor, disposition, actor, and disposal time. +`disposition_reason` is free text and is replaced with `[REDACTED]` by +default; the explicit `--include-reason` override is available only when an +operator has a separate review and storage boundary in place. +The census reports: + +- total alerts and every status count (including ``); +- closed-label count, class balance, and labelled/all-alert denominator; +- all-alert and labelled time ranges; +- labelled sensor, scoring model, threshold, distinct-model-state, and + missing-field counts, including per-status breakdowns; +- the SHA-256 of the immutable NDJSON snapshot and the canonical-content + SHA-256 of the JSON report; and +- an explicit zero/single-class calibration gate. + +The output begins with a non-negotiable warning: + +> **PRECISION-ONLY:** only above-threshold alerts are persisted, so this +> corpus cannot measure deployment recall or calibrate ordinary +> below-threshold traffic. + +An unlabelled alert is not a negative, and an absent below-threshold event is +not evidence that it was correctly rejected. #2986's calibration work may use +the snapshot only after the census shows enough labels and model-state +diversity; this command never performs that calibration. + +For an offline verification, pass a complete local JSON array of ES-shaped +hits to `--fixture`. The fixture includes open/legacy rows as well as closed +rows so the two denominators can be tested without a network. + ## Safety properties - Fixtures are the per-sensor documents from `ml-worker/tests/fixtures.py` — diff --git a/ml-worker/benchmarks/disposition_corpus.py b/ml-worker/benchmarks/disposition_corpus.py new file mode 100644 index 000000000..3d09a2fc6 --- /dev/null +++ b/ml-worker/benchmarks/disposition_corpus.py @@ -0,0 +1,592 @@ +"""Read-only export and census for operator-disposition alerts.""" + +from __future__ import annotations + +import hashlib +import json +import os +import sys +import tempfile +import urllib.error +import urllib.request +from collections import Counter +from copy import deepcopy +from dataclasses import dataclass, field +from datetime import datetime, timezone +from pathlib import Path +from typing import Any +from urllib.parse import urlencode, urlsplit + + +BENCHMARK_VERSION = "apiary-ml-worker-disposition-corpus-v1" +INDEX = "ml-anomalies" +STATUSES = ("open", "true_positive", "false_positive", "benign_known") +CLOSED_STATUSES = ("true_positive", "false_positive", "benign_known") +REASON_REDACTION = "[REDACTED]" +EXPORT_FIELDS = ( + "@timestamp", "composite_score", "model_scores", "contributing_detectors", + "alert_threshold", "model_state_id", "sensor", "status", + "disposition_reason", "disposition_by", "disposed_at", +) +PRECISION_ONLY = ( + "PRECISION-ONLY: only above-threshold alerts are persisted; deployment recall and " + "ordinary below-threshold calibration are unavailable from this corpus." +) + + +class CorpusError(RuntimeError): + """An actionable export/census failure.""" + + +class ElasticsearchUnavailable(CorpusError): + """The configured Elasticsearch endpoint is unreachable.""" + + +@dataclass +class Census: + total: int + status_counts: dict[str, int] + labelled_count: int + rows: list[dict[str, Any]] = field(default_factory=list) + time_range: dict[str, str | None] | None = None + + +def now_utc() -> str: + return datetime.now(timezone.utc).isoformat().replace("+00:00", "Z") + + +def present(value: Any) -> bool: + return value is not None and value not in ("", [], {}) + + +def model_name(value: Any) -> str | None: + if isinstance(value, dict): + for name, score in value.items(): + if score is not None: + return str(name) + return None + + +def source_value(source: dict[str, Any], field: str) -> Any: + if field == "sensor": + return source.get("sensor") or (source.get("event") or {}).get("sensor") + return source.get(field) + + +def normalise_row(hit: dict[str, Any], redact_reason: bool = True) -> dict[str, Any]: + if not isinstance(hit, dict) or not isinstance(hit.get("_id"), str): + raise CorpusError("each corpus hit must have a string _id") + source = hit.get("_source") + if not isinstance(source, dict): + raise CorpusError(f"corpus hit {hit['_id']} has no object _source") + if source.get("status") not in CLOSED_STATUSES: + raise CorpusError(f"corpus hit {hit['_id']} is not a closed disposition") + row = {"_id": hit["_id"]} + for field in EXPORT_FIELDS: + value = source_value(source, field) + if value is not None: + row[field] = deepcopy(value) + if redact_reason and "disposition_reason" in row: + row["disposition_reason"] = REASON_REDACTION + return row + + +def normalise_hits(value: Any) -> list[dict[str, Any]]: + while isinstance(value, dict) and "hits" in value: + value = value["hits"] + if not isinstance(value, list): + raise CorpusError("fixture must be a JSON list of Elasticsearch hits") + return value + + +def census_from_hits(hits: Any, redact_reason: bool = True) -> Census: + """Census a complete corpus, excluding open rows from the export.""" + hits = normalise_hits(hits) + counts: Counter[str] = Counter() + timestamps: list[str] = [] + for hit in hits: + if not isinstance(hit, dict) or not isinstance(hit.get("_source"), dict): + raise CorpusError("each corpus hit must have an object _source") + source = hit["_source"] + status = source.get("status") + counts[str(status) if status else ""] += 1 + timestamp = source.get("@timestamp") + if present(timestamp): + timestamps.append(str(timestamp)) + closed_rows = [ + normalise_row(hit, redact_reason=redact_reason) + for hit in hits if (hit.get("_source") or {}).get("status") in CLOSED_STATUSES + ] + status_counts = {status: counts.get(status, 0) for status in STATUSES} + for status, count in counts.items(): + if status not in status_counts: + status_counts[status] = count + labelled_count = sum(status_counts[status] for status in CLOSED_STATUSES) + if labelled_count != len(closed_rows): + raise CorpusError("closed status counts and exported rows disagree") + return Census( + total=len(hits), + status_counts=status_counts, + labelled_count=labelled_count, + rows=closed_rows, + time_range={"min": min(timestamps), "max": max(timestamps)} if timestamps else None, + ) + + +def _range(rows: list[dict[str, Any]]) -> dict[str, str | None] | None: + for field in ("@timestamp", "disposed_at"): + values = [str(row[field]) for row in rows if present(row.get(field))] + if values: + return {"min": min(values), "max": max(values)} + return None + + +def _missing(rows: list[dict[str, Any]]) -> dict[str, int]: + missing = { + field: sum(not present(source_value(row, field)) for row in rows) + for field in EXPORT_FIELDS + } + missing["model_scores"] = sum(model_name(row.get("model_scores")) is None for row in rows) + return missing + + +def _counter(values: list[Any]) -> dict[str, int]: + return dict(sorted(Counter(str(value) for value in values if present(value)).items())) + + +def _groups(rows: list[dict[str, Any]]) -> dict[str, Any]: + return { + "by_sensor": _counter([row.get("sensor") for row in rows]), + "by_model": _counter([model_name(row.get("model_scores")) for row in rows]), + "by_threshold": _counter([row.get("alert_threshold") for row in rows]), + "distinct_model_state_ids": len({ + str(row["model_state_id"]) for row in rows if present(row.get("model_state_id")) + }), + "missing_fields": _missing(rows), + } + + +def build_report(census: Census, *, source: dict[str, Any], export: dict[str, Any]) -> dict[str, Any]: + if sum(census.status_counts.values()) != census.total: + raise CorpusError("status counts do not sum to the total alert count") + if census.labelled_count != len(census.rows): + raise CorpusError("labelled count does not match exported rows") + labelled = census.rows + class_counts = {status: census.status_counts.get(status, 0) for status in CLOSED_STATUSES} + classes_present = sum(count > 0 for count in class_counts.values()) + if census.labelled_count == 0: + gate = "non_calibratable" + reason = "zero labelled dispositions" + elif classes_present < 2: + gate = "non_calibratable" + reason = "single labelled class" + else: + gate = "eligible_for_tier_2_calibration" + reason = "at least two labelled classes" + by_status = {status: _groups([ + row for row in labelled if row.get("status") == status + ]) for status in CLOSED_STATUSES} + return { + "generated_at": now_utc(), + "protocol": { + "version": BENCHMARK_VERSION, + "read_only": True, + "index": INDEX, + "closed_statuses": list(CLOSED_STATUSES), + "export_fields": list(EXPORT_FIELDS), + "reasons_redacted": export.get("reasons_redacted", True), + }, + "source": source, + "census": { + "alerts": { + "total": census.total, + "by_disposition_status": census.status_counts, + "labelled_count": census.labelled_count, + "labelled_denominator": census.total, + "labelled_fraction": census.labelled_count / census.total if census.total else None, + "class_balance": { + "counts": class_counts, + "denominator": census.labelled_count, + "fractions": { + status: count / census.labelled_count if census.labelled_count else None + for status, count in class_counts.items() + }, + }, + }, + "time_range": { + "all_alerts": census.time_range, + "labelled": _range(labelled), + "by_labelled_status": { + status: _range([row for row in labelled if row.get("status") == status]) + for status in CLOSED_STATUSES + }, + }, + "groupings": { + "labelled": _groups(labelled), + "by_labelled_status": by_status, + }, + }, + "calibration_gate": { + "status": gate, + "reason": reason, + "classes_present": classes_present, + "precision_only": True, + "precision_only_message": PRECISION_ONLY, + "deployment_recall_available": False, + }, + "export": export, + "caps": [ + PRECISION_ONLY, + "Unlabelled and below-threshold events are not negative examples.", + "This tool reports eligibility; it does not fit or score a calibrator.", + ], + } + + +def _json_bytes(value: Any) -> bytes: + return (json.dumps(value, indent=2, sort_keys=True, allow_nan=False) + "\n").encode() + + +def report_digest(report: dict[str, Any]) -> str: + """Hash the report content without its digest field to avoid self-reference.""" + content = {key: value for key, value in report.items() if key != "report_sha256"} + return hashlib.sha256(_json_bytes(content)).hexdigest() + + +def write_report(path: str | os.PathLike[str], report: dict[str, Any]) -> str: + report["report_sha256"] = report_digest(report) + return write_atomic(path, _json_bytes(report)) + + +def write_atomic(path: str | os.PathLike[str], payload: bytes) -> str: + destination = Path(path).expanduser().resolve() + destination.parent.mkdir(parents=True, exist_ok=True) + descriptor, temporary = tempfile.mkstemp(prefix=".dispositions-", dir=destination.parent) + try: + with os.fdopen(descriptor, "wb") as output: + output.write(payload) + output.flush() + os.fsync(output.fileno()) + os.replace(temporary, destination) + finally: + try: + os.unlink(temporary) + except FileNotFoundError: + pass + return hashlib.sha256(payload).hexdigest() + + +def export_rows(census: Census, output: str | os.PathLike[str], *, + reasons_redacted: bool = True) -> dict[str, Any]: + lines = [json.dumps(row, sort_keys=True, separators=(",", ":")) for row in census.rows] + payload = ("\n".join(lines) + ("\n" if lines else "")).encode() + digest = write_atomic(output, payload) + return { + "path": str(Path(output).expanduser().resolve()), + "format": "ndjson", + "rows": len(census.rows), + "bytes": len(payload), + "sha256": digest, + "reasons_redacted": reasons_redacted, + } + + +def print_report(report: dict[str, Any]) -> None: + alerts = report["census"]["alerts"] + print("=" * 78) + print(PRECISION_ONLY) + print("=" * 78) + if alerts["labelled_count"] == 0: + print("HEADLINE: zero labelled dispositions; calibration is blocked.") + else: + print("HEADLINE: labelled dispositions found; calibration is not run by this tool.") + print(f"Total alerts: {alerts['total']}") + print(f"Labelled subset: {alerts['labelled_count']} / {alerts['labelled_denominator']}") + for status, count in alerts["by_disposition_status"].items(): + print(f" {status}: {count}") + time_range = report["census"]["time_range"] + print(f"Alert time range: {time_range['all_alerts'] or 'unavailable'}") + print(f"Labelled time range: {time_range['labelled'] or 'unavailable'}") + print(f"Calibration gate: {report['calibration_gate']['status']}") + print(f"Export: {report['export']['rows']} rows -> {report['export']['path']}") + print(f"Snapshot SHA-256: {report['export']['sha256']}") + if report.get("report_sha256"): + print(f"Report SHA-256 (canonical content): {report['report_sha256']}") + + +def load_fixture(path: str | os.PathLike[str]) -> list[dict[str, Any]]: + fixture = Path(path).expanduser() + try: + with fixture.open(encoding="utf-8") as source: + value = json.load(source) + except (OSError, json.JSONDecodeError) as exc: + raise CorpusError(f"could not read fixture {fixture}: {exc}") from exc + return normalise_hits(value) + + +def _hits_total(body: dict[str, Any]) -> int: + total = body.get("hits", {}).get("total") + value = total.get("value") if isinstance(total, dict) else total + if isinstance(value, bool) or not isinstance(value, int) or value < 0: + raise CorpusError("Elasticsearch response has no valid hits.total") + return value + + +def _buckets(body: dict[str, Any], name: str) -> list[dict[str, Any]]: + value = body.get("aggregations", {}).get(name, {}).get("buckets") + if not isinstance(value, list): + raise CorpusError(f"Elasticsearch response has no {name} buckets") + return value + + +def _request(client: dict[str, Any], method: str, path: str, + body: Any = None) -> dict[str, Any]: + endpoint = str(client.get("endpoint") or "").rstrip("/") + if not endpoint: + raise CorpusError("an explicit Elasticsearch endpoint is required") + headers = {"Accept": "application/json", "Content-Type": "application/json", + **client.get("headers", {})} + request = urllib.request.Request( + endpoint + path, + data=json.dumps(body).encode() if body is not None else None, + headers=headers, + method=method, + ) + try: + with urllib.request.urlopen(request, timeout=client["timeout"]) as response: + raw = response.read() + except urllib.error.HTTPError as exc: + detail = exc.read().decode(errors="replace")[:300] + raise CorpusError(f"Elasticsearch {method} {path} failed ({exc.code}): {detail}") from exc + except (urllib.error.URLError, TimeoutError, OSError) as exc: + raise ElasticsearchUnavailable(f"could not reach Elasticsearch at {endpoint}: {exc}") from exc + try: + value = json.loads(raw) + except json.JSONDecodeError as exc: + raise CorpusError(f"Elasticsearch returned invalid JSON for {path}") from exc + if not isinstance(value, dict): + raise CorpusError(f"Elasticsearch returned a non-object for {path}") + return value + + +def validate_endpoint(value: str | None) -> str: + if not value: + raise CorpusError( + "an explicit Elasticsearch endpoint is required; pass --es-host or set ES_HOST") + parsed = urlsplit(value) + if parsed.scheme not in ("http", "https") or not parsed.netloc: + raise CorpusError("Elasticsearch endpoint must be an explicit http(s) URL") + if parsed.username or parsed.password: + raise CorpusError("do not put credentials in --es-host; use --api-key or environment variables") + return value.rstrip("/") + + +def _client(endpoint: str, api_key: str | None, username: str | None, + password: str | None, timeout: float) -> dict[str, Any]: + import base64 + headers = {} + if api_key: + headers["Authorization"] = f"ApiKey {api_key}" + elif username or password: + token = base64.b64encode(f"{username or ''}:{password or ''}".encode()).decode() + headers["Authorization"] = f"Basic {token}" + return {"endpoint": endpoint, "headers": headers, "timeout": timeout} + + +def _census_query() -> dict[str, Any]: + return { + "size": 0, + "track_total_hits": True, + "query": {"match_all": {}}, + "aggs": { + "statuses": {"terms": {"field": "status", "missing": "", "size": 100}}, + "alert_time": {"stats": {"field": "@timestamp"}}, + }, + } + + +def _stats_range(aggregations: dict[str, Any]) -> dict[str, str | None] | None: + alert_time = aggregations.get("alert_time") + if not isinstance(alert_time, dict): + raise CorpusError("Elasticsearch response has no alert-time stats") + stats = alert_time.get("stats", alert_time) + if not isinstance(stats, dict): + raise CorpusError("Elasticsearch response has invalid alert-time stats") + minimum = stats.get("min_as_string", stats.get("min")) + maximum = stats.get("max_as_string", stats.get("max")) + if minimum is None and maximum is None: + return None + return {"min": _text(minimum), "max": _text(maximum)} + + +def _text(value: Any) -> str | None: + return None if value is None else str(value) + + +def _census_from_elasticsearch(client: dict[str, Any], pit_id: str) -> Census: + query = _census_query() + query["pit"] = {"id": pit_id, "keep_alive": "1m"} + body = _request(client, "POST", "/_search", query) + counts: Counter[str] = Counter() + for bucket in _buckets(body, "statuses"): + key = bucket.get("key") + value = bucket.get("doc_count") + if not isinstance(key, str) or isinstance(value, bool) or not isinstance(value, int): + raise CorpusError("Elasticsearch returned an invalid status bucket") + counts[key] += value + total = _hits_total(body) + if sum(counts.values()) != total: + raise CorpusError("status aggregation does not cover every alert") + status_counts = {status: counts.get(status, 0) for status in STATUSES} + for status, count in counts.items(): + status_counts.setdefault(status, count) + labelled = sum(status_counts[status] for status in CLOSED_STATUSES) + status_counts = {status: status_counts.get(status, 0) for status in STATUSES} + for status, count in counts.items(): + status_counts.setdefault(status, count) + return Census(total, status_counts, labelled, [], _stats_range(body.get("aggregations", {}))) + + +def _open_pit(client: dict[str, Any]) -> str: + body = _request(client, "POST", f"/{INDEX}/_pit?{urlencode({'keep_alive': '1m'})}") + pit_id = body.get("id") + if not isinstance(pit_id, str) or not pit_id: + raise CorpusError("Elasticsearch did not return a PIT id") + return pit_id + + +def _close_pit(client: dict[str, Any], pit_id: str) -> None: + try: + _request(client, "DELETE", "/_pit", {"id": pit_id}) + except CorpusError: + pass + + +def _fetch_rows(client: dict[str, Any], page_size: int, + redact_reason: bool, pit_id: str) -> list[dict[str, Any]]: + rows: list[dict[str, Any]] = [] + search_after: list[Any] | None = None + try: + while True: + body: dict[str, Any] = { + "size": page_size, + "pit": {"id": pit_id, "keep_alive": "1m"}, + "query": {"terms": {"status": list(CLOSED_STATUSES)}}, + "sort": [{"_shard_doc": "asc"}], + "_source": list(EXPORT_FIELDS), + "track_total_hits": True, + } + if search_after is not None: + body["search_after"] = search_after + response = _request(client, "POST", "/_search", body) + pit_id = response.get("pit_id", pit_id) + hits = response.get("hits", {}).get("hits") + if not isinstance(hits, list): + raise CorpusError("Elasticsearch search response has no hits list") + rows.extend(normalise_row(hit, redact_reason) for hit in hits) + if len(hits) < page_size: + break + if not hits or "sort" not in hits[-1]: + raise CorpusError("Elasticsearch page is missing search_after sort values") + search_after = hits[-1]["sort"] + finally: + _close_pit(client, pit_id) + return rows + + +def fetch_elasticsearch_census(client: dict[str, Any], *, page_size: int = 1000, + redact_reason: bool = True) -> tuple[Census, list[dict[str, Any]]]: + """Open one PIT, then use it for census and immutable labelled export.""" + pit_id = _open_pit(client) + try: + census = _census_from_elasticsearch(client, pit_id) + rows = _fetch_rows(client, page_size, redact_reason, pit_id) + except Exception: + _close_pit(client, pit_id) + raise + return census, rows + + +def run_elasticsearch(*, endpoint: str, output: str, report_output: str, + api_key: str | None, username: str | None, + password: str | None, page_size: int, timeout: float, + include_reason: bool = False) -> tuple[Census, dict[str, Any]]: + client = _client(validate_endpoint(endpoint), api_key, username, password, timeout) + census, rows = fetch_elasticsearch_census(client, page_size=page_size, + redact_reason=not include_reason) + status_counts: Counter[str] = Counter() + for row in rows: + status = row.get("status") + if status not in CLOSED_STATUSES: + raise CorpusError("exported row is not a closed disposition") + status_counts[status] += 1 + if any(status_counts.get(status, 0) != census.status_counts.get(status, 0) + for status in CLOSED_STATUSES): + raise CorpusError("exported status counts differ from the census") + census.rows = rows + if len(census.rows) != census.labelled_count: + raise CorpusError("exported labelled rows do not match the census labelled count") + export = export_rows(census, output, reasons_redacted=not include_reason) + source = {"kind": "elasticsearch", "index": INDEX, "endpoint": endpoint} + report = build_report(census, source=source, export=export) + write_report(report_output, report) + return census, report + + +def run_fixture(path: str, output: str, report_output: str, + include_reason: bool = False) -> tuple[Census, dict[str, Any]]: + census = census_from_hits(load_fixture(path), redact_reason=not include_reason) + export = export_rows(census, output, reasons_redacted=not include_reason) + source = {"kind": "fixture", "path": str(Path(path).resolve())} + report = build_report(census, source=source, export=export) + write_report(report_output, report) + return census, report + + +def parse_args(argv: list[str] | None = None) -> Any: + import argparse + parser = argparse.ArgumentParser(description=__doc__) + source = parser.add_mutually_exclusive_group(required=True) + source.add_argument( + "--es-host", help="explicit Elasticsearch URL; env ES_HOST/ELASTICSEARCH_URL also works") + source.add_argument("--fixture", help="local JSON ES-hit fixture for offline verification") + parser.add_argument("--output", required=True, help="labelled NDJSON snapshot path") + parser.add_argument("--report", required=True, help="hashed JSON census report path") + parser.add_argument("--api-key", default=os.getenv("ES_API_KEY") or os.getenv("ELASTICSEARCH_API_KEY")) + parser.add_argument("--username", default=os.getenv("ES_USERNAME") or os.getenv("ELASTICSEARCH_USERNAME")) + parser.add_argument("--password", default=os.getenv("ES_PASSWORD") or os.getenv("ELASTICSEARCH_PASSWORD")) + parser.add_argument("--page-size", type=int, default=1000) + parser.add_argument("--timeout", type=float, default=30.0) + parser.add_argument( + "--include-reason", action="store_true", help="include free-text reasons (not default)") + args = parser.parse_args(argv) + if args.page_size < 1 or args.page_size > 10000: + parser.error("--page-size must be between 1 and 10000") + if args.timeout <= 0: + parser.error("--timeout must be positive") + return args + + +def main(argv: list[str] | None = None) -> int: + args = parse_args(argv) + try: + if args.fixture: + _, report = run_fixture(args.fixture, args.output, args.report, args.include_reason) + else: + endpoint = args.es_host or os.getenv("ES_HOST") or os.getenv("ELASTICSEARCH_URL") + _, report = run_elasticsearch( + endpoint=validate_endpoint(endpoint), output=args.output, + report_output=args.report, api_key=args.api_key, + username=args.username, password=args.password, + page_size=args.page_size, timeout=args.timeout, + include_reason=args.include_reason, + ) + except CorpusError as exc: + print(f"disposition corpus: {exc}", file=sys.stderr) + return 2 + print_report(report) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/ml-worker/tests/test_disposition_corpus.py b/ml-worker/tests/test_disposition_corpus.py new file mode 100644 index 000000000..e7c38ea6a --- /dev/null +++ b/ml-worker/tests/test_disposition_corpus.py @@ -0,0 +1,177 @@ +"""Tests for the read-only operator-disposition export and census.""" + +import io +import json +import sys +from contextlib import redirect_stdout +from pathlib import Path + +import pytest + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from benchmarks import disposition_corpus as corpus # noqa: E402 + + +def hit(identifier, status, *, sensor="cowrie-login", model="isolation_forest", + model_state="iso:1|hbos:1|lstm:1", threshold=0.85, reason="private note", + timestamp="2026-09-25T10:00:00Z"): + source = { + "@timestamp": timestamp, + "composite_score": 0.91, + "model_scores": {model: 0.91, "lstm_ae": 0.9, "hbos": 0.8}, + "contributing_detectors": [model, "lstm_ae", "hbos"], + "alert_threshold": threshold, + "model_state_id": model_state, + "sensor": sensor, + "status": status, + "disposition_reason": reason, + "disposition_by": "operator@example.test", + "disposed_at": timestamp, + } + if status is None: + source.pop("status") + return {"_id": identifier, "_source": source} + + +def run_local(tmp_path, hits): + fixture = tmp_path / "fixture.json" + fixture.write_text(json.dumps(hits), encoding="utf-8") + _, report = corpus.run_fixture( + str(fixture), str(tmp_path / "labels.ndjson"), str(tmp_path / "report.json")) + return report, (tmp_path / "labels.ndjson").read_text(encoding="utf-8") + + +def test_census_denominators_and_balanced_labels(tmp_path): + hits = [ + hit("tp-1", "true_positive", model_state="iso:1|hbos:1|lstm:1"), + hit("fp-1", "false_positive", sensor="dionaea-connection", model_state="iso:1|hbos:1|lstm:2"), + hit("bk-1", "benign_known", sensor="conpot-modbus", model="lstm_ae", + model_state="iso:1|hbos:1|lstm:3"), + hit("open-1", "open", timestamp="2026-09-24T10:00:00Z"), + hit("legacy-1", None, timestamp="2026-09-23T10:00:00Z"), + ] + report, exported = run_local(tmp_path, hits) + alerts = report["census"]["alerts"] + assert alerts["total"] == 5 + assert alerts["labelled_count"] == 3 + assert alerts["labelled_denominator"] == 5 + assert alerts["labelled_fraction"] == pytest.approx(0.6) + assert alerts["by_disposition_status"]["open"] == 1 + assert alerts["by_disposition_status"][""] == 1 + assert report["calibration_gate"]["status"] == "eligible_for_tier_2_calibration" + assert report["census"]["groupings"]["labelled"]["by_sensor"] == { + "conpot-modbus": 1, "cowrie-login": 1, "dionaea-connection": 1, + } + assert report["census"]["time_range"]["all_alerts"] == { + "min": "2026-09-23T10:00:00Z", "max": "2026-09-25T10:00:00Z", + } + assert len(exported.splitlines()) == 3 + assert "open-1" not in exported and "legacy-1" not in exported + + +def test_zero_labels_is_prominent_and_precision_only(tmp_path): + report, exported = run_local(tmp_path, [hit("open-1", "open")]) + assert report["census"]["alerts"]["labelled_count"] == 0 + assert report["calibration_gate"]["status"] == "non_calibratable" + assert report["calibration_gate"]["deployment_recall_available"] is False + output = io.StringIO() + with redirect_stdout(output): + corpus.print_report(report) + text = output.getvalue() + assert "PRECISION-ONLY" in text + assert "HEADLINE: zero labelled dispositions" in text + assert exported == "" + + +@pytest.mark.parametrize("status", corpus.CLOSED_STATUSES) +def test_single_class_is_non_calibratable(tmp_path, status): + report, _ = run_local(tmp_path, [hit("one", status)]) + assert report["calibration_gate"]["status"] == "non_calibratable" + assert report["calibration_gate"]["reason"] == "single labelled class" + + +def test_reasons_are_redacted_by_default_but_can_be_audited(tmp_path): + source = hit("tp-1", "true_positive", reason="do not publish this note") + redacted = corpus.census_from_hits([source]).rows[0] + audited = corpus.census_from_hits([source], redact_reason=False).rows[0] + assert redacted["disposition_reason"] == corpus.REASON_REDACTION + assert audited["disposition_reason"] == "do not publish this note" + + +def test_open_and_legacy_rows_are_never_exported(): + result = corpus.census_from_hits([ + hit("tp-1", "true_positive"), + hit("open-1", "open"), + hit("legacy-1", None), + ]) + assert result.labelled_count == 1 + assert [row["_id"] for row in result.rows] == ["tp-1"] + + +def test_missing_field_and_model_state_counts(): + row = hit("tp-1", "true_positive") + row["_source"].pop("sensor") + row["_source"].pop("model_state_id") + row["_source"]["model_scores"] = {"isolation_forest": None, "lstm_ae": None, "hbos": None} + result = corpus.census_from_hits([row]) + groups = corpus._groups(result.rows) + assert groups["missing_fields"]["sensor"] == 1 + assert groups["missing_fields"]["model_state_id"] == 1 + assert groups["missing_fields"]["model_scores"] == 1 + assert groups["distinct_model_state_ids"] == 0 + + +def test_report_hash_matches_export(tmp_path): + report, _ = run_local(tmp_path, [hit("tp-1", "true_positive")]) + export = tmp_path / "labels.ndjson" + assert corpus.hashlib.sha256(export.read_bytes()).hexdigest() == report["export"]["sha256"] + report_file = tmp_path / "report.json" + assert corpus.report_digest(json.loads(report_file.read_text(encoding="utf-8"))) == report["report_sha256"] + + +def test_endpoint_is_explicit_and_unreachable_is_actionable(monkeypatch): + with pytest.raises(corpus.CorpusError, match="explicit"): + corpus.validate_endpoint(None) + with pytest.raises(corpus.CorpusError, match="http"): + corpus.validate_endpoint("elasticsearch:9200") + + def unreachable(*args, **kwargs): + raise corpus.urllib.error.URLError("connection refused") + + monkeypatch.setattr(corpus.urllib.request, "urlopen", unreachable) + with pytest.raises(corpus.ElasticsearchUnavailable, match="could not reach"): + corpus.run_elasticsearch( + endpoint="http://127.0.0.1:1", output="/tmp/no-snapshot", report_output="/tmp/no-report", + api_key=None, username=None, password=None, page_size=10, timeout=0.1, + ) + + +def test_elasticsearch_client_only_sends_read_methods(monkeypatch): + calls = [] + + def fake_request(client, method, path, body=None): + calls.append((method, path, body)) + if method == "POST" and path == "/_search" and body.get("aggs"): + return { + "hits": {"total": {"value": 3, "relation": "eq"}, "hits": []}, + "aggregations": { + "statuses": {"buckets": [{"key": "open", "doc_count": 3}]}, + "alert_time": {"count": 3, "min": None, "max": None}, + }, + } + if "/_pit" in path and method == "POST": + return {"id": "pit-1"} + if path == "/_search": + return {"hits": {"hits": []}, "pit_id": "pit-1"} + if method == "DELETE": + return {"succeeded": True} + raise AssertionError((method, path)) + + monkeypatch.setattr(corpus, "_request", fake_request) + census, rows = corpus.fetch_elasticsearch_census( + {"endpoint": "http://es.invalid", "headers": {}, "timeout": 1}, page_size=2) + assert census.total == 3 and census.labelled_count == 0 and rows == [] + assert {method for method, _, _ in calls} <= {"POST", "DELETE"} + assert all(method in {"POST", "DELETE"} for method, _, _ in calls) + assert any("/_pit" in path for _, path, _ in calls)