File size: 11,617 Bytes
875e4af
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
"""
Tests for src/turn_detector/inference.py.

Honesty about what's actually tested here (this sandbox has no torch/
transformers and no network β€” see docs/INITIAL_ANALYSIS.md):

  - Audio validation, decision logic (threshold/debounce/hysteresis), and
    output schema: tested directly, no ML involved, fully real.
  - The classifier stage (loading models/whisper_classifier.joblib and
    running predict_proba): tested with the REAL trained classifier and
    REAL embeddings from the actual EXP-004 Colab run
    (artifacts/exp004/whisper_embeddings.npz) β€” not synthetic data.
  - The Whisper encoder stage itself (audio -> embedding): CANNOT be
    tested in this sandbox (no torch/transformers). Tests that would need
    it are skipped with a clear reason, not silently passed or faked.
"""

import sys
from pathlib import Path

sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))

import numpy as np

from turn_detector.inference import (
    TurnDetector, TurnDetectorConfig, TurnDecisionConfig, TurnDecisionState,
    validate_audio, decide, AudioValidationError, InferenceError,
    DEFAULT_CLASSIFIER_PATH,
)

REAL_EMBEDDINGS_PATH = Path(__file__).resolve().parents[1] / "artifacts" / "exp004" / "whisper_embeddings.npz"

SKIPPED = []


def skip(reason):
    SKIPPED.append(reason)
    print(f"  SKIPPED: {reason}")


# ---------------------------------------------------------------------------
# Audio validation
# ---------------------------------------------------------------------------

def test_validate_audio_accepts_normal_clip():
    audio = np.random.randn(16000).astype(np.float32) * 0.1
    validate_audio(audio, 16000)  # should not raise


def test_validate_audio_rejects_none():
    try:
        validate_audio(None, 16000)
        assert False, "should have raised"
    except AudioValidationError:
        pass


def test_validate_audio_rejects_empty():
    try:
        validate_audio(np.array([]), 16000)
        assert False, "should have raised"
    except AudioValidationError:
        pass


def test_validate_audio_rejects_nan():
    audio = np.array([0.1, float("nan"), 0.2], dtype=np.float32)
    try:
        validate_audio(audio, 16000)
        assert False, "should have raised"
    except AudioValidationError:
        pass


def test_validate_audio_rejects_bad_sample_rate():
    audio = np.random.randn(16000).astype(np.float32)
    try:
        validate_audio(audio, 0)
        assert False, "should have raised"
    except AudioValidationError:
        pass
    try:
        validate_audio(audio, -16000)
        assert False, "should have raised"
    except AudioValidationError:
        pass


def test_validate_audio_rejects_too_short():
    audio = np.random.randn(100).astype(np.float32)  # ~6ms at 16kHz
    try:
        validate_audio(audio, 16000)
        assert False, "should have raised"
    except AudioValidationError:
        pass


def test_validate_audio_accepts_2d_multichannel():
    audio = np.random.randn(16000, 2).astype(np.float32) * 0.1
    validate_audio(audio, 16000)  # should not raise


# ---------------------------------------------------------------------------
# Decision logic (threshold / min_confidence / debounce)
# ---------------------------------------------------------------------------

def test_decide_basic_threshold():
    cfg = TurnDecisionConfig(end_threshold=0.5)
    assert decide(0.6, cfg) == "END"
    assert decide(0.4, cfg) == "CONTINUE"
    assert decide(0.5, cfg) == "END"  # boundary: >= threshold


def test_decide_custom_threshold():
    cfg = TurnDecisionConfig(end_threshold=0.8)
    assert decide(0.7, cfg) == "CONTINUE"
    assert decide(0.85, cfg) == "END"


def test_decide_min_confidence_blocks_uncertain_end():
    cfg = TurnDecisionConfig(end_threshold=0.5, min_confidence=0.5)
    # p=0.6 -> confidence = |0.6-0.5|*2 = 0.2, below 0.5 -> forced CONTINUE
    assert decide(0.6, cfg) == "CONTINUE"
    # p=0.9 -> confidence = 0.8, above 0.5 -> allowed to be END
    assert decide(0.9, cfg) == "END"


