Skip to content

Commit a3e4b59

Browse files
committed
fix(eval_optimize_loop): address review feedback for pipeline orchestrator
- Fix Optional[callable] -> Optional[Callable] with typing.Callable import - Add tests for _write_eval_config_temp, _run_optimization injection hook, and run() trace mode - Remove unused imports (Path, EvalConfig, EvalSetAggregateResult, GateDecision from pipeline.py) - Remove unused imports (AgentEvaluator, EvalConfig, PipelineResult, GateConfig from tests)
1 parent 91e64a0 commit a3e4b59

2 files changed

Lines changed: 155 additions & 10 deletions

File tree

examples/optimization/eval_optimize_loop/pipeline.py

Lines changed: 2 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -4,26 +4,22 @@
44
import time
55
import tempfile
66
from datetime import datetime, timezone
7-
from pathlib import Path
8-
from typing import Optional
7+
from typing import Callable, Optional
98

109
from trpc_agent_sdk.evaluation import (
1110
AgentEvaluator,
1211
AgentOptimizer,
1312
CallAgent,
1413
EvalCaseResult,
15-
EvalConfig,
1614
EvalStatus,
1715
EvaluateResult,
18-
EvalSetAggregateResult,
1916
TargetPrompt,
2017
)
2118

2219
from .delta import compute_delta
2320
from .failure_attribution import attribute_failures
2421
from .gate import apply_gate
2522
from .models import (
26-
GateDecision,
2723
PerCaseResult,
2824
PipelineConfig,
2925
PipelineResult,
@@ -37,7 +33,7 @@ def __init__(self, config: PipelineConfig) -> None:
3733
self._config = config
3834
self._live_call_agent: Optional[CallAgent] = None
3935
self._live_target_prompt: Optional[TargetPrompt] = None
40-
self._optimizer_call: Optional[callable] = None # type: ignore[valid-type]
36+
self._optimizer_call: Optional[Callable] = None
4137

4238
@classmethod
4339
def from_config(

examples/optimization/eval_optimize_loop/tests/test_pipeline.py

Lines changed: 153 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
import sys
55
import tempfile
66
from pathlib import Path
7-
from unittest.mock import AsyncMock
7+
from unittest.mock import AsyncMock, patch
88

99
import pytest
1010

@@ -13,9 +13,7 @@
1313
if str(_EXAMPLE_ROOT.parents[1]) not in sys.path:
1414
sys.path.insert(0, str(_EXAMPLE_ROOT.parents[1]))
1515

16-
from trpc_agent_sdk.evaluation import AgentEvaluator, EvalConfig
17-
18-
from ..models import PipelineConfig, PipelineResult, GateConfig
16+
from ..models import PipelineConfig
1917
from ..pipeline import EvalOptimizePipeline
2018

2119

@@ -207,3 +205,154 @@ async def test_pipeline_build_split_result():
207205
assert sr.per_case["case_a"].passed is True
208206
assert sr.per_case["case_b"].passed is False
209207
assert "m1" in sr.metric_breakdown
208+
209+
210+
# ── Helpers ──────────────────────────────────────────────────────
211+
212+
213+
def _make_fake_eval_result(
214+
eval_set_id: str,
215+
case_ids: list[str],
216+
passed: list[bool],
217+
metric_name: str = "m1",
218+
) -> "EvaluateResult":
219+
from trpc_agent_sdk.evaluation import (
220+
EvalCaseResult,
221+
EvalMetricResult,
222+
EvalStatus,
223+
EvaluateResult,
224+
EvalSetAggregateResult,
225+
)
226+
227+
eval_results_by_eval_id: dict[str, list[EvalCaseResult]] = {}
228+
for case_id, is_pass in zip(case_ids, passed):
229+
status = EvalStatus.PASSED if is_pass else EvalStatus.FAILED
230+
score = 1.0 if is_pass else 0.0
231+
eval_results_by_eval_id[case_id] = [
232+
EvalCaseResult(
233+
eval_set_id=eval_set_id,
234+
eval_id=case_id,
235+
final_eval_status=status,
236+
overall_eval_metric_results=[
237+
EvalMetricResult(
238+
metric_name=metric_name,
239+
score=score,
240+
threshold=1.0,
241+
eval_status=status,
242+
),
243+
],
244+
eval_metric_result_per_invocation=[],
245+
session_id=f"s_{case_id}",
246+
)
247+
]
248+
249+
return EvaluateResult(
250+
results_by_eval_set_id={
251+
eval_set_id: EvalSetAggregateResult(
252+
eval_results_by_eval_id=eval_results_by_eval_id,
253+
num_runs=1,
254+
)
255+
}
256+
)
257+
258+
259+
def _make_pipeline(config_overrides: dict | None = None) -> EvalOptimizePipeline:
260+
base = {
261+
"mode": "trace",
262+
"output_dir": "/tmp/fake_outputs",
263+
"evaluate": {
264+
"metrics": [
265+
{
266+
"metric_name": "m1",
267+
"threshold": 1.0,
268+
"criterion": {"final_response": {"text": {"match": "contains"}}},
269+
}
270+
],
271+
"num_runs": 1,
272+
},
273+
"train_baseline_evalset": "/tmp/train_base.json",
274+
"val_baseline_evalset": "/tmp/val_base.json",
275+
"train_candidate_evalset": "/tmp/train_cand.json",
276+
"val_candidate_evalset": "/tmp/val_cand.json",
277+
"seed": 42,
278+
}
279+
if config_overrides:
280+
base.update(config_overrides)
281+
282+
pipeline = EvalOptimizePipeline.__new__(EvalOptimizePipeline)
283+
pipeline._config = PipelineConfig.model_validate(base)
284+
pipeline._live_call_agent = None
285+
pipeline._live_target_prompt = None
286+
pipeline._optimizer_call = None
287+
return pipeline
288+
289+
290+
# ── New high-priority tests ─────────────────────────────────────
291+
292+
293+
@pytest.mark.asyncio
294+
async def test_write_eval_config_temp_creates_and_returns_path():
295+
pipeline = _make_pipeline()
296+
path = await pipeline._write_eval_config_temp()
297+
try:
298+
assert os.path.isfile(path)
299+
import json
300+
with open(path) as f:
301+
data = json.load(f)
302+
assert "metrics" in data
303+
assert data["numRuns"] == 1
304+
finally:
305+
os.unlink(path)
306+
307+
308+
@pytest.mark.asyncio
309+
async def test_run_optimization_calls_injected_hook():
310+
pipeline = _make_pipeline({"mode": "live"})
311+
hook = AsyncMock()
312+
pipeline._optimizer_call = hook
313+
await pipeline._run_optimization()
314+
hook.assert_called_once_with(pipeline)
315+
316+
317+
@pytest.mark.asyncio
318+
async def test_run_trace_mode_orchestration():
319+
pipeline = _make_pipeline()
320+
321+
fake_train = _make_fake_eval_result("train", ["a", "b"], [True, False])
322+
fake_val = _make_fake_eval_result("val", ["c", "d"], [True, True])
323+
324+
async def _fake_run_eval(_path: str):
325+
return fake_train if "train" in _path else fake_val
326+
327+
with patch.object(pipeline, "_run_eval", side_effect=_fake_run_eval):
328+
with patch("examples.optimization.eval_optimize_loop.pipeline.write_reports"):
329+
result = await pipeline.run()
330+
331+
assert result.mode == "trace"
332+
assert result.seed == 42
333+
assert "train" in result.baseline
334+
assert "val" in result.baseline
335+
assert result.baseline["train"].pass_rate == 0.5
336+
assert result.baseline["val"].pass_rate == 1.0
337+
338+
339+
@pytest.mark.asyncio
340+
async def test_run_optimization_calls_agent_optimizer_when_no_hook():
341+
pipeline = _make_pipeline(
342+
{
343+
"mode": "live",
344+
"optimizer_config_path": "/tmp/opt.json",
345+
"live_train_evalset": "/tmp/train.json",
346+
"live_val_evalset": "/tmp/val.json",
347+
}
348+
)
349+
with patch(
350+
"examples.optimization.eval_optimize_loop.pipeline.AgentOptimizer.optimize",
351+
new_callable=AsyncMock,
352+
) as mock_optimize:
353+
await pipeline._run_optimization()
354+
355+
mock_optimize.assert_called_once()
356+
call_kwargs = mock_optimize.call_args.kwargs
357+
assert call_kwargs["config_path"] == "/tmp/opt.json"
358+
assert call_kwargs["train_dataset_path"] == "/tmp/train.json"

0 commit comments

Comments
 (0)