Skip to content

Commit 270f9a4

Browse files
committed
Harden eval optimize audit inputs
1 parent 227827c commit 270f9a4

5 files changed

Lines changed: 274 additions & 21 deletions

File tree

examples/optimization/eval_optimize_loop/eval_loop/backends.py

Lines changed: 23 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44

55
import asyncio
66
import importlib
7+
import math
78
from dataclasses import dataclass
89
from pathlib import Path
910
from typing import Any
@@ -170,6 +171,8 @@ def _load_call_agent(path: str):
170171
call_agent = getattr(module, function_name, None)
171172
if call_agent is None:
172173
raise ValueError(f"--sdk-call-agent target {path!r} was not found")
174+
if not callable(call_agent):
175+
raise ValueError(f"--sdk-call-agent target {path!r} was found but is not callable")
173176
return call_agent
174177

175178

@@ -184,13 +187,17 @@ def _has_running_loop() -> bool:
184187
def _summarize_sdk_result(result: Any) -> dict[str, Any]:
185188
return {
186189
"status": _safe_jsonable(getattr(result, "status", None)),
187-
"baseline_pass_rate": _safe_jsonable(getattr(result, "baseline_pass_rate", None)),
188-
"best_pass_rate": _safe_jsonable(getattr(result, "best_pass_rate", None)),
189-
"pass_rate_improvement": _safe_jsonable(getattr(result, "pass_rate_improvement", None)),
190+
"baseline_pass_rate": _safe_result_field(
191+
"baseline_pass_rate", getattr(result, "baseline_pass_rate", None)
192+
),
193+
"best_pass_rate": _safe_result_field("best_pass_rate", getattr(result, "best_pass_rate", None)),
194+
"pass_rate_improvement": _safe_result_field(
195+
"pass_rate_improvement", getattr(result, "pass_rate_improvement", None)
196+
),
190197
"baseline_metric_breakdown": _safe_jsonable(getattr(result, "baseline_metric_breakdown", {})),
191198
"best_metric_breakdown": _safe_jsonable(getattr(result, "best_metric_breakdown", {})),
192199
"metric_thresholds": _safe_jsonable(getattr(result, "metric_thresholds", {})),
193-
"total_llm_cost": _safe_jsonable(getattr(result, "total_llm_cost", 0.0)),
200+
"total_llm_cost": _safe_result_field("total_llm_cost", getattr(result, "total_llm_cost", 0.0)),
194201
"total_token_usage": _safe_jsonable(getattr(result, "total_token_usage", {})),
195202
"duration_seconds": _safe_jsonable(getattr(result, "duration_seconds", 0.0)),
196203
"started_at": _safe_jsonable(getattr(result, "started_at", None)),
@@ -211,6 +218,13 @@ def _summarize_sdk_result(result: Any) -> dict[str, Any]:
211218
}
212219

213220

221+
def _safe_result_field(field_name: str, value: Any) -> Any:
222+
try:
223+
return _safe_jsonable(value)
224+
except ValueError as exc:
225+
raise ValueError(f"SDK OptimizeResult field {field_name} must be a finite number") from exc
226+
227+
214228
def _safe_jsonable(value: Any) -> Any:
215229
if hasattr(value, "model_dump"):
216230
return _safe_jsonable(value.model_dump(mode="json"))
@@ -220,7 +234,11 @@ def _safe_jsonable(value: Any) -> Any:
220234
return {str(key): _safe_jsonable(item) for key, item in value.items()}
221235
if isinstance(value, (list, tuple)):
222236
return [_safe_jsonable(item) for item in value]
223-
if isinstance(value, (str, int, float, bool)) or value is None:
237+
if isinstance(value, float):
238+
if not math.isfinite(value):
239+
raise ValueError("value must be a finite number")
240+
return value
241+
if isinstance(value, (str, int, bool)) or value is None:
224242
return value
225243
return repr(value)
226244

examples/optimization/eval_optimize_loop/eval_loop/report.py

Lines changed: 21 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
from __future__ import annotations
44

55
import json
6+
import re
67
import shutil
78
from pathlib import Path
89
from typing import Any
@@ -25,6 +26,7 @@
2526
"--output-dir /tmp/eval-optimize-loop "
2627
"--fake-model --fake-judge --trace"
2728
)
29+
ARTIFACT_NAME_RE = re.compile(r"^[A-Za-z0-9_.-]+$")
2830

