Skip to content

Commit dae8692

Browse files
committed
fix(examples): keep contract change self-contained
1 parent d6d225c commit dae8692

5 files changed

Lines changed: 181 additions & 453 deletions

File tree

examples/optimization/eval_optimize_loop/eval_loop/loader.py

Lines changed: 3 additions & 90 deletions
Original file line numberDiff line numberDiff line change
@@ -30,16 +30,10 @@ def _reject_non_standard_json_constant(constant: str) -> None:
3030
def load_eval_cases(path: str | Path, split: str | None = None) -> list[EvalCase]:
3131
payload = read_json(path)
3232
cases = payload.get("cases")
33+
if not isinstance(cases, list):
34+
raise ValueError(f"evalset {path} must contain a cases list")
3335
effective_split = split or payload.get("split") or Path(path).name.split(".", 1)[0]
34-
if isinstance(cases, list):
35-
return [EvalCase.from_dict(case, str(effective_split)) for case in cases]
36-
37-
sdk_cases = payload.get("eval_cases") or payload.get("evalCases")
38-
if isinstance(sdk_cases, list):
39-
_validate_sdk_evalset(payload, path)
40-
return [_eval_case_from_sdk_dict(case, str(effective_split)) for case in sdk_cases]
41-
42-
raise ValueError(f"evalset {path} must contain a cases or evalCases list")
36+
return [EvalCase.from_dict(case, str(effective_split)) for case in cases]
4337

4438

4539
def load_optimizer_config(path: str | Path) -> OptimizerConfig:
@@ -66,84 +60,3 @@ def sha256_file(path: str | Path) -> str:
6660
for chunk in iter(lambda: file.read(1024 * 1024), b""):
6761
digest.update(chunk)
6862
return digest.hexdigest()
69-
70-
71-
def _validate_sdk_evalset(payload: dict[str, Any], path: str | Path) -> None:
72-
try:
73-
from trpc_agent_sdk.evaluation._eval_set import EvalSet
74-
75-
EvalSet.model_validate(payload)
76-
except Exception as exc:
77-
raise ValueError(f"evalset {path} is not a valid SDK EvalSet: {exc}") from exc
78-
79-
80-
def _eval_case_from_sdk_dict(payload: dict[str, Any], split: str) -> EvalCase:
81-
case_id = payload.get("eval_id") or payload.get("evalId")
82-
if not case_id:
83-
raise ValueError(f"SDK eval case is missing evalId/eval_id: {payload!r}")
84-
85-
session_input = payload.get("session_input") or payload.get("sessionInput") or {}
86-
state = session_input.get("state") if isinstance(session_input, dict) else {}
87-
state = state if isinstance(state, dict) else {}
88-
expectation = state.get("eval_optimize_expectation")
89-
if not isinstance(expectation, dict):
90-
expectation = _infer_expectation_from_sdk_case(payload)
91-
92-
return EvalCase(
93-
case_id=str(case_id),
94-
split=split,
95-
input=_first_user_text(payload),
96-
expectation=dict(expectation),
97-
tags=[str(item) for item in state.get("eval_optimize_tags", [])],
98-
protected=bool(state.get("eval_optimize_protected", False)),
99-
simulated_outputs=dict(state.get("eval_optimize_simulated_outputs") or expectation.get("simulated_outputs") or {}),
100-
expected_failure_category=state.get("eval_optimize_expected_failure_category")
101-
or expectation.get("expected_failure_category"),
102-
)
103-
104-
105-
def _infer_expectation_from_sdk_case(payload: dict[str, Any]) -> dict[str, Any]:
106-
expected = _first_final_response_text(payload)
107-
if expected:
108-
return {
109-
"type": "exact",
110-
"expected": expected,
111-
"expected_failure_category": "final_response_mismatch",
112-
}
113-
raise ValueError(
114-
"SDK eval case must put fake-mode metadata in sessionInput.state.eval_optimize_expectation "
115-
f"or provide a finalResponse that can be treated as an exact expectation: {payload!r}"
116-
)
117-
118-
119-
def _first_user_text(payload: dict[str, Any]) -> str:
120-
for invocation in _conversation(payload):
121-
content = invocation.get("user_content") or invocation.get("userContent") or {}
122-
text = _content_text(content)
123-
if text:
124-
return text
125-
return ""
126-
127-
128-
def _first_final_response_text(payload: dict[str, Any]) -> str:
129-
for invocation in _conversation(payload):
130-
content = invocation.get("final_response") or invocation.get("finalResponse") or {}
131-
text = _content_text(content)
132-
if text:
133-
return text
134-
return ""
135-
136-
137-
def _conversation(payload: dict[str, Any]) -> list[dict[str, Any]]:
138-
conversation = payload.get("conversation") or []
139-
return [item for item in conversation if isinstance(item, dict)]
140-
141-
142-
def _content_text(content: Any) -> str:
143-
if not isinstance(content, dict):
144-
return ""
145-
texts = []
146-
for part in content.get("parts") or []:
147-
if isinstance(part, dict) and part.get("text") is not None:
148-
texts.append(str(part["text"]))
149-
return "\n".join(texts)

examples/optimization/eval_optimize_loop/eval_loop/schemas.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -57,9 +57,9 @@ class CaseResult:
5757
score: float
5858
passed: bool
5959
output: str
60-
metrics: dict[str, float] = field(default_factory=dict)
60+
metrics: dict[str, float] = field(default_factory=dict, kw_only=True)
6161
trace: dict[str, Any] = field(default_factory=dict)
62-
trace_available: bool = False
62+
trace_available: bool = field(default=False, kw_only=True)
6363
failure_category: str | None = None
6464
failure_reason: str | None = None
6565
evidence: str | None = None

0 commit comments

Comments
 (0)