use std::time::{Duration, Instant};
use anyhow::{Context, Result, bail};
use crate::chem_env::{ChemEnv, RetroRule, default_rules, load_rules_from_file};
use crate::search::{self, Route, SearchConfig, SearchControl, SearchStats, SearchTermination};
#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize)]
#[serde(rename_all = "snake_case")]
pub enum SelectedStage {
Stage1,
Stage2,
}
#[derive(Debug)]
pub struct CoverageModeResult {
pub routes: Vec<Route>,
pub selected_stage: SelectedStage,
pub stats: SearchStats,
pub stage1_solved: bool,
pub stage2_invoked: bool,
pub stage1_elapsed_ms: f64,
pub stage2_elapsed_ms: Option<f64>,
pub total_elapsed_ms: f64,
pub stage1_timeout: bool,
pub stage2_timeout: bool,
pub reranker_failures: u64,
}
pub fn validate_coverage_mode_flags(
bond_index: bool,
ring_context_policy_active: bool,
onnx_scorer_active: bool,
) -> Result<()> {
if bond_index {
bail!(
"coverage mode does not support --bond-index in v0 -- Stage 2 would need its own, \
separately validated retrieval index against the coverage template set"
);
}
if onnx_scorer_active {
bail!(
"coverage mode does not support an ONNX --scorer in v0 -- Stage 2 would need its \
own, separately validated scorer vocabulary against the coverage template set"
);
}
if ring_context_policy_active {
bail!(
"coverage mode does not support an active --ring-context-policy in v0 -- Stage 2 \
would need its own, separately validated ring-context sidecar against the coverage \
template set"
);
}
Ok(())
}
pub fn validate_coverage_mode_config(config: &SearchConfig) -> Result<()> {
let ring_context_policy_active = !matches!(
config.ring_context,
crate::ring_context::RingContextConfig::Disabled
);
#[cfg(feature = "nn-scoring")]
let onnx_scorer_active = config.nn_scorer.is_some();
#[cfg(not(feature = "nn-scoring"))]
let onnx_scorer_active = false;
validate_coverage_mode_flags(
config.bond_index,
ring_context_policy_active,
onnx_scorer_active,
)
}
pub fn load_coverage_rules(coverage_templates_path: &str) -> Result<Vec<RetroRule>> {
let metadata = std::fs::metadata(coverage_templates_path).with_context(|| {
format!(
"--coverage-templates path does not exist or is not readable: \
{coverage_templates_path}"
)
})?;
if !metadata.is_file() {
bail!("--coverage-templates path is not a file: {coverage_templates_path}");
}
std::fs::read_to_string(coverage_templates_path).with_context(|| {
format!(
"--coverage-templates path exists but could not be read as valid UTF-8 text \
(permission error or binary/non-UTF-8 content): {coverage_templates_path}"
)
})?;
let extra = load_rules_from_file(coverage_templates_path);
if extra.is_empty() {
bail!("--coverage-templates file contains no valid templates: {coverage_templates_path}");
}
let mut rules = default_rules();
rules.extend(extra);
Ok(rules)
}
pub fn run_coverage_mode_with_configs(
target_smiles: &str,
env: &ChemEnv,
stage1_rules: &[RetroRule],
stage1_config: &SearchConfig,
stage2_rules: &[RetroRule],
stage2_config: &SearchConfig,
stage2_timeout: Option<Duration>,
) -> Result<CoverageModeResult> {
validate_coverage_mode_config(stage1_config)?;
let total_start = Instant::now();
let stage1_start = Instant::now();
let (stage1_routes, stage1_stats) =
search::find_routes(target_smiles, env, stage1_rules, stage1_config)?;
let stage1_elapsed_ms = stage1_start.elapsed().as_secs_f64() * 1000.0;
let stage1_solved = !stage1_routes.is_empty();
if stage1_solved {
let reranker_failures = stage1_stats.reranker_failures;
return Ok(CoverageModeResult {
routes: stage1_routes,
selected_stage: SelectedStage::Stage1,
stats: stage1_stats,
stage1_solved: true,
stage2_invoked: false,
stage1_elapsed_ms,
stage2_elapsed_ms: None,
total_elapsed_ms: total_start.elapsed().as_secs_f64() * 1000.0,
stage1_timeout: false,
stage2_timeout: false,
reranker_failures,
});
}
let stage2_start = Instant::now();
let control = match stage2_timeout {
Some(d) => SearchControl::with_timeout(d),
None => SearchControl::unlimited(),
};
let stage2_result = search::find_routes_with_control(
target_smiles,
env,
stage2_rules,
stage2_config,
&control,
)?;
let stage2_elapsed_ms = stage2_start.elapsed().as_secs_f64() * 1000.0;
Ok(stage2_outcome_to_result(
stage2_result,
stage1_stats.reranker_failures,
stage1_elapsed_ms,
stage2_elapsed_ms,
total_start.elapsed().as_secs_f64() * 1000.0,
))
}
fn stage2_outcome_to_result(
stage2_result: search::SearchRunResult,
stage1_reranker_failures: u64,
stage1_elapsed_ms: f64,
stage2_elapsed_ms: f64,
total_elapsed_ms: f64,
) -> CoverageModeResult {
let stage2_timed_out = stage2_result.termination == SearchTermination::DeadlineExceeded;
let reranker_failures = stage1_reranker_failures + stage2_result.stats.reranker_failures;
CoverageModeResult {
routes: stage2_result.routes,
selected_stage: SelectedStage::Stage2,
stats: stage2_result.stats,
stage1_solved: false,
stage2_invoked: true,
stage1_elapsed_ms,
stage2_elapsed_ms: Some(stage2_elapsed_ms),
total_elapsed_ms,
stage1_timeout: false,
stage2_timeout: stage2_timed_out,
reranker_failures,
}
}
pub fn run_coverage_mode(
target_smiles: &str,
env: &ChemEnv,
stage1_rules: &[RetroRule],
config: &SearchConfig,
coverage_rules: &[RetroRule],
stage2_timeout: Option<Duration>,
) -> Result<CoverageModeResult> {
run_coverage_mode_with_configs(
target_smiles,
env,
stage1_rules,
config,
coverage_rules,
config,
stage2_timeout,
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::chem_env::default_rules;
fn env() -> ChemEnv {
ChemEnv::load("data/building_blocks.smi").unwrap_or_else(|_| {
ChemEnv::in_memory(&["CC(=O)O", "Oc1ccccc1C(=O)O", "c1ccccc1C(=O)O", "C", "O"])
})
}
fn cfg() -> SearchConfig {
SearchConfig {
max_depth: 5,
max_routes: 5,
beam_width: 0,
..Default::default()
}
}
fn shallow_beam_limited_cfg() -> SearchConfig {
SearchConfig {
max_depth: 2,
max_routes: 5,
beam_width: 100,
..Default::default()
}
}
const ASPIRIN: &str = "CC(=O)Oc1ccccc1C(=O)O";
const SOLVABLE_BY_DEFAULT_RULES_ALONE: &str = "CCN(CC)C(=O)c1ccccc1F";
const SOLVABLE_ONLY_WITH_FIXTURE_TEMPLATES: &str = "O=C1CCC(=O)N1c1ccccc1";
const UNKNOWN: &str = "c1ccc2c(c1)c1ccccc1c1ccccc21";
fn fixture_rules() -> Vec<RetroRule> {
load_rules_from_file("tests/fixtures/coverage_mode_templates.smi")
}
struct PanicReranker;
impl crate::candidate::CandidateReranker for PanicReranker {
fn score_pool(
&self,
_target: &str,
_candidates: &mut [crate::candidate::ReactionCandidate],
) -> anyhow::Result<()> {
panic!("PanicReranker::score_pool was called -- Stage 2 must not have run");
}
}
struct FailingReranker(std::sync::Arc<std::sync::atomic::AtomicUsize>);
impl crate::candidate::CandidateReranker for FailingReranker {
fn score_pool(
&self,
_target: &str,
_candidates: &mut [crate::candidate::ReactionCandidate],
) -> anyhow::Result<()> {
self.0.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
anyhow::bail!("FailingReranker: deliberate failure for aggregation test")
}
}
fn synthetic_route() -> Route {
Route {
steps: vec![],
depth: 0,
score: 0.0,
building_blocks: vec![],
confidence: 1.0,
convergency: 1.0,
success_probability: 1.0,
route_cost: 0.0,
}
}
#[test]
fn stage1_solved_never_invokes_stage2() {
let env = env();
let stage1_rules = default_rules();
let stage2_rules = default_rules();
let stage1_config = cfg();
let stage2_config = SearchConfig {
reranker: Some(std::sync::Arc::new(PanicReranker)),
..cfg()
};
let result = run_coverage_mode_with_configs(
"CC(=O)O",
&env,
&stage1_rules,
&stage1_config,
&stage2_rules,
&stage2_config,
None,
)
.unwrap();
assert_eq!(result.selected_stage, SelectedStage::Stage1);
assert!(result.stage1_solved);
assert!(!result.stage2_invoked);
assert!(!result.routes.is_empty());
assert!(result.stage2_elapsed_ms.is_none());
}
#[test]
fn stage1_unsolved_invokes_stage2() {
let env = env();
let stage1_rules: Vec<RetroRule> = vec![]; let stage2_rules = default_rules();
let result =
run_coverage_mode(ASPIRIN, &env, &stage1_rules, &cfg(), &stage2_rules, None).unwrap();
assert!(!result.stage1_solved);
assert!(result.stage2_invoked);
assert_eq!(result.selected_stage, SelectedStage::Stage2);
assert!(result.stage2_elapsed_ms.is_some());
assert!(!result.routes.is_empty());
}
#[test]
fn stage1_valid_route_never_overwritten() {
let env = env();
let stage1_rules = default_rules();
let mut stage2_rules = default_rules();
stage2_rules.extend(fixture_rules());
let baseline =
search::find_routes(SOLVABLE_BY_DEFAULT_RULES_ALONE, &env, &stage1_rules, &cfg())
.unwrap();
assert!(
!baseline.0.is_empty(),
"fixture target must be Stage-1-solvable"
);
let result = run_coverage_mode(
SOLVABLE_BY_DEFAULT_RULES_ALONE,
&env,
&stage1_rules,
&cfg(),
&stage2_rules,
None,
)
.unwrap();
assert_eq!(result.selected_stage, SelectedStage::Stage1);
assert!(!result.stage2_invoked);
assert_eq!(
serde_json::to_string(&result.routes).unwrap(),
serde_json::to_string(&baseline.0).unwrap(),
"coverage mode's Stage-1 result must be byte-identical to a direct find_routes call \
with the same Stage-1 rules -- anything else means it was altered on the way through"
);
}
#[test]
fn stage2_uses_its_own_rules_not_stage1_rules() {
let env = env();
let stage1_rules = default_rules(); let mut stage2_rules = default_rules();
stage2_rules.extend(fixture_rules());
let stage1_only = search::find_routes(
SOLVABLE_ONLY_WITH_FIXTURE_TEMPLATES,
&env,
&stage1_rules,
&shallow_beam_limited_cfg(),
)
.unwrap();
assert!(
stage1_only.0.is_empty(),
"fixture target must NOT be solvable by Stage 1's rules alone"
);
let result = run_coverage_mode(
SOLVABLE_ONLY_WITH_FIXTURE_TEMPLATES,
&env,
&stage1_rules,
&shallow_beam_limited_cfg(),
&stage2_rules,
None,
)
.unwrap();
assert!(result.stage2_invoked);
assert_eq!(result.selected_stage, SelectedStage::Stage2);
assert!(
!result.routes.is_empty(),
"Stage 2 must have used its own (larger) rule set, not Stage 1's"
);
}
#[test]
fn stage2_timeout_is_surfaced() {
let env = env();
let stage1_rules: Vec<RetroRule> = vec![];
let stage2_rules = default_rules();
let result = run_coverage_mode(
ASPIRIN,
&env,
&stage1_rules,
&cfg(),
&stage2_rules,
Some(Duration::from_nanos(1)),
)
.unwrap();
assert!(result.stage2_invoked);
assert!(result.stage2_timeout);
}
#[test]
fn stage2_outcome_conversion_retains_partial_routes_on_timeout() {
let synthetic = search::SearchRunResult {
routes: vec![synthetic_route()],
stats: SearchStats::default(),
termination: SearchTermination::DeadlineExceeded,
};
let result = stage2_outcome_to_result(synthetic, 0, 10.0, 20.0, 30.0);
assert_eq!(result.routes.len(), 1);
assert!(result.stage2_timeout);
assert_eq!(result.selected_stage, SelectedStage::Stage2);
assert!(result.stage2_invoked);
assert_eq!(result.stage1_elapsed_ms, 10.0);
assert_eq!(result.stage2_elapsed_ms, Some(20.0));
assert_eq!(result.total_elapsed_ms, 30.0);
}
#[test]
fn stage2_outcome_conversion_preserves_completed_routes_too() {
let synthetic = search::SearchRunResult {
routes: vec![synthetic_route(), synthetic_route()],
stats: SearchStats::default(),
termination: SearchTermination::Completed,
};
let result = stage2_outcome_to_result(synthetic, 0, 1.0, 2.0, 3.0);
assert_eq!(result.routes.len(), 2);
assert!(!result.stage2_timeout);
}
#[test]
fn reranker_failures_summed_across_invoked_stages() {
let env = env();
let stage1_rules: Vec<RetroRule> = vec![]; let stage2_rules = default_rules(); let call_count = std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0));
let stage_config = SearchConfig {
reranker: Some(std::sync::Arc::new(FailingReranker(call_count.clone()))),
..cfg()
};
let result = run_coverage_mode(
ASPIRIN,
&env,
&stage1_rules,
&stage_config,
&stage2_rules,
None,
)
.unwrap();
assert!(result.stage2_invoked);
assert_eq!(
result.stats.reranker_failures, 1,
"the selected (Stage 2) stage's own count must be exactly 1"
);
assert_eq!(
result.reranker_failures, 2,
"summed across both stages (1 + 1), not just the selected stage's own count"
);
assert!(
call_count.load(std::sync::atomic::Ordering::SeqCst) >= 2,
"the reranker must have actually been invoked by both stages, not just one"
);
}
#[test]
fn elapsed_and_stage_fields_are_self_consistent() {
let env = env();
let stage1_rules: Vec<RetroRule> = vec![];
let stage2_rules = default_rules();
let result =
run_coverage_mode(ASPIRIN, &env, &stage1_rules, &cfg(), &stage2_rules, None).unwrap();
assert!(result.stage1_elapsed_ms >= 0.0);
assert!(result.stage2_elapsed_ms.unwrap() >= 0.0);
assert!(result.total_elapsed_ms >= result.stage1_elapsed_ms);
assert!(result.total_elapsed_ms >= result.stage2_elapsed_ms.unwrap());
assert_eq!(result.selected_stage, SelectedStage::Stage2);
assert!(result.stage2_invoked);
}
#[test]
fn deterministic_repeated_output_with_sufficient_budget() {
let env = env();
let stage1_rules: Vec<RetroRule> = vec![];
let stage2_rules = default_rules();
let r1 =
run_coverage_mode(ASPIRIN, &env, &stage1_rules, &cfg(), &stage2_rules, None).unwrap();
let r2 =
run_coverage_mode(ASPIRIN, &env, &stage1_rules, &cfg(), &stage2_rules, None).unwrap();
assert_eq!(
serde_json::to_string(&r1.routes).unwrap(),
serde_json::to_string(&r2.routes).unwrap()
);
assert_eq!(r1.selected_stage, r2.selected_stage);
assert_eq!(r1.reranker_failures, r2.reranker_failures);
}
#[test]
fn validate_config_rejects_bond_index() {
let config = SearchConfig {
bond_index: true,
..Default::default()
};
assert!(validate_coverage_mode_config(&config).is_err());
}
#[test]
fn validate_flags_rejects_bond_index() {
assert!(validate_coverage_mode_flags(true, false, false).is_err());
}
#[test]
fn validate_flags_rejects_ring_context_active() {
assert!(validate_coverage_mode_flags(false, true, false).is_err());
}
#[test]
fn validate_flags_rejects_onnx_scorer_active() {
assert!(validate_coverage_mode_flags(false, false, true).is_err());
}
#[test]
fn validate_flags_accepts_all_inactive() {
assert!(validate_coverage_mode_flags(false, false, false).is_ok());
}
#[test]
fn load_coverage_rules_missing_path_fails_loud() {
let result = load_coverage_rules("/nonexistent/path/does_not_exist.smi");
assert!(result.is_err());
}
#[test]
fn load_coverage_rules_directory_path_fails_loud() {
let result = load_coverage_rules("data");
assert!(result.is_err());
}
#[test]
fn load_coverage_rules_empty_file_fails_loud() {
let dir = std::env::temp_dir();
let path = dir.join(format!(
"renkin_coverage_mode_test_empty_{}.smi",
std::process::id()
));
std::fs::write(&path, "# only a comment, no templates\n").unwrap();
let result = load_coverage_rules(path.to_str().unwrap());
let _ = std::fs::remove_file(&path);
assert!(result.is_err());
}
#[test]
fn unreadable_coverage_templates_path_reports_a_read_failure_not_missing_templates() {
let dir = std::env::temp_dir();
let path = dir.join(format!(
"renkin_coverage_mode_test_invalid_utf8_{}.smi",
std::process::id()
));
std::fs::write(&path, [0xFF, 0xFE, 0x00, 0xFF, 0xD8, 0x00]).unwrap();
let result = load_coverage_rules(path.to_str().unwrap());
let _ = std::fs::remove_file(&path);
let err = result.expect_err("invalid UTF-8 content must fail loud");
let msg = format!("{err:#}");
assert!(
msg.contains("could not be read as valid UTF-8"),
"error must describe a read failure, got: {msg}"
);
assert!(
!msg.contains("contains no valid templates"),
"must not be misreported as the empty-templates case: {msg}"
);
}
#[test]
fn unknown_target_via_coverage_mode_does_not_panic() {
let env = env();
let stage1_rules: Vec<RetroRule> = vec![];
let stage2_rules: Vec<RetroRule> = vec![]; let result = run_coverage_mode(UNKNOWN, &env, &stage1_rules, &cfg(), &stage2_rules, None);
assert!(result.is_ok());
}
}