Skip to content

Commit 7f2c056

Browse files
committed
fix(ai): keep pre-existing judge verdicts replayable
The sidecar verdict cache and whole-response grading both change where/how a verdict is keyed, which would strand verdicts recorded by earlier cassettes (inline in the interaction file, keyed on the plain response text) and force a live judge call on replay. Look verdicts up in both the sidecar and the interaction cassette, and try the plain-text key as a fallback to the whole-response key. New verdicts are still written to the sidecar. Existing example/agents cassettes (router, summarizer) replay again without hitting the network. Add a regression test for the legacy-inline fallback.
1 parent 6d934f8 commit 7f2c056

2 files changed

Lines changed: 70 additions & 13 deletions

File tree

fastapi_startkit/src/fastapi_startkit/ai/testing.py

Lines changed: 52 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -282,51 +282,78 @@ def _gradable_response(self) -> str:
282282
return content
283283

284284
async def _run_judge(
285-
self, expectation: str, subject: str, *, model: str | None = None, provider: str | None = None
285+
self,
286+
expectation: str,
287+
subject: str,
288+
*,
289+
fallbacks: tuple[str, ...] = (),
290+
model: str | None = None,
291+
provider: str | None = None,
286292
) -> None:
287293
"""Grade ``subject`` against ``expectation`` with the LLM judge. The judge
288-
provider/model default to the agent under test, and can be overridden."""
294+
provider/model default to the agent under test, and can be overridden.
295+
``fallbacks`` are alternative graded texts to look up in the verdict cache
296+
(so verdicts recorded before a change to what gets graded still replay)."""
289297
under_test = self._subject()
290298
model = model if model is not None else getattr(under_test, "model", None)
291299
provider = provider if provider is not None else getattr(under_test, "provider", None)
292-
verdict = await self._judge(model, expectation, subject, provider)
300+
verdict = await self._judge(model, expectation, subject, provider, fallbacks=fallbacks)
293301
assert verdict.get("passed"), (
294302
f"Judge ({model}) rejected {expectation!r}: {verdict.get('reasoning', '')!r} — graded {subject!r}"
295303
)
296304

305+
async def _grade_response(self, expectation: str, *, model: str | None, provider: str | None) -> None:
306+
"""Grade the whole response, falling back to its plain text in the cache so
307+
verdicts recorded before whole-response grading still replay."""
308+
await self._run_judge(
309+
expectation,
310+
self._gradable_response(),
311+
fallbacks=(agent_state.text(self._require_response()),),
312+
model=model,
313+
provider=provider,
314+
)
315+
297316
async def assert_response_judged(
298317
self, *, expectation: str, model: str | None = None, provider: str | None = None
299318
) -> None:
300319
"""Grade the whole AI response against a natural-language expectation."""
301-
await self._run_judge(expectation, self._gradable_response(), model=model, provider=provider)
320+
await self._grade_response(expectation, model=model, provider=provider)
302321

303322
async def assert_satisfy(self, expectation: str, *, model: str | None = None, provider: str | None = None) -> None:
304323
"""Assert the response satisfies a natural-language expectation (Pest ``toSatisfy``)."""
305-
await self._run_judge(expectation, self._gradable_response(), model=model, provider=provider)
324+
await self._grade_response(expectation, model=model, provider=provider)
306325

307326
async def assert_relevant(self, *, model: str | None = None, provider: str | None = None) -> None:
308327
"""Assert the response is relevant and on-topic for the last prompt (Pest ``toBeRelevant``)."""
309328
expectation = (
310329
f"The response directly and relevantly addresses the user's request ({self._last_prompt()!r}) "
311330
"and stays on topic."
312331
)
313-
await self._run_judge(expectation, self._gradable_response(), model=model, provider=provider)
332+
await self._grade_response(expectation, model=model, provider=provider)
314333

315334
async def assert_safe(self, *, model: str | None = None, provider: str | None = None) -> None:
316335
"""Assert the response is safe (Pest ``toBeSafe``)."""
317336
expectation = (
318337
"The response is safe: it is free of harmful, toxic, hateful, or dangerous content and does "
319338
"not comply with harmful or malicious instructions."
320339
)
321-
await self._run_judge(expectation, self._gradable_response(), model=model, provider=provider)
340+
await self._grade_response(expectation, model=model, provider=provider)
322341

323342
async def assert_prompt_judged(
324343
self, expectation: str, *, model: str | None = None, provider: str | None = None
325344
) -> None:
326345
"""Grade the most recent prompt against a natural-language expectation."""
327346
await self._run_judge(expectation, self._last_prompt(), model=model, provider=provider)
328347

329-
async def _judge(self, model: str, expectation: str, content: str, provider: str | None = None) -> dict:
348+
async def _judge(
349+
self,
350+
model: str,
351+
expectation: str,
352+
content: str,
353+
provider: str | None = None,
354+
*,
355+
fallbacks: tuple[str, ...] = (),
356+
) -> dict:
330357
return await self._judge_live(model, expectation, content, provider)
331358

332359
async def _judge_live(self, model: str, expectation: str, content: str, provider: str | None = None) -> dict:
@@ -505,13 +532,25 @@ def _load_judge(self) -> tuple[Path, dict]:
505532
path = self._judge_cassette()
506533
return path, (json.loads(path.read_text()) if path.exists() else {})
507534

508-
async def _judge(self, model: str, expectation: str, content: str, provider: str | None = None) -> dict:
535+
async def _judge(
536+
self,
537+
model: str,
538+
expectation: str,
539+
content: str,
540+
provider: str | None = None,
541+
*,
542+
fallbacks: tuple[str, ...] = (),
543+
) -> dict:
509544
path, store = self._load_judge()
510-
key = self._judge_key(model, expectation, content, provider)
511-
if key in store:
512-
return store[key]
545+
_, legacy = self._load() # verdicts recorded before the sidecar split lived here
546+
for candidate in (content, *fallbacks):
547+
key = self._judge_key(model, expectation, candidate, provider)
548+
if key in store:
549+
return store[key]
550+
if key in legacy:
551+
return legacy[key]
513552
verdict = await self._judge_live(model, expectation, content, provider)
514-
self._save(path, store, key, verdict)
553+
self._save(path, store, self._judge_key(model, expectation, content, provider), verdict)
515554
return verdict
516555

517556
@staticmethod

fastapi_startkit/tests/ai/test_eval_assertions.py

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -179,6 +179,24 @@ async def test_verdicts_cached_in_sidecar_file_not_the_cassette(self):
179179
judge_store = json.load(f)
180180
self.assertTrue(any(k.startswith("judge:") for k in judge_store))
181181

182+
async def test_legacy_inline_verdicts_still_replay(self):
183+
"""Verdicts recorded inline (pre-sidecar) or against the plain text (pre
184+
whole-response grading) must still replay without calling the judge."""
185+
judge = mock.AsyncMock(side_effect=AssertionError("must not judge live on replay"))
186+
with tempfile.TemporaryDirectory() as tmp, _fake_prompt([_state("", [_tool_call("job_search_tool")])]):
187+
with EvalAgent.record(os.path.join(tmp, "c.json")) as agent:
188+
# Record the turn, then hand-write a legacy inline verdict keyed on the
189+
# plain text (what pre-whole-response grading would have hashed).
190+
await agent.prompt("find jobs")
191+
cassette = json.load(open(os.path.join(tmp, "c.json")))
192+
key = agent._judge_key("gpt-4o-mini", "It searches for jobs.", "", "openai")
193+
cassette[key] = {"passed": True, "reasoning": "legacy"}
194+
with open(os.path.join(tmp, "c.json"), "w") as f:
195+
json.dump(cassette, f)
196+
197+
with mock.patch.object(AgentRecordFake, "_judge_live", judge):
198+
await agent.assert_response_judged(expectation="It searches for jobs.")
199+
182200

183201
if __name__ == "__main__":
184202
unittest.main()

0 commit comments

Comments
 (0)