Skip to content

Commit d6d225c

Browse files
committed
refactor(examples): define safe optimization contracts
1 parent 9642a85 commit d6d225c

6 files changed

Lines changed: 802 additions & 182 deletions

File tree

examples/optimization/eval_optimize_loop/eval_loop/config.py

Lines changed: 49 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22

33
from __future__ import annotations
44

5+
import math
56
from dataclasses import dataclass
67
from dataclasses import field
78
from pathlib import Path
@@ -16,7 +17,7 @@ class GateConfig:
1617
allow_new_hard_fail: bool = False
1718
protected_case_ids: list[str] = field(default_factory=list)
1819
max_score_drop_per_case: float = 0.0
19-
max_total_cost: float = 1.0
20+
max_total_cost: float | None = 1.0
2021
extras: dict[str, Any] = field(default_factory=dict)
2122

2223
def to_dict(self) -> dict[str, Any]:
@@ -125,9 +126,12 @@ def _parse_gate_config(payload: dict[str, Any], *, path: str) -> GateConfig:
125126
}
126127
extras = {key: value for key, value in payload.items() if key not in allowed}
127128

128-
min_val = payload.get("min_val_score_improvement", 0.01)
129-
if not isinstance(min_val, (int, float)) or min_val < 0 or min_val > 1:
130-
raise ValueError(f"{path}: field 'gate.min_val_score_improvement' must be a number between 0 and 1")
129+
min_val = _finite_number(
130+
payload.get("min_val_score_improvement", 0.01),
131+
f"{path}: field 'gate.min_val_score_improvement'",
132+
0.0,
133+
1.0,
134+
)
131135

132136
allow_new_hard_fail = payload.get("allow_new_hard_fail", False)
133137
if not isinstance(allow_new_hard_fail, bool):
@@ -137,24 +141,56 @@ def _parse_gate_config(payload: dict[str, Any], *, path: str) -> GateConfig:
137141
if not isinstance(protected_case_ids, list) or not all(isinstance(item, str) for item in protected_case_ids):
138142
raise ValueError(f"{path}: field 'gate.protected_case_ids' must be a list of strings")
139143

140-
max_drop = payload.get("max_score_drop_per_case", 0.0)
141-
if not isinstance(max_drop, (int, float)) or max_drop < 0:
142-
raise ValueError(f"{path}: field 'gate.max_score_drop_per_case' must be a non-negative number")
144+
max_drop = _finite_number(
145+
payload.get("max_score_drop_per_case", 0.0),
146+
f"{path}: field 'gate.max_score_drop_per_case'",
147+
0.0,
148+
)
143149

144-
max_total_cost = payload.get("max_total_cost", 1.0)
145-
if not isinstance(max_total_cost, (int, float)) or max_total_cost < 0:
146-
raise ValueError(f"{path}: field 'gate.max_total_cost' must be a non-negative number")
150+
max_total_cost_value = payload.get("max_total_cost", 1.0)
151+
max_total_cost = (
152+
None
153+
if max_total_cost_value is None
154+
else _finite_number(
155+
max_total_cost_value,
156+
f"{path}: field 'gate.max_total_cost'",
157+
0.0,
158+
)
159+
)
147160

148161
return GateConfig(
149-
min_val_score_improvement=float(min_val),
162+
min_val_score_improvement=min_val,
150163
allow_new_hard_fail=allow_new_hard_fail,
151164
protected_case_ids=list(protected_case_ids),
152-
max_score_drop_per_case=float(max_drop),
153-
max_total_cost=float(max_total_cost),
165+
max_score_drop_per_case=max_drop,
166+
max_total_cost=max_total_cost,
154167
extras=extras,
155168
)
156169

157170

171+
def _finite_number(
172+
value: Any,
173+
field_name: str,
174+
minimum: float,
175+
maximum: float | None = None,
176+
) -> float:
177+
if isinstance(value, bool) or not isinstance(value, (int, float)):
178+
raise ValueError(f"{field_name} must be a finite number")
179+
try:
180+
number = float(value)
181+
except OverflowError as exc:
182+
raise ValueError(f"{field_name} must be a finite number") from exc
183+
if not math.isfinite(number):
184+
raise ValueError(f"{field_name} must be a finite number")
185+
if number < minimum:
186+
raise ValueError(f"{field_name} must be a finite number greater than or equal to {minimum:g}")
187+
if maximum is not None and number > maximum:
188+
raise ValueError(
189+
f"{field_name} must be a finite number between {minimum:g} and {maximum:g}"
190+
)
191+
return number
192+
193+
158194
def _validate_cases(cases: list[EvalCase], *, split: str, path: str | Path) -> None:
159195
seen: set[str] = set()
160196
for case in cases:

examples/optimization/eval_optimize_loop/eval_loop/loader.py

Lines changed: 99 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -13,20 +13,33 @@
1313

1414
def read_json(path: str | Path) -> dict[str, Any]:
1515
resolved = Path(path)
16-
with resolved.open("r", encoding="utf-8") as file:
17-
payload = json.load(file)
16+
try:
17+
with resolved.open("r", encoding="utf-8") as file:
18+
payload = json.load(file, parse_constant=_reject_non_standard_json_constant)
19+
except (json.JSONDecodeError, ValueError) as exc:
20+
raise ValueError(f"{resolved}: invalid JSON: {exc}") from exc
1821
if not isinstance(payload, dict):
1922
raise ValueError(f"expected JSON object in {resolved}")
2023
return payload
2124

