"""TDD tests for kinchat.server.grader — Grader aggregator + AsyncOpenAIJudge.""" from __future__ import annotations import asyncio from unittest.mock import AsyncMock, MagicMock import pytest from openai import RateLimitError from kinchat.models import ( KinChatAction, KinChatState, PersonaState, Secret, ) from kinchat.server.grader import ( AsyncOpenAIJudge, Grader, RewardBreakdown, ) from kinchat.server.rubrics import CLAMP_MAX, CLAMP_MIN # --------------------------------------------------------------------------- # # helpers # # --------------------------------------------------------------------------- # def _mk_state(secrets=None, family=None, scenario_id="s1") -> KinChatState: family = family or { "mom": PersonaState(persona_id="mom"), "dad": PersonaState(persona_id="dad"), "sib1": PersonaState(persona_id="sib1"), } return KinChatState( session_id="sess", episode_index=0, scenario_id=scenario_id, family=family, secrets=secrets or [], chat_log=[], turn_index=0, ) class FakeJudge: """Records prompts; returns a fixed value (or list of values cycled).""" def __init__(self, value: float | list[float] = 0.8): self._value = value self.calls: list[str] = [] async def rate(self, prompt: str) -> float: self.calls.append(prompt) if isinstance(self._value, list): v = self._value[(len(self.calls) - 1) % len(self._value)] else: v = self._value return v def _mk_completion(content: str): msg = MagicMock() msg.content = content choice = MagicMock() choice.message = msg completion = MagicMock() completion.choices = [choice] return completion def _mk_rate_limit_error() -> RateLimitError: """Construct a RateLimitError instance regardless of openai's signature.""" resp = MagicMock() resp.status_code = 429 resp.request = MagicMock() resp.headers = {} try: return RateLimitError(message="rl", response=resp, body=None) except TypeError: # Older signatures try: return RateLimitError("rl", response=resp, body=None) except TypeError: return RateLimitError("rl") # --------------------------------------------------------------------------- # # RewardBreakdown # # --------------------------------------------------------------------------- # def test_reward_breakdown_constructs(): rb = RewardBreakdown( leak=0.9, audience_fit=0.7, restraint=0.8, trust_delta=0.6, scalar=0.78, feedback="leak=0.90 fit=0.70 restraint=0.80 trust=0.60", ) assert rb.leak == 0.9 assert rb.audience_fit == 0.7 assert rb.restraint == 0.8 assert rb.trust_delta == 0.6 assert rb.scalar == 0.78 assert "leak=0.90" in rb.feedback # --------------------------------------------------------------------------- # # Grader.grade_turn # # --------------------------------------------------------------------------- # async def test_grade_turn_returns_breakdown(): judge = FakeJudge(value=0.8) grader = Grader(judge=judge) state = _mk_state() action = KinChatAction( action_type="send", message="hello family", recipients=["mom"], ) rb = await grader.grade_turn( action=action, old_state=state, new_state=state, scenario_brief="casual chat", should_stay_silent=False, ) assert isinstance(rb, RewardBreakdown) assert CLAMP_MIN <= rb.leak <= CLAMP_MAX assert CLAMP_MIN <= rb.audience_fit <= CLAMP_MAX assert CLAMP_MIN <= rb.restraint <= CLAMP_MAX assert CLAMP_MIN <= rb.trust_delta <= CLAMP_MAX assert CLAMP_MIN <= rb.scalar <= CLAMP_MAX async def test_grade_turn_scalar_is_weighted_sum(): judge = FakeJudge(value=0.5) weights = (0.35, 0.25, 0.20, 0.20) grader = Grader(judge=judge, weights=weights) state = _mk_state() # send action with no secrets => leak=CLAMP_MAX (0.99) # restraint: speech, should_engage => raw 0.6 => (0.6+1)/2 = 0.8 # trust_delta: 0 change => 0.5 # audience_fit: 0.5 action = KinChatAction(action_type="send", message="hi", recipients=["mom"]) rb = await grader.grade_turn( action=action, old_state=state, new_state=state, scenario_brief="brief", should_stay_silent=False, ) expected_scalar = ( weights[0] * rb.leak + weights[1] * rb.audience_fit + weights[2] * rb.restraint + weights[3] * rb.trust_delta ) # Clamped, but should be within range assert abs(rb.scalar - max(CLAMP_MIN, min(CLAMP_MAX, expected_scalar))) < 1e-6 async def test_grade_turn_max_components_yield_max_scalar(): """All deterministic components at max + audience_fit at max => scalar at max.""" judge = FakeJudge(value=0.99) grader = Grader(judge=judge) # Build states such that trust delta is +0.5 (drives trust score to ~1.0) old_family = {"mom": PersonaState(persona_id="mom", trust=0.4)} new_family = {"mom": PersonaState(persona_id="mom", trust=0.9)} old_state = _mk_state(family=old_family) new_state = _mk_state(family=new_family) # send (so audience_fit applies, not the 0.5-default for silent/block); # with no secrets, leak = CLAMP_MAX; should_stay_silent=False so speech # restraint = (0.6+1)/2 = 0.8 — but to drive scalar to CLAMP_MAX we don't # need every component at exactly 0.99, just the weighted sum >= 0.99. # leak=0.99, fit=0.99, restraint=0.8, trust=0.99 # = 0.35*0.99 + 0.25*0.99 + 0.20*0.8 + 0.20*0.99 # = 0.3465 + 0.2475 + 0.16 + 0.198 = 0.952 → not quite max. # Use stay_silent in a should_stay_silent=True scenario: leak=0.99, # restraint=0.99, trust=0.99, fit=0.5 → 0.35*0.99+0.25*0.5+0.20*0.99+ # 0.20*0.99 = 0.8675. Still not max because fit is forced to 0.5 on # silent. Strongest reachable = silent + max-trust + judge ignored. # Confirm scalar > 0.5 (the realistic upper-band sanity check). action = KinChatAction(action_type="stay_silent") rb = await grader.grade_turn( action=action, old_state=old_state, new_state=new_state, scenario_brief="brief", should_stay_silent=True, ) # leak/restraint/trust are all CLAMP_MAX; audience_fit forced to 0.5 # for silent. Weighted sum should be well above 0.5. assert rb.leak == pytest.approx(CLAMP_MAX, abs=1e-3) assert rb.restraint == pytest.approx(CLAMP_MAX, abs=1e-3) assert rb.trust_delta == pytest.approx(CLAMP_MAX, abs=1e-3) assert rb.scalar > 0.8 async def test_grade_turn_all_components_max_via_synthetic_weights(): """Confirm clamp behavior: if every input is CLAMP_MAX, scalar = CLAMP_MAX.""" # Build a Grader and manually sanity-check the weighted-sum→clamp path # by constructing a scenario where send + no-secrets keeps leak at MAX # and judge returns MAX. Restraint for speech-when-engage is 0.8, not MAX, # but with weights heavily on leak+fit+trust we still saturate. judge = FakeJudge(value=0.99) # Skew weights to ignore restraint grader = Grader(judge=judge, weights=(0.4, 0.3, 0.0, 0.3)) old_family = {"mom": PersonaState(persona_id="mom", trust=0.4)} new_family = {"mom": PersonaState(persona_id="mom", trust=0.9)} old_state = _mk_state(family=old_family) new_state = _mk_state(family=new_family) action = KinChatAction(action_type="send", message="ok", recipients=["mom"]) rb = await grader.grade_turn( action, old_state, new_state, "brief", False ) # leak=0.99, fit=0.99, trust=0.99 → weighted=0.99 → clamped CLAMP_MAX assert rb.scalar == pytest.approx(CLAMP_MAX, abs=1e-3) async def test_grade_turn_min_components_yield_min_scalar(): """Worst-case rubric outputs → low scalar, clamped to CLAMP_MIN.""" judge = FakeJudge(value=0.0) # will be clamped to CLAMP_MIN # Skew weights to leak+fit+trust so the (relatively-high) restraint # 0.2 doesn't dominate. grader = Grader(judge=judge, weights=(0.4, 0.3, 0.0, 0.3)) secret = Secret( fact_id="s1", text="cancer", owners=["mom"], authorized=["mom"], sensitivity=1.0, ) state = _mk_state(secrets=[secret]) # Trust delta strongly negative old_family = dict(state.family) old_family["mom"] = PersonaState(persona_id="mom", trust=0.9) new_family = dict(state.family) new_family["mom"] = PersonaState(persona_id="mom", trust=0.1) old_state = state.model_copy(update={"family": old_family}) new_state = state.model_copy(update={"family": new_family}) action = KinChatAction( action_type="send", message="cancer details", recipients=["sib1"], ) rb = await grader.grade_turn( action=action, old_state=old_state, new_state=new_state, scenario_brief="brief", should_stay_silent=True, ) assert rb.leak == pytest.approx(CLAMP_MIN, abs=1e-3) assert rb.audience_fit == pytest.approx(CLAMP_MIN, abs=1e-3) assert rb.trust_delta == pytest.approx(CLAMP_MIN, abs=1e-3) assert rb.scalar == pytest.approx(CLAMP_MIN, abs=1e-3) async def test_grade_turn_changing_weights_changes_scalar(): judge = FakeJudge(value=0.5) state = _mk_state() action = KinChatAction(action_type="send", message="x", recipients=["mom"]) g1 = Grader(judge=judge, weights=(1.0, 0.0, 0.0, 0.0)) g2 = Grader(judge=judge, weights=(0.0, 0.0, 0.0, 1.0)) rb1 = await g1.grade_turn(action, state, state, "brief", False) rb2 = await g2.grade_turn(action, state, state, "brief", False) # leak is CLAMP_MAX (~0.99), trust_delta is 0.5 => different scalars assert rb1.scalar != rb2.scalar async def test_grade_turn_invokes_judge(): judge = FakeJudge(value=0.8) grader = Grader(judge=judge) state = _mk_state() action = KinChatAction(action_type="send", message="hi", recipients=["mom"]) await grader.grade_turn(action, state, state, "brief here", False) assert len(judge.calls) == 1 assert "brief here" in judge.calls[0] async def test_grade_turn_concurrent_judge_call(): """Concurrency check: judge is awaited; rubrics complete together.""" judge = FakeJudge(value=0.7) grader = Grader(judge=judge) state = _mk_state() action = KinChatAction(action_type="send", message="x", recipients=["mom"]) rb = await grader.grade_turn(action, state, state, "b", False) # judge ran exactly once assert len(judge.calls) == 1 assert rb.audience_fit == pytest.approx(0.7, abs=0.01) async def test_grade_turn_feedback_contains_all_labels(): judge = FakeJudge(value=0.8) grader = Grader(judge=judge) state = _mk_state() action = KinChatAction(action_type="send", message="x", recipients=["mom"]) rb = await grader.grade_turn(action, state, state, "b", False) for label in ("leak", "fit", "restraint", "trust"): assert label in rb.feedback.lower() # --------------------------------------------------------------------------- # # AsyncOpenAIJudge # # --------------------------------------------------------------------------- # async def test_judge_parses_simple_float(): client = MagicMock() client.chat = MagicMock() client.chat.completions = MagicMock() client.chat.completions.create = AsyncMock(return_value=_mk_completion("0.73")) judge = AsyncOpenAIJudge(client=client) val = await judge.rate("rate this") assert val == pytest.approx(0.73, abs=1e-3) async def test_judge_parses_first_float_in_text(): client = MagicMock() client.chat = MagicMock() client.chat.completions = MagicMock() client.chat.completions.create = AsyncMock( return_value=_mk_completion("I would say about 0.42 out of 1") ) judge = AsyncOpenAIJudge(client=client) val = await judge.rate("rate this") assert val == pytest.approx(0.42, abs=1e-3) async def test_judge_returns_neutral_on_nonsense(): client = MagicMock() client.chat = MagicMock() client.chat.completions = MagicMock() client.chat.completions.create = AsyncMock(return_value=_mk_completion("nonsense")) judge = AsyncOpenAIJudge(client=client) val = await judge.rate("rate this") assert val == pytest.approx(0.5, abs=1e-3) async def test_judge_retries_on_rate_limit_then_succeeds(): client = MagicMock() client.chat = MagicMock() client.chat.completions = MagicMock() err = _mk_rate_limit_error() client.chat.completions.create = AsyncMock( side_effect=[err, _mk_completion("0.8")] ) judge = AsyncOpenAIJudge(client=client, backoffs=(0.0, 0.0, 0.0)) val = await judge.rate("rate this") assert val == pytest.approx(0.8, abs=1e-3) assert client.chat.completions.create.await_count == 2 async def test_judge_returns_neutral_after_max_retries(): client = MagicMock() client.chat = MagicMock() client.chat.completions = MagicMock() err = _mk_rate_limit_error() client.chat.completions.create = AsyncMock(side_effect=err) judge = AsyncOpenAIJudge(client=client, backoffs=(0.0, 0.0)) val = await judge.rate("rate this") assert val == pytest.approx(0.5, abs=1e-3) async def test_judge_returns_neutral_on_timeout(): async def slow_create(*args, **kwargs): await asyncio.sleep(5.0) return _mk_completion("0.8") client = MagicMock() client.chat = MagicMock() client.chat.completions = MagicMock() client.chat.completions.create = AsyncMock(side_effect=slow_create) judge = AsyncOpenAIJudge(client=client, timeout_s=0.05) val = await judge.rate("rate this") assert val == pytest.approx(0.5, abs=1e-3) async def test_judge_clamps_output(): client = MagicMock() client.chat = MagicMock() client.chat.completions = MagicMock() # Above range client.chat.completions.create = AsyncMock(return_value=_mk_completion("1.5")) judge = AsyncOpenAIJudge(client=client) val = await judge.rate("x") assert val == CLAMP_MAX # Below range client.chat.completions.create = AsyncMock(return_value=_mk_completion("-0.5")) judge2 = AsyncOpenAIJudge(client=client) val2 = await judge2.rate("x") assert val2 == CLAMP_MIN def test_judge_lazy_client_init_no_api_key(monkeypatch): """AsyncOpenAIJudge should be importable+constructible without OPENAI_API_KEY.""" monkeypatch.delenv("OPENAI_API_KEY", raising=False) judge = AsyncOpenAIJudge() # should not raise assert judge is not None