From 48c030ed13ed30440a2dbc5901daca5a6de36672 Mon Sep 17 00:00:00 2001 From: Adam Moussa <166072409+amoussa1229@users.noreply.github.com> Date: Fri, 29 May 2026 14:52:53 -0400 Subject: [PATCH] Add tests for the classifier handler transforms --- tests/test_classifier_handler.py | 523 +++++++++++++++++++++++++++++++ 1 file changed, 523 insertions(+) create mode 100644 tests/test_classifier_handler.py diff --git a/tests/test_classifier_handler.py b/tests/test_classifier_handler.py new file mode 100644 index 0000000..37ef4d1 --- /dev/null +++ b/tests/test_classifier_handler.py @@ -0,0 +1,523 @@ +"""Tests for the classifier handler: pure transforms and the handler() entrypoint. + +Uses the importlib trick to avoid an ambiguous bare ``import handler`` (both +lambdas/classifier/handler.py and lambdas/slack_post/handler.py are on +pythonpath). ``awswrangler`` is stubbed at the sys.modules level before the +module is loaded, so no real AWS/network calls are made. +""" + +from __future__ import annotations + +import csv +import io +import json +import sys +import tempfile +import types +from pathlib import Path +from unittest.mock import MagicMock + +import pytest + +# --------------------------------------------------------------------------- +# Load the classifier handler under a unique module name. +# awswrangler must be stubbed before exec_module() runs the top-level imports. +# --------------------------------------------------------------------------- + +# Build a minimal awswrangler stub so handler.py's top-level `import awswrangler as wr` +# succeeds without the real package installed. +_wr_stub = types.ModuleType("awswrangler") +_wr_stub.s3 = types.ModuleType("awswrangler.s3") +_wr_stub.s3.to_parquet = MagicMock() +sys.modules.setdefault("awswrangler", _wr_stub) +sys.modules.setdefault("awswrangler.s3", _wr_stub.s3) + +import importlib.util # noqa: E402 + +_HANDLER_PATH = ( + Path(__file__).resolve().parents[1] / "lambdas" / "classifier" / "handler.py" +) +_spec = importlib.util.spec_from_file_location("classifier_handler", _HANDLER_PATH) +handler = importlib.util.module_from_spec(_spec) +_spec.loader.exec_module(handler) + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +_HEADER = [ + "WO Number", + "WO Description", + "Equipment Code", + "Organization", + "Due Date", + "Department", + "WO Status", + "Hold Reason", + "Last Comment", + "Last Comment By", + "Last Comment Date", + "Contractor", + "Contractor Description", +] + + +def _make_row( + wo_number="WO-001", + wo_description="Fix HVAC", + equipment_code="HVAC-01", + site="ABQ5", + due_date="2026-05-30", + department="SSP", + wo_status="IP", + hold_reason="", + last_comment="WO schedule confirmed with vendor.", + last_comment_by="tech@example.com", + last_comment_date="2026-05-28", + contractor="ABC HVAC", + contractor_description="HVAC Services", +): + return [ + wo_number, + wo_description, + equipment_code, + site, + due_date, + department, + wo_status, + hold_reason, + last_comment, + last_comment_by, + last_comment_date, + contractor, + contractor_description, + ] + + +def _write_csv(rows: list[list], header: list[str] = _HEADER) -> str: + """Write header + rows to a temp CSV file and return the path.""" + with tempfile.NamedTemporaryFile( + mode="w", suffix=".csv", delete=False, newline="" + ) as fh: + writer = csv.writer(fh) + writer.writerow(header) + writer.writerows(rows) + return fh.name + + +def _write_xlsx(rows: list[list], header: list[str] = _HEADER) -> str: + """Write header + rows to a temp xlsx file and return the path.""" + import openpyxl + + wb = openpyxl.Workbook() + ws = wb.active + ws.append(header) + for row in rows: + ws.append(row) + with tempfile.NamedTemporaryFile(suffix=".xlsx", delete=False) as fh: + path = fh.name + wb.save(path) + return path + + +# --------------------------------------------------------------------------- +# _resolve_columns +# --------------------------------------------------------------------------- + + +class TestResolveColumns: + def test_normal_13_col_header(self): + result = handler._resolve_columns(_HEADER) + assert result["wo_number"] == 0 + assert result["site"] == 3 # "Organization" column + assert result["wo_status"] == 6 + assert result["hold_reason"] == 7 + assert result["last_comment"] == 8 + assert result["last_comment_by"] == 9 + assert result["last_comment_date"] == 10 + + def test_header_drift_extra_spaces_and_case(self): + drifted = [ + " WO NUMBER ", + "WO DESCRIPTION", + "Equipment Code", + "Organization", + "Due Date", + "Department", + "WO Status ", + "Hold Reason", + "Last Comment", + "Last Comment By", + "Last Comment Date", + "Contractor", + "Contractor Description", + ] + result = handler._resolve_columns(drifted) + assert result["wo_number"] == 0 + assert result["wo_status"] == 6 + assert result["last_comment"] == 8 + + def test_prefix_collision_last_comment_resolves_to_exact_column(self): + """'last comment' must resolve to col 8 (Last Comment), not col 9 or 10.""" + result = handler._resolve_columns(_HEADER) + col_idx = result["last_comment"] + assert col_idx == 8 + assert _HEADER[col_idx] == "Last Comment" + + # last_comment_by and last_comment_date must not steal last_comment's slot + assert result["last_comment_by"] == 9 + assert result["last_comment_date"] == 10 + + +# --------------------------------------------------------------------------- +# _read_rows +# --------------------------------------------------------------------------- + + +class TestReadRows: + def test_csv_header_and_rows(self): + data = [_make_row(wo_number="WO-001"), _make_row(wo_number="WO-002")] + path = _write_csv(data) + header, rows = handler._read_rows(path, "raw/export.csv") + assert header[0] == "WO Number" + assert len(rows) == 2 + assert rows[0][0] == "WO-001" + assert rows[1][0] == "WO-002" + + def test_xlsx_header_and_rows(self): + data = [_make_row(wo_number="WO-003"), _make_row(wo_number="WO-004")] + path = _write_xlsx(data) + header, rows = handler._read_rows(path, "raw/export.xlsx") + assert header[0] == "WO Number" + assert len(rows) == 2 + assert rows[0][0] == "WO-003" + assert rows[1][0] == "WO-004" + + +# --------------------------------------------------------------------------- +# _build_snapshot +# --------------------------------------------------------------------------- + + +class TestBuildSnapshot: + def test_blank_comment_rows_excluded(self): + rows = [ + _make_row( + wo_number="WO-010", last_comment="Schedule confirmed." + ), + _make_row(wo_number="WO-011", last_comment=""), # blank — excluded + _make_row( + wo_number="WO-012", last_comment="" + ), # strips to "" — excluded + ] + df, blank = handler._build_snapshot(_HEADER, rows) + assert blank == 2 + assert len(df) == 1 + assert df.iloc[0]["wo_number"] == "WO-010" + + def test_missing_last_comment_column_raises(self): + # Remove all "last comment" variants so last_comment cannot resolve via + # substring match either — only columns with no "last comment" remain. + bad_header = [ + "WO Number", + "WO Description", + "Equipment Code", + "Organization", + "Due Date", + "Department", + "WO Status", + "Hold Reason", + "Contractor", + "Contractor Description", + ] + rows = [ + [ + "WO-001", + "Fix HVAC", + "HVAC", + "ABQ5", + "2026-05-30", + "SSP", + "IP", + "", + "ABC", + "HVAC", + ] + ] + with pytest.raises(ValueError, match="required columns"): + handler._build_snapshot(bad_header, rows) + + def test_missing_wo_status_column_raises(self): + bad_header = [ + "WO Number", + "WO Description", + "Equipment Code", + "Organization", + "Due Date", + "Department", + # "WO Status" missing + "Hold Reason", + "Last Comment", + "Last Comment By", + "Last Comment Date", + "Contractor", + "Contractor Description", + ] + rows = [ + [ + "WO-001", + "Fix HVAC", + "HVAC", + "ABQ5", + "2026-05-30", + "SSP", + "", + "Schedule confirmed.", + "tech@example.com", + "2026-05-28", + "ABC", + "HVAC", + ] + ] + with pytest.raises(ValueError, match="required columns"): + handler._build_snapshot(bad_header, rows) + + def test_is_escalation_and_is_action_derived(self): + import classify as clf + + rows = [ + _make_row( + wo_number="WO-020", + wo_status="H", + hold_reason="REPORT", + last_comment="3rd attempt process for schedule confirmation.", + ), + _make_row( + wo_number="WO-021", + wo_status="IP", + hold_reason="", + last_comment="WO schedule confirmed with vendor.", + ), + ] + df, _ = handler._build_snapshot(_HEADER, rows) + + esc_row = df[df["wo_number"] == "WO-020"].iloc[0] + assert bool(esc_row["is_escalation"]) is True + assert esc_row["category"] in clf.ESCALATION_CATEGORIES + + routine_row = df[df["wo_number"] == "WO-021"].iloc[0] + assert bool(routine_row["is_escalation"]) is False + assert routine_row["category"] == "Schedule Confirmed" + + +# --------------------------------------------------------------------------- +# _build_summary +# --------------------------------------------------------------------------- + + +class TestBuildSummary: + def _make_df(self): + """Build a small DataFrame via _build_snapshot.""" + rows = [ + _make_row( + wo_number="WO-030", + wo_status="H", + hold_reason="REPORT", + last_comment="3rd attempt process for schedule confirmation.", + site="ABQ5", + ), + _make_row( + wo_number="WO-031", + wo_status="H", + hold_reason="REPORT", + last_comment="3rd attempt process for schedule confirmation.", + site="ACY9", + ), + _make_row( + wo_number="WO-032", + wo_status="IP", + hold_reason="", + last_comment="WO schedule confirmed with vendor.", + site="ABQ5", + ), + # This row produces a mismatch: completion comment + IP status + _make_row( + wo_number="WO-033", + wo_status="IP", + hold_reason="REPORT", + last_comment="Vendor arrived and performed task.", + site="ABQ5", + ), + ] + df, blank = handler._build_snapshot(_HEADER, rows) + return df, blank + + def test_third_escalation_count(self): + df, blank = self._make_df() + summary = handler._build_summary(df, "2026-05-28", "raw/export.csv", blank) + assert summary["third_escalation_count"] == 2 + + def test_category_counts_present(self): + df, blank = self._make_df() + summary = handler._build_summary(df, "2026-05-28", "raw/export.csv", blank) + assert "3rd Escalation" in summary["category_counts"] + assert summary["category_counts"]["3rd Escalation"] == 2 + + def test_escalation_total(self): + df, blank = self._make_df() + summary = handler._build_summary(df, "2026-05-28", "raw/export.csv", blank) + assert summary["escalation_total"] == 2 + + def test_action_needed_and_routine_sum_to_classified_total(self): + df, blank = self._make_df() + summary = handler._build_summary(df, "2026-05-28", "raw/export.csv", blank) + assert ( + summary["action_needed"] + summary["routine"] == summary["classified_total"] + ) + + def test_top_sites_shape(self): + df, blank = self._make_df() + summary = handler._build_summary(df, "2026-05-28", "raw/export.csv", blank) + assert isinstance(summary["top_sites"], list) + for entry in summary["top_sites"]: + assert "site" in entry + assert "count" in entry + # ABQ5 appears 3 times, should be first + assert summary["top_sites"][0]["site"] == "ABQ5" + + def test_mismatches_list(self): + df, blank = self._make_df() + summary = handler._build_summary(df, "2026-05-28", "raw/export.csv", blank) + assert isinstance(summary["mismatches"], list) + # WO-033 has a mismatch (completion comment + REPORT hold) + assert len(summary["mismatches"]) >= 1 + mismatch_wos = [m["wo_number"] for m in summary["mismatches"]] + assert "WO-033" in mismatch_wos + + +# --------------------------------------------------------------------------- +# _event_dt +# --------------------------------------------------------------------------- + + +class TestEventDt: + def test_event_time_extracted(self): + record = {"eventTime": "2026-05-28T22:23:40.123Z"} + assert handler._event_dt(record) == "2026-05-28" + + def test_no_event_time_falls_back_to_today(self): + from datetime import datetime, timezone + + record = {} + result = handler._event_dt(record) + today = datetime.now(timezone.utc).strftime("%Y-%m-%d") + # Result must look like a YYYY-MM-DD date string + assert len(result) == 10 + assert result[4] == "-" and result[7] == "-" + assert result == today + + +# --------------------------------------------------------------------------- +# handler() entrypoint — two S3 records, monkeypatched AWS boundaries +# --------------------------------------------------------------------------- + + +class TestHandlerEntrypoint: + def _build_csv_bytes( + self, + wo_status="IP", + last_comment="Schedule confirmed with vendor.", + ): + """Return CSV bytes for a single-row export.""" + buf = io.StringIO() + writer = csv.writer(buf) + writer.writerow(_HEADER) + writer.writerow(_make_row(wo_status=wo_status, last_comment=last_comment)) + return buf.getvalue().encode("utf-8") + + def test_two_records_processed_slack_invoked_once_per_dt( + self, tmp_path, monkeypatch + ): + csv_bytes_a = self._build_csv_bytes() + csv_bytes_b = self._build_csv_bytes( + last_comment="3rd attempt process for schedule confirmation." + ) + + # Track put_object and lambda.invoke calls + put_calls: list[dict] = [] + invoke_calls: list[dict] = [] + + def fake_download_fileobj(bucket, key, fh): + if "file_a" in key: + fh.write(csv_bytes_a) + else: + fh.write(csv_bytes_b) + + fake_s3 = MagicMock() + fake_s3.download_fileobj.side_effect = fake_download_fileobj + fake_s3.put_object.side_effect = lambda **kw: put_calls.append(kw) + + fake_lambda = MagicMock() + fake_lambda.invoke.side_effect = lambda **kw: invoke_calls.append(kw) + + monkeypatch.setattr(handler, "_s3", fake_s3) + monkeypatch.setattr(handler, "_lambda", fake_lambda) + monkeypatch.setattr(handler.wr.s3, "to_parquet", MagicMock()) + + # Set the SLACK_POST_FUNCTION_NAME env var so the lambda invoke fires + monkeypatch.setenv("SLACK_POST_FUNCTION_NAME", "apm-slack-post") + + event = { + "Records": [ + { + "eventTime": "2026-05-28T10:00:00.000Z", + "s3": { + "bucket": {"name": "test-bucket"}, + "object": {"key": "raw/file_a.csv"}, + }, + }, + { + "eventTime": "2026-05-28T11:00:00.000Z", + "s3": { + "bucket": {"name": "test-bucket"}, + "object": {"key": "raw/file_b.csv"}, + }, + }, + ] + } + + result = handler.handler(event, None) + + # Both records processed + assert len(result["processed"]) == 2 + + # summary.json and details.json written for each record (2 put_object calls each = 4) + assert len(put_calls) == 4 + + # Slack post invoked exactly once (both records share the same dt "2026-05-28") + assert len(invoke_calls) == 1 + assert invoke_calls[0]["FunctionName"] == "apm-slack-post" + payload = json.loads(invoke_calls[0]["Payload"]) + assert payload["dt"] == "2026-05-28" + + def test_non_export_key_skipped(self, monkeypatch): + fake_s3 = MagicMock() + fake_lambda = MagicMock() + monkeypatch.setattr(handler, "_s3", fake_s3) + monkeypatch.setattr(handler, "_lambda", fake_lambda) + + event = { + "Records": [ + { + "eventTime": "2026-05-28T10:00:00.000Z", + "s3": { + "bucket": {"name": "test-bucket"}, + "object": {"key": "raw/not-an-export.txt"}, + }, + } + ] + } + result = handler.handler(event, None) + assert result["processed"] == [] + fake_s3.download_fileobj.assert_not_called()