2225

26+
def _reject_non_standard_json_constant(constant: str) -> None:
27+
raise ValueError(f"non-standard JSON constant {constant!r}")
28+
29+
2330
def load_eval_cases(path: str | Path, split: str | None = None) -> list[EvalCase]:
2431
payload = read_json(path)
2532
cases = payload.get("cases")
26-
if not isinstance(cases, list):
27-
raise ValueError(f"evalset {path} must contain a cases list")
2833
effective_split = split or payload.get("split") or Path(path).name.split(".", 1)[0]
29-
return [EvalCase.from_dict(case, str(effective_split)) for case in cases]
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")
3043

3144

3245
def load_optimizer_config(path: str | Path) -> OptimizerConfig:
@@ -53,3 +66,84 @@ def sha256_file(path: str | Path) -> str:
5366
for chunk in iter(lambda: file.read(1024 * 1024), b""):
5467
digest.update(chunk)
5568
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: 74 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@
77
from dataclasses import field
88
from dataclasses import is_dataclass
99
from typing import Any
10+
from typing import Literal
1011

1112

1213
@dataclass(frozen=True)
@@ -27,12 +28,16 @@ def from_dict(cls, payload: dict[str, Any], split: str) -> "EvalCase":
2728
case_id = payload.get("case_id") or payload.get("id")
2829
if not case_id:
2930
raise ValueError(f"eval case is missing id/case_id: {payload!r}")
31+
if "split" in payload and str(payload["split"]) != str(split):
32+
raise ValueError(
33+
f"eval case {case_id!r} split mismatch: payload has {payload['split']!r}, expected {split!r}"
34+
)
3035
expectation = payload.get("expectation")
3136
if not isinstance(expectation, dict):
3237
raise ValueError(f"eval case {case_id!r} is missing expectation object")
3338
return cls(
3439
case_id=str(case_id),
35-
split=str(payload.get("split") or split),
40+
split=str(split),
3641
input=str(payload.get("input") or payload.get("user_input") or ""),
3742
expectation=dict(expectation),
3843
tags=list(payload.get("tags") or []),
@@ -52,7 +57,9 @@ class CaseResult:
5257
score: float
5358
passed: bool
5459
output: str
60+
metrics: dict[str, float] = field(default_factory=dict)
5561
trace: dict[str, Any] = field(default_factory=dict)
62+
trace_available: bool = False
5663
failure_category: str | None = None
5764
failure_reason: str | None = None
5865
evidence: str | None = None
@@ -84,6 +91,67 @@ class CandidatePrompt:
8491
prompt: str
8592
rationale: str
8693
prompt_diff: str
94+
prompt_fields: dict[str, str] = field(default_factory=dict)
95+
96+
def bundle(self) -> dict[str, str]:
97+
"""Return this candidate's complete prompt bundle."""
98+
99+
if self.prompt_fields:
100+
return dict(self.prompt_fields)
101+
return {"system_prompt": self.prompt}
102+
103+
104+
@dataclass(frozen=True)
105+
class CostSummary:
106+
"""Cost attribution for an optimization run."""
107+
108+
optimizer: float = 0.0
109+
evaluator: float = 0.0
110+
agent: float = 0.0
111+
total: float = 0.0
112+
complete: bool = True
113+
114+
115+
@dataclass(frozen=True)
116+
class OptimizationRound:
117+
"""One auditable optimizer round."""
118+
119+
round_id: int
120+
candidate_id: str
121+
prompts: dict[str, str]
122+
rationale: str
123+
metrics: dict[str, float]
124+
cost: CostSummary
125+
duration_seconds: float
126+
127+
128+
WritebackStatus = Literal[
129+
"rejected",
130+
"not_requested",
131+
"applied",
132+
"rolled_back",
133+
"rollback_failed",
134+
]
135+
136+
137+
@dataclass(frozen=True)
138+
class WritebackResult:
139+
"""Outcome of an optional source prompt writeback."""
140+
141+
status: WritebackStatus
142+
before_hashes: dict[str, str] = field(default_factory=dict)
143+
after_hashes: dict[str, str] = field(default_factory=dict)
144+
error: str | None = None
145+
146+
147+
@dataclass(frozen=True)
148+
class OptimizationResult:
149+
"""Backend-neutral optimization output."""
150+
151+
candidates: list[CandidatePrompt]
152+
rounds: list[OptimizationRound]
153+
cost: CostSummary
154+
raw_summary: dict[str, Any] = field(default_factory=dict)
87155

88156

89157
@dataclass(frozen=True)
@@ -141,6 +209,11 @@ class OptimizationReport:
141209
gate_decisions: list[GateDecision]
142210
selected_candidate: str | None
143211
audit: dict[str, Any]
212+
rounds: list[OptimizationRound] = field(default_factory=list)
213+
cost_summary: CostSummary = field(default_factory=CostSummary)
214+
writeback: WritebackResult = field(
215+
default_factory=lambda: WritebackResult(status="not_requested")
216+
)
144217

145218

146219
def to_jsonable(value: Any) -> Any:

0 commit comments

Comments
 (0)