Add tests for the Haiku fallback

This commit is contained in:
Adam Moussa 2026-05-29 14:54:47 -04:00
parent 48c030ed13
commit e555a75d97

271
tests/test_haiku.py Normal file
View file

@ -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 = "<html>Uplift request submitted pending management approval.</html>"
_TRIVIAL_COMMENT = "<html>Ok</html>" # 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", "", "<html>WO schedule confirmed with vendor.</html>"
)
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", "<html>xyz nothing here</html>"
)
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"