""" 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)