Skip to content

Commit c6b5172

Browse files
committed
fix(examples): parse sdk evalset expectations
1 parent 717a413 commit c6b5172

2 files changed

Lines changed: 280 additions & 11 deletions

File tree

examples/optimization/eval_optimize_loop/eval_loop/backends.py

Lines changed: 135 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,7 @@
1616
from .fake_judge import FakeJudge
1717
from .fake_model import FakeModel
1818
from .loader import load_eval_cases
19+
from .loader import read_json
1920
from .optimizer import FakeOptimizer
2021
from .schemas import CandidatePrompt
2122
from .schemas import CaseResult
@@ -421,7 +422,7 @@ async def evaluate(
421422
target_paths,
422423
context=f"cannot evaluate {prompt_id}",
423424
)
424-
expected_cases = load_eval_cases(dataset_path, split=split)
425+
expected_cases = _load_sdk_expected_cases(dataset_path, split=split)
425426
snapshot = snapshot_prompt_files(target_paths)
426427
result: Any | None = None
427428
with temporary_prompt_bundle(snapshot, candidate_prompts):
@@ -447,7 +448,7 @@ async def evaluate(
447448
result,
448449
prompt_id=prompt_id,
449450
split=split,
450-
expected_cases=expected_cases,
451+
expected_cases=expected_cases.values(),
451452
)
452453

453454
def _load_required_call_agent(self, *, for_evaluation: bool):
@@ -721,6 +722,138 @@ def _safe_jsonable(value: Any) -> Any:
721722
return repr(value)
722723

723724

725+
def _load_sdk_expected_cases(
726+
dataset_path: str | Path,
727+
*,
728+
split: str,
729+
) -> dict[str, EvalCase]:
730+
"""Load wrapper metadata from an SDK EvalSet without changing the SDK file."""
731+
732+
payload = read_json(dataset_path)
733+
if "eval_cases" not in payload:
734+
if "cases" in payload:
735+
return _eval_cases_by_id(
736+
load_eval_cases(dataset_path, split=split),
737+
context=f"legacy evalset {dataset_path}",
738+
)
739+
raise ValueError(f"SDK evalset {dataset_path} must contain an eval_cases list")
740+
741+
eval_set_id = payload.get("eval_set_id")
742+
if not isinstance(eval_set_id, str) or not eval_set_id.strip():
743+
raise ValueError(f"SDK evalset {dataset_path} is missing non-empty eval_set_id")
744+
raw_cases = payload["eval_cases"]
745+
if not isinstance(raw_cases, list):
746+
raise ValueError(f"SDK evalset {dataset_path} eval_cases must be a list")
747+
748+
expected_cases: dict[str, EvalCase] = {}
749+
for index, raw_case in enumerate(raw_cases):
750+
if not isinstance(raw_case, dict):
751+
raise ValueError(
752+
f"SDK evalset {dataset_path} eval_cases[{index}] must be an object"
753+
)
754+
eval_id = raw_case.get("eval_id")
755+
if not isinstance(eval_id, str) or not eval_id.strip():
756+
raise ValueError(
757+
f"SDK evalset {dataset_path} eval_cases[{index}] is missing non-empty eval_id"
758+
)
759+
if eval_id in expected_cases:
760+
raise ValueError(f"SDK evalset {dataset_path} contains duplicate eval_id {eval_id!r}")
761+
expected_cases[eval_id] = _expected_case_from_sdk_case(
762+
raw_case,
763+
eval_id=eval_id,
764+
split=split,
765+
dataset_path=dataset_path,
766+
)
767+
return expected_cases
768+
769+
770+
def _expected_case_from_sdk_case(
771+
raw_case: dict[str, Any],
772+
*,
773+
eval_id: str,
774+
split: str,
775+
dataset_path: str | Path,
776+
) -> EvalCase:
777+
context = f"SDK evalset {dataset_path} case {eval_id!r}"
778+
conversation = raw_case.get("conversation")
779+
if not isinstance(conversation, list) or not conversation:
780+
raise ValueError(f"{context} must contain a non-empty conversation list")
781+
782+
input_text = ""
783+
for turn_index, invocation in reversed(list(enumerate(conversation))):
784+
if not isinstance(invocation, dict):
785+
raise ValueError(f"{context} conversation[{turn_index}] must be an object")
786+
user_content = invocation.get("user_content")
787+
if user_content is None:
788+
continue
789+
if not isinstance(user_content, dict):
790+
raise ValueError(
791+
f"{context} conversation[{turn_index}].user_content must be an object"
792+
)
793+
candidate_text = _content_text(user_content)
794+
if candidate_text.strip():
795+
input_text = candidate_text
796+
break
797+
if not input_text:
798+
raise ValueError(f"{context} conversation has no user_content text")
799+
800+
session_input = raw_case.get("session_input")
801+
if not isinstance(session_input, dict):
802+
raise ValueError(f"{context} must contain a session_input object")
803+
state = session_input.get("state")
804+
if not isinstance(state, dict):
805+
raise ValueError(f"{context} session_input.state must be an object")
806+
expectation = state.get("eval_optimize_expectation")
807+
if not isinstance(expectation, dict):
808+
raise ValueError(
809+
f"{context} session_input.state must contain eval_optimize_expectation object"
810+
)
811+
812+
tags = state.get("eval_optimize_tags", [])
813+
if not isinstance(tags, list) or any(not isinstance(tag, str) for tag in tags):
814+
raise ValueError(f"{context} eval_optimize_tags must be a list of strings")
815+
protected = state.get("eval_optimize_protected", False)
816+
if not isinstance(protected, bool):
817+
raise ValueError(f"{context} eval_optimize_protected must be a boolean")
818+
819+
expected_failure_category = state.get("eval_optimize_expected_failure_category")
820+
if expected_failure_category is None:
821+
expected_failure_category = state.get("expected_failure_category")
822+
if expected_failure_category is None:
823+
expected_failure_category = expectation.get("expected_failure_category")
824+
if (
825+
expected_failure_category is not None
826+
and (
827+
not isinstance(expected_failure_category, str)
828+
or not expected_failure_category.strip()
829+
)
830+
):
831+
raise ValueError(f"{context} expected_failure_category must be a non-empty string")
832+
833+
return EvalCase(
834+
case_id=eval_id,
835+
split=split,
836+
input=input_text,
837+
expectation=dict(expectation),
838+
tags=list(tags),
839+
protected=protected,
840+
expected_failure_category=expected_failure_category,
841+
)
842+
843+
844+
def _eval_cases_by_id(
845+
cases: Iterable[EvalCase],
846+
*,
847+
context: str,
848+
) -> dict[str, EvalCase]:
849+
by_id: dict[str, EvalCase] = {}
850+
for case in cases:
851+
if case.case_id in by_id:
852+
raise ValueError(f"{context} contains duplicate case id {case.case_id!r}")
853+
by_id[case.case_id] = case
854+
return by_id
855+
856+
724857
def _eval_result_from_sdk_result(
725858
result: Any,
726859
*,

examples/optimization/eval_optimize_loop/tests/test_sdk_backend.py

Lines changed: 145 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -328,24 +328,114 @@ def test_sdk_result_mapping_rejects_non_finite_metric_scores():
328328
)
329329