def test_decision_state_debounce_requires_consecutive_calls():
    cfg = TurnDecisionConfig(end_threshold=0.5, debounce_consecutive_calls=3)
    state = TurnDecisionState(cfg)
    r1 = state.update(0.9)
    r2 = state.update(0.9)
    r3 = state.update(0.9)
    assert r1["decision"] == "CONTINUE"
    assert r2["decision"] == "CONTINUE"
    assert r3["decision"] == "END"


def test_decision_state_debounce_resets_on_low_probability():
    cfg = TurnDecisionConfig(end_threshold=0.5, debounce_consecutive_calls=2)
    state = TurnDecisionState(cfg)
    state.update(0.9)
    r_reset = state.update(0.1)  # breaks the streak
    r_after = state.update(0.9)  # streak restarts, only 1 consecutive so far
    assert r_reset["decision"] == "CONTINUE"
    assert r_after["decision"] == "CONTINUE"


def test_decision_state_reset():
    cfg = TurnDecisionConfig(debounce_consecutive_calls=2)
    state = TurnDecisionState(cfg)
    state.update(0.9)
    state.reset()
    r = state.update(0.9)
    assert r["decision"] == "CONTINUE"  # streak was reset, only 1 consecutive


# ---------------------------------------------------------------------------
# Classifier loading β€” REAL artifact
# ---------------------------------------------------------------------------

def test_classifier_artifact_exists():
    assert DEFAULT_CLASSIFIER_PATH.exists(), (
        f"Expected trained classifier at {DEFAULT_CLASSIFIER_PATH} "
        f"(models/whisper_classifier.joblib) β€” real artifact from EXP-004."
    )


def test_turndetector_loads_real_classifier():
    detector = TurnDetector()
    assert detector._classifier is not None
    # Real sklearn Pipeline structure, not a mock
    assert hasattr(detector._classifier, "predict_proba")


def test_missing_classifier_raises_clear_error(tmp_path=None):
    import tempfile
    fake_path = Path(tempfile.gettempdir()) / "definitely_does_not_exist.joblib"
    if fake_path.exists():
        fake_path.unlink()
    cfg = TurnDetectorConfig(classifier_path=fake_path)
    try:
        TurnDetector(config=cfg)
        assert False, "should have raised InferenceError"
    except InferenceError as e:
        assert "not found" in str(e).lower()


# ---------------------------------------------------------------------------
# End-to-end predict() β€” REAL classifier + REAL embeddings (injected,
# since the Whisper encoder itself can't run in this sandbox)
# ---------------------------------------------------------------------------

def test_predict_end_to_end_with_real_classifier_and_real_embedding():
    if not REAL_EMBEDDINGS_PATH.exists():
        skip(f"real embeddings not found at {REAL_EMBEDDINGS_PATH} β€” cannot run this test")
        return

    data = np.load(REAL_EMBEDDINGS_PATH)
    real_embedding = data["embeddings"][0]
    assert real_embedding.shape == (384,), f"unexpected embedding shape: {real_embedding.shape}"

    def real_embed_fn(audio, sr):
        return real_embedding  # bypasses Whisper (unavailable here), uses a REAL saved embedding

    detector = TurnDetector(embed_fn=real_embed_fn)
    fake_audio = np.random.randn(16000).astype(np.float32) * 0.1
    result = detector.predict(fake_audio, sr=16000)

    # Output schema
    assert set(result.keys()) == {"decision", "end_probability", "continue_probability", "latency_ms"}
    assert result["decision"] in ("END", "CONTINUE")
    assert 0.0 <= result["end_probability"] <= 1.0
    assert 0.0 <= result["continue_probability"] <= 1.0
    assert abs(result["end_probability"] + result["continue_probability"] - 1.0) < 1e-9
    assert result["latency_ms"] >= 0
    assert result["latency_ms"] < 100, (
        f"latency suspiciously high ({result['latency_ms']}ms) for a classifier-only call "
        f"with an injected embedding β€” likely counting disk I/O that shouldn't be in the hot path"
    )


