Skip to content

Commit 6a6a368

Browse files
committed
examples: harden eval optimization regression loop
Fail closed on ambiguous replay data, unsafe artifact labels, non-finite metrics, key-case pass regressions, and common tool error payloads. Isolate optimizer prompt writes from source files, reject unsafe in-process concurrency, and add adversarial regression coverage for gate, attribution, path, apply, and concurrency behavior. Fixes #91 RELEASE NOTES: Harden the evaluation and optimization example against ambiguous data, unsafe artifacts, and regression-gate edge cases.
1 parent 8f9a323 commit 6a6a368

8 files changed

Lines changed: 482 additions & 109 deletions

File tree

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,3 @@
11
# 设计说明(约 400 字)
22

3-
闭环以 `EvalOptimizePipeline.run()` 作为唯一外部接口,把数据校验、基线评估、提示词优化、候选复评、差异计算、门禁和报告隐藏在同一深模块内。优化阶段直接复用 `AgentOptimizer`、`TargetPrompt` 与真实 GEPA;无密钥模式仅通过 `ModelRegistry` 注入离线 agent、judge 和 reflector,因而仍会经过 `LlmAgent`、内置 rubric 解析及优化器回调,不依赖 monkeypatch。外层将每个候选运行物化为 trace,再用 `AgentEvaluator` 独立评测 train 与 validation,避免采信优化器自己的聚合分数。候选集合按完整提示词哈希去重,即使优化器内部提前拒绝某次提案,外层仍会捕获并独立复评,防止遗漏验证回退。逐 case 差异区分新增通过、新增失败、升分、降分和不变;归因按工具执行、参数、格式、知识召回、rubric、最终回答的根因优先级输出证据。SDK 聚合器偶尔不保留逐 rubric 详情,格式归因依次采用详情、回放证据和请求中明确的 JSON、单行或 Markdown 约束,否定式要求不触发检查;judge 无有效 verdict 时单列 evaluation_error,避免把系统错误或内容失败误报为格式失败。Gate 对验证增益、配对 bootstrap 下界、新增 hard fail、关键 case、过拟合及资源预算做可配置 AND,并先过滤不安全候选,再计算质量、token、P95 时延 Pareto。运行前检查跨集合重复、近重复和答案泄漏;运行后保存完整 prompt、seed、哈希、耗时、成本、trace 与 JSON/Markdown 报告。默认始终恢复源 prompt,只有显式允许且 Gate 全通过才原子回写。
3+
闭环以 `EvalOptimizePipeline.run()` 作为唯一外部接口,把数据校验、基线评估、提示词优化、候选复评、差异计算、门禁和报告隐藏在同一深模块内。优化阶段直接复用 `AgentOptimizer`、`TargetPrompt` 与真实 GEPA;无密钥模式仅通过 `ModelRegistry` 注入离线 agent、judge 和 reflector,因而仍会经过 `LlmAgent`、内置 rubric 解析及优化器回调,不依赖 monkeypatch。外层将每个候选运行物化为 trace,再用 `AgentEvaluator` 独立评测 train 与 validation,避免采信优化器自己的聚合分数。候选集合按完整提示词哈希去重,即使优化器内部提前拒绝某次提案,外层仍会捕获并独立复评,防止遗漏验证回退。逐 case 差异区分新增通过、新增失败、升分、降分和不变;归因按工具执行、参数、格式、知识召回、rubric、最终回答的根因优先级输出证据。SDK 聚合器偶尔不保留逐 rubric 详情,格式归因依次采用详情、回放证据和请求中明确的 JSON、单行或 Markdown 约束,否定式要求不触发检查;judge 无有效 verdict 时单列 evaluation_error,避免把系统错误或内容失败误报为格式失败。Gate 对验证增益、配对 bootstrap 下界、新增 hard fail、关键 case、过拟合及资源预算做可配置 AND,并先过滤不安全候选,再计算质量、token、P95 时延 Pareto。运行前检查跨集合重复、近重复和答案泄漏;运行后保存完整 prompt、seed、哈希、耗时、成本、trace 与 JSON/Markdown 报告。默认仅在隔离 workspace 内改写 prompt,不触碰源文件,只有显式允许且 Gate 全通过才原子回写。