330330

331+
def test_sdk_expected_cases_parse_standard_evalset_metadata(tmp_path: Path):
332+
dataset_path = tmp_path / "validation.evalset.json"
333+
case_payload = _sdk_eval_case_payload(
334+
"case-1",
335+
query="query",
336+
expected="expected",
337+
tags=["x"],
338+
protected=True,
339+
)
340+
case_payload["conversation"].insert(
341+
0,
342+
{
343+
"invocation_id": "case-1-turn-0",
344+
"user_content": {
345+
"role": "user",
346+
"parts": [{"text": "earlier query"}],
347+
},
348+
"final_response": {
349+
"role": "model",
350+
"parts": [{"text": "earlier expected"}],
351+
},
352+
},
353+
)
354+
dataset_path.write_text(
355+
json.dumps(
356+
_sdk_evalset_payload([case_payload])
357+
),
358+
encoding="utf-8",
359+
)
360+
361+
expected_cases = backend_module._load_sdk_expected_cases(
362+
dataset_path,
363+
split="validation",
364+
)
365+
366+
assert set(expected_cases) == {"case-1"}
367+
expected_case = expected_cases["case-1"]
368+
assert expected_case.input == "query"
369+
assert expected_case.expectation == {
370+
"type": "exact",
371+
"expected": "expected",
372+
"expected_failure_category": "final_response_mismatch",
373+
}
374+
assert expected_case.tags == ["x"]
375+
assert expected_case.protected is True
376+
assert expected_case.expected_failure_category == "final_response_mismatch"
377+
assert expected_case.split == "validation"
378+
379+
380+
def test_sdk_expected_cases_reject_duplicate_eval_ids(tmp_path: Path):
381+
dataset_path = tmp_path / "duplicate.evalset.json"
382+
dataset_path.write_text(
383+
json.dumps(
384+
_sdk_evalset_payload([
385+
_sdk_eval_case_payload("case-1"),
386+
_sdk_eval_case_payload("case-1"),
387+
])
388+
),
389+
encoding="utf-8",
390+
)
391+
392+
with pytest.raises(ValueError, match="duplicate.*case-1"):
393+
backend_module._load_sdk_expected_cases(dataset_path, split="validation")
394+
395+
396+
def test_sdk_expected_cases_require_expectation_metadata(tmp_path: Path):
397+
payload = _sdk_evalset_payload([_sdk_eval_case_payload("case-1")])
398+
del payload["eval_cases"][0]["session_input"]["state"]["eval_optimize_expectation"]
399+
dataset_path = tmp_path / "missing_expectation.evalset.json"
400+
dataset_path.write_text(json.dumps(payload), encoding="utf-8")
401+
402+
with pytest.raises(ValueError, match="case-1.*eval_optimize_expectation"):
403+
backend_module._load_sdk_expected_cases(dataset_path, split="validation")
404+
405+
406+
def test_sdk_expected_cases_reject_invalid_eval_cases_shape(tmp_path: Path):
407+
dataset_path = tmp_path / "invalid_shape.evalset.json"
408+
dataset_path.write_text(
409+
json.dumps({"eval_set_id": "set", "eval_cases": {}}),
410+
encoding="utf-8",
411+
)
412+
413+
with pytest.raises(ValueError, match="eval_cases.*list"):
414+
backend_module._load_sdk_expected_cases(dataset_path, split="validation")
415+
416+
331417
@pytest.mark.asyncio
332418
async def test_sdk_backend_evaluate_temporarily_installs_and_restores_prompt_bytes(
333419
tmp_path: Path,
334420
monkeypatch,
335421
):
336422
dataset_path = tmp_path / "validation.evalset.json"
337423
dataset_path.write_text(
338-
json.dumps({
339-
"split": "validation",
340-
"cases": [{
341-
"case_id": "case_a",
342-
"input": "question",
343-
"expectation": {"answer": "answer"},
344-
"expected_failure_category": "format_violation",
345-
}],
346-
}),
424+
json.dumps(
425+
_sdk_evalset_payload([
426+
_sdk_eval_case_payload(
427+
"case_a",
428+
query="question",
429+
expected="answer",
430+
expected_failure_category="format_violation",
431+
)
432+
])
433+
),
347434
encoding="utf-8",
348435
)
436+
from trpc_agent_sdk.evaluation import EvalSet
437+
438+
assert EvalSet.model_validate_json(dataset_path.read_text(encoding="utf-8")).eval_cases[0].eval_id == "case_a"
349439
prompt_path = tmp_path / "prompt.txt"
350440
original_bytes = b"original prompt\r\n"
351441
prompt_path.write_bytes(original_bytes)
@@ -385,6 +475,7 @@ async def test_sdk_backend_evaluate_temporarily_installs_and_restores_prompt_byt
385475
assert calls["eval_result_output_dir"] == str(tmp_path / "sdk_eval")
386476
assert prompt_path.read_bytes() == original_bytes
387477
assert mapped.cases[0].trace_available is True
478+
assert mapped.cases[0].expected_failure_category == "format_violation"
388479

