Skip to main content

somatize_runtime/executors/
study.rs

1//! Study runner — orchestrates hyperparameter optimization.
2//!
3//! Iterates over trials: samples parameters, calls the executor with a
4//! [`TrialContext`] handle, records metrics, feeds results back to the
5//! sampler (ask/tell), consults the pruner on intermediate metrics,
6//! and persists the study through an optional tracker after every
7//! trial. Resume-aware: a study with N recorded trials continues at
8//! trial N.
9
10use crate::event_bus::EventBus;
11use crate::pruner::{MedianPruner, PercentilePruner, Pruner, TrialMetricHistory};
12use crate::sampler::Sampler;
13use chrono::Utc;
14use somatize_core::error::Result;
15use somatize_core::event::{Event, MetricRecord};
16use somatize_core::study::{Direction, PruningStrategy, Study, Trial, TrialState};
17use somatize_core::tracking::Tracker;
18use std::collections::HashMap;
19use std::sync::{Arc, Mutex};
20use std::time::Instant;
21
22/// Result of executing a trial. Separates control flow (pruning) from errors.
23#[derive(Debug, Clone)]
24pub enum TrialOutcome {
25    /// Trial completed successfully with final metrics.
26    Completed(Vec<MetricRecord>),
27    /// Trial was pruned (stopped early) at the given step.
28    Pruned {
29        /// Step at which the pruner stopped the trial.
30        step: usize,
31        /// The pruner's explanation, recorded on the trial.
32        reason: String,
33    },
34}
35
36/// Handle passed to a trial: reports intermediate metrics and asks
37/// whether the pruner wants the trial stopped.
38///
39/// [`report`](Self::report) is the single channel for per-step metrics:
40/// it records the value on the trial, emits [`Event::TrialMetric`], and
41/// — when the metric is the study's objective — consults the pruner
42/// against the completed trials' histories. Values are compared on a
43/// maximize scale (the runner pre-normalizes for `Minimize`).
44///
45/// The handle is cheaply cloneable (`Arc`-backed shared state) so it
46/// can outlive the borrow stack — e.g. cross into a Python callback.
47#[derive(Clone)]
48pub struct TrialContext {
49    study_id: String,
50    trial_id: String,
51    /// Objective metric name + direction used for pruning decisions.
52    objective: Option<(String, Direction)>,
53    pruner: Option<Arc<dyn Pruner>>,
54    /// Completed trials' metric histories, direction-normalized.
55    history: Arc<Vec<TrialMetricHistory>>,
56    event_bus: Arc<EventBus>,
57    shared: Arc<Mutex<TrialShared>>,
58}
59
60#[derive(Default)]
61struct TrialShared {
62    metrics: Vec<MetricRecord>,
63    pruned: Option<(usize, String)>,
64}
65
66impl TrialContext {
67    /// Record an intermediate metric at `step`. Returns `true` when the
68    /// trial should stop (pruned) — the executor should then return
69    /// early; the runner marks the trial pruned regardless.
70    pub fn report(&self, name: &str, value: f64, step: usize) -> bool {
71        let record = MetricRecord {
72            name: name.to_string(),
73            value,
74            step,
75            timestamp: Utc::now(),
76        };
77        {
78            let mut shared = self.lock_shared();
79            shared.metrics.push(record.clone());
80        }
81        self.event_bus.emit(Event::TrialMetric {
82            study_id: self.study_id.clone(),
83            trial_id: self.trial_id.clone(),
84            metric: record,
85        });
86
87        if self.should_prune() {
88            return true;
89        }
90        if let (Some(pruner), Some((obj_name, direction))) = (&self.pruner, &self.objective)
91            && name == obj_name
92            && let Some(reason) =
93                pruner.should_prune(obj_name, direction.normalize(value), step, &self.history)
94        {
95            self.lock_shared().pruned = Some((step, reason));
96            return true;
97        }
98        false
99    }
100
101    /// Whether the pruner has decided to stop this trial.
102    pub fn should_prune(&self) -> bool {
103        self.lock_shared().pruned.is_some()
104    }
105
106    /// Metrics reported so far.
107    pub fn metrics(&self) -> Vec<MetricRecord> {
108        self.lock_shared().metrics.clone()
109    }
110
111    /// Id of the trial this handle belongs to.
112    pub fn trial_id(&self) -> &str {
113        &self.trial_id
114    }
115
116    fn lock_shared(&self) -> std::sync::MutexGuard<'_, TrialShared> {
117        match self.shared.lock() {
118            Ok(guard) => guard,
119            Err(poisoned) => poisoned.into_inner(),
120        }
121    }
122
123    fn take_results(&self) -> (Vec<MetricRecord>, Option<(usize, String)>) {
124        let mut shared = self.lock_shared();
125        (std::mem::take(&mut shared.metrics), shared.pruned.take())
126    }
127}
128
129/// Callback that executes a trial given sampled parameters.
130///
131/// Returns `Ok(TrialOutcome)` for normal completion or pruning,
132/// `Err(SomaError)` only for unexpected failures.
133pub trait TrialExecutor: Send + Sync {
134    /// Run one trial with the sampled `params`, reporting intermediate
135    /// metrics through `ctx` and honouring its pruning verdicts.
136    fn execute_trial(
137        &self,
138        params: &HashMap<String, serde_json::Value>,
139        ctx: &TrialContext,
140    ) -> Result<TrialOutcome>;
141}
142
143/// Function-based trial executor for convenience.
144pub struct FnTrialExecutor<F>(pub F);
145
146impl<F> TrialExecutor for FnTrialExecutor<F>
147where
148    F: Fn(&HashMap<String, serde_json::Value>, &TrialContext) -> Result<TrialOutcome> + Send + Sync,
149{
150    fn execute_trial(
151        &self,
152        params: &HashMap<String, serde_json::Value>,
153        ctx: &TrialContext,
154    ) -> Result<TrialOutcome> {
155        (self.0)(params, ctx)
156    }
157}
158
159/// Runs a Study: samples parameters, executes trials, records results.
160pub struct StudyRunner {
161    event_bus: Arc<EventBus>,
162    tracker: Option<Arc<dyn Tracker>>,
163}
164
165impl StudyRunner {
166    /// A runner emitting trial events on `event_bus`, with no persistence
167    /// until [`Self::with_tracker`] adds it.
168    pub fn new(event_bus: Arc<EventBus>) -> Self {
169        Self {
170            event_bus,
171            tracker: None,
172        }
173    }
174
175    /// Persist the study (`study.json`, atomic) after every trial.
176    pub fn with_tracker(mut self, tracker: Arc<dyn Tracker>) -> Self {
177        self.tracker = Some(tracker);
178        self
179    }
180
181    /// Run the study to completion.
182    ///
183    /// Resume-aware: starts at `trial_index = study.trials.len()` and
184    /// replays already-completed trials into the sampler's history
185    /// first, so model-based samplers continue informed and grids don't
186    /// repeat configurations.
187    pub fn run(
188        &self,
189        study: &mut Study,
190        sampler: &mut dyn Sampler,
191        executor: &dyn TrialExecutor,
192    ) -> Result<()> {
193        sampler.prepare(&study.search_space);
194        // With experiment seeds, every sampled config runs once per seed.
195        let total = sampler.n_trials().unwrap_or(0) * study.seeds.len().max(1);
196        if total > 0 {
197            study.planned_trials = Some(total);
198        }
199        let direction = study.primary_direction().unwrap_or(Direction::Maximize);
200        if study.created_at.is_none() {
201            study.created_at = Some(Utc::now());
202        }
203
204        // Replay prior trials (resume) into the sampler's history.
205        for trial in &study.trials {
206            if let Some(value) = study.objective_value(trial) {
207                sampler.record_result(&trial.params, direction.normalize(value));
208            }
209        }
210
211        self.event_bus.emit(Event::StudyStarted {
212            study_id: study.id.clone(),
213            name: study.name.clone(),
214            total_trials: total,
215        });
216
217        let pruner = build_pruner(&study.pruning);
218        let mut trial_index = study.trials.len();
219
220        // Experiment seeds: each sampled configuration runs once per
221        // seed (params carry "seed"), so every seed is an independent,
222        // resumable trial with its own cache line. trial_index enumerates
223        // config-major: config 0 × all seeds, config 1 × all seeds, …
224        let seeds = study.seeds.clone();
225        let n_seeds = seeds.len().max(1);
226        let mut current_config: Option<HashMap<String, serde_json::Value>> = None;
227
228        loop {
229            let config_index = trial_index / n_seeds;
230            let seed_slot = trial_index % n_seeds;
231
232            let base = if seed_slot == 0 || current_config.is_none() {
233                if seed_slot > 0 {
234                    // Resuming mid-seed-block: recover the block's config
235                    // from the previous (persisted) trial.
236                    let mut prev = study.trials[trial_index - 1].params.clone();
237                    prev.remove("seed");
238                    Some(prev)
239                } else {
240                    sampler.sample(&study.search_space, config_index)?
241                }
242            } else {
243                current_config.clone()
244            };
245            let Some(base) = base else { break };
246            current_config = Some(base.clone());
247
248            let mut params = base;
249            if !seeds.is_empty() {
250                params.insert("seed".to_string(), serde_json::json!(seeds[seed_slot]));
251            }
252            // Frozen parameters are fixed values excluded from the
253            // search space — inject them into every configuration.
254            for (name, value) in &study.frozen {
255                params.insert(name.clone(), value.clone());
256            }
257
258            let trial_id = format!("trial_{trial_index:04}");
259            let mut trial = Trial::new(trial_id.clone(), params.clone());
260            trial.state = TrialState::Running;
261            trial.started_at = Some(Utc::now());
262
263            self.event_bus.emit(Event::TrialStarted {
264                study_id: study.id.clone(),
265                trial_id: trial_id.clone(),
266                params: serde_json::json!(params),
267            });
268
269            // Histories the pruner compares against, on a maximize scale.
270            let ctx = TrialContext {
271                study_id: study.id.clone(),
272                trial_id: trial_id.clone(),
273                objective: objective_metric(study).map(|name| (name, direction)),
274                pruner: pruner.clone(),
275                history: Arc::new(normalized_histories(study, direction)),
276                event_bus: self.event_bus.clone(),
277                shared: Arc::new(Mutex::new(TrialShared::default())),
278            };
279
280            let start = Instant::now();
281            let outcome = executor.execute_trial(&params, &ctx);
282            let (reported, pruned) = ctx.take_results();
283            trial.duration_ms = Some(start.elapsed().as_millis() as u64);
284            trial.finished_at = Some(Utc::now());
285            trial.metrics = reported;
286
287            match (outcome, pruned) {
288                // The pruner's verdict wins over a Completed return —
289                // an executor may not notice report() returned true.
290                (Ok(_), Some((step, reason)))
291                | (Ok(TrialOutcome::Pruned { step, reason }), None) => {
292                    trial.state = TrialState::Pruned {
293                        step,
294                        reason: reason.clone(),
295                    };
296                    self.event_bus.emit(Event::TrialPruned {
297                        study_id: study.id.clone(),
298                        trial_id: trial_id.clone(),
299                        step,
300                        reason,
301                    });
302                }
303                (Ok(TrialOutcome::Completed(final_metrics)), None) => {
304                    for metric in &final_metrics {
305                        self.event_bus.emit(Event::TrialMetric {
306                            study_id: study.id.clone(),
307                            trial_id: trial_id.clone(),
308                            metric: metric.clone(),
309                        });
310                    }
311                    trial.metrics.extend(final_metrics.clone());
312                    trial.state = TrialState::Completed;
313
314                    self.event_bus.emit(Event::TrialCompleted {
315                        study_id: study.id.clone(),
316                        trial_id: trial_id.clone(),
317                        final_metrics,
318                    });
319                }
320                (Err(e), _) => {
321                    trial.state = TrialState::Failed {
322                        error: e.to_string(),
323                    };
324                    self.event_bus.emit(Event::TrialFailed {
325                        study_id: study.id.clone(),
326                        trial_id: trial_id.clone(),
327                        error: e.to_string(),
328                    });
329                }
330            }
331
332            study.trials.push(trial);
333
334            // Tell the sampler (ask/tell feedback for TPE and future BO).
335            if let Some(value) = study.objective_value(study.trials.last().unwrap()) {
336                sampler.record_result(&params, direction.normalize(value));
337            }
338
339            // Check if we have a new best
340            if let Some(best) = study.best_trial()
341                && best.id == trial_id
342                && let Some(val) = study.best_value()
343            {
344                self.event_bus.emit(Event::BestUpdated {
345                    study_id: study.id.clone(),
346                    trial_id: trial_id.clone(),
347                    value: val,
348                    params: serde_json::json!(params),
349                });
350            }
351
352            let completed = study.trials.iter().filter(|t| t.is_terminal()).count();
353            self.event_bus.emit(Event::StudyProgress {
354                study_id: study.id.clone(),
355                completed,
356                total,
357                best_value: study.best_value().unwrap_or(f64::NAN),
358            });
359
360            study.updated_at = Some(Utc::now());
361            self.save_study(study);
362            trial_index += 1;
363        }
364
365        let best_trial_id = study.best_trial().map(|t| t.id.clone()).unwrap_or_default();
366        let best_value = study.best_value().unwrap_or(f64::NAN);
367
368        self.event_bus.emit(Event::StudyCompleted {
369            study_id: study.id.clone(),
370            best_trial_id,
371            best_value,
372        });
373        study.updated_at = Some(Utc::now());
374        self.save_study(study);
375        // Make the event log durable before returning — a caller may
376        // read the run directory immediately after run() completes.
377        self.event_bus.flush_sinks();
378
379        Ok(())
380    }
381
382    fn save_study(&self, study: &Study) {
383        if let Some(tracker) = &self.tracker
384            && let Err(e) = tracker.save_study(study)
385        {
386            tracing::warn!("tracking: failed to persist study: {e}");
387        }
388    }
389}
390
391/// Metric name the pruner watches: the composite's recorded name is not
392/// a raw metric, so pruning tracks the first declared objective (or the
393/// first composite term as a fallback).
394fn objective_metric(study: &Study) -> Option<String> {
395    study
396        .objectives
397        .first()
398        .map(|o| o.metric.clone())
399        .or_else(|| {
400            study
401                .composite
402                .as_ref()
403                .and_then(|c| c.terms.first().map(|(name, _)| name.clone()))
404        })
405}
406
407fn build_pruner(strategy: &PruningStrategy) -> Option<Arc<dyn Pruner>> {
408    match strategy {
409        PruningStrategy::None | PruningStrategy::Hyperband => None,
410        PruningStrategy::Median { n_warmup_steps } => {
411            Some(Arc::new(MedianPruner::new(*n_warmup_steps)))
412        }
413        PruningStrategy::Percentile {
414            percentile,
415            n_warmup_steps,
416        } => Some(Arc::new(PercentilePruner::new(
417            *percentile,
418            *n_warmup_steps,
419        ))),
420    }
421}
422
423/// Completed trials' metric histories with values mapped onto a
424/// maximize scale, so pruners can always assume higher-is-better.
425fn normalized_histories(study: &Study, direction: Direction) -> Vec<TrialMetricHistory> {
426    study
427        .trials
428        .iter()
429        .filter(|t| t.is_complete())
430        .map(|t| TrialMetricHistory {
431            trial_id: t.id.clone(),
432            metrics: t
433                .metrics
434                .iter()
435                .map(|m| MetricRecord {
436                    name: m.name.clone(),
437                    value: direction.normalize(m.value),
438                    step: m.step,
439                    timestamp: m.timestamp,
440                })
441                .collect(),
442        })
443        .collect()
444}
445
446#[cfg(test)]
447mod tests {
448    use super::*;
449    use crate::sampler::{BayesianSampler, GridSampler, RandomSampler};
450    use crate::study_io::StudyIo;
451    use chrono::Utc;
452    use somatize_core::error::SomaError;
453    use somatize_core::search::{Scale, SearchDimension, SearchSpace};
454    use somatize_core::study::{Direction, Objective, SearchStrategy};
455
456    fn sample_space() -> SearchSpace {
457        let mut space = SearchSpace::new();
458        space.add(SearchDimension::Float {
459            name: "lr".into(),
460            low: 0.001,
461            high: 0.1,
462            scale: Scale::Log,
463            default: None,
464        });
465        space.add(SearchDimension::Categorical {
466            name: "activation".into(),
467            choices: vec![serde_json::json!("relu"), serde_json::json!("tanh")],
468        });
469        space
470    }
471
472    /// Simple executor: f1 = 1.0 - |lr - 0.01| * 10
473    fn make_executor() -> FnTrialExecutor<
474        impl Fn(&HashMap<String, serde_json::Value>, &TrialContext) -> Result<TrialOutcome>,
475    > {
476        FnTrialExecutor(
477            |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
478                let lr = params["lr"].as_f64().unwrap();
479                let f1 = (1.0 - (lr - 0.01).abs() * 10.0).max(0.0);
480                Ok(TrialOutcome::Completed(vec![MetricRecord {
481                    name: "f1".into(),
482                    value: f1,
483                    step: 0,
484                    timestamp: Utc::now(),
485                }]))
486            },
487        )
488    }
489
490    #[test]
491    fn study_runner_grid_search() {
492        let bus = Arc::new(EventBus::new(256));
493        let mut rx = bus.subscribe();
494        let runner = StudyRunner::new(bus);
495
496        let space = sample_space();
497        let mut study = Study::new(
498            "grid_test",
499            space,
500            SearchStrategy::Grid { points_per_dim: 3 },
501            vec![Objective {
502                metric: "f1".into(),
503                direction: Direction::Maximize,
504            }],
505        );
506
507        let mut sampler = GridSampler::new(3);
508        let executor = make_executor();
509
510        runner.run(&mut study, &mut sampler, &executor).unwrap();
511
512        // 3 lr points * 2 activations = 6 trials
513        assert_eq!(study.trials.len(), 6);
514        assert!(study.trials.iter().all(|t| t.is_complete()));
515
516        // Best trial should have lr closest to 0.01
517        let best = study.best_trial().unwrap();
518        let best_lr = best.params["lr"].as_f64().unwrap();
519        assert!(
520            (best_lr - 0.01).abs() < 0.05,
521            "best lr should be near 0.01, got {best_lr}"
522        );
523
524        // Check events were emitted
525        let mut events = Vec::new();
526        while let Ok(e) = rx.try_recv() {
527            events.push(e);
528        }
529        assert!(
530            events
531                .iter()
532                .any(|e| matches!(e, Event::StudyStarted { .. }))
533        );
534        assert!(
535            events
536                .iter()
537                .any(|e| matches!(e, Event::TrialStarted { .. }))
538        );
539        assert!(
540            events
541                .iter()
542                .any(|e| matches!(e, Event::TrialCompleted { .. }))
543        );
544        assert!(
545            events
546                .iter()
547                .any(|e| matches!(e, Event::BestUpdated { .. }))
548        );
549        assert!(
550            events
551                .iter()
552                .any(|e| matches!(e, Event::StudyCompleted { .. }))
553        );
554    }
555
556    #[test]
557    fn study_runner_random_search() {
558        let bus = Arc::new(EventBus::new(256));
559        let runner = StudyRunner::new(bus);
560
561        let space = sample_space();
562        let mut study = Study::new(
563            "random_test",
564            space,
565            SearchStrategy::Random {
566                n_trials: 20,
567                seed: Some(42),
568            },
569            vec![Objective {
570                metric: "f1".into(),
571                direction: Direction::Maximize,
572            }],
573        );
574
575        let mut sampler = RandomSampler::new(20, Some(42));
576        let executor = make_executor();
577
578        runner.run(&mut study, &mut sampler, &executor).unwrap();
579
580        assert_eq!(study.trials.len(), 20);
581        assert!(study.best_trial().is_some());
582    }
583
584    #[test]
585    fn study_runner_handles_failed_trials() {
586        let bus = Arc::new(EventBus::new(256));
587        let runner = StudyRunner::new(bus);
588
589        let mut space = SearchSpace::new();
590        space.add(SearchDimension::Float {
591            name: "x".into(),
592            low: 0.0,
593            high: 1.0,
594            scale: Scale::Linear,
595            default: None,
596        });
597
598        let mut study = Study::new(
599            "fail_test",
600            space,
601            SearchStrategy::Random {
602                n_trials: 5,
603                seed: None,
604            },
605            vec![Objective {
606                metric: "f1".into(),
607                direction: Direction::Maximize,
608            }],
609        );
610
611        // Executor that fails on even trials
612        let executor = FnTrialExecutor(
613            |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
614                let x = params["x"].as_f64().unwrap();
615                if x > 0.5 {
616                    Err(SomaError::Other("too high".into()))
617                } else {
618                    Ok(TrialOutcome::Completed(vec![MetricRecord {
619                        name: "f1".into(),
620                        value: x,
621                        step: 0,
622                        timestamp: Utc::now(),
623                    }]))
624                }
625            },
626        );
627
628        let mut sampler = RandomSampler::new(5, Some(42));
629        runner.run(&mut study, &mut sampler, &executor).unwrap();
630
631        assert_eq!(study.trials.len(), 5);
632        // Some should be Failed
633        let failed = study
634            .trials
635            .iter()
636            .filter(|t| matches!(t.state, TrialState::Failed { .. }))
637            .count();
638        assert!(failed > 0, "should have some failed trials");
639    }
640
641    #[test]
642    fn study_runner_handles_pruned_trials() {
643        let bus = Arc::new(EventBus::new(256));
644        let runner = StudyRunner::new(bus);
645
646        let mut space = SearchSpace::new();
647        space.add(SearchDimension::Float {
648            name: "x".into(),
649            low: 0.0,
650            high: 1.0,
651            scale: Scale::Linear,
652            default: None,
653        });
654
655        let mut study = Study::new(
656            "prune_test",
657            space,
658            SearchStrategy::Random {
659                n_trials: 3,
660                seed: None,
661            },
662            vec![Objective {
663                metric: "f1".into(),
664                direction: Direction::Maximize,
665            }],
666        );
667
668        // Executor that prunes every trial
669        let executor = FnTrialExecutor(
670            |_params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
671                Ok(TrialOutcome::Pruned {
672                    step: 5,
673                    reason: "below median".into(),
674                })
675            },
676        );
677
678        let mut sampler = RandomSampler::new(3, Some(42));
679        runner.run(&mut study, &mut sampler, &executor).unwrap();
680
681        assert!(
682            study
683                .trials
684                .iter()
685                .all(|t| matches!(t.state, TrialState::Pruned { .. }))
686        );
687    }
688
689    fn one_dim_space() -> SearchSpace {
690        let mut space = SearchSpace::new();
691        space.add(SearchDimension::Float {
692            name: "x".into(),
693            low: 0.0,
694            high: 1.0,
695            scale: Scale::Linear,
696            default: None,
697        });
698        space
699    }
700
701    #[test]
702    fn grid_study_started_reports_real_total() {
703        let bus = Arc::new(EventBus::new(256));
704        let mut rx = bus.subscribe();
705        let runner = StudyRunner::new(bus);
706
707        let mut study = Study::new(
708            "grid_total",
709            sample_space(),
710            SearchStrategy::Grid { points_per_dim: 3 },
711            vec![Objective {
712                metric: "f1".into(),
713                direction: Direction::Maximize,
714            }],
715        );
716        let mut sampler = GridSampler::new(3);
717        runner
718            .run(&mut study, &mut sampler, &make_executor())
719            .unwrap();
720
721        let mut started_total = None;
722        while let Ok(e) = rx.try_recv() {
723            if let Event::StudyStarted { total_trials, .. } = e {
724                started_total = Some(total_trials);
725            }
726        }
727        // 3 lr points × 2 activations — known BEFORE the first sample.
728        assert_eq!(started_total, Some(6));
729    }
730
731    /// Sampler spy that counts record_result calls.
732    struct SpySampler {
733        inner: RandomSampler,
734        recorded: Vec<f64>,
735    }
736
737    impl Sampler for SpySampler {
738        fn sample(
739            &mut self,
740            space: &SearchSpace,
741            trial_index: usize,
742        ) -> Result<Option<HashMap<String, serde_json::Value>>> {
743            self.inner.sample(space, trial_index)
744        }
745        fn n_trials(&self) -> Option<usize> {
746            self.inner.n_trials()
747        }
748        fn record_result(&mut self, _params: &HashMap<String, serde_json::Value>, value: f64) {
749            self.recorded.push(value);
750        }
751    }
752
753    #[test]
754    fn sampler_receives_feedback_per_completed_trial() {
755        let bus = Arc::new(EventBus::new(256));
756        let runner = StudyRunner::new(bus);
757
758        let mut study = Study::new(
759            "feedback",
760            one_dim_space(),
761            SearchStrategy::Random {
762                n_trials: 4,
763                seed: Some(1),
764            },
765            vec![Objective {
766                metric: "loss".into(),
767                direction: Direction::Minimize,
768            }],
769        );
770        let executor = FnTrialExecutor(
771            |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
772                Ok(TrialOutcome::Completed(vec![MetricRecord {
773                    name: "loss".into(),
774                    value: params["x"].as_f64().unwrap(),
775                    step: 0,
776                    timestamp: Utc::now(),
777                }]))
778            },
779        );
780
781        let mut sampler = SpySampler {
782            inner: RandomSampler::new(4, Some(1)),
783            recorded: Vec::new(),
784        };
785        runner.run(&mut study, &mut sampler, &executor).unwrap();
786
787        assert_eq!(
788            sampler.recorded.len(),
789            4,
790            "one feedback per completed trial"
791        );
792        // Minimize → values arrive negated (maximize scale).
793        assert!(sampler.recorded.iter().all(|v| *v <= 0.0));
794    }
795
796    fn pruning_study(direction: Direction) -> Study {
797        let metric = match direction {
798            Direction::Maximize => "f1",
799            Direction::Minimize => "loss",
800        };
801        Study::new(
802            "pruning",
803            one_dim_space(),
804            SearchStrategy::Random {
805                n_trials: 4,
806                seed: Some(7),
807            },
808            vec![Objective {
809                metric: metric.into(),
810                direction,
811            }],
812        )
813        .with_pruning(PruningStrategy::Median { n_warmup_steps: 2 })
814    }
815
816    /// First trial completes with a good curve; every later (bad)
817    /// trial must get pruned mid-way once it reports below the median.
818    fn run_pruning_case(direction: Direction) -> Study {
819        use std::sync::atomic::{AtomicUsize, Ordering};
820
821        let bus = Arc::new(EventBus::new(256));
822        let runner = StudyRunner::new(bus);
823        let mut study = pruning_study(direction);
824        let counter = AtomicUsize::new(0);
825
826        let executor = FnTrialExecutor(
827            move |_params: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
828                let good = counter.fetch_add(1, Ordering::SeqCst) == 0;
829                let metric = match direction {
830                    Direction::Maximize => "f1",
831                    Direction::Minimize => "loss",
832                };
833                for step in 0..10 {
834                    let value = match (direction, good) {
835                        (Direction::Maximize, true) => 0.5 + step as f64 * 0.05,
836                        (Direction::Maximize, false) => 0.01,
837                        (Direction::Minimize, true) => 1.0 - step as f64 * 0.05,
838                        (Direction::Minimize, false) => 10.0,
839                    };
840                    if ctx.report(metric, value, step) {
841                        return Ok(TrialOutcome::Pruned {
842                            step,
843                            reason: "stopped by pruner".into(),
844                        });
845                    }
846                }
847                Ok(TrialOutcome::Completed(vec![]))
848            },
849        );
850
851        let mut sampler = RandomSampler::new(4, Some(7));
852        runner.run(&mut study, &mut sampler, &executor).unwrap();
853        study
854    }
855
856    fn assert_bad_trials_pruned(study: &Study) {
857        let pruned: Vec<_> = study
858            .trials
859            .iter()
860            .filter(|t| matches!(t.state, TrialState::Pruned { .. }))
861            .collect();
862        assert_eq!(pruned.len(), 3, "all trials after the first get pruned");
863        for t in pruned {
864            if let TrialState::Pruned { step, .. } = &t.state {
865                assert_eq!(*step, 2, "pruned right after warmup, not at the end");
866            }
867        }
868        assert!(study.trials[0].is_complete());
869    }
870
871    #[test]
872    fn median_pruner_stops_bad_trials_maximize() {
873        assert_bad_trials_pruned(&run_pruning_case(Direction::Maximize));
874    }
875
876    #[test]
877    fn median_pruner_stops_bad_trials_minimize() {
878        // Direction normalization: a HIGH loss must read as "bad".
879        assert_bad_trials_pruned(&run_pruning_case(Direction::Minimize));
880    }
881
882    #[test]
883    fn resume_continues_without_repeating_grid_params() {
884        let executor = make_executor();
885
886        // Full reference run: 6 grid trials.
887        let bus = Arc::new(EventBus::new(256));
888        let runner = StudyRunner::new(bus);
889        let mut full = Study::new(
890            "full",
891            sample_space(),
892            SearchStrategy::Grid { points_per_dim: 3 },
893            vec![Objective {
894                metric: "f1".into(),
895                direction: Direction::Maximize,
896            }],
897        );
898        runner
899            .run(&mut full, &mut GridSampler::new(3), &executor)
900            .unwrap();
901        assert_eq!(full.trials.len(), 6);
902
903        // Interrupted run: only the first 3 trials happened.
904        let mut resumed = full.clone();
905        resumed.trials.truncate(3);
906
907        let bus = Arc::new(EventBus::new(256));
908        let mut rx = bus.subscribe();
909        let runner = StudyRunner::new(bus);
910        runner
911            .run(&mut resumed, &mut GridSampler::new(3), &executor)
912            .unwrap();
913
914        assert_eq!(resumed.trials.len(), 6);
915        // The resumed half reproduces exactly the reference tail — no
916        // repeated or skipped configurations, ids continue.
917        for i in 3..6 {
918            assert_eq!(resumed.trials[i].params, full.trials[i].params);
919            assert_eq!(resumed.trials[i].id, format!("trial_{i:04}"));
920        }
921        // Only 3 TrialStarted events in the resumed run.
922        let mut started = 0;
923        while let Ok(e) = rx.try_recv() {
924            if matches!(e, Event::TrialStarted { .. }) {
925                started += 1;
926            }
927        }
928        assert_eq!(started, 3);
929    }
930
931    #[test]
932    fn frozen_params_reach_every_trial() {
933        let bus = Arc::new(EventBus::new(256));
934        let runner = StudyRunner::new(bus);
935
936        let mut study = Study::new(
937            "frozen",
938            one_dim_space(),
939            SearchStrategy::Random {
940                n_trials: 3,
941                seed: Some(2),
942            },
943            vec![Objective {
944                metric: "f1".into(),
945                direction: Direction::Maximize,
946            }],
947        );
948        study
949            .frozen
950            .insert("batch_size".into(), serde_json::json!(64));
951
952        let executor = FnTrialExecutor(
953            |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
954                assert_eq!(params["batch_size"], serde_json::json!(64));
955                Ok(TrialOutcome::Completed(vec![MetricRecord {
956                    name: "f1".into(),
957                    value: 0.5,
958                    step: 0,
959                    timestamp: Utc::now(),
960                }]))
961            },
962        );
963        let mut sampler = RandomSampler::new(3, Some(2));
964        runner.run(&mut study, &mut sampler, &executor).unwrap();
965
966        for t in &study.trials {
967            assert_eq!(t.params["batch_size"], serde_json::json!(64));
968        }
969    }
970
971    #[test]
972    fn composite_objective_selects_best_trial() {
973        use somatize_core::study::{CompositeObjective, Scalarizer};
974
975        let bus = Arc::new(EventBus::new(256));
976        let runner = StudyRunner::new(bus);
977
978        let mut study = Study::new(
979            "composite",
980            one_dim_space(),
981            SearchStrategy::Grid { points_per_dim: 5 },
982            vec![],
983        )
984        .with_composite(CompositeObjective {
985            // Maximize x while penalizing x² → optimum at x = 0.5.
986            terms: vec![("x".into(), 1.0), ("x_sq".into(), -1.0)],
987            direction: Direction::Maximize,
988            scalarizer: Scalarizer::WeightedSum,
989        });
990
991        let executor = FnTrialExecutor(
992            |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
993                let x = params["x"].as_f64().unwrap();
994                let now = Utc::now();
995                Ok(TrialOutcome::Completed(vec![
996                    MetricRecord {
997                        name: "x".into(),
998                        value: x,
999                        step: 0,
1000                        timestamp: now,
1001                    },
1002                    MetricRecord {
1003                        name: "x_sq".into(),
1004                        value: x * x,
1005                        step: 0,
1006                        timestamp: now,
1007                    },
1008                ]))
1009            },
1010        );
1011        let mut sampler = GridSampler::new(5);
1012        runner.run(&mut study, &mut sampler, &executor).unwrap();
1013
1014        let best_x = study.best_trial().unwrap().params["x"].as_f64().unwrap();
1015        assert!(
1016            (best_x - 0.5).abs() < 1e-9,
1017            "optimum of x - x² is 0.5, got {best_x}"
1018        );
1019    }
1020
1021    #[test]
1022    fn tracker_persists_study_after_every_trial() {
1023        use somatize_core::tracking::{RunKind, Tracker as _};
1024
1025        let root = tempfile::tempdir().unwrap();
1026        let tracker = Arc::new(
1027            crate::tracking::LocalTracker::create(root.path(), RunKind::Study, "t").unwrap(),
1028        );
1029        let run_dir = tracker.run_dir().to_path_buf();
1030
1031        let bus = Arc::new(EventBus::new(256));
1032        bus.add_sink(tracker.sink());
1033        let runner = StudyRunner::new(bus).with_tracker(tracker);
1034
1035        let mut study = Study::new(
1036            "persisted",
1037            one_dim_space(),
1038            SearchStrategy::Random {
1039                n_trials: 3,
1040                seed: Some(3),
1041            },
1042            vec![Objective {
1043                metric: "f1".into(),
1044                direction: Direction::Maximize,
1045            }],
1046        );
1047        let mut sampler = RandomSampler::new(3, Some(3));
1048        runner
1049            .run(&mut study, &mut sampler, &make_simple_executor())
1050            .unwrap();
1051
1052        // study.json is a complete, loadable study (crash-safe resume).
1053        let loaded = Study::load(run_dir.join("study.json")).unwrap();
1054        assert_eq!(loaded.trials.len(), 3);
1055        assert!(loaded.updated_at.is_some());
1056        // Trial events reached the run's events.jsonl through the sink.
1057        let events = std::fs::read_to_string(run_dir.join("events.jsonl")).unwrap();
1058        assert!(events.contains("TrialCompleted"));
1059    }
1060
1061    fn make_simple_executor() -> FnTrialExecutor<
1062        impl Fn(&HashMap<String, serde_json::Value>, &TrialContext) -> Result<TrialOutcome>,
1063    > {
1064        FnTrialExecutor(
1065            |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1066                Ok(TrialOutcome::Completed(vec![MetricRecord {
1067                    name: "f1".into(),
1068                    value: params["x"].as_f64().unwrap_or(0.5),
1069                    step: 0,
1070                    timestamp: Utc::now(),
1071                }]))
1072            },
1073        )
1074    }
1075
1076    // ── TrialContext unit tests (direct construction) ──
1077
1078    use crate::pruner::TrialMetricHistory as History;
1079    use std::sync::atomic::{AtomicUsize, Ordering};
1080
1081    /// Spy pruner: counts consultations, prunes everything after warmup.
1082    struct SpyPruner {
1083        calls: AtomicUsize,
1084        prune_from_step: usize,
1085    }
1086
1087    impl Pruner for SpyPruner {
1088        fn should_prune(
1089            &self,
1090            _metric: &str,
1091            _value: f64,
1092            step: usize,
1093            _history: &[History],
1094        ) -> Option<String> {
1095            self.calls.fetch_add(1, Ordering::SeqCst);
1096            (step >= self.prune_from_step).then(|| "spy".to_string())
1097        }
1098    }
1099
1100    fn make_ctx(pruner: Option<Arc<SpyPruner>>) -> (TrialContext, Arc<EventBus>) {
1101        let bus = Arc::new(EventBus::new(64));
1102        let ctx = TrialContext {
1103            study_id: "s".into(),
1104            trial_id: "trial_0000".into(),
1105            objective: Some(("f1".to_string(), Direction::Maximize)),
1106            pruner: pruner.map(|p| p as Arc<dyn Pruner>),
1107            history: Arc::new(Vec::new()),
1108            event_bus: bus.clone(),
1109            shared: Arc::new(Mutex::new(TrialShared::default())),
1110        };
1111        (ctx, bus)
1112    }
1113
1114    #[test]
1115    fn report_after_prune_is_sticky_and_skips_the_pruner() {
1116        let pruner = Arc::new(SpyPruner {
1117            calls: AtomicUsize::new(0),
1118            prune_from_step: 0,
1119        });
1120        let (ctx, _bus) = make_ctx(Some(pruner.clone()));
1121
1122        assert!(!ctx.should_prune());
1123        assert!(ctx.report("f1", 0.1, 0), "pruned immediately");
1124        assert!(ctx.should_prune());
1125        let calls_after_prune = pruner.calls.load(Ordering::SeqCst);
1126
1127        // Later reports return true WITHOUT consulting the pruner…
1128        assert!(ctx.report("f1", 0.9, 1));
1129        assert!(ctx.report("f1", 0.9, 2));
1130        assert_eq!(pruner.calls.load(Ordering::SeqCst), calls_after_prune);
1131        // …but the metrics are still recorded (pinned: push happens
1132        // before the prune check).
1133        assert_eq!(ctx.metrics().len(), 3);
1134    }
1135
1136    #[test]
1137    fn non_objective_metric_never_consults_the_pruner() {
1138        let pruner = Arc::new(SpyPruner {
1139            calls: AtomicUsize::new(0),
1140            prune_from_step: 0,
1141        });
1142        let (ctx, _bus) = make_ctx(Some(pruner.clone()));
1143
1144        assert!(!ctx.report("train_loss", 99.0, 0));
1145        assert!(!ctx.report("lr", 0.001, 0));
1146        assert_eq!(pruner.calls.load(Ordering::SeqCst), 0);
1147        assert!(ctx.report("f1", 0.0, 0), "objective metric does consult");
1148        assert_eq!(pruner.calls.load(Ordering::SeqCst), 1);
1149    }
1150
1151    #[test]
1152    fn report_emits_trial_metric_events() {
1153        let (ctx, bus) = make_ctx(None);
1154        let mut rx = bus.subscribe();
1155        ctx.report("f1", 0.5, 3);
1156        ctx.report("aux", 1.5, 3);
1157
1158        let mut got = Vec::new();
1159        while let Ok(e) = rx.try_recv() {
1160            if let Event::TrialMetric {
1161                study_id,
1162                trial_id,
1163                metric,
1164            } = e
1165            {
1166                assert_eq!(study_id, "s");
1167                assert_eq!(trial_id, "trial_0000");
1168                got.push((metric.name, metric.value, metric.step));
1169            }
1170        }
1171        assert_eq!(
1172            got,
1173            vec![("f1".to_string(), 0.5, 3), ("aux".to_string(), 1.5, 3)]
1174        );
1175    }
1176
1177    #[test]
1178    fn trial_context_accessors_and_cross_thread_clone() {
1179        let (ctx, _bus) = make_ctx(None);
1180        assert_eq!(ctx.trial_id(), "trial_0000");
1181
1182        // The documented property: the handle can cross threads (e.g.
1183        // into a Python callback) and all reports land on the trial.
1184        let clone = ctx.clone();
1185        let handle = std::thread::spawn(move || {
1186            clone.report("from_thread", 1.0, 0);
1187        });
1188        ctx.report("from_main", 2.0, 0);
1189        handle.join().unwrap();
1190
1191        let names: Vec<String> = ctx.metrics().into_iter().map(|m| m.name).collect();
1192        assert_eq!(names.len(), 2);
1193        assert!(names.contains(&"from_thread".to_string()));
1194        assert!(names.contains(&"from_main".to_string()));
1195    }
1196
1197    // ── Runner behavior ──
1198
1199    #[test]
1200    fn percentile_pruning_works_through_the_runner() {
1201        // Same shape as the median cases but exercising the Percentile
1202        // arm of build_pruner (dead code until now).
1203        use std::sync::atomic::AtomicUsize;
1204        let bus = Arc::new(EventBus::new(256));
1205        let runner = StudyRunner::new(bus);
1206        let mut study = Study::new(
1207            "percentile",
1208            one_dim_space(),
1209            SearchStrategy::Random {
1210                n_trials: 4,
1211                seed: Some(7),
1212            },
1213            vec![Objective {
1214                metric: "f1".into(),
1215                direction: Direction::Maximize,
1216            }],
1217        )
1218        .with_pruning(PruningStrategy::Percentile {
1219            percentile: 50.0,
1220            n_warmup_steps: 2,
1221        });
1222        let counter = AtomicUsize::new(0);
1223        let executor = FnTrialExecutor(
1224            move |_p: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
1225                let good = counter.fetch_add(1, Ordering::SeqCst) == 0;
1226                for step in 0..10 {
1227                    let value = if good { 0.5 + step as f64 * 0.05 } else { 0.01 };
1228                    if ctx.report("f1", value, step) {
1229                        return Ok(TrialOutcome::Pruned {
1230                            step,
1231                            reason: "stopped".into(),
1232                        });
1233                    }
1234                }
1235                Ok(TrialOutcome::Completed(vec![]))
1236            },
1237        );
1238        let mut sampler = RandomSampler::new(4, Some(7));
1239        runner.run(&mut study, &mut sampler, &executor).unwrap();
1240        assert_bad_trials_pruned(&study);
1241    }
1242
1243    #[test]
1244    fn pruner_verdict_wins_over_completed_outcome() {
1245        // An executor that IGNORES report()'s stop signal and returns
1246        // Completed anyway: the runner must mark the trial pruned with
1247        // the pruner's step/reason and drop the returned metrics.
1248        use std::sync::atomic::AtomicUsize;
1249        let bus = Arc::new(EventBus::new(256));
1250        let mut rx = bus.subscribe();
1251        let runner = StudyRunner::new(bus);
1252        let mut study = pruning_study(Direction::Maximize);
1253
1254        let counter = AtomicUsize::new(0);
1255        let executor = FnTrialExecutor(
1256            move |_p: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
1257                let good = counter.fetch_add(1, Ordering::SeqCst) == 0;
1258                for step in 0..10 {
1259                    let value = if good { 0.5 + step as f64 * 0.05 } else { 0.01 };
1260                    ctx.report("f1", value, step); // return value ignored!
1261                }
1262                Ok(TrialOutcome::Completed(vec![MetricRecord {
1263                    name: "sneaky".into(),
1264                    value: 1.0,
1265                    step: 0,
1266                    timestamp: Utc::now(),
1267                }]))
1268            },
1269        );
1270        let mut sampler = RandomSampler::new(4, Some(7));
1271        runner.run(&mut study, &mut sampler, &executor).unwrap();
1272
1273        let pruned: Vec<_> = study
1274            .trials
1275            .iter()
1276            .filter(|t| matches!(t.state, TrialState::Pruned { .. }))
1277            .collect();
1278        assert_eq!(pruned.len(), 3);
1279        for t in &pruned {
1280            if let TrialState::Pruned { step, reason } = &t.state {
1281                assert_eq!(*step, 2, "pruner's step, not the executor's");
1282                assert!(reason.contains("median"), "pruner's reason: {reason}");
1283            }
1284            assert!(
1285                !t.metrics.iter().any(|m| m.name == "sneaky"),
1286                "final metrics of an overridden Completed are dropped"
1287            );
1288        }
1289        // TrialPruned events emitted, and no TrialCompleted for them.
1290        let mut completed = 0;
1291        let mut pruned_events = 0;
1292        while let Ok(e) = rx.try_recv() {
1293            match e {
1294                Event::TrialCompleted { .. } => completed += 1,
1295                Event::TrialPruned { .. } => pruned_events += 1,
1296                _ => {}
1297            }
1298        }
1299        assert_eq!(completed, 1);
1300        assert_eq!(pruned_events, 3);
1301    }
1302
1303    /// Spy tracker recording every save_study call.
1304    #[derive(Default)]
1305    struct SpyTracker {
1306        saves: Mutex<Vec<usize>>, // trials.len() at each save
1307        fail: bool,
1308    }
1309
1310    impl somatize_core::tracking::Tracker for SpyTracker {
1311        fn run_id(&self) -> &str {
1312            "spy"
1313        }
1314        fn run_dir(&self) -> &std::path::Path {
1315            std::path::Path::new("/nonexistent")
1316        }
1317        fn sink(&self) -> Arc<dyn somatize_core::tracking::EventSink> {
1318            struct Null;
1319            impl somatize_core::tracking::EventSink for Null {
1320                fn record(&self, _event: &Event) {}
1321            }
1322            Arc::new(Null)
1323        }
1324        fn save_manifest(&self, _m: &somatize_core::tracking::RunManifest) -> Result<()> {
1325            Ok(())
1326        }
1327        fn save_artifact(&self, _p: &str, _b: &[u8]) -> Result<()> {
1328            Ok(())
1329        }
1330        fn save_study(&self, study: &Study) -> Result<()> {
1331            if self.fail {
1332                return Err(somatize_core::SomaError::Other("disk gone".into()));
1333            }
1334            self.saves.lock().unwrap().push(study.trials.len());
1335            Ok(())
1336        }
1337        fn heartbeat(&self) -> Result<()> {
1338            Ok(())
1339        }
1340        fn finalize(&self, _s: somatize_core::tracking::RunState) -> Result<()> {
1341            Ok(())
1342        }
1343    }
1344
1345    #[test]
1346    fn study_is_saved_after_every_trial_with_monotonic_growth() {
1347        let bus = Arc::new(EventBus::new(256));
1348        let tracker = Arc::new(SpyTracker::default());
1349        let runner = StudyRunner::new(bus).with_tracker(tracker.clone());
1350
1351        let mut study = Study::new(
1352            "persist",
1353            one_dim_space(),
1354            SearchStrategy::Random {
1355                n_trials: 3,
1356                seed: Some(3),
1357            },
1358            vec![Objective {
1359                metric: "f1".into(),
1360                direction: Direction::Maximize,
1361            }],
1362        );
1363        let mut sampler = RandomSampler::new(3, Some(3));
1364        runner
1365            .run(&mut study, &mut sampler, &make_simple_executor())
1366            .unwrap();
1367
1368        // One save per trial plus the final save: a runner that saved
1369        // only once at the end would fail this.
1370        let saves = tracker.saves.lock().unwrap().clone();
1371        assert_eq!(saves, vec![1, 2, 3, 3]);
1372    }
1373
1374    #[test]
1375    fn failing_tracker_never_fails_the_study() {
1376        let bus = Arc::new(EventBus::new(256));
1377        let tracker = Arc::new(SpyTracker {
1378            fail: true,
1379            ..Default::default()
1380        });
1381        let runner = StudyRunner::new(bus).with_tracker(tracker);
1382
1383        let mut study = Study::new(
1384            "resilient",
1385            one_dim_space(),
1386            SearchStrategy::Random {
1387                n_trials: 3,
1388                seed: Some(3),
1389            },
1390            vec![Objective {
1391                metric: "f1".into(),
1392                direction: Direction::Maximize,
1393            }],
1394        );
1395        let mut sampler = RandomSampler::new(3, Some(3));
1396        runner
1397            .run(&mut study, &mut sampler, &make_simple_executor())
1398            .unwrap();
1399        assert_eq!(study.trials.len(), 3);
1400    }
1401
1402    /// Sampler spy recording the ORDER of record_result vs sample calls.
1403    struct OrderSpySampler {
1404        inner: GridSampler,
1405        log: Vec<String>,
1406    }
1407
1408    impl Sampler for OrderSpySampler {
1409        fn prepare(&mut self, space: &SearchSpace) {
1410            self.inner.prepare(space);
1411        }
1412        fn sample(
1413            &mut self,
1414            space: &SearchSpace,
1415            trial_index: usize,
1416        ) -> Result<Option<HashMap<String, serde_json::Value>>> {
1417            self.log.push(format!("sample:{trial_index}"));
1418            self.inner.sample(space, trial_index)
1419        }
1420        fn n_trials(&self) -> Option<usize> {
1421            self.inner.n_trials()
1422        }
1423        fn record_result(&mut self, _params: &HashMap<String, serde_json::Value>, value: f64) {
1424            self.log.push(format!("record:{value:.2}"));
1425        }
1426    }
1427
1428    #[test]
1429    fn resume_replays_history_into_the_sampler_before_sampling() {
1430        let executor = make_executor();
1431
1432        // Reference run to harvest 3 completed trials.
1433        let bus = Arc::new(EventBus::new(256));
1434        let runner = StudyRunner::new(bus);
1435        let mut study = Study::new(
1436            "replay",
1437            sample_space(),
1438            SearchStrategy::Grid { points_per_dim: 3 },
1439            vec![Objective {
1440                metric: "f1".into(),
1441                direction: Direction::Maximize,
1442            }],
1443        );
1444        runner
1445            .run(&mut study, &mut GridSampler::new(3), &executor)
1446            .unwrap();
1447        study.trials.truncate(3);
1448
1449        let bus = Arc::new(EventBus::new(256));
1450        let runner = StudyRunner::new(bus);
1451        let mut spy = OrderSpySampler {
1452            inner: GridSampler::new(3),
1453            log: Vec::new(),
1454        };
1455        runner.run(&mut study, &mut spy, &executor).unwrap();
1456
1457        // The first 3 log entries are replayed history; sampling only
1458        // starts afterwards, at the resume index.
1459        assert_eq!(
1460            spy.log.iter().filter(|e| e.starts_with("record:")).count(),
1461            3 + 3
1462        );
1463        assert!(
1464            spy.log[..3].iter().all(|e| e.starts_with("record:")),
1465            "history replay must precede sampling: {:?}",
1466            &spy.log[..4]
1467        );
1468        assert_eq!(spy.log[3], "sample:3", "sampling resumes at the next index");
1469    }
1470
1471    #[test]
1472    fn frozen_param_overrides_a_sampled_dimension() {
1473        let bus = Arc::new(EventBus::new(256));
1474        let runner = StudyRunner::new(bus);
1475
1476        // "x" is IN the search space and also frozen: frozen wins.
1477        let mut study = Study::new(
1478            "frozen-collision",
1479            one_dim_space(),
1480            SearchStrategy::Random {
1481                n_trials: 3,
1482                seed: Some(2),
1483            },
1484            vec![Objective {
1485                metric: "f1".into(),
1486                direction: Direction::Maximize,
1487            }],
1488        );
1489        study.frozen.insert("x".into(), serde_json::json!(0.75));
1490
1491        let executor = FnTrialExecutor(
1492            |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1493                assert_eq!(params["x"], serde_json::json!(0.75));
1494                Ok(TrialOutcome::Completed(vec![MetricRecord {
1495                    name: "f1".into(),
1496                    value: 0.5,
1497                    step: 0,
1498                    timestamp: Utc::now(),
1499                }]))
1500            },
1501        );
1502        let mut sampler = RandomSampler::new(3, Some(2));
1503        runner.run(&mut study, &mut sampler, &executor).unwrap();
1504
1505        for t in &study.trials {
1506            assert_eq!(t.params["x"], serde_json::json!(0.75));
1507        }
1508    }
1509
1510    #[test]
1511    fn pruning_watches_first_composite_term_when_no_objectives() {
1512        use somatize_core::study::{CompositeObjective, Scalarizer};
1513        use std::sync::atomic::AtomicUsize;
1514
1515        let bus = Arc::new(EventBus::new(256));
1516        let runner = StudyRunner::new(bus);
1517        let mut study = Study::new(
1518            "composite-pruning",
1519            one_dim_space(),
1520            SearchStrategy::Random {
1521                n_trials: 4,
1522                seed: Some(7),
1523            },
1524            vec![], // no objectives — pruning falls back to terms[0]
1525        )
1526        .with_composite(CompositeObjective {
1527            terms: vec![("f1".into(), 1.0), ("aux".into(), 0.1)],
1528            direction: Direction::Maximize,
1529            scalarizer: Scalarizer::WeightedSum,
1530        })
1531        .with_pruning(PruningStrategy::Median { n_warmup_steps: 2 });
1532
1533        let counter = AtomicUsize::new(0);
1534        let executor = FnTrialExecutor(
1535            move |_p: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
1536                let good = counter.fetch_add(1, Ordering::SeqCst) == 0;
1537                for step in 0..10 {
1538                    let value = if good { 0.5 + step as f64 * 0.05 } else { 0.01 };
1539                    if ctx.report("f1", value, step) {
1540                        return Ok(TrialOutcome::Pruned {
1541                            step,
1542                            reason: "stopped".into(),
1543                        });
1544                    }
1545                }
1546                Ok(TrialOutcome::Completed(vec![MetricRecord {
1547                    name: "aux".into(),
1548                    value: 0.0,
1549                    step: 0,
1550                    timestamp: Utc::now(),
1551                }]))
1552            },
1553        );
1554        let mut sampler = RandomSampler::new(4, Some(7));
1555        runner.run(&mut study, &mut sampler, &executor).unwrap();
1556        assert_bad_trials_pruned(&study);
1557    }
1558
1559    #[test]
1560    fn planned_trials_is_stamped_and_progress_completes() {
1561        let bus = Arc::new(EventBus::new(256));
1562        let runner = StudyRunner::new(bus);
1563        let mut study = Study::new(
1564            "stamped",
1565            sample_space(),
1566            SearchStrategy::Grid { points_per_dim: 3 },
1567            vec![Objective {
1568                metric: "f1".into(),
1569                direction: Direction::Maximize,
1570            }],
1571        );
1572        assert_eq!(study.total_trials(), None);
1573        runner
1574            .run(&mut study, &mut GridSampler::new(3), &make_executor())
1575            .unwrap();
1576        assert_eq!(study.planned_trials, Some(6));
1577        assert_eq!(study.total_trials(), Some(6));
1578        assert!((study.progress() - 1.0).abs() < f64::EPSILON);
1579    }
1580
1581    #[test]
1582    fn all_failed_study_completes_with_nan_best() {
1583        let bus = Arc::new(EventBus::new(256));
1584        let mut rx = bus.subscribe();
1585        let runner = StudyRunner::new(bus);
1586        let mut study = Study::new(
1587            "doomed",
1588            one_dim_space(),
1589            SearchStrategy::Random {
1590                n_trials: 3,
1591                seed: Some(1),
1592            },
1593            vec![Objective {
1594                metric: "f1".into(),
1595                direction: Direction::Maximize,
1596            }],
1597        );
1598        let executor = FnTrialExecutor(
1599            |_p: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1600                Err(SomaError::Other("cuda out of memory".into()))
1601            },
1602        );
1603        let mut sampler = RandomSampler::new(3, Some(1));
1604        runner.run(&mut study, &mut sampler, &executor).unwrap();
1605
1606        // Error strings preserved, ids contiguous, best is empty/NaN.
1607        for (i, t) in study.trials.iter().enumerate() {
1608            assert_eq!(t.id, format!("trial_{i:04}"));
1609            match &t.state {
1610                TrialState::Failed { error } => assert!(error.contains("cuda out of memory")),
1611                other => panic!("expected Failed, got {other:?}"),
1612            }
1613        }
1614        assert!(study.best_trial().is_none());
1615
1616        let mut failed_events = 0;
1617        let mut completed_nan = false;
1618        while let Ok(e) = rx.try_recv() {
1619            match e {
1620                Event::TrialFailed { error, .. } => {
1621                    assert!(error.contains("cuda out of memory"));
1622                    failed_events += 1;
1623                }
1624                Event::StudyCompleted {
1625                    best_trial_id,
1626                    best_value,
1627                    ..
1628                } => {
1629                    assert!(best_trial_id.is_empty());
1630                    assert!(best_value.is_nan());
1631                    completed_nan = true;
1632                }
1633                _ => {}
1634            }
1635        }
1636        assert_eq!(failed_events, 3);
1637        assert!(completed_nan);
1638    }
1639
1640    #[test]
1641    fn best_updated_fires_once_when_trials_worsen() {
1642        use std::sync::atomic::AtomicUsize;
1643        let bus = Arc::new(EventBus::new(256));
1644        let mut rx = bus.subscribe();
1645        let runner = StudyRunner::new(bus);
1646        let mut study = Study::new(
1647            "worsening",
1648            one_dim_space(),
1649            SearchStrategy::Random {
1650                n_trials: 3,
1651                seed: Some(1),
1652            },
1653            vec![Objective {
1654                metric: "f1".into(),
1655                direction: Direction::Maximize,
1656            }],
1657        );
1658        let counter = AtomicUsize::new(0);
1659        let executor = FnTrialExecutor(
1660            move |_p: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1661                let i = counter.fetch_add(1, Ordering::SeqCst);
1662                Ok(TrialOutcome::Completed(vec![MetricRecord {
1663                    name: "f1".into(),
1664                    value: 0.9 - i as f64 * 0.2, // 0.9, 0.7, 0.5
1665                    step: 0,
1666                    timestamp: Utc::now(),
1667                }]))
1668            },
1669        );
1670        let mut sampler = RandomSampler::new(3, Some(1));
1671        runner.run(&mut study, &mut sampler, &executor).unwrap();
1672
1673        let mut best_events = Vec::new();
1674        while let Ok(e) = rx.try_recv() {
1675            if let Event::BestUpdated { value, .. } = e {
1676                best_events.push(value);
1677            }
1678        }
1679        assert_eq!(best_events, vec![0.9], "only the first trial is ever best");
1680    }
1681
1682    /// CONTRACT (pinned): metrics reported via ctx AND returned in
1683    /// Completed(final_metrics) are BOTH kept — the same name appears
1684    /// twice and the TrialMetric event fires twice. De-dup is the
1685    /// executor's responsibility.
1686    #[test]
1687    fn reported_and_final_metrics_are_concatenated_not_deduped() {
1688        let bus = Arc::new(EventBus::new(256));
1689        let mut rx = bus.subscribe();
1690        let runner = StudyRunner::new(bus);
1691        let mut study = Study::new(
1692            "dup",
1693            one_dim_space(),
1694            SearchStrategy::Random {
1695                n_trials: 1,
1696                seed: Some(1),
1697            },
1698            vec![Objective {
1699                metric: "f1".into(),
1700                direction: Direction::Maximize,
1701            }],
1702        );
1703        let executor = FnTrialExecutor(
1704            |_p: &HashMap<String, serde_json::Value>, ctx: &TrialContext| {
1705                ctx.report("f1", 0.4, 0);
1706                Ok(TrialOutcome::Completed(vec![MetricRecord {
1707                    name: "f1".into(),
1708                    value: 0.6,
1709                    step: 1,
1710                    timestamp: Utc::now(),
1711                }]))
1712            },
1713        );
1714        let mut sampler = RandomSampler::new(1, Some(1));
1715        runner.run(&mut study, &mut sampler, &executor).unwrap();
1716
1717        let f1_records: Vec<f64> = study.trials[0]
1718            .metrics
1719            .iter()
1720            .filter(|m| m.name == "f1")
1721            .map(|m| m.value)
1722            .collect();
1723        assert_eq!(f1_records, vec![0.4, 0.6]);
1724
1725        let mut metric_events = 0;
1726        while let Ok(e) = rx.try_recv() {
1727            if matches!(e, Event::TrialMetric { .. }) {
1728                metric_events += 1;
1729            }
1730        }
1731        assert_eq!(metric_events, 2);
1732    }
1733
1734    #[test]
1735    fn sampler_feedback_values_are_exact_and_skip_non_completed() {
1736        let bus = Arc::new(EventBus::new(256));
1737        let runner = StudyRunner::new(bus);
1738        let mut study = Study::new(
1739            "feedback-exact",
1740            one_dim_space(),
1741            SearchStrategy::Random {
1742                n_trials: 3,
1743                seed: Some(1),
1744            },
1745            vec![Objective {
1746                metric: "loss".into(),
1747                direction: Direction::Minimize,
1748            }],
1749        );
1750        use std::sync::atomic::AtomicUsize;
1751        let counter = AtomicUsize::new(0);
1752        let executor = FnTrialExecutor(
1753            move |_p: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| match counter
1754                .fetch_add(1, Ordering::SeqCst)
1755            {
1756                0 => Ok(TrialOutcome::Completed(vec![MetricRecord {
1757                    name: "loss".into(),
1758                    value: 0.25,
1759                    step: 0,
1760                    timestamp: Utc::now(),
1761                }])),
1762                1 => Ok(TrialOutcome::Pruned {
1763                    step: 1,
1764                    reason: "bad".into(),
1765                }),
1766                _ => Err(SomaError::Other("boom".into())),
1767            },
1768        );
1769        let mut spy = SpySampler {
1770            inner: RandomSampler::new(3, Some(1)),
1771            recorded: Vec::new(),
1772        };
1773        runner.run(&mut study, &mut spy, &executor).unwrap();
1774
1775        // Only the completed trial fed back, negated for Minimize.
1776        assert_eq!(spy.recorded, vec![-0.25]);
1777    }
1778
1779    #[test]
1780    fn timestamps_backfilled_and_monotonic() {
1781        let bus = Arc::new(EventBus::new(256));
1782        let tracker = Arc::new(SpyTracker::default());
1783        let runner = StudyRunner::new(bus).with_tracker(tracker);
1784        let mut study = Study::new(
1785            "ts",
1786            one_dim_space(),
1787            SearchStrategy::Random {
1788                n_trials: 2,
1789                seed: Some(1),
1790            },
1791            vec![Objective {
1792                metric: "f1".into(),
1793                direction: Direction::Maximize,
1794            }],
1795        );
1796        study.created_at = None; // e.g. loaded from a pre-timestamp JSON
1797        let mut sampler = RandomSampler::new(2, Some(1));
1798        runner
1799            .run(&mut study, &mut sampler, &make_simple_executor())
1800            .unwrap();
1801
1802        assert!(study.created_at.is_some(), "backfilled");
1803        assert!(study.updated_at.is_some());
1804        for t in &study.trials {
1805            assert!(t.started_at.is_some());
1806            assert!(t.finished_at.unwrap() >= t.started_at.unwrap());
1807        }
1808    }
1809
1810    #[test]
1811    fn bayesian_through_runner_improves_over_time() {
1812        // End-to-end proof that the ask/tell wiring feeds TPE: on a
1813        // unimodal objective the later trials must beat the early ones.
1814        // Deterministic (fixed seed) — not statistical.
1815        let bus = Arc::new(EventBus::new(1024));
1816        let runner = StudyRunner::new(bus);
1817        let mut study = Study::new(
1818            "tpe",
1819            one_dim_space(),
1820            SearchStrategy::Bayesian {
1821                n_trials: 30,
1822                n_startup: 8,
1823                seed: Some(42),
1824            },
1825            vec![Objective {
1826                metric: "score".into(),
1827                direction: Direction::Maximize,
1828            }],
1829        );
1830        let executor = FnTrialExecutor(
1831            |params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1832                let x = params["x"].as_f64().unwrap();
1833                Ok(TrialOutcome::Completed(vec![MetricRecord {
1834                    name: "score".into(),
1835                    value: 1.0 - (x - 0.7).abs(), // peak at x = 0.7
1836                    step: 0,
1837                    timestamp: Utc::now(),
1838                }]))
1839            },
1840        );
1841        let mut sampler = BayesianSampler::new(30, 8, Some(42));
1842        runner.run(&mut study, &mut sampler, &executor).unwrap();
1843
1844        let scores: Vec<f64> = study
1845            .trials
1846            .iter()
1847            .map(|t| t.last_metric("score").unwrap())
1848            .collect();
1849        let head: f64 = scores[..8].iter().sum::<f64>() / 8.0; // startup = random
1850        let tail: f64 = scores[20..].iter().sum::<f64>() / 10.0;
1851        assert!(
1852            tail > head,
1853            "TPE with feedback must beat its random startup: head={head:.3} tail={tail:.3}"
1854        );
1855    }
1856
1857    #[test]
1858    fn study_progress_tracking() {
1859        let bus = Arc::new(EventBus::new(256));
1860        let mut rx = bus.subscribe();
1861        let runner = StudyRunner::new(bus);
1862
1863        let mut space = SearchSpace::new();
1864        space.add(SearchDimension::Float {
1865            name: "x".into(),
1866            low: 0.0,
1867            high: 1.0,
1868            scale: Scale::Linear,
1869            default: None,
1870        });
1871
1872        let mut study = Study::new(
1873            "progress_test",
1874            space,
1875            SearchStrategy::Random {
1876                n_trials: 3,
1877                seed: None,
1878            },
1879            vec![Objective {
1880                metric: "f1".into(),
1881                direction: Direction::Maximize,
1882            }],
1883        );
1884
1885        let executor = FnTrialExecutor(
1886            |_params: &HashMap<String, serde_json::Value>, _ctx: &TrialContext| {
1887                Ok(TrialOutcome::Completed(vec![MetricRecord {
1888                    name: "f1".into(),
1889                    value: 0.5,
1890                    step: 0,
1891                    timestamp: Utc::now(),
1892                }]))
1893            },
1894        );
1895
1896        let mut sampler = RandomSampler::new(3, Some(42));
1897        runner.run(&mut study, &mut sampler, &executor).unwrap();
1898
1899        // Collect progress events
1900        let mut progress_events = Vec::new();
1901        while let Ok(e) = rx.try_recv() {
1902            if let Event::StudyProgress {
1903                completed, total, ..
1904            } = e
1905            {
1906                progress_events.push((completed, total));
1907            }
1908        }
1909
1910        assert_eq!(progress_events.len(), 3);
1911        assert_eq!(progress_events[0], (1, 3));
1912        assert_eq!(progress_events[1], (2, 3));
1913        assert_eq!(progress_events[2], (3, 3));
1914    }
1915}