examples/optimization/eval_optimize_loop/README.md

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@ python examples/optimization/eval_optimize_loop/run_pipeline.py \
1818
--output-dir /tmp/eval-optimize-loop
1919
```
2020

21-
命令正常完成时输出 `decision=accepted` 与报告路径。默认不会覆盖源 prompt;只有显式传入 `--apply-if-accepted` 且所有必选 Gate 均通过,才会原子回写最终候选。
21+
命令正常完成时输出 `decision=accepted` 与报告路径。优化器只改写输出目录中的隔离 prompt workspace,默认运行不会触碰源 prompt;只有显式传入 `--apply-if-accepted` 且所有必选 Gate 均通过,才会原子回写最终候选。
2222

2323
## 为什么不是“伪造优化”
2424

@@ -71,7 +71,7 @@ reflection
7171
2. validation pass-rate 增益达到 `min_validation_gain`
7272
3. paired bootstrap 区间下界达到配置值;
7373
4. 不新增 hard failure;
74-
5. key case 不降分
74+
5. key case 不允许从通过变失败,也不允许降分
7575
6. train 提升而 validation 下降时判定 overfit 并拒绝;趋势以 pass-rate 为主,持平时再比较 average score;
7676
7. metric calls、token、耗时和成本均在预算内。
7777

@@ -82,6 +82,7 @@ reflection
8282
- baseline 与 candidate 在同一 case 上配对,使用固定 seed 的 2,000 次 bootstrap,报告 validation pass-rate delta 的 95% 区间。
8383
- 先过安全 Gate,再在合格候选中按 validation 质量、token 与 P95 latency 标记 Pareto 前沿。
8484
- 运行前检查重复 id、train/validation 精确重复、去空白后的重复、相似度 ≥0.92 的近重复,以及 baseline/candidate prompt 中直接出现 validation reference answer;命中即 fail-closed。
85+
- artifact label 只允许有限长度的字母、数字、下划线和连字符,候选 id 必须唯一;离线回放遇到空问题或同一问题映射多个 case 时会拒绝运行,避免路径逃逸和静默错配。
8586
- baseline、每个候选和每个 split 的 trace 均保存到 artifact 目录,可由 `AgentEvaluator` 重新播放。
8687

8788
## 输出
@@ -112,4 +113,4 @@ pytest -q tests/evaluation/test_eval_optimize_loop_example.py
112113

113114
## 接入真实业务
114115

115-
保留 `EvalOptimizePipeline.run()` 这个外部 interface,将 `offline.py` 的三个 adapter 替换为业务实现:`call_agent` 驱动真实 Agent,reflection/judge 的 `provider_name` 改为实际 provider,trace 物化器读取业务运行日志。`optimizer.json``TargetPrompt`、逐候选复评、diff、Gate 和报告模型无需改变。生产环境建议每次使用唯一 output directory;离线 registry 使用进程级固定 replay 配置,不支持在同一 Python 进程中并发启动两个 pipeline
116+
保留 `EvalOptimizePipeline.run()` 这个外部 interface,将 `offline.py` 的三个 adapter 替换为业务实现:`call_agent` 驱动真实 Agent,reflection/judge 的 `provider_name` 改为实际 provider,trace 物化器读取业务运行日志。`optimizer.json``TargetPrompt`、逐候选复评、diff、Gate 和报告模型无需改变。生产环境应为每次运行使用唯一 output directory;不同进程会在各自 workspace 优化而不触碰源 prompt,同一 Python 进程中的并发离线运行会明确拒绝,避免共享 registry 串扰

examples/optimization/eval_optimize_loop/loop/analysis.py

Lines changed: 17 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -83,6 +83,13 @@ def validate_data_quality(
8383

8484
all_ids = [case.eval_id for case in train_set.eval_cases + validation_set.eval_cases]
8585
duplicate_ids = sorted(case_id for case_id, count in Counter(all_ids).items() if count > 1)
86+
query_groups: dict[str, list[str]] = {}
87+
for case in train_set.eval_cases + validation_set.eval_cases:
88+
replay_query = self._case_replay_query(case)
89+
if not replay_query:
90+
raise ValueError(f"case {case.eval_id!r} has an empty user query")
91+
query_groups.setdefault(replay_query, []).append(case.eval_id)
92+
duplicate_queries = sorted(" == ".join(case_ids) for case_ids in query_groups.values() if len(case_ids) > 1)
8693
train_fingerprints = {self._case_fingerprint(case): case.eval_id for case in train_set.eval_cases}
8794
validation_fingerprints = {self._case_fingerprint(case): case.eval_id for case in validation_set.eval_cases}
8895
overlap = sorted(set(train_fingerprints) & set(validation_fingerprints))
@@ -106,9 +113,10 @@ def validate_data_quality(
106113
normalized = _normalize(expected)
107114
if len(normalized) >= 12 and normalized in prompt_normalized:
108115
leakage.append(case.eval_id)
109-
if duplicate_ids or cross_split or near_cross_split or leakage:
116+
if duplicate_ids or duplicate_queries or cross_split or near_cross_split or leakage:
110117
raise ValueError("data quality check failed: "
111-
f"duplicate_ids={duplicate_ids}, cross_split={cross_split}, "
118+
f"duplicate_ids={duplicate_ids}, duplicate_queries={duplicate_queries}, "
119+
f"cross_split={cross_split}, "
112120
f"near_cross_split={near_cross_split}, prompt_leakage={leakage}")
113121
return DataQualityAudit(
114122
passed=True,
@@ -131,6 +139,11 @@ def _case_normalized_content(case: EvalCase) -> str:
131139
for invocation in conversation)
132140
return _normalize(payload)
133141

142+
@staticmethod
143+
def _case_replay_query(case: EvalCase) -> str:
144+
conversation = case.conversation or []
145+
return _text(conversation[0].user_content) if conversation else ""
146+
134147
@staticmethod
135148
def _expected_response(case: EvalCase) -> str:
136149
conversation = case.conversation or []
@@ -324,7 +337,8 @@ def gate(
324337
baseline_validation_by_id = {case.case_id: case for case in baseline.validation.cases}
325338
key_regressions = sorted(
326339
case.case_id for case in candidate_validation.cases
327-
if case.key_case and baseline_validation_by_id[case.case_id].score - case.score > epsilon)
340+
if case.key_case and ((baseline_validation_by_id[case.case_id].passed and not case.passed)
341+
or baseline_validation_by_id[case.case_id].score - case.score > epsilon))
328342

329343
def _trend(split_delta: SplitDelta) -> int:
330344
if split_delta.pass_rate_delta > epsilon:

examples/optimization/eval_optimize_loop/loop/models.py

Lines changed: 94 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@
1010
from __future__ import annotations
1111

1212
import json
13+
import re
1314
from pathlib import Path
1415
from typing import Any
1516
from typing import Literal
@@ -18,6 +19,39 @@
1819
from pydantic import BaseModel
1920
from pydantic import ConfigDict
2021
from pydantic import Field
22+
from pydantic import field_validator
23+
from pydantic import model_validator
24+
25+
_SAFE_ARTIFACT_LABEL = r"^[A-Za-z0-9][A-Za-z0-9_-]{0,63}$"
26+
_WINDOWS_DEVICE_NAMES = {
27+
"aux",
28+
"con",
29+
"nul",
30+
"prn",
31+
*(f"com{index}" for index in range(1, 10)),
32+
*(f"lpt{index}" for index in range(1, 10)),
33+
}
34+
35+
36+
def _validate_artifact_label(
37+
value: str,
38+
*,
39+
subject: str,
40+
reserved: set[str] | None = None,
41+
) -> str:
42+
if re.fullmatch(_SAFE_ARTIFACT_LABEL, value) is None:
43+
raise ValueError(f"{subject} must be a safe artifact label")
44+
folded = value.casefold()
45+
if folded in _WINDOWS_DEVICE_NAMES or folded in (reserved or set()):
46+
raise ValueError(f"{subject} is a reserved artifact label: {value!r}")
47+
return value
48+
49+
50+
def _casefold_duplicates(values: list[str]) -> list[str]:
51+
groups: dict[str, list[str]] = {}
52+
for value in values:
53+
groups.setdefault(value.casefold(), []).append(value)
54+
return sorted(" == ".join(group) for group in groups.values() if len(group) > 1)
2155

2256

2357
class _StrictModel(BaseModel):
@@ -27,9 +61,18 @@ class _StrictModel(BaseModel):
2761
class CandidateSource(_StrictModel):
2862
"""A deterministic prompt proposal used by the offline reflection model."""
2963

30-
candidate_id: str
64+
candidate_id: str = Field(pattern=_SAFE_ARTIFACT_LABEL)
3165
path: Path
3266

67+
@field_validator("candidate_id")
68+
@classmethod
69+
def _candidate_id_is_artifact_safe(cls, value: str) -> str:
70+
return _validate_artifact_label(
71+
value,
72+
subject="candidate_id",
73+
reserved={"baseline"},
74+
)
75+
3376

3477
class PipelineSpec(_StrictModel):
3578
"""All paths and run controls needed by :class:`EvalOptimizePipeline`."""
@@ -40,14 +83,29 @@ class PipelineSpec(_StrictModel):
4083
gate_config: Path
4184
train_dataset: Path
4285
validation_dataset: Path
43-
target_prompts: dict[str, Path]
44-
candidate_sources: list[CandidateSource]
86+
target_prompts: dict[str, Path] = Field(min_length=1)
87+
candidate_sources: list[CandidateSource] = Field(min_length=1)
4588
output_dir: Path
4689
seed: int = 91
4790
bootstrap_samples: int = Field(default=2000, ge=100)
4891
confidence_level: float = Field(default=0.95, gt=0.0, lt=1.0)
4992
apply_if_accepted: bool = False
5093

94+
@model_validator(mode="after")
95+
def _validate_artifact_labels(self) -> "PipelineSpec":
96+
candidate_ids = [source.candidate_id for source in self.candidate_sources]
97+
duplicate_ids = _casefold_duplicates(candidate_ids)
98+
if duplicate_ids:
99+
raise ValueError(f"candidate_ids must be case-insensitively unique: {duplicate_ids}")
100+
prompt_names = list(self.target_prompts)
101+
duplicate_prompt_names = _casefold_duplicates(prompt_names)
102+
if duplicate_prompt_names:
103+
raise ValueError("target prompt names must be case-insensitively unique: "
104+
f"{duplicate_prompt_names}")
105+
for name in prompt_names:
106+
_validate_artifact_label(name, subject="target prompt name")
107+
return self
108+
51109
@classmethod
52110
def from_file(
53111
cls,
@@ -94,8 +152,8 @@ def _resolve(value: str | Path) -> Path:
94152

95153
class MetricOutcome(_StrictModel):
96154
metric_name: str
97-
score: Optional[float] = None
98-
threshold: float
155+
score: Optional[float] = Field(default=None, allow_inf_nan=False)
156+
threshold: float = Field(allow_inf_nan=False)
99157
passed: bool
100158
reason: str = ""
101159

@@ -115,7 +173,7 @@ class TrajectoryStep(_StrictModel):
115173
class CaseEvaluation(_StrictModel):
116174
case_id: str
117175
passed: bool
118-
score: float
176+
score: float = Field(allow_inf_nan=False)
119177
key_case: bool = False
120178
hard_fail: bool = False
121179
metrics: list[MetricOutcome] = Field(default_factory=list)
@@ -128,8 +186,8 @@ class CaseEvaluation(_StrictModel):
128186

129187
class SplitEvaluation(_StrictModel):
130188
split: Literal["train", "validation"]
131-
pass_rate: float
132-
average_score: float
189+
pass_rate: float = Field(ge=0.0, le=1.0, allow_inf_nan=False)
190+
average_score: float = Field(allow_inf_nan=False)
133191
cases: list[CaseEvaluation]
134192

135193

@@ -152,24 +210,24 @@ class CaseDelta(_StrictModel):
152210
status: DeltaStatus
153211
baseline_passed: bool
154212
candidate_passed: bool
155-
baseline_score: float
156-
candidate_score: float
157-
score_delta: float
213+
baseline_score: float = Field(allow_inf_nan=False)
214+
candidate_score: float = Field(allow_inf_nan=False)
215+
score_delta: float = Field(allow_inf_nan=False)
158216

159217

160218
class PairedConfidenceInterval(_StrictModel):
161-
point_estimate: float
162-
lower: float
163-
upper: float
164-
confidence_level: float
165-
bootstrap_samples: int
219+
point_estimate: float = Field(allow_inf_nan=False)
220+
lower: float = Field(allow_inf_nan=False)
221+
upper: float = Field(allow_inf_nan=False)
222+
confidence_level: float = Field(gt=0.0, lt=1.0, allow_inf_nan=False)
223+
bootstrap_samples: int = Field(ge=1)
166224
seed: int
167225

168226

169227
class SplitDelta(_StrictModel):
170228
split: Literal["train", "validation"]
171-
pass_rate_delta: float
172-
average_score_delta: float
229+
pass_rate_delta: float = Field(allow_inf_nan=False)
230+
average_score_delta: float = Field(allow_inf_nan=False)
173231
paired_pass_rate_ci: PairedConfidenceInterval
174232
newly_passed: list[str] = Field(default_factory=list)
175233
newly_failed: list[str] = Field(default_factory=list)
@@ -200,15 +258,15 @@ class GateDecision(_StrictModel):
200258

201259

202260
class ResourceUsage(_StrictModel):
203-
metric_calls: int = 0
204-
reflection_calls: int = 0
205-
judge_calls: Optional[int] = None
206-
prompt_tokens: int = 0
207-
completion_tokens: int = 0
208-
total_tokens: int = 0
209-
cost_usd: Optional[float] = None
210-
duration_seconds: float = 0.0
211-
p95_latency_ms: Optional[float] = None
261+
metric_calls: int = Field(default=0, ge=0)
262+
reflection_calls: int = Field(default=0, ge=0)
263+
judge_calls: Optional[int] = Field(default=None, ge=0)
264+
prompt_tokens: int = Field(default=0, ge=0)
265+
completion_tokens: int = Field(default=0, ge=0)
266+
total_tokens: int = Field(default=0, ge=0)
267+
cost_usd: Optional[float] = Field(default=None, ge=0.0, allow_inf_nan=False)
268+
duration_seconds: float = Field(default=0.0, ge=0.0, allow_inf_nan=False)
269+
p95_latency_ms: Optional[float] = Field(default=None, ge=0.0, allow_inf_nan=False)
212270
cost_measurement: str = "unavailable"
213271

214272

@@ -236,17 +294,17 @@ class OptimizerAudit(_StrictModel):
236294
status: str
237295
stop_reason: Optional[str] = None
238296
used_agent_optimizer: bool
239-
baseline_pass_rate: float
240-
best_pass_rate: float
241-
rounds: int
297+
baseline_pass_rate: float = Field(ge=0.0, le=1.0, allow_inf_nan=False)
298+
best_pass_rate: float = Field(ge=0.0, le=1.0, allow_inf_nan=False)
299+
rounds: int = Field(ge=0)
242300
resources: ResourceUsage
243301
artifact_dir: str
244302

245303

246304
class DataQualityAudit(_StrictModel):
247305
passed: bool
248-
train_cases: int
249-
validation_cases: int
306+
train_cases: int = Field(ge=0)
307+
validation_cases: int = Field(ge=0)
250308
duplicate_ids: list[str] = Field(default_factory=list)
251309
cross_split_duplicates: list[str] = Field(default_factory=list)
252310
near_cross_split_duplicates: list[str] = Field(default_factory=list)
@@ -257,7 +315,7 @@ class RunAudit(_StrictModel):
257315
run_id: str
258316
started_at: str
259317
finished_at: str
260-
duration_seconds: float
318+
duration_seconds: float = Field(ge=0.0, allow_inf_nan=False)
261319
seed: int
262320
config_sha256: str
263321
train_sha256: str
@@ -267,9 +325,9 @@ class RunAudit(_StrictModel):
267325

268326

269327
class FailureAttributionSummary(_StrictModel):
270-
explained_failed_cases: int
271-
total_failed_cases: int
272-
coverage_rate: float
328+
explained_failed_cases: int = Field(ge=0)
329+
total_failed_cases: int = Field(ge=0)
330+
coverage_rate: float = Field(ge=0.0, le=1.0, allow_inf_nan=False)
273331
category_counts: dict[str, int] = Field(default_factory=dict)
274332
by_case: dict[str, list[FailureReason]] = Field(default_factory=dict)
275333

examples/optimization/eval_optimize_loop/loop/offline.py

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -93,6 +93,12 @@ def configure(
9393
if not conversation:
9494
continue
9595
query = _content_text(conversation[0].user_content)
96+
if not query:
97+
raise ValueError(f"case {case.eval_id!r} has an empty offline replay query")
98+
if query in catalog:
99+
previous = catalog[query]
100+
raise ValueError("offline replay queries must be unique; "
101+
f"cases {previous.eval_id!r} and {case.eval_id!r} share the same query")
96102
catalog[query] = case
97103
with cls._configuration_lock:
98104
cls._catalog = catalog

0 commit comments

Comments
 (0)