389480

390481
def test_sdk_backend_requires_call_agent_path(tmp_path: Path):
@@ -1397,6 +1488,51 @@ def _sdk_evaluate_result(runs_by_case_id: dict[str, list[object]]):
13971488
)
13981489

13991490

1491+
def _sdk_evalset_payload(eval_cases: list[dict[str, object]]) -> dict[str, object]:
1492+
return {
1493+
"eval_set_id": "set",
1494+
"eval_cases": eval_cases,
1495+
}
1496+
1497+
1498+
def _sdk_eval_case_payload(
1499+
eval_id: str,
1500+
*,
1501+
query: str = "query",
1502+
expected: str = "expected",
1503+
expected_failure_category: str = "final_response_mismatch",
1504+
tags: list[str] | None = None,
1505+
protected: bool = False,
1506+
) -> dict[str, object]:
1507+
return {
1508+
"eval_id": eval_id,
1509+
"conversation": [{
1510+
"invocation_id": f"{eval_id}-turn-1",
1511+
"user_content": {
1512+
"role": "user",
1513+
"parts": [{"text": query}],
1514+
},
1515+
"final_response": {
1516+
"role": "model",
1517+
"parts": [{"text": expected}],
1518+
},
1519+
}],
1520+
"session_input": {
1521+
"app_name": "eval_optimize_loop",
1522+
"user_id": "test-user",
1523+
"state": {
1524+
"eval_optimize_expectation": {
1525+
"type": "exact",
1526+
"expected": expected,
1527+
"expected_failure_category": expected_failure_category,
1528+
},
1529+
"eval_optimize_tags": list(tags or []),
1530+
"eval_optimize_protected": protected,
1531+
},
1532+
},
1533+
}
1534+
1535+
14001536
def _install_fake_agent_evaluator(
14011537
monkeypatch,
14021538
*,

0 commit comments

Comments
 (0)