Bhargav
Initial KinChat env: models, personas, scenarios, rubrics, grader, env loop, FastAPI app, client, dashboard, baseline inference (377 tests passing)
2e8387b | """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 | |