kinchat / tests /test_grader.py
Bhargav
Initial KinChat env: models, personas, scenarios, rubrics, grader, env loop, FastAPI app, client, dashboard, baseline inference (377 tests passing)
2e8387b
Raw
History Blame Contribute Delete
15 kB
"""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