def test_predict_matches_real_classifier_output_directly():
    """Cross-check: TurnDetector.predict()'s probability should exactly
    match calling the loaded classifier's predict_proba() directly on the
    same real embedding β€” catches any class-index-ordering bugs.
    """
    if not REAL_EMBEDDINGS_PATH.exists():
        skip(f"real embeddings not found at {REAL_EMBEDDINGS_PATH} β€” cannot run this test")
        return

    import joblib
    data = np.load(REAL_EMBEDDINGS_PATH)
    real_embedding = data["embeddings"][3]

    clf = joblib.load(DEFAULT_CLASSIFIER_PATH)
    classes = list(clf.named_steps["clf"].classes_)
    end_idx = classes.index(1) if 1 in classes else classes.index(True)
    expected_end_proba = float(clf.predict_proba(real_embedding.reshape(1, -1))[0][end_idx])

    def embed_fn(audio, sr):
        return real_embedding

    detector = TurnDetector(embed_fn=embed_fn)
    result = detector.predict(np.random.randn(16000).astype(np.float32) * 0.1, sr=16000)

    assert abs(result["end_probability"] - expected_end_proba) < 1e-9, (
        f"TurnDetector's end_probability ({result['end_probability']}) doesn't match "
        f"direct classifier output ({expected_end_proba}) β€” class-index bug likely"
    )


def test_predict_streaming_step_applies_debounce():
    if not REAL_EMBEDDINGS_PATH.exists():
        skip(f"real embeddings not found at {REAL_EMBEDDINGS_PATH} β€” cannot run this test")
        return

    data = np.load(REAL_EMBEDDINGS_PATH)
    real_embedding = data["embeddings"][0]

    def embed_fn(audio, sr):
        return real_embedding

    cfg = TurnDetectorConfig(decision=TurnDecisionConfig(debounce_consecutive_calls=2))
    detector = TurnDetector(config=cfg, embed_fn=embed_fn)
    fake_audio = np.random.randn(16000).astype(np.float32) * 0.1

    r1 = detector.predict_streaming_step(fake_audio)
    r2 = detector.predict_streaming_step(fake_audio)
    # Same underlying probability both times (same injected embedding) β€”
    # decision should only commit to END (if end_probability warrants it)
    # on the 2nd call due to debounce_consecutive_calls=2.
    if r1["end_probability"] >= 0.5:
        assert r1["decision"] == "CONTINUE"  # debouncing on call 1
        assert r2["decision"] == "END"       # committed on call 2
    detector.reset_stream()


def test_load_whisper_without_transformers_raises_clear_error():
    """Confirms the graceful-failure path for the Whisper stage itself,
    which genuinely cannot run in this sandbox β€” this test's PASS
    condition is that it fails informatively, not that Whisper works.
    """
    try:
        import transformers  # noqa
        skip("transformers IS installed in this environment β€” this test's premise doesn't apply here")
        return
    except ImportError:
        pass

    detector = TurnDetector()
    try:
        detector.load_whisper()
        assert False, "expected InferenceError since transformers is not installed"
    except InferenceError as e:
        assert "torch/transformers" in str(e) or "transformers" in str(e).lower()


if __name__ == "__main__":
    import sys as _sys
    test_fns = [v for k, v in list(globals().items()) if k.startswith("test_") and callable(v)]
    passed, failed = 0, []
    for fn in test_fns:
        try:
            fn()
            passed += 1
            print(f"PASS  {fn.__name__}")
        except Exception as e:
            failed.append((fn.__name__, e))
            print(f"FAIL  {fn.__name__}: {e}")
    print()
    print(f"{passed}/{len(test_fns)} passed, {len(SKIPPED)} skipped (real reasons logged above)")
    if failed:
        print("FAILURES:")
        for name, e in failed:
            print(f"  {name}: {e}")
        _sys.exit(1)