|
4 | 4 | import sys |
5 | 5 | import tempfile |
6 | 6 | from pathlib import Path |
7 | | -from unittest.mock import AsyncMock |
| 7 | +from unittest.mock import AsyncMock, patch |
8 | 8 |
|
9 | 9 | import pytest |
10 | 10 |
|
|
13 | 13 | if str(_EXAMPLE_ROOT.parents[1]) not in sys.path: |
14 | 14 | sys.path.insert(0, str(_EXAMPLE_ROOT.parents[1])) |
15 | 15 |
|
16 | | -from trpc_agent_sdk.evaluation import AgentEvaluator, EvalConfig |
17 | | - |
18 | | -from ..models import PipelineConfig, PipelineResult, GateConfig |
| 16 | +from ..models import PipelineConfig |
19 | 17 | from ..pipeline import EvalOptimizePipeline |
20 | 18 |
|
21 | 19 |
|
@@ -207,3 +205,154 @@ async def test_pipeline_build_split_result(): |
207 | 205 | assert sr.per_case["case_a"].passed is True |
208 | 206 | assert sr.per_case["case_b"].passed is False |
209 | 207 | 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