Skip to main content

assay_core/engine/
runner.rs

1use crate::cache::vcr::VcrCache;
2use crate::metrics_api::Metric;
3use crate::model::{EvalConfig, LlmResponse, TestCase, TestResultRow, TestStatus};
4use crate::providers::llm::LlmClient;
5use crate::quarantine::QuarantineMode;
6use crate::report::progress::ProgressSink;
7use crate::report::RunArtifacts;
8use crate::storage::store::Store;
9use std::sync::Arc;
10
11#[path = "runner_next/mod.rs"]
12mod runner_next;
13
14#[derive(Debug, Clone)]
15pub struct RunPolicy {
16    pub rerun_failures: u32,
17    pub quarantine_mode: QuarantineMode,
18    pub replay_strict: bool,
19}
20
21impl Default for RunPolicy {
22    fn default() -> Self {
23        Self {
24            rerun_failures: 1,
25            quarantine_mode: QuarantineMode::Warn,
26            replay_strict: false,
27        }
28    }
29}
30
31pub struct Runner {
32    pub store: Store,
33    pub cache: VcrCache,
34    pub client: Arc<dyn LlmClient>,
35    pub metrics: Vec<Arc<dyn Metric>>,
36    pub policy: RunPolicy,
37    pub _network_guard: Option<crate::providers::network::NetworkPolicyGuard>,
38    pub embedder: Option<Arc<dyn crate::providers::embedder::Embedder>>,
39    pub refresh_embeddings: bool,
40    pub incremental: bool,
41    pub refresh_cache: bool,
42    pub judge: Option<crate::judge::JudgeService>,
43    pub baseline: Option<crate::baseline::Baseline>,
44}
45
46impl Runner {
47    /// Run the suite; results are collected in completion order internally but returned
48    /// sorted by test_id for deterministic output. If `progress` is set, it is called
49    /// after each test completes (E4.3 realtime progress).
50    pub async fn run_suite(
51        &self,
52        cfg: &EvalConfig,
53        progress: Option<ProgressSink>,
54    ) -> anyhow::Result<RunArtifacts> {
55        runner_next::execute::run_suite_impl(self, cfg, progress).await
56    }
57
58    fn apply_agent_assertions(
59        &self,
60        run_id: i64,
61        tc: &TestCase,
62        final_row: &mut TestResultRow,
63    ) -> anyhow::Result<()> {
64        runner_next::assertions::apply_agent_assertions_impl(self, run_id, tc, final_row)
65    }
66
67    async fn run_test_once(
68        &self,
69        cfg: &EvalConfig,
70        tc: &TestCase,
71    ) -> anyhow::Result<(TestResultRow, LlmResponse)> {
72        runner_next::single::run_test_once_impl(self, cfg, tc).await
73    }
74
75    async fn call_llm(&self, cfg: &EvalConfig, tc: &TestCase) -> anyhow::Result<LlmResponse> {
76        runner_next::execute::call_llm_impl(self, cfg, tc).await
77    }
78
79    fn check_baseline_regressions(
80        &self,
81        tc: &TestCase,
82        cfg: &EvalConfig,
83        details: &serde_json::Value,
84        metrics: &[Arc<dyn Metric>],
85        baseline: &crate::baseline::Baseline,
86    ) -> Option<(TestStatus, String)> {
87        runner_next::baseline::check_baseline_regressions_impl(
88            self, tc, cfg, details, metrics, baseline,
89        )
90    }
91
92    // Embeddings logic
93    async fn enrich_semantic(
94        &self,
95        _cfg: &EvalConfig,
96        tc: &TestCase,
97        resp: &mut LlmResponse,
98    ) -> anyhow::Result<()> {
99        runner_next::scoring::enrich_semantic_impl(self, _cfg, tc, resp).await
100    }
101
102    pub async fn embed_text(
103        &self,
104        model_id: &str,
105        embedder: &dyn crate::providers::embedder::Embedder,
106        text: &str,
107    ) -> anyhow::Result<(Vec<f32>, &'static str)> {
108        runner_next::cache::embed_text_impl(self, model_id, embedder, text).await
109    }
110
111    async fn enrich_judge(
112        &self,
113        cfg: &EvalConfig,
114        tc: &TestCase,
115        resp: &mut LlmResponse,
116    ) -> anyhow::Result<()> {
117        runner_next::scoring::enrich_judge_impl(self, cfg, tc, resp).await
118    }
119}
120
121#[cfg(test)]
122mod tests {
123    use super::*;
124    use crate::metrics_api::{Metric, MetricResult};
125    use crate::model::{Expected, Settings, TestInput};
126    use crate::on_error::ErrorPolicy;
127    use crate::providers::llm::fake::FakeClient;
128    use crate::providers::llm::LlmClient;
129    use async_trait::async_trait;
130    use std::sync::atomic::{AtomicUsize, Ordering};
131
132    #[derive(Clone, Copy)]
133    enum MetricMode {
134        FailThenPass,
135        AlwaysFail,
136        AlwaysPass,
137    }
138
139    struct ScriptedMetric {
140        mode: MetricMode,
141        calls: AtomicUsize,
142    }
143
144    impl ScriptedMetric {
145        fn fail_then_pass() -> Self {
146            Self {
147                mode: MetricMode::FailThenPass,
148                calls: AtomicUsize::new(0),
149            }
150        }
151
152        fn always_fail() -> Self {
153            Self {
154                mode: MetricMode::AlwaysFail,
155                calls: AtomicUsize::new(0),
156            }
157        }
158
159        fn always_pass() -> Self {
160            Self {
161                mode: MetricMode::AlwaysPass,
162                calls: AtomicUsize::new(0),
163            }
164        }
165    }
166
167    #[async_trait]
168    impl Metric for ScriptedMetric {
169        fn name(&self) -> &'static str {
170            "scripted"
171        }
172
173        async fn evaluate(
174            &self,
175            _tc: &TestCase,
176            _expected: &Expected,
177            _resp: &LlmResponse,
178        ) -> anyhow::Result<MetricResult> {
179            let n = self.calls.fetch_add(1, Ordering::SeqCst);
180            match self.mode {
181                MetricMode::FailThenPass => {
182                    if n == 0 {
183                        Ok(MetricResult::fail(0.0, "scripted_fail_once"))
184                    } else {
185                        Ok(MetricResult::pass(1.0))
186                    }
187                }
188                MetricMode::AlwaysFail => Ok(MetricResult::fail(0.0, "scripted_fail")),
189                MetricMode::AlwaysPass => Ok(MetricResult::pass(1.0)),
190            }
191        }
192    }
193
194    struct ErrorClient;
195
196    #[async_trait]
197    impl LlmClient for ErrorClient {
198        async fn complete(
199            &self,
200            _prompt: &str,
201            _context: Option<&[String]>,
202        ) -> anyhow::Result<LlmResponse> {
203            Err(anyhow::anyhow!("scripted provider error"))
204        }
205
206        fn provider_name(&self) -> &'static str {
207            "error_client"
208        }
209    }
210
211    fn runner_for_contract_tests(
212        client: Arc<dyn LlmClient>,
213        metrics: Vec<Arc<dyn Metric>>,
214        rerun_failures: u32,
215    ) -> Runner {
216        let store = Store::memory().expect("in-memory store");
217        store.init_schema().expect("schema init");
218        Runner {
219            store: store.clone(),
220            cache: VcrCache::new(store),
221            client,
222            metrics,
223            policy: RunPolicy {
224                rerun_failures,
225                quarantine_mode: QuarantineMode::Off,
226                replay_strict: false,
227            },
228            _network_guard: None,
229            embedder: None,
230            refresh_embeddings: false,
231            incremental: false,
232            refresh_cache: false,
233            judge: None,
234            baseline: None,
235        }
236    }
237
238    fn single_test_config(on_error: ErrorPolicy) -> EvalConfig {
239        EvalConfig {
240            version: 1,
241            suite: "runner-contract".to_string(),
242            model: "fake-model".to_string(),
243            settings: Settings {
244                parallel: Some(1),
245                cache: Some(false),
246                seed: Some(1234),
247                on_error,
248                ..Default::default()
249            },
250            thresholds: Default::default(),
251            otel: Default::default(),
252            tests: vec![TestCase {
253                id: "t1".to_string(),
254                input: TestInput {
255                    prompt: "contract prompt".to_string(),
256                    context: None,
257                },
258                // Expected payload is not used by scripted metrics, but keeps test case valid.
259                expected: Expected::MustContain {
260                    must_contain: vec!["ok".to_string()],
261                },
262                assertions: None,
263                on_error: None,
264                tags: vec![],
265                metadata: None,
266            }],
267        }
268    }
269
270    fn config_with_test_ids(ids: &[&str], on_error: ErrorPolicy) -> EvalConfig {
271        EvalConfig {
272            version: 1,
273            suite: "runner-contract".to_string(),
274            model: "fake-model".to_string(),
275            settings: Settings {
276                parallel: Some(1),
277                cache: Some(false),
278                seed: Some(1234),
279                on_error,
280                ..Default::default()
281            },
282            thresholds: Default::default(),
283            otel: Default::default(),
284            tests: ids
285                .iter()
286                .map(|id| TestCase {
287                    id: (*id).to_string(),
288                    input: TestInput {
289                        prompt: format!("prompt-{id}"),
290                        context: None,
291                    },
292                    expected: Expected::MustContain {
293                        must_contain: vec!["ok".to_string()],
294                    },
295                    assertions: None,
296                    on_error: None,
297                    tags: vec![],
298                    metadata: None,
299                })
300                .collect(),
301        }
302    }
303
304    #[tokio::test]
305    async fn runner_contract_flake_fail_then_pass_classified_flaky() -> anyhow::Result<()> {
306        let cfg = single_test_config(ErrorPolicy::Block);
307        let client = Arc::new(FakeClient::new("fake-model".to_string()).with_response("ok".into()));
308        let metric = Arc::new(ScriptedMetric::fail_then_pass());
309        let runner = runner_for_contract_tests(client, vec![metric], 1);
310
311        let artifacts = runner.run_suite(&cfg, None).await?;
312        let row = artifacts
313            .results
314            .iter()
315            .find(|r| r.test_id == "t1")
316            .expect("result for t1");
317
318        assert_eq!(row.status, TestStatus::Flaky);
319        assert_eq!(row.message, "flake detected (rerun passed)");
320        let attempts = row.attempts.as_ref().expect("attempts");
321        assert_eq!(attempts.len(), 2);
322        assert_eq!(attempts[0].status, TestStatus::Fail);
323        assert_eq!(attempts[1].status, TestStatus::Pass);
324        Ok(())
325    }
326
327    #[tokio::test]
328    async fn runner_contract_fail_after_retries_stays_fail() -> anyhow::Result<()> {
329        let cfg = single_test_config(ErrorPolicy::Block);
330        let client = Arc::new(FakeClient::new("fake-model".to_string()).with_response("ok".into()));
331        let metric = Arc::new(ScriptedMetric::always_fail());
332        let runner = runner_for_contract_tests(client, vec![metric], 1);
333
334        let artifacts = runner.run_suite(&cfg, None).await?;
335        let row = artifacts
336            .results
337            .iter()
338            .find(|r| r.test_id == "t1")
339            .expect("result for t1");
340
341        assert_eq!(row.status, TestStatus::Fail);
342        assert!(
343            row.message.contains("failed: scripted"),
344            "expected stable failure reason, got: {}",
345            row.message
346        );
347        let attempts = row.attempts.as_ref().expect("attempts");
348        assert_eq!(attempts.len(), 2);
349        assert_eq!(attempts[0].status, TestStatus::Fail);
350        assert_eq!(attempts[1].status, TestStatus::Fail);
351        Ok(())
352    }
353
354    #[tokio::test]
355    async fn runner_contract_on_error_allow_marks_allowed_and_policy_applied() -> anyhow::Result<()>
356    {
357        let cfg = single_test_config(ErrorPolicy::Allow);
358        let client = Arc::new(ErrorClient);
359        let runner = runner_for_contract_tests(client, vec![], 2);
360
361        let artifacts = runner.run_suite(&cfg, None).await?;
362        let row = artifacts
363            .results
364            .iter()
365            .find(|r| r.test_id == "t1")
366            .expect("result for t1");
367
368        assert_eq!(row.status, TestStatus::AllowedOnError);
369        assert_eq!(row.error_policy_applied, Some(ErrorPolicy::Allow));
370        assert_eq!(row.details["policy_applied"], serde_json::json!("allow"));
371        let attempts = row.attempts.as_ref().expect("attempts");
372        assert_eq!(attempts.len(), 1);
373        assert_eq!(attempts[0].status, TestStatus::AllowedOnError);
374        Ok(())
375    }
376
377    #[tokio::test]
378    async fn runner_contract_results_sorted_by_test_id() -> anyhow::Result<()> {
379        let mut cfg = config_with_test_ids(&["t3", "t1", "t2"], ErrorPolicy::Block);
380        cfg.settings.parallel = Some(3);
381        let client = Arc::new(FakeClient::new("fake-model".to_string()).with_response("ok".into()));
382        let metric = Arc::new(ScriptedMetric::always_pass());
383        let runner = runner_for_contract_tests(client, vec![metric], 0);
384
385        let artifacts = runner.run_suite(&cfg, None).await?;
386        let ids: Vec<_> = artifacts
387            .results
388            .iter()
389            .map(|r| r.test_id.as_str())
390            .collect();
391        assert_eq!(ids, vec!["t1", "t2", "t3"]);
392        Ok(())
393    }
394
395    #[tokio::test]
396    async fn runner_contract_progress_sink_reports_done_total() -> anyhow::Result<()> {
397        let cfg = config_with_test_ids(&["p1", "p2", "p3"], ErrorPolicy::Block);
398        let client = Arc::new(FakeClient::new("fake-model".to_string()).with_response("ok".into()));
399        let metric = Arc::new(ScriptedMetric::always_pass());
400        let runner = runner_for_contract_tests(client, vec![metric], 0);
401
402        let events = Arc::new(std::sync::Mutex::new(Vec::<(usize, usize)>::new()));
403        let sink = {
404            let events = Arc::clone(&events);
405            Arc::new(move |ev: crate::report::progress::ProgressEvent| {
406                events
407                    .lock()
408                    .expect("progress lock")
409                    .push((ev.done, ev.total));
410            }) as crate::report::progress::ProgressSink
411        };
412
413        let artifacts = runner.run_suite(&cfg, Some(sink)).await?;
414        assert_eq!(artifacts.results.len(), 3);
415
416        let observed = events.lock().expect("progress lock");
417        assert_eq!(observed.len(), 3);
418        assert_eq!(observed.last(), Some(&(3, 3)));
419        assert!(observed.windows(2).all(|w| w[0].0 < w[1].0));
420        Ok(())
421    }
422
423    #[tokio::test]
424    async fn runner_contract_relative_baseline_missing_warns_in_helper() -> anyhow::Result<()> {
425        let mut cfg = single_test_config(ErrorPolicy::Block);
426        cfg.settings.thresholding = Some(crate::model::ThresholdingSettings {
427            mode: Some("relative".to_string()),
428            max_drop: Some(0.05),
429            min_floor: None,
430        });
431
432        let client = Arc::new(FakeClient::new("fake-model".to_string()).with_response("ok".into()));
433        let metric = Arc::new(ScriptedMetric::always_pass());
434        let runner = runner_for_contract_tests(client, vec![], 0);
435        let baseline = crate::baseline::Baseline {
436            schema_version: 1,
437            suite: "runner-contract".to_string(),
438            assay_version: env!("CARGO_PKG_VERSION").to_string(),
439            created_at: "2026-01-01T00:00:00Z".to_string(),
440            config_fingerprint: "md5:test".to_string(),
441            git_info: None,
442            entries: vec![],
443        };
444        let tc = cfg.tests.first().cloned().expect("single test case");
445        let details = serde_json::json!({
446            "metrics": {
447                "scripted": {
448                    "score": 1.0,
449                    "passed": true,
450                    "unstable": false,
451                    "details": {}
452                }
453            }
454        });
455
456        let verdict = runner.check_baseline_regressions(&tc, &cfg, &details, &[metric], &baseline);
457        let (status, message) = verdict.expect("relative baseline decision");
458        assert_eq!(status, TestStatus::Warn);
459        assert_eq!(message, "missing baseline for t1/scripted");
460        Ok(())
461    }
462}