"""test_teacher_replay.py — unit test for DPO-pair extraction. We DON'T hit OpenRouter in unit tests (cost + flakiness). We test the deterministic local logic: given fake teacher results, extract_dpo_pairs should produce the right (chosen, rejected) pairs. Run: pytest spikes/005-integrated-trainer-skeleton/tests/test_teacher_replay.py -v """ from __future__ import annotations import sys from pathlib import Path sys.path.insert(0, str(Path(__file__).resolve().parents[1])) from teacher_replay import extract_dpo_pairs # noqa: E402 # ---------------------------------------------------------------------------- # Helpers # ---------------------------------------------------------------------------- def _state(state_id: str, student_action: str) -> dict: return { "state_id": state_id, "messages": [{"role": "user", "content": f"task for {state_id}"}], "student_action": student_action, } def _teacher_call(state_id: str, slug: str, response: str) -> dict: return { "state_id": state_id, "teacher_slug": slug, "response_text": response, "latency_s": 1.0, "prompt_tokens": 100, "completion_tokens": 20, "cost_usd": 0.001, "error": None, } # ---------------------------------------------------------------------------- # Tests # ---------------------------------------------------------------------------- def test_consensus_against_student_yields_pair(): """All 3 teachers agree on X, student picked Y → emit (X, Y) pair.""" states = [_state("s1", student_action="option B")] teacher_calls = [ _teacher_call("s1", "anthropic/opus", "Option A"), _teacher_call("s1", "openai/gpt5", "option a"), # case-insensitive normalize _teacher_call("s1", "deepseek/v4", "Option A"), ] pairs = extract_dpo_pairs(states, teacher_calls, agreement_threshold=2) assert len(pairs) == 1 p = pairs[0] assert p["state_id"] == "s1" assert p["chosen"].lower().strip() == "option a" assert p["rejected"] == "option B" assert p["n_teachers_agreeing"] == 3 def test_no_pair_when_student_matches_consensus(): """All teachers agree with student → no pair (no signal).""" states = [_state("s1", student_action="option A")] teacher_calls = [ _teacher_call("s1", "anthropic/opus", "Option A"), _teacher_call("s1", "openai/gpt5", "Option A"), _teacher_call("s1", "deepseek/v4", "Option A"), ] pairs = extract_dpo_pairs(states, teacher_calls, agreement_threshold=2) assert len(pairs) == 0 def test_no_pair_when_all_teachers_disagree(): """All 3 teachers disagree with each other AND none meets threshold → no pair.""" states = [_state("s1", student_action="option D")] teacher_calls = [ _teacher_call("s1", "anthropic/opus", "Option A"), _teacher_call("s1", "openai/gpt5", "Option B"), _teacher_call("s1", "deepseek/v4", "Option C"), ] pairs = extract_dpo_pairs(states, teacher_calls, agreement_threshold=2) assert len(pairs) == 0 def test_threshold_2_with_2_of_3_consensus(): """2 teachers agree on X, third disagrees, student picked Y → emit (X, Y).""" states = [_state("s1", student_action="option C")] teacher_calls = [ _teacher_call("s1", "anthropic/opus", "Option A"), _teacher_call("s1", "openai/gpt5", "Option A"), _teacher_call("s1", "deepseek/v4", "Option B"), ] pairs = extract_dpo_pairs(states, teacher_calls, agreement_threshold=2) assert len(pairs) == 1 assert pairs[0]["chosen"].lower().strip() == "option a" assert pairs[0]["n_teachers_agreeing"] == 2 def test_strict_threshold_3_filters_2of3(): """With agreement_threshold=3, only unanimous consensus counts.""" states = [_state("s1", student_action="option C")] teacher_calls = [ _teacher_call("s1", "anthropic/opus", "Option A"), _teacher_call("s1", "openai/gpt5", "Option A"), _teacher_call("s1", "deepseek/v4", "Option B"), ] pairs = extract_dpo_pairs(states, teacher_calls, agreement_threshold=3) assert len(pairs) == 0 # only 2/3 agree, threshold is 3 def test_errored_teacher_calls_excluded(): """Failed API calls (error != None) should be ignored when computing consensus.""" states = [_state("s1", student_action="option C")] teacher_calls = [ _teacher_call("s1", "anthropic/opus", "Option A"), {**_teacher_call("s1", "openai/gpt5", "Option A"), "error": "rate limit"}, _teacher_call("s1", "deepseek/v4", "Option A"), ] # Only 2 valid responses, both agree → meets threshold=2 pairs = extract_dpo_pairs(states, teacher_calls, agreement_threshold=2) assert len(pairs) == 1 assert pairs[0]["n_teachers_agreeing"] == 2 def test_multiple_states_independent(): """Each state's pair extraction is independent of other states.""" states = [ _state("s1", student_action="picked X"), # consensus is "picked Y" _state("s2", student_action="picked Z"), # all teachers agree with student ] teacher_calls = [ _teacher_call("s1", "t1", "picked Y"), _teacher_call("s1", "t2", "picked Y"), _teacher_call("s1", "t3", "picked Y"), _teacher_call("s2", "t1", "picked Z"), _teacher_call("s2", "t2", "picked Z"), _teacher_call("s2", "t3", "picked Z"), ] pairs = extract_dpo_pairs(states, teacher_calls, agreement_threshold=2) assert len(pairs) == 1 assert pairs[0]["state_id"] == "s1"