2931

3032
def compute_case_deltas(
@@ -106,7 +108,7 @@ def write_reports(report: OptimizationReport, output_dir: str | Path) -> tuple[P
106108

107109

108110
def report_to_json(report: OptimizationReport) -> str:
109-
return json.dumps(to_jsonable(report), indent=2, ensure_ascii=False, sort_keys=True) + "\n"
111+
return json.dumps(to_jsonable(report), indent=2, ensure_ascii=False, sort_keys=True, allow_nan=False) + "\n"
110112

111113

112114
def render_markdown(report: OptimizationReport) -> str:
@@ -248,7 +250,7 @@ def render_markdown(report: OptimizationReport) -> str:
248250

249251

250252
def write_audit_artifacts(report: OptimizationReport, output_path: Path) -> None:
251-
run_id = str(report.run.get("run_id") or "run")
253+
run_id = _safe_artifact_name(str(report.run.get("run_id") or "run"))
252254
run_dir = output_path / "runs" / run_id
253255
if run_dir.exists() and report.run.get("mode") == "fake":
254256
shutil.rmtree(run_dir)
@@ -275,20 +277,26 @@ def write_audit_artifacts(report: OptimizationReport, output_path: Path) -> None
275277

276278
for record in report.candidates:
277279
candidate: CandidatePrompt = record["candidate"]
278-
candidate_dir = prompt_dir / candidate.candidate_id
280+
candidate_name = _safe_artifact_name(candidate.candidate_id)
281+
candidate_dir = prompt_dir / candidate_name
279282
candidate_dir.mkdir(exist_ok=True)
280283
best_prompts = report.audit.get("sdk_result_summary", {}).get("best_prompts", {})
281284
if report.run.get("mode") == "sdk" and isinstance(best_prompts, dict) and best_prompts:
282285
for field_name, prompt_text in best_prompts.items():
283-
(candidate_dir / f"{field_name}.txt").write_text(str(prompt_text), encoding="utf-8")
286+
field_artifact = _safe_artifact_name(str(field_name))
287+
(candidate_dir / f"{field_artifact}.txt").write_text(str(prompt_text), encoding="utf-8")
284288
(candidate_dir / "bundle.txt").write_text(candidate.prompt, encoding="utf-8")
285289
else:
286290
(candidate_dir / "system_prompt.txt").write_text(candidate.prompt, encoding="utf-8")
287-
(diffs_dir / f"{candidate.candidate_id}.diff").write_text(candidate.prompt_diff, encoding="utf-8")
291+
(diffs_dir / f"{candidate_name}.diff").write_text(candidate.prompt_diff, encoding="utf-8")
288292
for split_name in ("train_result", "validation_result"):
289293
split_result = record[split_name]
290-
path = results_dir / f"{candidate.candidate_id}_{split_result.split}.json"
291-
path.write_text(json.dumps(to_jsonable(split_result), indent=2, sort_keys=True) + "\n", encoding="utf-8")
294+
split_artifact = _safe_artifact_name(str(split_result.split))
295+
path = results_dir / f"{candidate_name}_{split_artifact}.json"
296+
path.write_text(
297+
json.dumps(to_jsonable(split_result), indent=2, sort_keys=True, allow_nan=False) + "\n",
298+
encoding="utf-8",
299+
)
292300

293301

294302
def _delta_type(*, baseline_passed: bool, candidate_passed: bool, delta: float) -> str:
@@ -301,3 +309,9 @@ def _delta_type(*, baseline_passed: bool, candidate_passed: bool, delta: float)
301309
if delta < 0:
302310
return "score_down"
303311
return "unchanged"
312+
313+
314+
def _safe_artifact_name(name: str) -> str:
315+
if name in {"", ".", ".."} or not ARTIFACT_NAME_RE.fullmatch(name):
316+
raise ValueError(f"unsafe audit artifact name: {name!r}")
317+
return name

examples/optimization/eval_optimize_loop/run_pipeline.py

Lines changed: 51 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
import hashlib
77
import json
88
import math
9+
import re
910
import shlex
1011
import sys
1112
import tempfile
@@ -42,6 +43,8 @@
4243
DEFAULT_OPTIMIZER_CONFIG = HERE / "data" / "optimizer.json"
4344
DEFAULT_PROMPT = HERE / "prompts" / "baseline_system_prompt.txt"
4445
DEFAULT_OUTPUT_DIR = Path(tempfile.gettempdir()) / "eval-optimize-loop"
46+
TARGET_PROMPT_FIELD_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
47+
RUN_ID_RE = re.compile(r"^[A-Za-z0-9_.-]+$")
4548

4649

4750
def run_pipeline(
@@ -65,6 +68,8 @@ def run_pipeline(
6568

6669
if mode not in {"fake", "sdk"}:
6770
raise ValueError("field 'mode' must be one of: fake, sdk")
71+
if run_id is not None:
72+
run_id = validate_run_id(run_id)
6873
if mode == "fake" and (not fake_model or not fake_judge):
6974
raise ValueError(
7075
"fake mode requires fake_model=True and fake_judge=True. Pass --fake-model --fake-judge "
@@ -109,6 +114,8 @@ def run_pipeline(
109114
sdk_call_agent=sdk_call_agent,
110115
run_id=run_id,
111116
)
117+
if run_id is None:
118+
_resolve_default_sdk_run_id_collision(report, output_dir)
112119
write_reports(report, output_dir)
113120
return report
114121

@@ -351,14 +358,15 @@ def _build_sdk_report(
351358
prompt_path=prompt_path,
352359
)
353360
sdk_summary = sdk_backend.last_result_summary or {}
354-
baseline_pass_rate = _summary_float(sdk_summary, "baseline_pass_rate", 0.0)
355-
best_pass_rate = _summary_float(sdk_summary, "best_pass_rate", baseline_pass_rate)
361+
baseline_pass_rate = _summary_float(sdk_summary, "baseline_pass_rate", 0.0, required=True)
362+
best_pass_rate = _summary_float(sdk_summary, "best_pass_rate", baseline_pass_rate, required=True)
356363
pass_rate_improvement = _summary_float(
357364
sdk_summary,
358365
"pass_rate_improvement",
359366
best_pass_rate - baseline_pass_rate,
367+
required=True,
360368
)
361-
total_llm_cost = _summary_float(sdk_summary, "total_llm_cost", 0.0)
369+
total_llm_cost = _summary_float(sdk_summary, "total_llm_cost", 0.0, required=True)
362370
duration_seconds = _summary_float(sdk_summary, "duration_seconds", 0.0)
363371
effective_run_id = run_id or _default_sdk_run_id(sdk_summary)
364372
target_prompt_hashes = {
@@ -562,8 +570,8 @@ def _sdk_gate_decision(
562570
gate_config: dict[str, float],
563571
) -> GateDecision:
564572
status = str(sdk_summary.get("status") or "UNKNOWN")
565-
improvement = _summary_float(sdk_summary, "pass_rate_improvement", 0.0)
566-
total_cost = _summary_float(sdk_summary, "total_llm_cost", 0.0)
573+
improvement = _summary_float(sdk_summary, "pass_rate_improvement", 0.0, required=True)
574+
total_cost = _summary_float(sdk_summary, "total_llm_cost", 0.0, required=True)
567575
min_improvement = gate_config["min_val_score_improvement"]
568576
max_cost = gate_config["max_total_cost"]
569577

@@ -659,14 +667,25 @@ def _is_non_negative_finite_number(value: Any) -> bool:
659667
)
660668

661669

662-
def _summary_float(summary: dict[str, Any], key: str, default: float) -> float:
670+
def _summary_float(summary: dict[str, Any], key: str, default: float, *, required: bool = False) -> float:
663671
value = summary.get(key, default)
664672
if value is None:
665673
return default
674+
if isinstance(value, bool):
675+
if required:
676+
raise ValueError(f"SDK OptimizeResult field {key} must be a finite number")
677+
return default
666678
try:
667-
return float(value)
679+
parsed = float(value)
668680
except (TypeError, ValueError):
681+
if required:
682+
raise ValueError(f"SDK OptimizeResult field {key} must be a finite number")
683+
return default
684+
if not math.isfinite(parsed):
685+
if required:
686+
raise ValueError(f"SDK OptimizeResult field {key} must be a finite number")
669687
return default
688+
return parsed
670689

671690

672691
def _default_sdk_run_id(sdk_summary: dict[str, Any]) -> str:
@@ -704,16 +723,39 @@ def _parse_target_prompt_paths(
704723
if "=" not in item:
705724
raise ValueError("--target-prompt must use name=path format")
706725
name, path = item.split("=", 1)
707-
name = name.strip()
708726
path = path.strip()
709-
if not name or not path:
727+
if not TARGET_PROMPT_FIELD_RE.fullmatch(name):
728+
raise ValueError(
729+
f"--target-prompt field name {name!r} is invalid; use /^[A-Za-z_][A-Za-z0-9_]*$/"
730+
)
731+
if not path:
710732
raise ValueError("--target-prompt must use non-empty name=path values")
711733
if name in parsed:
712734
raise ValueError(f"--target-prompt duplicate field name {name!r}")
713735
parsed[name] = Path(path)
714736
return parsed
715737

716738

739+
def validate_run_id(run_id: str) -> str:
740+
if not isinstance(run_id, str):
741+
raise ValueError(f"--run-id value {run_id!r} must be a string")
742+
if run_id in {"", ".", ".."} or not RUN_ID_RE.fullmatch(run_id):
743+
raise ValueError(f"--run-id value {run_id!r} is invalid")
744+
return run_id
745+
746+
747+
def _resolve_default_sdk_run_id_collision(report: OptimizationReport, output_dir: str | Path) -> None:
748+
base_run_id = str(report.run.get("run_id") or "")
749+
validate_run_id(base_run_id)
750+
run_root = Path(output_dir) / "runs"
751+
candidate = base_run_id
752+
suffix = 1
753+
while (run_root / candidate).exists():
754+
candidate = f"{base_run_id}-{suffix}"
755+
suffix += 1
756+
report.run["run_id"] = candidate
757+
758+
717759
def _candidate_prompt_hashes_by_field(
718760
candidates: list[CandidatePrompt],
719761
sdk_summary: dict[str, Any],

examples/optimization/eval_optimize_loop/tests/test_pipeline_fake_mode.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
from examples.optimization.eval_optimize_loop.run_pipeline import DEFAULT_TRAIN
1111
from examples.optimization.eval_optimize_loop.run_pipeline import DEFAULT_VAL
1212
from examples.optimization.eval_optimize_loop.run_pipeline import run_pipeline
13+
from examples.optimization.eval_optimize_loop.eval_loop.report import report_to_json
1314

1415

1516
def test_fake_mode_pipeline_generates_json_and_markdown_reports(tmp_path: Path):
@@ -69,6 +70,7 @@ def test_fake_mode_pipeline_generates_json_and_markdown_reports(tmp_path: Path):
6970
"total_run_cost",
7071
}
7172
assert payload["run"]["reproducibility_command"].startswith("python examples/optimization")
73+
assert payload["run"]["run_id"] == "eval_optimize_loop_seed_91"
7274
assert payload["audit"]["total_run_cost"] == payload["audit"]["cost"]["total"]
7375
assert "candidate_001_overfit" in payload["audit"]["prompt_diffs"]
7476
assert "candidate_002_safe" in payload["audit"]["prompt_diffs"]
@@ -157,3 +159,24 @@ def test_cli_mode_fake_and_legacy_fake_flags_both_run(tmp_path: Path):
157159

158160
assert (first / "optimization_report.json").is_file()
159161
assert (second / "optimization_report.md").is_file()
162+
163+
164+
def test_report_to_json_rejects_nan_values(tmp_path: Path):
165+
report = run_pipeline(output_dir=tmp_path / "run", mode="fake", trace=True)
166+
report.audit["bad_float"] = float("nan")
167+
168+
try:
169+
report_to_json(report)
170+
except ValueError as exc:
171+
assert "Out of range float values are not JSON compliant" in str(exc)
172+
else:
173+
raise AssertionError("report_to_json should reject NaN")
174+
175+
176+
def test_fake_report_json_remains_strict_json(tmp_path: Path):
177+
report = run_pipeline(output_dir=tmp_path / "run", mode="fake", trace=True)
178+
payload = report_to_json(report)
179+
180+
assert "NaN" not in payload
181+
assert "Infinity" not in payload
182+
assert json.loads(payload)["selected_candidate"] == "candidate_002_safe"

0 commit comments

Comments
 (0)