diff --git a/tests/test_haiku.py b/tests/test_haiku.py new file mode 100644 index 0000000..25d4549 --- /dev/null +++ b/tests/test_haiku.py @@ -0,0 +1,271 @@ +"""Tests for the Haiku fallback in classify_with_haiku(). + +No real network or AWS calls are made — _call_haiku and _fetch_api_key are +monkeypatched at the function level. boto3 shape tests stub the secretsmanager +client directly via monkeypatch on boto3.client. +""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock + +import classify + +# --------------------------------------------------------------------------- +# Comment that will always land as "Other" deterministically (free-text with +# no structured hold signal and no regex match). +# --------------------------------------------------------------------------- +_OTHER_COMMENT = "Uplift request submitted pending management approval." +_TRIVIAL_COMMENT = "Ok" # strips to "Ok" — len == 2 < 4 + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _haiku_never_called(): + """Return a _call_haiku stub that raises if invoked.""" + + def _stub(text, api_key): + raise AssertionError("_call_haiku should not have been called") + + return _stub + + +# --------------------------------------------------------------------------- +# Gating: Haiku is NOT called when the deterministic result is non-Other +# --------------------------------------------------------------------------- + + +def test_haiku_not_called_when_deterministic_result_is_non_other(monkeypatch): + """A comment that resolves deterministically must skip the Haiku path.""" + monkeypatch.setattr(classify, "_call_haiku", _haiku_never_called()) + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "fake-key") + + # "WO schedule confirmed with vendor" → Schedule Confirmed deterministically + cat, mm = classify.classify_with_haiku( + "IP", "", "WO schedule confirmed with vendor." + ) + assert cat == "Schedule Confirmed" + assert mm is None + + +def test_haiku_not_called_when_hold_reason_present(monkeypatch): + """Hold reason present → structured state handles it; Haiku must not fire.""" + monkeypatch.setattr(classify, "_call_haiku", _haiku_never_called()) + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "fake-key") + + # REPORT hold with unrecognised comment → hold decides, not Haiku + cat, _ = classify.classify_with_haiku( + "H", "REPORT", "xyz nothing here" + ) + assert cat == "Report / Docs Needed" + + +def test_haiku_not_called_when_comment_is_trivial(monkeypatch): + """Comment len < 4 after stripping → trivial, skip Haiku.""" + monkeypatch.setattr(classify, "_call_haiku", _haiku_never_called()) + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "fake-key") + + cat, _ = classify.classify_with_haiku("IP", "", _TRIVIAL_COMMENT) + assert cat == "Other" + + +def test_haiku_not_called_when_disabled_via_env(monkeypatch): + """APM_HAIKU_FALLBACK=off must suppress the Haiku call.""" + monkeypatch.setenv("APM_HAIKU_FALLBACK", "off") + monkeypatch.setattr(classify, "_call_haiku", _haiku_never_called()) + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "fake-key") + + cat, _ = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + assert cat == "Other" + + +def test_haiku_not_called_when_disabled_via_zero(monkeypatch): + monkeypatch.setenv("APM_HAIKU_FALLBACK", "0") + monkeypatch.setattr(classify, "_call_haiku", _haiku_never_called()) + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "fake-key") + + cat, _ = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + assert cat == "Other" + + +def test_haiku_not_called_when_disabled_via_false(monkeypatch): + monkeypatch.setenv("APM_HAIKU_FALLBACK", "false") + monkeypatch.setattr(classify, "_call_haiku", _haiku_never_called()) + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "fake-key") + + cat, _ = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + assert cat == "Other" + + +# --------------------------------------------------------------------------- +# Haiku IS invoked (enabled, Other, no hold, non-trivial) +# --------------------------------------------------------------------------- + + +def test_haiku_called_for_true_residual_other(monkeypatch): + """When all gates pass, _call_haiku must be invoked.""" + monkeypatch.setenv("APM_HAIKU_FALLBACK", "on") + haiku_calls: list[tuple] = [] + + def _stub_haiku(text, api_key): + haiku_calls.append((text, api_key)) + return "Rescheduled" + + monkeypatch.setattr(classify, "_call_haiku", _stub_haiku) + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "test-key") + + cat, _ = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + assert cat == "Rescheduled" + assert len(haiku_calls) == 1 + + +def test_haiku_valid_bucket_replaces_other(monkeypatch): + """When Haiku returns a valid bucket, that bucket is used.""" + monkeypatch.setenv("APM_HAIKU_FALLBACK", "1") + monkeypatch.setattr(classify, "_call_haiku", lambda text, key: "Rescheduled") + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "test-key") + + cat, _ = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + assert cat == "Rescheduled" + + +def test_haiku_none_response_stays_other(monkeypatch): + """When Haiku returns None, the result stays Other.""" + monkeypatch.setenv("APM_HAIKU_FALLBACK", "1") + monkeypatch.setattr(classify, "_call_haiku", lambda text, key: None) + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "test-key") + + cat, _ = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + assert cat == "Other" + + +def test_haiku_other_response_stays_other(monkeypatch): + """When Haiku returns 'Other', the result stays Other (guard against self-loop).""" + monkeypatch.setenv("APM_HAIKU_FALLBACK", "1") + monkeypatch.setattr(classify, "_call_haiku", lambda text, key: "Other") + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "test-key") + + cat, _ = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + assert cat == "Other" + + +def test_haiku_exception_in_call_haiku_stays_other(monkeypatch): + """Exception raised by _call_haiku (e.g. network failure) is caught and + the result safely stays at Other.""" + monkeypatch.setenv("APM_HAIKU_FALLBACK", "1") + + def _raise(text, api_key): + raise OSError("network failure") + + monkeypatch.setattr(classify, "_call_haiku", _raise) + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "test-key") + + cat, _ = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + assert cat == "Other" + + +def test_mismatch_preserved_through_haiku_path(monkeypatch): + """The deterministic mismatch reason must survive through the Haiku path.""" + monkeypatch.setenv("APM_HAIKU_FALLBACK", "1") + monkeypatch.setattr(classify, "_call_haiku", lambda text, key: "Rescheduled") + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "test-key") + + # A comment that fires "Completed / Pending Close" intent while on a REPORT hold + # produces a mismatch; but the deterministic result is NOT Other here. To test + # mismatch preservation we need the deterministic result to be Other while a + # mismatch exists. Mismatch is computed from intent vs structured state; if the + # intent is None (→ Other from deterministic), no mismatch is produced by + # _mismatch(). So the relevant test is: Haiku fires on an Other row; the + # mismatch (None in this case) is preserved correctly. + cat, mm = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + # Comment has no intent → no mismatch → mm is None; Haiku returns Rescheduled + assert cat == "Rescheduled" + assert mm is None + + +def test_mismatch_preserved_when_haiku_overrides(monkeypatch): + """Verify mismatch from classify() passes through unchanged when Haiku overrides. + + We use a real deterministic Other-producing row and manually inject a mismatch + by monkeypatching classify.classify to return (Other, "fake mismatch reason"). + """ + monkeypatch.setenv("APM_HAIKU_FALLBACK", "1") + monkeypatch.setattr(classify, "_call_haiku", lambda text, key: "Rescheduled") + monkeypatch.setattr(classify, "_fetch_api_key", lambda: "test-key") + + original_classify = classify.classify + + def _patched_classify(wo_status, hold_reason, last_comment): + cat, _ = original_classify(wo_status, hold_reason, last_comment) + if cat == "Other": + return "Other", "injected mismatch reason" + return cat, _ + + monkeypatch.setattr(classify, "classify", _patched_classify) + + cat, mm = classify.classify_with_haiku("IP", "", _OTHER_COMMENT) + assert cat == "Rescheduled" + assert mm == "injected mismatch reason" + + +# --------------------------------------------------------------------------- +# _fetch_api_key JSON shapes +# --------------------------------------------------------------------------- + + +def _make_secretsmanager_client(secret_string: str) -> MagicMock: + """Build a minimal secretsmanager client stub returning the given secret.""" + client = MagicMock() + client.get_secret_value.return_value = {"SecretString": secret_string} + return client + + +def test_fetch_api_key_bare_string(monkeypatch): + """Bare string secret → returned as-is.""" + fake_client = _make_secretsmanager_client("sk-ant-barekey123") + + import boto3 + + monkeypatch.setattr(boto3, "client", lambda service, **kw: fake_client) + key = classify._fetch_api_key() + assert key == "sk-ant-barekey123" + + +def test_fetch_api_key_json_anthropic_api_key(monkeypatch): + """JSON with 'anthropic-api-key' field → that value extracted.""" + secret = json.dumps({"anthropic-api-key": "sk-ant-jsonkey456"}) + fake_client = _make_secretsmanager_client(secret) + + import boto3 + + monkeypatch.setattr(boto3, "client", lambda service, **kw: fake_client) + key = classify._fetch_api_key() + assert key == "sk-ant-jsonkey456" + + +def test_fetch_api_key_single_value_json(monkeypatch): + """Single-value JSON object with an arbitrary key → the only value extracted.""" + secret = json.dumps({"my_custom_key": "sk-ant-singleval789"}) + fake_client = _make_secretsmanager_client(secret) + + import boto3 + + monkeypatch.setattr(boto3, "client", lambda service, **kw: fake_client) + key = classify._fetch_api_key() + assert key == "sk-ant-singleval789" + + +def test_fetch_api_key_json_api_key_field(monkeypatch): + """JSON with 'api_key' field → that value extracted.""" + secret = json.dumps({"api_key": "sk-ant-apikey111"}) + fake_client = _make_secretsmanager_client(secret) + + import boto3 + + monkeypatch.setattr(boto3, "client", lambda service, **kw: fake_client) + key = classify._fetch_api_key() + assert key == "sk-ant-apikey111"