use std::collections::HashMap;
#[cfg(feature = "perf-instrumentation")]
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Mutex, OnceLock};
#[cfg(feature = "perf-instrumentation")]
use std::time::Instant;
#[cfg(not(target_arch = "wasm32"))]
use rayon::prelude::*;
use serde::Serialize;
use sha2::{Digest, Sha256};
#[cfg(test)]
use crate::chem_env::apply_retro;
use crate::chem_env::{
Molecule, PrecursorMol, RetroRule, TemplateBondIndex, mol_from_smiles, to_canonical,
};
use crate::score::step_cost;
#[cfg(test)]
use crate::search::is_extracted_template;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum UpstreamScoreStatus {
Available,
NotApplicable,
ModelNotConfigured,
TargetParseFailed,
InferenceFailed,
OutputShapeMismatch,
}
pub struct ScoredRuleRef<'a> {
pub rule: &'a RetroRule,
pub source_rank: usize,
pub upstream_score: Option<f32>,
pub upstream_score_status: UpstreamScoreStatus,
}
pub enum ProposalMode {
Exhaustive,
BondIndexed { top_k: usize },
ScorerConditioned {
input: ScorerConditionedInput,
top_k: usize,
},
}
#[derive(Debug, Clone)]
pub struct ScorerConditionedInput {
pub scores: Vec<(usize, f32, usize)>,
pub status: UpstreamScoreStatus,
pub rules_offset: usize,
pub scorer_identity: String,
pub scorer_model_sha256: String,
}
pub struct ProposalConfig {
pub mode: ProposalMode,
}
impl Default for ProposalConfig {
fn default() -> Self {
Self {
mode: ProposalMode::Exhaustive,
}
}
}
#[derive(Debug, Clone, Default)]
pub struct CandidateFeatures {
pub values: Vec<f32>,
pub missing: Vec<bool>,
}
#[derive(Debug, Clone)]
pub struct CandidateSource {
pub template_id: String,
pub rule_name: String,
pub original_rank: usize,
pub upstream_score: Option<f32>,
pub upstream_score_status: UpstreamScoreStatus,
pub template_log_frequency_raw: Option<f32>,
pub base_step_cost: f64,
}
#[derive(Debug, Clone)]
pub struct ReactionCandidate {
pub candidate_id: String,
pub target_smiles: String,
pub precursor_smiles: Vec<String>,
pub sources: Vec<CandidateSource>,
pub source_template_count: usize,
pub best_upstream_score: Option<f32>,
pub best_upstream_rank: usize,
pub min_base_step_cost: f64,
pub max_template_frequency: Option<f32>,
pub mean_template_frequency: Option<f32>,
pub features: CandidateFeatures,
pub reranker_score: Option<f64>,
}
pub struct CandidatePool {
pub group_id: String,
pub target_id: String,
pub target_smiles: String,
pub candidates: Vec<ReactionCandidate>,
}
pub trait CandidateReranker: Send + Sync {
fn score_pool(&self, target: &str, candidates: &mut [ReactionCandidate]) -> anyhow::Result<()>;
}
#[derive(Debug, Clone, Copy, Default)]
pub struct TemplateTransformationFeatures {
pub mapped_atom_count: u32,
pub unmapped_atom_count: u32,
pub deleted_bond_count: u32,
pub added_bond_count: u32,
pub changed_bond_order_count: u32,
pub reaction_center_atom_count: u32,
pub extractable: bool,
}
type TransformationCacheKey = (String, String);
fn transformation_cache()
-> &'static Mutex<HashMap<TransformationCacheKey, TemplateTransformationFeatures>> {
static CACHE: OnceLock<Mutex<HashMap<TransformationCacheKey, TemplateTransformationFeatures>>> =
OnceLock::new();
CACHE.get_or_init(|| Mutex::new(HashMap::new()))
}
fn smirks_hash(smirks: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(smirks.as_bytes());
crate::sha256_hex(hasher.finalize())
}
struct AtomMapIndex {
by_map: rustc_hash::FxHashMap<u16, (usize, chematic::core::AtomIdx)>,
duplicate: bool,
}
fn index_atoms_by_map(mols: &[Molecule]) -> AtomMapIndex {
let mut by_map = rustc_hash::FxHashMap::default();
let mut duplicate = false;
for (mol_idx, mol) in mols.iter().enumerate() {
for (atom_idx, atom) in mol.atoms() {
if let Some(map_num) = atom.atom_map
&& by_map.insert(map_num, (mol_idx, atom_idx)).is_some()
{
duplicate = true;
}
}
}
AtomMapIndex { by_map, duplicate }
}
fn bond_orders_by_atom_map(
mols: &[Molecule],
index: &rustc_hash::FxHashMap<u16, (usize, chematic::core::AtomIdx)>,
) -> rustc_hash::FxHashMap<(u16, u16), chematic::core::BondOrder> {
let mut bonds = rustc_hash::FxHashMap::default();
for (&map_a, &(mol_idx, atom_idx)) in index {
let mol = &mols[mol_idx];
for (neighbor_idx, bond_idx) in mol.neighbors(atom_idx) {
let Some(map_b) = mol.atom(neighbor_idx).atom_map else {
continue;
};
if map_b <= map_a {
continue; }
bonds.insert((map_a, map_b), mol.bond(bond_idx).order);
}
}
bonds
}
#[derive(Debug, Clone, Copy, Default)]
struct ReactionCenterDiff {
deleted_bond_count: u32,
added_bond_count: u32,
changed_bond_order_count: u32,
reaction_center_atom_count: u32,
extractable: bool,
}
fn compute_reaction_center(rxn: &chematic::rxn::Reaction) -> ReactionCenterDiff {
let reactant_index = index_atoms_by_map(&rxn.reactants);
let product_index = index_atoms_by_map(&rxn.products);
if reactant_index.by_map.is_empty() || reactant_index.duplicate || product_index.duplicate {
return ReactionCenterDiff::default();
}
let reactant_bonds = bond_orders_by_atom_map(&rxn.reactants, &reactant_index.by_map);
let product_bonds = bond_orders_by_atom_map(&rxn.products, &product_index.by_map);
let mut deleted = 0u32;
let mut changed_order = 0u32;
let mut center_maps: std::collections::HashSet<u16> = std::collections::HashSet::new();
for (&key, r_order) in &reactant_bonds {
match product_bonds.get(&key) {
None => {
deleted += 1;
center_maps.insert(key.0);
center_maps.insert(key.1);
}
Some(p_order) if p_order != r_order => {
changed_order += 1;
center_maps.insert(key.0);
center_maps.insert(key.1);
}
_ => {}
}
}
let mut added = 0u32;
for &key in product_bonds.keys() {
if !reactant_bonds.contains_key(&key) {
added += 1;
center_maps.insert(key.0);
center_maps.insert(key.1);
}
}
for (&map_num, &(r_mol_idx, r_atom_idx)) in &reactant_index.by_map {
if let Some(&(p_mol_idx, p_atom_idx)) = product_index.by_map.get(&map_num) {
let r_atom = rxn.reactants[r_mol_idx].atom(r_atom_idx);
let p_atom = rxn.products[p_mol_idx].atom(p_atom_idx);
if r_atom.element != p_atom.element
|| r_atom.charge != p_atom.charge
|| r_atom.aromatic != p_atom.aromatic
{
center_maps.insert(map_num);
}
}
}
ReactionCenterDiff {
deleted_bond_count: deleted,
added_bond_count: added,
changed_bond_order_count: changed_order,
reaction_center_atom_count: center_maps.len() as u32,
extractable: true,
}
}
pub fn template_transformation_features(rule: &RetroRule) -> TemplateTransformationFeatures {
let cache_key = (rule.template_id.clone(), smirks_hash(&rule.smirks));
if let Some(cached) = transformation_cache().lock().unwrap().get(&cache_key) {
return *cached;
}
let features = if rule.smirks.is_empty() {
TemplateTransformationFeatures {
extractable: false,
..Default::default()
}
} else {
match chematic::rxn::parse_reaction(&rule.smirks) {
Ok(rxn) => {
let has_atom_map = rxn
.reactants
.iter()
.any(|m| m.atoms().any(|(_, a)| a.atom_map.is_some()));
if !has_atom_map {
TemplateTransformationFeatures {
extractable: false,
..Default::default()
}
} else {
let mapped_atom_count = rxn
.reactants
.iter()
.flat_map(|m| m.atoms())
.filter(|(_, a)| a.atom_map.is_some())
.count() as u32;
let unmapped_atom_count = rxn
.reactants
.iter()
.flat_map(|m| m.atoms())
.filter(|(_, a)| a.atom_map.is_none())
.count() as u32;
let center = compute_reaction_center(&rxn);
if !center.extractable {
TemplateTransformationFeatures {
mapped_atom_count,
unmapped_atom_count,
extractable: false,
..Default::default()
}
} else {
TemplateTransformationFeatures {
mapped_atom_count,
unmapped_atom_count,
deleted_bond_count: center.deleted_bond_count,
added_bond_count: center.added_bond_count,
changed_bond_order_count: center.changed_bond_order_count,
reaction_center_atom_count: center.reaction_center_atom_count,
extractable: true,
}
}
}
}
Err(_) => TemplateTransformationFeatures {
extractable: false,
..Default::default()
},
}
};
transformation_cache()
.lock()
.unwrap()
.insert(cache_key, features);
features
}
#[derive(Debug, Clone, Copy, Default)]
pub struct TransformationFeatureAggregate {
pub reaction_center_atom_count_min: u32,
pub reaction_center_atom_count_max: u32,
pub reaction_center_atom_count_mean: f32,
pub reaction_center_extractable_fraction: f32,
}
pub fn aggregate_transformation_features(
features: &[TemplateTransformationFeatures],
) -> TransformationFeatureAggregate {
if features.is_empty() {
return TransformationFeatureAggregate::default();
}
let extractable: Vec<&TemplateTransformationFeatures> =
features.iter().filter(|f| f.extractable).collect();
let (min, max, mean) = if extractable.is_empty() {
(0, 0, 0.0)
} else {
let counts: Vec<u32> = extractable
.iter()
.map(|f| f.reaction_center_atom_count)
.collect();
let min = *counts.iter().min().unwrap();
let max = *counts.iter().max().unwrap();
let mean = counts.iter().sum::<u32>() as f32 / counts.len() as f32;
(min, max, mean)
};
TransformationFeatureAggregate {
reaction_center_atom_count_min: min,
reaction_center_atom_count_max: max,
reaction_center_atom_count_mean: mean,
reaction_center_extractable_fraction: extractable.len() as f32 / features.len() as f32,
}
}
pub const FEATURE_SCHEMA_VERSION: u32 = 1;
pub const FEATURE_NAMES_V1: &[&str] = &[
"num_precursors",
"target_heavy_atom_count",
"precursor_heavy_atom_count_sum",
"precursor_heavy_atom_count_max",
"heavy_atom_retention_ratio",
"net_charge_balanced",
"no_heavy_atom_gain",
"source_template_count",
"reaction_center_atom_count_min",
"reaction_center_atom_count_max",
"reaction_center_atom_count_mean",
"reaction_center_extractable_fraction",
"min_base_step_cost",
"best_upstream_score",
"fraction_precursors_in_stock",
"all_precursors_in_stock",
"max_template_log_frequency",
"mean_template_log_frequency",
];
pub const FEATURE_GROUP1_LEN: usize = 14;
pub fn feature_index(name: &str) -> Option<usize> {
FEATURE_NAMES_V1.iter().position(|&n| n == name)
}
pub fn feature_schema_hash() -> String {
let mut hasher = Sha256::new();
hasher.update(b"renkin-retrospect-feature-schema-v1\0");
hasher.update(FEATURE_SCHEMA_VERSION.to_be_bytes());
hasher.update((FEATURE_NAMES_V1.len() as u64).to_be_bytes());
for name in FEATURE_NAMES_V1 {
hasher.update((name.len() as u64).to_be_bytes());
hasher.update(name.as_bytes());
}
format!("sha256:{}", crate::sha256_hex(hasher.finalize()))
}
fn heavy_atom_count_and_charge(mol: &Molecule) -> (u32, i64) {
let mut heavy = 0u32;
let mut charge = 0i64;
for (_, atom) in mol.atoms() {
if atom.element != chematic::core::Element::H {
heavy += 1;
}
charge += i64::from(atom.charge);
}
(heavy, charge)
}
pub fn extract_features(
candidate: &ReactionCandidate,
target_mol: &Molecule,
templates_by_id: &HashMap<String, &RetroRule>,
stock: Option<&crate::chem_env::ChemEnv>,
) -> CandidateFeatures {
let n = FEATURE_NAMES_V1.len();
let mut values = vec![0.0f32; n];
let mut missing = vec![false; n];
values[0] = candidate.precursor_smiles.len() as f32;
let mut precursor_mols: Vec<Molecule> = Vec::with_capacity(candidate.precursor_smiles.len());
let mut reparse_failed = false;
for smi in &candidate.precursor_smiles {
match mol_from_smiles(smi) {
Ok(m) => precursor_mols.push(m),
Err(_) => reparse_failed = true,
}
}
if reparse_failed || precursor_mols.is_empty() {
for m in missing.iter_mut().take(7).skip(1) {
*m = true;
}
} else {
let (target_heavy, target_charge) = heavy_atom_count_and_charge(target_mol);
let per_precursor: Vec<(u32, i64)> = precursor_mols
.iter()
.map(heavy_atom_count_and_charge)
.collect();
let precursor_heavy_sum: u32 = per_precursor.iter().map(|(h, _)| h).sum();
let precursor_charge_sum: i64 = per_precursor.iter().map(|(_, c)| c).sum();
let precursor_heavy_max = per_precursor.iter().map(|(h, _)| *h).max().unwrap_or(0);
values[1] = target_heavy as f32;
values[2] = precursor_heavy_sum as f32;
values[3] = precursor_heavy_max as f32;
if precursor_heavy_sum > 0 {
values[4] = target_heavy as f32 / precursor_heavy_sum as f32;
} else {
missing[4] = true;
}
values[5] = if precursor_charge_sum == target_charge {
1.0
} else {
0.0
};
values[6] = if precursor_heavy_sum >= target_heavy {
1.0
} else {
0.0
};
}
values[7] = candidate.source_template_count as f32;
let per_source_template_features: Vec<TemplateTransformationFeatures> = candidate
.sources
.iter()
.filter_map(|s| templates_by_id.get(&s.template_id).copied())
.map(template_transformation_features)
.collect();
let agg = aggregate_transformation_features(&per_source_template_features);
values[8] = agg.reaction_center_atom_count_min as f32;
values[9] = agg.reaction_center_atom_count_max as f32;
values[10] = agg.reaction_center_atom_count_mean;
values[11] = agg.reaction_center_extractable_fraction;
values[12] = candidate.min_base_step_cost as f32;
match candidate.best_upstream_score {
Some(s) => values[13] = s,
None => missing[13] = true,
}
match stock {
Some(stock) => {
let in_stock: Vec<bool> = candidate
.precursor_smiles
.iter()
.map(|smi| stock.is_building_block_smiles(smi))
.collect();
let n_in_stock = in_stock.iter().filter(|b| **b).count();
values[14] = n_in_stock as f32 / in_stock.len().max(1) as f32;
values[15] = if in_stock.iter().all(|b| *b) {
1.0
} else {
0.0
};
}
None => {
missing[14] = true;
missing[15] = true;
}
}
missing[16] = true;
missing[17] = true;
CandidateFeatures { values, missing }
}
pub fn index_rules_by_template_id(
rules: &[RetroRule],
) -> anyhow::Result<HashMap<String, &RetroRule>> {
let mut by_id: HashMap<String, &RetroRule> = HashMap::new();
for rule in rules {
if let Some(existing) = by_id.get(&rule.template_id) {
let conflicting = existing.name != rule.name
|| existing.smirks != rule.smirks
|| existing.weight != rule.weight
|| existing.required_elements != rule.required_elements;
if conflicting {
anyhow::bail!(
"template_id {:?} maps to two different rules: \
{{name: {:?}, smirks: {:?}, weight: {}}} vs \
{{name: {:?}, smirks: {:?}, weight: {}}}",
rule.template_id,
existing.name,
existing.smirks,
existing.weight,
rule.name,
rule.smirks,
rule.weight
);
}
continue;
}
by_id.insert(rule.template_id.clone(), rule);
}
Ok(by_id)
}
fn select_active_rules<'a>(
target_mol: &Molecule,
rules: &'a [RetroRule],
mode: &ProposalMode,
bond_index: Option<&TemplateBondIndex>,
) -> anyhow::Result<Vec<ScoredRuleRef<'a>>> {
match mode {
ProposalMode::Exhaustive => Ok(rules
.iter()
.enumerate()
.map(|(i, rule)| ScoredRuleRef {
rule,
source_rank: i,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
})
.collect()),
ProposalMode::BondIndexed { top_k } => {
let idx = bond_index.ok_or_else(|| {
anyhow::anyhow!(
"ProposalMode::BondIndexed requires a prepared TemplateBondIndex -- \
construct CandidateProposalContext::new(rules, true), or use the \
propose_one_step free function (which always prepares one for this \
mode); this never silently falls back to Exhaustive"
)
})?;
Ok(idx
.retrieve(target_mol, *top_k, rules)
.into_iter()
.enumerate()
.filter_map(|(rank, i)| {
rules.get(i).map(|rule| ScoredRuleRef {
rule,
source_rank: rank,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
})
})
.collect())
}
ProposalMode::ScorerConditioned { input, top_k } => {
if input.status != UpstreamScoreStatus::Available {
anyhow::bail!(
"ScorerConditioned proposal mode requires a successful \
scorer (status: Available), got {:?} -- failing closed \
rather than silently narrowing to zero file templates \
as if the scorer had succeeded",
input.status
);
}
let offset = input.rules_offset.min(rules.len());
let mut result: Vec<ScoredRuleRef> = rules
.iter()
.enumerate()
.take(offset)
.map(|(i, rule)| ScoredRuleRef {
rule,
source_rank: i,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
})
.collect();
let mut seen_indices = std::collections::HashSet::new();
let mut seen_ranks = std::collections::HashSet::new();
for &(rule_index, raw_logit, rank) in &input.scores {
if rule_index < offset || rule_index >= rules.len() {
anyhow::bail!(
"ScorerConditioned scores entry has rule_index {rule_index} \
out of bounds [{offset}, {})",
rules.len()
);
}
if !seen_indices.insert(rule_index) {
anyhow::bail!(
"ScorerConditioned scores contain duplicate rule_index {rule_index}"
);
}
if !seen_ranks.insert(rank) {
anyhow::bail!("ScorerConditioned scores contain duplicate rank {rank}");
}
if !raw_logit.is_finite() {
anyhow::bail!(
"ScorerConditioned scores entry for rule_index {rule_index} \
has a non-finite raw_logit ({raw_logit})"
);
}
}
let mut by_rank: Vec<&(usize, f32, usize)> = input.scores.iter().collect();
by_rank.sort_by_key(|s| s.2);
let mut rank_counter = offset;
for &(rule_index, raw_logit, _) in by_rank.into_iter().take(*top_k) {
if let Some(rule) = rules.get(rule_index) {
result.push(ScoredRuleRef {
rule,
source_rank: rank_counter,
upstream_score: Some(raw_logit),
upstream_score_status: UpstreamScoreStatus::Available,
});
rank_counter += 1;
}
}
Ok(result)
}
}
}
pub struct RawCandidate {
pub rule_name: String,
pub template_id: String,
pub rule_weight: f64,
pub original_rank: usize,
pub upstream_score: Option<f32>,
pub upstream_score_status: UpstreamScoreStatus,
pub precursors: Vec<PrecursorMol>,
}
type PerRuleProposal = (
Vec<RawCandidate>,
crate::ring_context::RingContextDiagnostics,
Vec<crate::spectator_bond::SpectatorBondLossFinding>,
Vec<crate::spectator_bond::GatedCandidateRecord>,
);
pub(crate) fn raw_propose(
target_mol: &Molecule,
target_smi: &str,
active_rules: &[ScoredRuleRef<'_>],
ring: crate::ring_context::RingContextArgs,
spectator_bond_policy: crate::spectator_bond::SpectatorBondPolicy,
) -> (
Vec<RawCandidate>,
crate::ring_context::RingContextDiagnostics,
Vec<crate::spectator_bond::SpectatorBondLossFinding>,
Vec<crate::spectator_bond::GatedCandidateRecord>,
) {
use crate::spectator_bond::SpectatorBondPolicy;
let target_elem_mask: u64 = crate::search::elem_mask_from_smiles(target_smi);
#[cfg(not(target_arch = "wasm32"))]
let iter = active_rules.par_iter();
#[cfg(target_arch = "wasm32")]
let iter = active_rules.iter();
let per_rule: Vec<PerRuleProposal> = iter
.filter(|r| {
r.rule.required_elements == 0
|| (target_elem_mask & r.rule.required_elements == r.rule.required_elements)
})
.map(|r| {
let mut diag = crate::ring_context::RingContextDiagnostics::default();
let mut candidates = crate::ring_context::apply_retro_with_policy(
target_mol,
r.rule,
&ring.config,
&mut diag,
)
.into_iter()
.filter(|precs| !precs.is_empty() && !precs.iter().any(|p| p.smiles == target_smi))
.map(|precs| RawCandidate {
rule_name: r.rule.name.to_string(),
template_id: r.rule.template_id.clone(),
rule_weight: r.rule.weight,
original_rank: r.source_rank,
upstream_score: r.upstream_score,
upstream_score_status: r.upstream_score_status,
precursors: precs,
})
.collect::<Vec<_>>();
let sbl_findings = if spectator_bond_policy != SpectatorBondPolicy::Off {
let mut findings = crate::spectator_bond::detect_case_a(target_mol, r.rule);
findings.extend(crate::spectator_bond::detect_case_b(target_mol, r.rule));
findings
} else {
Vec::new()
};
let mut gated_out = Vec::new();
if spectator_bond_policy == SpectatorBondPolicy::Gated {
let signatures: Vec<Vec<String>> = candidates
.iter()
.map(|c| {
let mut sig: Vec<String> =
c.precursors.iter().map(|p| p.smiles.clone()).collect();
sig.sort_unstable();
sig
})
.collect();
let verdicts =
crate::spectator_bond::gate_candidates(target_mol, r.rule, &signatures);
let mut kept = Vec::with_capacity(candidates.len());
for (candidate, verdict) in candidates.into_iter().zip(verdicts) {
match verdict {
crate::spectator_bond::SpectatorBondGateVerdict::Rejected { findings } => {
gated_out.push(crate::spectator_bond::GatedCandidateRecord {
rule_name: candidate.rule_name.clone(),
template_id: candidate.template_id.clone(),
precursor_smiles: candidate
.precursors
.iter()
.map(|p| p.smiles.clone())
.collect(),
findings,
});
}
_ => kept.push(candidate),
}
}
candidates = kept;
}
(candidates, diag, sbl_findings, gated_out)
})
.collect();
let mut merged_diag = crate::ring_context::RingContextDiagnostics::default();
let mut merged_sbl_findings = Vec::new();
let mut merged_gated_out = Vec::new();
let raw: Vec<RawCandidate> = per_rule
.into_iter()
.flat_map(|(candidates, diag, sbl_findings, gated_out)| {
merged_diag.merge(&diag);
merged_sbl_findings.extend(sbl_findings);
merged_gated_out.extend(gated_out);
candidates
})
.collect();
(raw, merged_diag, merged_sbl_findings, merged_gated_out)
}
fn hash_string_sequence(hasher: &mut Sha256, values: &[String]) {
hasher.update((values.len() as u64).to_be_bytes());
for value in values {
let bytes = value.as_bytes();
hasher.update((bytes.len() as u64).to_be_bytes());
hasher.update(bytes);
}
}
pub(crate) fn candidate_id_for(canonical_target: &str, precursor_smiles: &[String]) -> String {
let mut hasher = Sha256::new();
hasher.update(b"renkin-retrospect-candidate-v1\0");
hash_string_sequence(&mut hasher, &[canonical_target.to_string()]);
hasher.update(b"\0precursors\0");
hash_string_sequence(&mut hasher, precursor_smiles);
format!("sha256:{}", crate::sha256_hex(hasher.finalize()))
}
fn merge_duplicate_sources(sources: Vec<CandidateSource>) -> anyhow::Result<Vec<CandidateSource>> {
let mut order: Vec<(String, String)> = Vec::new();
let mut by_key: HashMap<(String, String), CandidateSource> = HashMap::new();
for s in sources {
let key = (s.template_id.clone(), s.rule_name.clone());
match by_key.get_mut(&key) {
None => {
order.push(key.clone());
by_key.insert(key, s);
}
Some(existing) => {
if existing.template_log_frequency_raw != s.template_log_frequency_raw {
anyhow::bail!(
"duplicate source (template_id={:?}, rule_name={:?}) reports \
inconsistent template_log_frequency_raw ({:?} vs {:?}) for what \
must be the same rule",
key.0,
key.1,
existing.template_log_frequency_raw,
s.template_log_frequency_raw
);
}
if existing.upstream_score_status != s.upstream_score_status {
anyhow::bail!(
"duplicate source (template_id={:?}, rule_name={:?}) reports \
inconsistent upstream_score_status ({:?} vs {:?}) for what must \
be the same rule",
key.0,
key.1,
existing.upstream_score_status,
s.upstream_score_status
);
}
existing.upstream_score = match (existing.upstream_score, s.upstream_score) {
(Some(a), Some(b)) => Some(a.max(b)),
(Some(a), None) | (None, Some(a)) => Some(a),
(None, None) => None,
};
existing.original_rank = existing.original_rank.min(s.original_rank);
existing.base_step_cost = existing.base_step_cost.min(s.base_step_cost);
}
}
}
Ok(order
.into_iter()
.map(|key| by_key.remove(&key).expect("key was just inserted above"))
.collect())
}
pub(crate) fn merge_into_candidates(
canonical_target: &str,
raw: &[RawCandidate],
) -> anyhow::Result<Vec<ReactionCandidate>> {
let mut order: Vec<String> = Vec::new();
let mut precursors_by_id: HashMap<String, Vec<String>> = HashMap::new();
let mut sources_by_id: HashMap<String, Vec<CandidateSource>> = HashMap::new();
for proposal in raw {
let mut precursor_smiles: Vec<String> = proposal
.precursors
.iter()
.map(|p| p.smiles.clone())
.collect();
precursor_smiles.sort_unstable();
let candidate_id = candidate_id_for(canonical_target, &precursor_smiles);
let base_step_cost = step_cost(
&proposal
.precursors
.iter()
.map(|p| &p.mol)
.collect::<Vec<_>>(),
);
let source = CandidateSource {
template_id: proposal.template_id.clone(),
rule_name: proposal.rule_name.clone(),
original_rank: proposal.original_rank,
upstream_score: proposal.upstream_score,
upstream_score_status: proposal.upstream_score_status,
template_log_frequency_raw: Some(proposal.rule_weight as f32),
base_step_cost,
};
if !sources_by_id.contains_key(&candidate_id) {
order.push(candidate_id.clone());
precursors_by_id.insert(candidate_id.clone(), precursor_smiles);
}
sources_by_id.entry(candidate_id).or_default().push(source);
}
order
.into_iter()
.map(|id| {
let raw_sources = sources_by_id
.remove(&id)
.expect("id was just inserted above");
let precursor_smiles = precursors_by_id
.remove(&id)
.expect("id was just inserted above");
let mut sources = merge_duplicate_sources(raw_sources)?;
sources.sort_by(|a, b| {
b.upstream_score
.partial_cmp(&a.upstream_score)
.unwrap_or(std::cmp::Ordering::Equal)
.then(
a.base_step_cost
.partial_cmp(&b.base_step_cost)
.unwrap_or(std::cmp::Ordering::Equal),
)
.then(a.original_rank.cmp(&b.original_rank))
.then(a.template_id.cmp(&b.template_id))
.then(a.rule_name.cmp(&b.rule_name))
});
let best_upstream_score = sources
.iter()
.filter_map(|s| s.upstream_score)
.fold(None, |acc: Option<f32>, v| {
Some(acc.map_or(v, |a| a.max(v)))
});
let best_upstream_rank = match best_upstream_score {
Some(best) => sources
.iter()
.filter(|s| s.upstream_score == Some(best))
.map(|s| s.original_rank)
.min()
.unwrap_or(0),
None => sources.iter().map(|s| s.original_rank).min().unwrap_or(0),
};
let min_base_step_cost = sources
.iter()
.map(|s| s.base_step_cost)
.fold(f64::INFINITY, f64::min);
let freqs: Vec<f32> = sources
.iter()
.filter_map(|s| s.template_log_frequency_raw)
.collect();
let max_template_frequency = freqs.iter().copied().fold(None, |acc: Option<f32>, v| {
Some(acc.map_or(v, |a| a.max(v)))
});
let mean_template_frequency = if freqs.is_empty() {
None
} else {
Some(freqs.iter().sum::<f32>() / freqs.len() as f32)
};
Ok(ReactionCandidate {
candidate_id: id,
target_smiles: canonical_target.to_string(),
precursor_smiles,
source_template_count: sources.len(),
sources,
best_upstream_score,
best_upstream_rank,
min_base_step_cost,
max_template_frequency,
mean_template_frequency,
features: CandidateFeatures::default(),
reranker_score: None,
})
})
.collect()
}
#[cfg(feature = "perf-instrumentation")]
static PHASE_SELECT_NANOS: AtomicU64 = AtomicU64::new(0);
#[cfg(feature = "perf-instrumentation")]
static PHASE_RAW_PROPOSE_NANOS: AtomicU64 = AtomicU64::new(0);
#[cfg(feature = "perf-instrumentation")]
static PHASE_MERGE_NANOS: AtomicU64 = AtomicU64::new(0);
#[derive(Debug, Clone, Copy, Default, serde::Serialize)]
pub struct ProposePhaseNanos {
pub select: u64,
pub raw_propose: u64,
pub merge: u64,
}
#[cfg(feature = "perf-instrumentation")]
pub fn propose_phase_nanos() -> ProposePhaseNanos {
ProposePhaseNanos {
select: PHASE_SELECT_NANOS.load(Ordering::Relaxed),
raw_propose: PHASE_RAW_PROPOSE_NANOS.load(Ordering::Relaxed),
merge: PHASE_MERGE_NANOS.load(Ordering::Relaxed),
}
}
#[cfg(not(feature = "perf-instrumentation"))]
pub fn propose_phase_nanos() -> ProposePhaseNanos {
ProposePhaseNanos::default()
}
#[cfg(feature = "perf-instrumentation")]
pub fn reset_propose_phase_nanos() {
PHASE_SELECT_NANOS.store(0, Ordering::Relaxed);
PHASE_RAW_PROPOSE_NANOS.store(0, Ordering::Relaxed);
PHASE_MERGE_NANOS.store(0, Ordering::Relaxed);
}
#[cfg(not(feature = "perf-instrumentation"))]
pub fn reset_propose_phase_nanos() {}
pub struct CandidateProposalContext<'a> {
rules: &'a [RetroRule],
bond_index: Option<TemplateBondIndex>,
}
impl<'a> CandidateProposalContext<'a> {
pub fn new(rules: &'a [RetroRule], prepare_bond_index: bool) -> Self {
Self {
rules,
bond_index: prepare_bond_index.then(|| TemplateBondIndex::build(rules)),
}
}
pub fn propose_one_step(
&self,
group_id: &str,
target_smiles: &str,
config: &ProposalConfig,
) -> anyhow::Result<CandidatePool> {
let target_mol = mol_from_smiles(target_smiles)?;
let canonical_target = to_canonical(&target_mol);
#[cfg(feature = "perf-instrumentation")]
let t0 = Instant::now();
let active_rules = select_active_rules(
&target_mol,
self.rules,
&config.mode,
self.bond_index.as_ref(),
)?;
#[cfg(feature = "perf-instrumentation")]
PHASE_SELECT_NANOS.fetch_add(t0.elapsed().as_nanos() as u64, Ordering::Relaxed);
#[cfg(feature = "perf-instrumentation")]
let t1 = Instant::now();
let (raw, _ring_diag, _sbl_findings, _gated_out) = raw_propose(
&target_mol,
&canonical_target,
&active_rules,
crate::ring_context::RingContextArgs::default(),
crate::spectator_bond::SpectatorBondPolicy::Off,
);
#[cfg(feature = "perf-instrumentation")]
PHASE_RAW_PROPOSE_NANOS.fetch_add(t1.elapsed().as_nanos() as u64, Ordering::Relaxed);
#[cfg(feature = "perf-instrumentation")]
let t2 = Instant::now();
let candidates = merge_into_candidates(&canonical_target, &raw)?;
#[cfg(feature = "perf-instrumentation")]
PHASE_MERGE_NANOS.fetch_add(t2.elapsed().as_nanos() as u64, Ordering::Relaxed);
Ok(CandidatePool {
group_id: group_id.to_string(),
target_id: canonical_target.clone(),
target_smiles: canonical_target,
candidates,
})
}
}
pub fn propose_one_step(
group_id: &str,
target_smiles: &str,
rules: &[RetroRule],
config: &ProposalConfig,
) -> anyhow::Result<CandidatePool> {
let prepare_bond_index = matches!(config.mode, ProposalMode::BondIndexed { .. });
CandidateProposalContext::new(rules, prepare_bond_index).propose_one_step(
group_id,
target_smiles,
config,
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::chem_env::default_rules;
fn rule(name: &str, smirks: &str) -> RetroRule {
RetroRule {
name: name.to_string(),
template_id: format!("rule:{name}"),
smirks: smirks.to_string(),
weight: 1.0,
required_elements: 0,
}
}
fn extracted_rule(idx: usize, smirks: &str, weight: f64) -> RetroRule {
RetroRule {
name: format!("extracted_{idx}"),
template_id: format!("smirks-sha256:fake{idx}"),
smirks: smirks.to_string(),
weight,
required_elements: 0,
}
}
fn scorer_input(
scores: Vec<(usize, f32, usize)>,
rules_offset: usize,
) -> ScorerConditionedInput {
ScorerConditionedInput {
scores,
status: UpstreamScoreStatus::Available,
rules_offset,
scorer_identity: "test-scorer".to_string(),
scorer_model_sha256: "sha256:test".to_string(),
}
}
#[test]
fn bond_index_is_built_once_and_reused_across_two_targets() {
let rules = default_rules();
let ctx = CandidateProposalContext::new(&rules, true);
let index_ptr_before = ctx
.bond_index
.as_ref()
.map(|idx| idx as *const TemplateBondIndex);
assert!(
index_ptr_before.is_some(),
"prepare_bond_index: true must build the index up front"
);
let config = ProposalConfig {
mode: ProposalMode::BondIndexed { top_k: 0 },
};
let _pool_a = ctx.propose_one_step("g1", "CCO", &config).unwrap();
let _pool_b = ctx.propose_one_step("g2", "c1ccccc1", &config).unwrap();
let index_ptr_after = ctx
.bond_index
.as_ref()
.map(|idx| idx as *const TemplateBondIndex);
assert_eq!(
index_ptr_before, index_ptr_after,
"the same TemplateBondIndex instance must persist across both targets"
);
}
#[test]
fn bond_indexed_mode_without_a_prepared_index_is_a_hard_error() {
let rules = default_rules();
let ctx = CandidateProposalContext::new(&rules, false);
assert!(ctx.bond_index.is_none());
let config = ProposalConfig {
mode: ProposalMode::BondIndexed { top_k: 0 },
};
assert!(ctx.propose_one_step("g1", "CCO", &config).is_err());
}
#[test]
fn exhaustive_mode_does_not_need_a_bond_index() {
let rules = default_rules();
let ctx = CandidateProposalContext::new(&rules, false);
assert!(ctx.bond_index.is_none());
let pool = ctx
.propose_one_step("g1", "CCO", &ProposalConfig::default())
.unwrap();
assert!(!pool.candidates.is_empty());
}
#[test]
fn scorer_conditioned_mode_is_unaffected_by_context_index_preparation() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(extracted_rule(0, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
let config = ProposalConfig {
mode: ProposalMode::ScorerConditioned {
input: scorer_input(vec![(n_handcrafted, 0.9, 0)], n_handcrafted),
top_k: 1,
},
};
let without_index = CandidateProposalContext::new(&rules, false)
.propose_one_step("g1", "CCCC", &config)
.unwrap();
let with_index = CandidateProposalContext::new(&rules, true)
.propose_one_step("g1", "CCCC", &config)
.unwrap();
let mut without_ids: Vec<&str> = without_index
.candidates
.iter()
.map(|c| c.candidate_id.as_str())
.collect();
let mut with_ids: Vec<&str> = with_index
.candidates
.iter()
.map(|c| c.candidate_id.as_str())
.collect();
without_ids.sort();
with_ids.sort();
assert_eq!(without_ids, with_ids);
}
fn assert_context_api_matches_legacy_api(
rules: &[RetroRule],
target: &str,
config: &ProposalConfig,
) {
let legacy_pool = propose_one_step("g1", target, rules, config).unwrap();
let context_pool = CandidateProposalContext::new(rules, true)
.propose_one_step("g1", target, config)
.unwrap();
let summarize = |pool: &CandidatePool| {
let mut summary: Vec<(String, Vec<String>)> = pool
.candidates
.iter()
.map(|c| {
let mut rule_names: Vec<String> =
c.sources.iter().map(|s| s.rule_name.clone()).collect();
rule_names.sort();
(c.candidate_id.clone(), rule_names)
})
.collect();
summary.sort_by(|a, b| a.0.cmp(&b.0));
summary
};
assert_eq!(
summarize(&legacy_pool),
summarize(&context_pool),
"the legacy single-call API and the context-reuse API must produce \
identical candidate IDs and source provenance for the same input"
);
}
#[test]
fn context_api_matches_legacy_api_for_bond_indexed_top_k_zero() {
let rules = default_rules();
let config = ProposalConfig {
mode: ProposalMode::BondIndexed { top_k: 0 },
};
assert_context_api_matches_legacy_api(&rules, "CC(=O)c1ccccc1", &config);
}
#[test]
fn context_api_matches_legacy_api_for_bond_indexed_top_k_positive() {
let rules = default_rules();
let config = ProposalConfig {
mode: ProposalMode::BondIndexed { top_k: 2 },
};
assert_context_api_matches_legacy_api(&rules, "CC(=O)c1ccccc1", &config);
}
#[test]
fn exhaustive_mode_tries_all_rules() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let active =
select_active_rules(&target_mol, &rules, &ProposalMode::Exhaustive, None).unwrap();
assert_eq!(active.len(), rules.len());
for r in &active {
assert_eq!(r.upstream_score_status, UpstreamScoreStatus::NotApplicable);
assert!(r.upstream_score.is_none());
}
}
#[test]
fn scorer_conditioned_mode_tries_only_top_k_file_templates() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(extracted_rule(0, "[C:1][C:2]>>[C:1].[C:2]", 3.0));
rules.push(extracted_rule(1, "[C:1][C:2]>>[C:1].[C:2]", 2.0));
rules.push(extracted_rule(2, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
let scores = vec![
(n_handcrafted, 0.9, 0),
(n_handcrafted + 1, 0.5, 1),
(n_handcrafted + 2, 0.1, 2),
];
let mode = ProposalMode::ScorerConditioned {
input: scorer_input(scores, n_handcrafted),
top_k: 1,
};
let target = "CCCC";
let target_mol = mol_from_smiles(target).unwrap();
let active = select_active_rules(&target_mol, &rules, &mode, None).unwrap();
assert_eq!(active.len(), n_handcrafted + 1);
let extracted_in_active: Vec<&ScoredRuleRef> = active
.iter()
.filter(|r| is_extracted_template(&r.rule.name))
.collect();
assert_eq!(extracted_in_active.len(), 1);
assert_eq!(extracted_in_active[0].rule.name, "extracted_0");
assert_eq!(extracted_in_active[0].upstream_score, Some(0.9));
assert_eq!(
extracted_in_active[0].upstream_score_status,
UpstreamScoreStatus::Available
);
}
#[test]
fn handcrafted_rules_always_included_regardless_of_scorer() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(extracted_rule(0, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
let mode_zero_k = ProposalMode::ScorerConditioned {
input: scorer_input(vec![(n_handcrafted, 0.9, 0)], n_handcrafted),
top_k: 0, };
let target_mol = mol_from_smiles("CCCC").unwrap();
let active = select_active_rules(&target_mol, &rules, &mode_zero_k, None).unwrap();
assert_eq!(
active.len(),
n_handcrafted,
"hand-crafted rules must all still be present"
);
assert!(active.iter().all(|r| !is_extracted_template(&r.rule.name)));
}
#[test]
fn scorer_conditioned_classifies_by_rules_offset_position_not_name_prefix() {
let handcrafted = rule("totally_handcrafted", "[C:1][C:2]>>[C:1].[C:2]");
let file_template_with_plain_name = rule("not_prefixed_at_all", "[C:1][C:2]>>[C:1].[C:2]");
let rules = vec![handcrafted, file_template_with_plain_name];
let rules_offset = 1;
let mode = ProposalMode::ScorerConditioned {
input: scorer_input(vec![(1, 0.5, 0)], rules_offset),
top_k: 1,
};
let target_mol = mol_from_smiles("CCCC").unwrap();
let active = select_active_rules(&target_mol, &rules, &mode, None).unwrap();
assert_eq!(
active.len(),
2,
"handcrafted-by-position + 1 scored file template"
);
let scored: Vec<&ScoredRuleRef> = active
.iter()
.filter(|r| r.upstream_score_status == UpstreamScoreStatus::Available)
.collect();
assert_eq!(scored.len(), 1);
assert_eq!(scored[0].rule.name, "not_prefixed_at_all");
}
#[test]
fn exhaustive_and_scorer_conditioned_candidate_sets_can_differ() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(extracted_rule(0, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
rules.push(extracted_rule(1, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
let target = "CCCC";
let target_mol = mol_from_smiles(target).unwrap();
let exhaustive =
select_active_rules(&target_mol, &rules, &ProposalMode::Exhaustive, None).unwrap();
let conditioned = select_active_rules(
&target_mol,
&rules,
&ProposalMode::ScorerConditioned {
input: scorer_input(
vec![(n_handcrafted, 0.9, 0), (n_handcrafted + 1, 0.1, 1)],
n_handcrafted,
),
top_k: 1,
},
None,
)
.unwrap();
assert!(
conditioned.len() < exhaustive.len(),
"scorer-conditioned selection must be a strict subset here, not just reordered"
);
}
#[test]
fn original_rank_matches_within_mode_rank() {
let rules = default_rules();
let target_mol = mol_from_smiles("CC(=O)c1ccccc1").unwrap();
let active =
select_active_rules(&target_mol, &rules, &ProposalMode::Exhaustive, None).unwrap();
for (i, r) in active.iter().enumerate() {
assert_eq!(r.source_rank, i);
}
}
#[test]
fn scorer_failure_status_is_not_applicable_not_frequency() {
let rules = default_rules();
let target_mol = mol_from_smiles("CC(=O)c1ccccc1").unwrap();
let active =
select_active_rules(&target_mol, &rules, &ProposalMode::Exhaustive, None).unwrap();
for r in &active {
assert!(r.upstream_score.is_none());
assert_eq!(r.upstream_score_status, UpstreamScoreStatus::NotApplicable);
}
}
#[test]
fn scorer_conditioned_with_empty_scores_but_available_status_selects_only_handcrafted() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(extracted_rule(0, "[C:1][C:2]>>[C:1].[C:2]", 3.0));
rules.push(extracted_rule(1, "[C:1][C:2]>>[C:1].[C:2]", 2.0));
let mode = ProposalMode::ScorerConditioned {
input: scorer_input(Vec::new(), n_handcrafted),
top_k: 10,
};
let target_mol = mol_from_smiles("CCCC").unwrap();
let active = select_active_rules(&target_mol, &rules, &mode, None).unwrap();
assert_eq!(
active.len(),
n_handcrafted,
"must select only hand-crafted rules, zero file templates"
);
assert!(active.iter().all(|r| !is_extracted_template(&r.rule.name)));
assert!(
active.iter().all(|r| r.upstream_score.is_none()),
"no frequency-derived value may appear in upstream_score when the scorer produced no scores"
);
}
#[test]
fn scorer_conditioned_fails_closed_when_status_is_not_available() {
let rules = default_rules();
let n_handcrafted = rules.len();
for status in [
UpstreamScoreStatus::ModelNotConfigured,
UpstreamScoreStatus::TargetParseFailed,
UpstreamScoreStatus::InferenceFailed,
UpstreamScoreStatus::OutputShapeMismatch,
] {
let mode = ProposalMode::ScorerConditioned {
input: ScorerConditionedInput {
scores: Vec::new(),
status,
rules_offset: n_handcrafted,
scorer_identity: "test-scorer".to_string(),
scorer_model_sha256: "sha256:test".to_string(),
},
top_k: 10,
};
let target_mol = mol_from_smiles("CCCC").unwrap();
assert!(
select_active_rules(&target_mol, &rules, &mode, None).is_err(),
"status {status:?} must fail closed, not silently succeed with zero file templates"
);
}
}
#[test]
fn scorer_conditioned_rejects_out_of_bounds_rule_index() {
let rules = default_rules();
let n_handcrafted = rules.len();
let mode = ProposalMode::ScorerConditioned {
input: scorer_input(vec![(rules.len() + 5, 0.5, 0)], n_handcrafted),
top_k: 10,
};
let target_mol = mol_from_smiles("CCCC").unwrap();
assert!(select_active_rules(&target_mol, &rules, &mode, None).is_err());
}
#[test]
fn scorer_conditioned_rejects_rule_index_inside_handcrafted_prefix() {
let rules = default_rules();
let n_handcrafted = rules.len();
assert!(n_handcrafted > 0, "fixture assumption");
let mode = ProposalMode::ScorerConditioned {
input: scorer_input(vec![(0, 0.5, 0)], n_handcrafted),
top_k: 10,
};
let target_mol = mol_from_smiles("CCCC").unwrap();
assert!(
select_active_rules(&target_mol, &rules, &mode, None).is_err(),
"a scored rule_index inside [0, rules_offset) must be rejected"
);
}
#[test]
fn scorer_conditioned_rejects_duplicate_rule_index() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(extracted_rule(0, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
let mode = ProposalMode::ScorerConditioned {
input: scorer_input(
vec![(n_handcrafted, 0.5, 0), (n_handcrafted, 0.6, 1)],
n_handcrafted,
),
top_k: 10,
};
let target_mol = mol_from_smiles("CCCC").unwrap();
assert!(select_active_rules(&target_mol, &rules, &mode, None).is_err());
}
#[test]
fn scorer_conditioned_rejects_duplicate_rank() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(extracted_rule(0, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
rules.push(extracted_rule(1, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
let mode = ProposalMode::ScorerConditioned {
input: scorer_input(
vec![(n_handcrafted, 0.5, 0), (n_handcrafted + 1, 0.6, 0)],
n_handcrafted,
),
top_k: 10,
};
let target_mol = mol_from_smiles("CCCC").unwrap();
assert!(select_active_rules(&target_mol, &rules, &mode, None).is_err());
}
#[test]
fn scorer_conditioned_rejects_non_finite_raw_logit() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(extracted_rule(0, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
for bad in [f32::NAN, f32::INFINITY, f32::NEG_INFINITY] {
let mode = ProposalMode::ScorerConditioned {
input: scorer_input(vec![(n_handcrafted, bad, 0)], n_handcrafted),
top_k: 10,
};
let target_mol = mol_from_smiles("CCCC").unwrap();
assert!(
select_active_rules(&target_mol, &rules, &mode, None).is_err(),
"raw_logit {bad} must be rejected"
);
}
}
#[test]
fn scorer_conditioned_tie_break_is_rank_ascending() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(extracted_rule(0, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
rules.push(extracted_rule(1, "[C:1][C:2]>>[C:1].[C:2]", 1.0));
let mode = ProposalMode::ScorerConditioned {
input: scorer_input(
vec![(n_handcrafted + 1, 0.1, 1), (n_handcrafted, 0.9, 0)],
n_handcrafted,
),
top_k: 1,
};
let target_mol = mol_from_smiles("CCCC").unwrap();
let active = select_active_rules(&target_mol, &rules, &mode, None).unwrap();
let scored: Vec<&ScoredRuleRef> = active
.iter()
.filter(|r| r.upstream_score_status == UpstreamScoreStatus::Available)
.collect();
assert_eq!(scored.len(), 1);
assert_eq!(
scored[0].upstream_score,
Some(0.9),
"rank 0 (lowest rank) must be selected by top_k=1, not insertion order"
);
}
#[test]
fn raw_propose_matches_independent_reference_filter() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1"; let target_mol = mol_from_smiles(target).unwrap();
let canon_target = to_canonical(&target_mol);
let target_elem_mask = crate::search::elem_mask_from_smiles(&canon_target);
let mut reference: Vec<(String, Vec<String>)> = Vec::new();
for rule in &rules {
if !(rule.required_elements == 0
|| (target_elem_mask & rule.required_elements == rule.required_elements))
{
continue;
}
for precs in apply_retro(&target_mol, rule) {
if precs.is_empty() || precs.iter().any(|p| p.smiles == canon_target) {
continue;
}
reference.push((
rule.name.clone(),
precs.iter().map(|p| p.smiles.clone()).collect(),
));
}
}
let active_rules =
select_active_rules(&target_mol, &rules, &ProposalMode::Exhaustive, None).unwrap();
let (got, _ring_diag, _sbl_findings, _gated_out) = raw_propose(
&target_mol,
&canon_target,
&active_rules,
crate::ring_context::RingContextArgs::default(),
crate::spectator_bond::SpectatorBondPolicy::Off,
);
let mut got_pairs: Vec<(String, Vec<String>)> = got
.into_iter()
.map(|p| {
(
p.rule_name,
p.precursors.iter().map(|pm| pm.smiles.clone()).collect(),
)
})
.collect();
let mut reference_sorted = reference;
reference_sorted.sort();
got_pairs.sort();
assert_eq!(got_pairs, reference_sorted);
assert!(
!got_pairs.is_empty(),
"expected at least one match for acetophenone"
);
}
fn scored_rule_ref(rule: &RetroRule) -> ScoredRuleRef<'_> {
ScoredRuleRef {
rule,
source_rank: 0,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
}
}
fn extracted_824_rule() -> RetroRule {
RetroRule {
name: "extracted_824".to_string(),
template_id: "rule:extracted_824".to_string(),
smirks: "[C:5]-[O:6]-[C:3](=[O:4])-[NH:2]-[C:1]>>[C:1]-[N:2]=[C:3]=[O:4].[C:5]-[OH:6]"
.to_string(),
..Default::default()
}
}
#[test]
fn raw_propose_spectator_bond_policy_off_by_default_produces_none() {
let target_mol = mol_from_smiles("O=C2NCC(O2)Cc1ccccc1").unwrap();
let rule = extracted_824_rule();
let active_rules = [scored_rule_ref(&rule)];
let (_raw, _ring_diag, sbl_findings, gated_out) = raw_propose(
&target_mol,
"O=C2NCC(O2)Cc1ccccc1",
&active_rules,
crate::ring_context::RingContextArgs::default(),
crate::spectator_bond::SpectatorBondPolicy::Off,
);
assert!(
sbl_findings.is_empty(),
"policy Off must produce zero findings, even for a rule/target pair the detector \
would flag if enabled: {sbl_findings:?}"
);
assert!(gated_out.is_empty());
}
#[test]
fn raw_propose_spectator_bond_policy_diagnostics_only_detects_known_positive_control() {
let target_mol = mol_from_smiles("O=C2NCC(O2)Cc1ccccc1").unwrap();
let rule = extracted_824_rule();
let active_rules = [scored_rule_ref(&rule)];
let (raw, _ring_diag, sbl_findings, gated_out) = raw_propose(
&target_mol,
"O=C2NCC(O2)Cc1ccccc1",
&active_rules,
crate::ring_context::RingContextArgs::default(),
crate::spectator_bond::SpectatorBondPolicy::DiagnosticsOnly,
);
assert!(
!sbl_findings.is_empty(),
"policy DiagnosticsOnly must surface the real extracted_824 finding through \
raw_propose's own wiring, not just the detector functions directly"
);
assert_eq!(
sbl_findings[0].rule_name, "extracted_824",
"finding must be attributed to the rule that produced it"
);
assert!(
!raw.is_empty(),
"DiagnosticsOnly must never exclude a candidate"
);
assert!(gated_out.is_empty(), "DiagnosticsOnly must never gate");
}
#[test]
fn raw_propose_spectator_bond_policy_diagnostics_only_stays_empty_for_clean_rule() {
let target_mol = mol_from_smiles("CC(=O)NC").unwrap();
let rule = RetroRule {
name: "amide_formation_retro".to_string(),
template_id: "rule:amide_formation_retro".to_string(),
smirks: "[C:1](=[O:2])-[N:3]>>[C:1](=[O:2])[OH].[N:3]".to_string(),
..Default::default()
};
let active_rules = [scored_rule_ref(&rule)];
let (_raw, _ring_diag, sbl_findings, gated_out) = raw_propose(
&target_mol,
"CC(=O)NC",
&active_rules,
crate::ring_context::RingContextArgs::default(),
crate::spectator_bond::SpectatorBondPolicy::DiagnosticsOnly,
);
assert!(sbl_findings.is_empty());
assert!(gated_out.is_empty());
}
#[test]
fn raw_propose_spectator_bond_policy_gated_excludes_known_positive_control() {
let target_mol = mol_from_smiles("O=C2NCC(O2)Cc1ccccc1").unwrap();
let rule = extracted_824_rule();
let active_rules = [scored_rule_ref(&rule)];
let (raw, _ring_diag, sbl_findings, gated_out) = raw_propose(
&target_mol,
"O=C2NCC(O2)Cc1ccccc1",
&active_rules,
crate::ring_context::RingContextArgs::default(),
crate::spectator_bond::SpectatorBondPolicy::Gated,
);
assert!(
raw.iter().all(|c| c.rule_name != "extracted_824"),
"Gated must exclude the known-defective extracted_824 candidate, not just report it: \
{:?}",
raw.iter().map(|c| &c.rule_name).collect::<Vec<_>>()
);
assert!(
!sbl_findings.is_empty(),
"Gated must still record the finding -- policy changes the verdict, never the \
finding set"
);
assert_eq!(
gated_out.len(),
1,
"exactly one candidate excluded: {gated_out:?}"
);
assert_eq!(gated_out[0].rule_name, "extracted_824");
assert!(!gated_out[0].findings.is_empty());
}
#[test]
fn raw_propose_spectator_bond_policy_gated_keeps_clean_rule() {
let target_mol = mol_from_smiles("CC(=O)NC").unwrap();
let rule = RetroRule {
name: "amide_formation_retro".to_string(),
template_id: "rule:amide_formation_retro".to_string(),
smirks: "[C:1](=[O:2])-[N:3]>>[C:1](=[O:2])[OH].[N:3]".to_string(),
..Default::default()
};
let active_rules = [scored_rule_ref(&rule)];
let (raw, _ring_diag, _sbl_findings, gated_out) = raw_propose(
&target_mol,
"CC(=O)NC",
&active_rules,
crate::ring_context::RingContextArgs::default(),
crate::spectator_bond::SpectatorBondPolicy::Gated,
);
assert!(
!raw.is_empty(),
"an ordinary, non-defective disconnection must survive Gated policy"
);
assert!(gated_out.is_empty());
}
fn precs(smi: &str) -> PrecursorMol {
PrecursorMol {
smiles: smi.to_string(),
mol: mol_from_smiles(smi).unwrap(),
}
}
#[test]
fn candidate_id_join_ambiguity_is_resolved() {
let id_a = candidate_id_for("target", &["C.C".to_string(), "N".to_string()]);
let id_b = candidate_id_for("target", &["C".to_string(), "C.N".to_string()]);
assert_ne!(id_a, id_b, "join ambiguity must not collide");
}
#[test]
fn candidate_id_is_stable_sha256_prefixed() {
let id = candidate_id_for("target", &["CC".to_string()]);
assert!(id.starts_with("sha256:"));
assert_eq!(candidate_id_for("target", &["CC".to_string()]), id);
}
#[test]
fn duplicate_precursor_fragment_within_one_split_is_not_collapsed() {
let raw = vec![RawCandidate {
rule_name: "symmetric_split".to_string(),
template_id: "rule:symmetric_split".to_string(),
rule_weight: 1.0,
original_rank: 0,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
precursors: vec![precs("CC"), precs("CC")],
}];
let merged = merge_into_candidates("target", &raw).unwrap();
assert_eq!(merged.len(), 1);
assert_eq!(
merged[0].precursor_smiles,
vec!["CC".to_string(), "CC".to_string()],
"duplicate precursor fragment multiplicity must be preserved, not deduplicated"
);
}
#[test]
fn duplicate_precursor_set_from_two_rules_merges_with_provenance_retained() {
let raw = vec![
RawCandidate {
rule_name: "rule_a".to_string(),
template_id: "rule:rule_a".to_string(),
rule_weight: 1.0,
original_rank: 0,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
precursors: vec![precs("CC"), precs("O")],
},
RawCandidate {
rule_name: "rule_b".to_string(),
template_id: "rule:rule_b".to_string(),
rule_weight: 1.0,
original_rank: 5,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
precursors: vec![precs("O"), precs("CC")], },
];
let merged = merge_into_candidates("target", &raw).unwrap();
assert_eq!(
merged.len(),
1,
"identical precursor sets must merge into one candidate"
);
let c = &merged[0];
assert_eq!(c.source_template_count, 2);
assert_eq!(c.sources.len(), 2);
assert_eq!(
c.best_upstream_rank, 0,
"must keep the best (lowest) original_rank"
);
let names: Vec<&str> = c.sources.iter().map(|s| s.rule_name.as_str()).collect();
assert!(names.contains(&"rule_a"));
assert!(names.contains(&"rule_b"));
}
#[test]
fn duplicate_same_template_outcomes_do_not_inflate_source_count() {
let raw = vec![
RawCandidate {
rule_name: "symmetric_rule".to_string(),
template_id: "rule:symmetric_rule".to_string(),
rule_weight: 1.0,
original_rank: 3,
upstream_score: Some(0.4),
upstream_score_status: UpstreamScoreStatus::Available,
precursors: vec![precs("CC"), precs("O")],
},
RawCandidate {
rule_name: "symmetric_rule".to_string(),
template_id: "rule:symmetric_rule".to_string(),
rule_weight: 1.0,
original_rank: 1,
upstream_score: Some(0.4),
upstream_score_status: UpstreamScoreStatus::Available,
precursors: vec![precs("O"), precs("CC")], },
];
let merged = merge_into_candidates("target", &raw).unwrap();
assert_eq!(merged.len(), 1);
assert_eq!(
merged[0].source_template_count, 1,
"two applications of the same rule must merge into one source"
);
assert_eq!(merged[0].sources.len(), 1);
assert_eq!(
merged[0].sources[0].original_rank, 1,
"merged source must keep the min original_rank across duplicates"
);
}
#[test]
fn duplicate_same_template_outcomes_reject_inconsistent_frequency() {
let raw = vec![
RawCandidate {
rule_name: "r".to_string(),
template_id: "rule:r".to_string(),
rule_weight: 1.0,
original_rank: 0,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
precursors: vec![precs("CC")],
},
RawCandidate {
rule_name: "r".to_string(),
template_id: "rule:r".to_string(),
rule_weight: 2.0, original_rank: 1,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
precursors: vec![precs("CC")],
},
];
assert!(merge_into_candidates("target", &raw).is_err());
}
#[test]
fn best_upstream_score_and_min_cost_retained_on_merge() {
let mut a = RawCandidate {
rule_name: "rule_a".to_string(),
template_id: "rule:rule_a".to_string(),
rule_weight: 1.0,
original_rank: 3,
upstream_score: Some(0.2),
upstream_score_status: UpstreamScoreStatus::Available,
precursors: vec![precs("CC")],
};
let mut b = RawCandidate {
rule_name: "rule_b".to_string(),
template_id: "rule:rule_b".to_string(),
rule_weight: 1.0,
original_rank: 1,
upstream_score: Some(0.9),
upstream_score_status: UpstreamScoreStatus::Available,
precursors: vec![precs("CC")],
};
a.precursors = vec![precs("CC")];
b.precursors = vec![precs("CC")];
let merged = merge_into_candidates("target", &[a, b]).unwrap();
assert_eq!(merged.len(), 1);
assert_eq!(
merged[0].best_upstream_score,
Some(0.9),
"must keep the best (max) upstream score"
);
assert_eq!(
merged[0].best_upstream_rank, 1,
"must keep the best (min) original_rank"
);
assert!(merged[0].min_base_step_cost.is_finite());
}
#[test]
fn best_upstream_rank_is_the_best_scoring_sources_rank_not_the_global_minimum() {
let raw = vec![
RawCandidate {
rule_name: "rule_low_rank".to_string(),
template_id: "rule:rule_low_rank".to_string(),
rule_weight: 1.0,
original_rank: 0,
upstream_score: Some(0.1),
upstream_score_status: UpstreamScoreStatus::Available,
precursors: vec![precs("CC")],
},
RawCandidate {
rule_name: "rule_best_score".to_string(),
template_id: "rule:rule_best_score".to_string(),
rule_weight: 1.0,
original_rank: 2,
upstream_score: Some(0.9),
upstream_score_status: UpstreamScoreStatus::Available,
precursors: vec![precs("CC")],
},
];
let merged = merge_into_candidates("target", &raw).unwrap();
assert_eq!(merged.len(), 1);
assert_eq!(merged[0].best_upstream_score, Some(0.9));
assert_eq!(
merged[0].best_upstream_rank, 2,
"must be the rank of the source that achieved best_upstream_score, not the global minimum rank (0)"
);
}
#[test]
fn sources_sorted_deterministically_representative_unambiguous() {
let raw = vec![
RawCandidate {
rule_name: "z_rule".to_string(),
template_id: "z".to_string(),
rule_weight: 1.0,
original_rank: 0,
upstream_score: Some(0.5),
upstream_score_status: UpstreamScoreStatus::Available,
precursors: vec![precs("CC")],
},
RawCandidate {
rule_name: "a_rule".to_string(),
template_id: "a".to_string(),
rule_weight: 1.0,
original_rank: 0,
upstream_score: Some(0.5),
upstream_score_status: UpstreamScoreStatus::Available,
precursors: vec![precs("CC")],
},
];
let merged = merge_into_candidates("target", &raw).unwrap();
assert_eq!(merged.len(), 1);
assert_eq!(merged[0].sources[0].template_id, "a");
}
#[test]
fn merge_into_candidates_output_is_independent_of_input_order() {
fn one(rule_name: &str, template_id: &str, rank: usize, precursor: &str) -> RawCandidate {
RawCandidate {
rule_name: rule_name.to_string(),
template_id: template_id.to_string(),
rule_weight: 1.0,
original_rank: rank,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
precursors: vec![precs(precursor)],
}
}
fn summarize(candidates: Vec<ReactionCandidate>) -> Vec<(String, Vec<String>)> {
let mut summary: Vec<(String, Vec<String>)> = candidates
.into_iter()
.map(|c| {
let mut rule_names: Vec<String> =
c.sources.iter().map(|s| s.rule_name.clone()).collect();
rule_names.sort();
(c.candidate_id, rule_names)
})
.collect();
summary.sort_by(|a, b| a.0.cmp(&b.0));
summary
}
let forward = vec![
one("rule_a", "rule:a", 0, "CC"),
one("rule_b", "rule:b", 1, "CCO"),
one("rule_c", "rule:c", 2, "CC"),
];
let reversed = vec![
one("rule_c", "rule:c", 2, "CC"),
one("rule_b", "rule:b", 1, "CCO"),
one("rule_a", "rule:a", 0, "CC"),
];
let forward_summary = summarize(merge_into_candidates("target", &forward).unwrap());
let reversed_summary = summarize(merge_into_candidates("target", &reversed).unwrap());
assert_eq!(
forward_summary, reversed_summary,
"merged candidate set/content must not depend on RawCandidate input order"
);
assert_eq!(
forward_summary.len(),
2,
"CC merges rule_a+rule_c; CCO stays separate"
);
}
#[test]
fn no_precursors_produces_no_self_loop_candidate() {
let target = "CCO";
let rules = vec![rule("noop", "[C:1]>>[C:1]")];
let config = ProposalConfig::default();
let pool = propose_one_step("group:1", target, &rules, &config).unwrap();
for c in &pool.candidates {
assert_ne!(c.precursor_smiles, vec![pool.target_smiles.clone()]);
}
}
#[test]
fn same_target_different_group_shares_target_id_not_group_id() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let config = ProposalConfig::default();
let pool_a = propose_one_step("rxn-example-001", target, &rules, &config).unwrap();
let pool_b = propose_one_step("rxn-example-002", target, &rules, &config).unwrap();
assert_eq!(pool_a.target_id, pool_b.target_id);
assert_ne!(pool_a.group_id, pool_b.group_id);
assert_eq!(pool_a.group_id, "rxn-example-001");
assert_eq!(pool_b.group_id, "rxn-example-002");
}
#[test]
fn graph_based_rule_reaction_center_is_missing() {
let graph_rule = rule("cbz_deprotection_retro", "");
let f = template_transformation_features(&graph_rule);
assert!(!f.extractable);
assert_eq!(f.reaction_center_atom_count, 0);
}
#[test]
fn mapped_smirks_reaction_center_is_deterministic() {
let mapped_rule = rule(
"ester_hydrolysis_retro",
"[C:1](=[O:2])-[O:3]-[C:4]>>[C:1](=[O:2])-[OH:3].[OH]-[C:4]",
);
let a = template_transformation_features(&mapped_rule);
let b = template_transformation_features(&mapped_rule);
assert_eq!(a.mapped_atom_count, b.mapped_atom_count);
assert_eq!(a.deleted_bond_count, b.deleted_bond_count);
assert_eq!(a.extractable, b.extractable);
assert!(
a.extractable,
"a properly atom-mapped SMIRKS must be extractable"
);
assert_eq!(a.mapped_atom_count, 4, "C:1, O:2, O:3, C:4");
assert_eq!(a.deleted_bond_count, 1, "the O:3-C:4 ester bond is broken");
assert!(a.reaction_center_atom_count > 0);
}
#[test]
fn unmapped_smirks_reaction_center_is_missing_not_guessed() {
let unmapped_rule = rule("fake_unmapped", "CC>>C.C");
let f = template_transformation_features(&unmapped_rule);
assert!(!f.extractable);
}
#[test]
fn partially_mapped_smirks_is_extractable_over_the_mapped_atoms_only() {
let partial_rule = rule("partial_map", "[C:1]CC>>[C:1]C.C");
let f = template_transformation_features(&partial_rule);
assert!(
f.extractable,
"a partially-mapped SMIRKS (>=1 mapped atom, no duplicates) must still be extractable"
);
assert_eq!(
f.mapped_atom_count, 1,
"only atom map 1 is annotated on either side"
);
}
#[test]
fn changed_bond_order_is_detected_without_adding_or_deleting_a_bond() {
let bond_order_change_rule = rule("retro_reduction", "[C:1]=[O:2]>>[C:1][O:2]");
let f = template_transformation_features(&bond_order_change_rule);
assert!(f.extractable);
assert_eq!(f.deleted_bond_count, 0, "the C:1-O:2 bond is never deleted");
assert_eq!(
f.added_bond_count, 0,
"the C:1-O:2 bond is never newly formed"
);
assert_eq!(
f.changed_bond_order_count, 1,
"the C:1-O:2 bond order changes from double to single"
);
assert!(f.reaction_center_atom_count > 0);
}
#[test]
fn multi_component_reaction_center_does_not_collide_on_local_atom_idx() {
let cross_rule = rule(
"synthetic_cross_metathesis",
"[C:1][C:2].[C:3][C:4]>>[C:1][C:3].[C:2][C:4]",
);
let f = template_transformation_features(&cross_rule);
assert!(f.extractable);
assert_eq!(f.deleted_bond_count, 2, "both original bonds are broken");
assert_eq!(f.added_bond_count, 2, "both new cross-bonds are formed");
assert_eq!(
f.reaction_center_atom_count, 4,
"all four atoms (map1..map4) participate in the reaction center -- \
a raw-AtomIdx collision would undercount this to 2"
);
}
#[test]
fn duplicate_atom_map_within_one_side_is_not_extractable() {
let ambiguous_rule = rule("ambiguous_duplicate_map", "[C:1][C:1]>>[C:1].[C:1]");
let f = template_transformation_features(&ambiguous_rule);
assert!(
!f.extractable,
"a duplicate atom_map number on one side must not be extractable"
);
}
#[test]
fn transformation_cache_does_not_collide_on_reused_template_id_with_different_smirks() {
let mut a = rule("shared_id", "[C:1][C:2]>>[C:1].[C:2]");
a.template_id = "rule:shared_id".to_string();
let mut b = rule("shared_id", "[C:1][C:2][C:3]>>[C:1].[C:2].[C:3]");
b.template_id = "rule:shared_id".to_string();
let fa = template_transformation_features(&a);
let fb = template_transformation_features(&b);
assert!(fa.extractable);
assert!(fb.extractable);
assert_ne!(
fa.mapped_atom_count, fb.mapped_atom_count,
"these two SMIRKS have a different mapped atom count -- if the cache \
collided on template_id alone, one of these would incorrectly read \
back the other's cached result"
);
}
#[test]
fn index_rules_by_template_id_rejects_conflicting_duplicate() {
let mut a = rule("dup", "[C:1]>>[C:1]");
a.template_id = "rule:dup".to_string();
let mut b = rule("dup", "[N:1]>>[N:1]"); b.template_id = "rule:dup".to_string();
assert!(index_rules_by_template_id(&[a, b]).is_err());
}
#[test]
fn index_rules_by_template_id_succeeds_on_real_extracted_corpus() {
let path = concat!(env!("CARGO_MANIFEST_DIR"), "/data/templates_extracted.smi");
let rules = crate::chem_env::load_rules_from_file(path);
assert_eq!(
rules.len(),
500,
"load_rules_from_file must return exactly one RetroRule per raw template line"
);
index_rules_by_template_id(&rules)
.expect("candidate-pool export must not hard-error on the real extracted corpus");
}
#[test]
fn index_rules_by_template_id_tolerates_exact_duplicate() {
let mut a = rule("dup", "[C:1]>>[C:1]");
a.template_id = "rule:dup".to_string();
let b = a.clone();
let rules = [a, b];
let index = index_rules_by_template_id(&rules).unwrap();
assert_eq!(index.len(), 1);
}
#[test]
fn aggregate_transformation_features_no_nan_or_inf() {
let features = vec![
TemplateTransformationFeatures {
extractable: true,
reaction_center_atom_count: 4,
..Default::default()
},
TemplateTransformationFeatures {
extractable: false,
..Default::default()
},
];
let agg = aggregate_transformation_features(&features);
assert!(agg.reaction_center_atom_count_mean.is_finite());
assert!(agg.reaction_center_extractable_fraction.is_finite());
assert!(!agg.reaction_center_atom_count_mean.is_nan());
let empty_agg = aggregate_transformation_features(&[]);
assert!(!empty_agg.reaction_center_atom_count_mean.is_nan());
}
#[test]
fn feature_schema_v1_names_and_group_boundary_are_consistent() {
assert_eq!(FEATURE_NAMES_V1.len(), 18);
assert!(FEATURE_GROUP1_LEN < FEATURE_NAMES_V1.len());
assert_eq!(
FEATURE_NAMES_V1[FEATURE_GROUP1_LEN],
"fraction_precursors_in_stock"
);
assert_eq!(
feature_index("num_precursors"),
Some(0),
"feature_index must find a real schema-v1 name"
);
assert_eq!(
feature_index("not_a_real_feature"),
None,
"feature_index must return None for an unknown name"
);
}
#[test]
fn feature_schema_hash_is_stable_and_pinned_for_cross_language_verification() {
let a = feature_schema_hash();
let b = feature_schema_hash();
assert_eq!(a, b, "the hash must be deterministic across calls");
assert!(a.starts_with("sha256:"));
assert_eq!(
a,
"sha256:756404c59bbee9a65e194f92df3530e1b801028f333e01c67214917977061df1"
);
}
fn candidate_for(target: &str, rules: &[RetroRule], mode: ProposalMode) -> ReactionCandidate {
let pool = propose_one_step("group:1", target, rules, &ProposalConfig { mode }).unwrap();
pool.candidates
.into_iter()
.next()
.expect("expected at least one candidate for this fixture")
}
#[test]
fn extract_features_group2_missing_without_stock() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let candidate = candidate_for(target, &rules, ProposalMode::Exhaustive);
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let features = extract_features(&candidate, &target_mol, &templates_by_id, None);
assert_eq!(features.values.len(), FEATURE_NAMES_V1.len());
assert_eq!(features.missing.len(), FEATURE_NAMES_V1.len());
for name in ["fraction_precursors_in_stock", "all_precursors_in_stock"] {
let i = feature_index(name).unwrap();
assert!(
features.missing[i],
"{name} must be missing without a stock"
);
}
for name in ["max_template_log_frequency", "mean_template_log_frequency"] {
let i = feature_index(name).unwrap();
assert!(
features.missing[i],
"{name} must always be missing until split-aware recomputation lands"
);
}
let best_upstream_i = feature_index("best_upstream_score").unwrap();
for (i, name) in FEATURE_NAMES_V1.iter().enumerate().take(FEATURE_GROUP1_LEN) {
if i == best_upstream_i {
assert!(
features.missing[i],
"best_upstream_score must be missing under Exhaustive mode (no scorer involved)"
);
continue;
}
assert!(
!features.missing[i],
"group-1 feature {i} ({name}) must be computed for a normal candidate"
);
}
}
#[test]
fn extract_features_availability_reflects_stock_membership() {
use crate::chem_env::ChemEnv;
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let candidate = candidate_for(target, &rules, ProposalMode::Exhaustive);
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let full_stock = ChemEnv::in_memory(
&candidate
.precursor_smiles
.iter()
.map(|s| s.as_str())
.collect::<Vec<_>>(),
);
let f_full = extract_features(&candidate, &target_mol, &templates_by_id, Some(&full_stock));
let all_i = feature_index("all_precursors_in_stock").unwrap();
let frac_i = feature_index("fraction_precursors_in_stock").unwrap();
assert!(!f_full.missing[all_i]);
assert_eq!(f_full.values[all_i], 1.0);
assert_eq!(f_full.values[frac_i], 1.0);
let empty_stock = ChemEnv::in_memory(&[]);
let f_empty = extract_features(
&candidate,
&target_mol,
&templates_by_id,
Some(&empty_stock),
);
assert!(!f_empty.missing[all_i]);
assert_eq!(f_empty.values[all_i], 0.0);
assert_eq!(f_empty.values[frac_i], 0.0);
}
#[test]
fn extract_features_no_heavy_atom_gain_and_charge_balance_hold_for_real_reaction() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let candidate = candidate_for(target, &rules, ProposalMode::Exhaustive);
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let features = extract_features(&candidate, &target_mol, &templates_by_id, None);
let no_gain_i = feature_index("no_heavy_atom_gain").unwrap();
let charge_i = feature_index("net_charge_balanced").unwrap();
assert_eq!(
features.values[no_gain_i], 1.0,
"a real retro disconnection must never gain heavy atoms in the target"
);
assert_eq!(
features.values[charge_i], 1.0,
"a real retro disconnection on a neutral target/precursors must be charge-balanced"
);
}
#[test]
fn extract_features_num_precursors_survives_reparse_failure() {
let target = "CCO";
let target_mol = mol_from_smiles(target).unwrap();
let candidate = ReactionCandidate {
candidate_id: "sha256:fake".to_string(),
target_smiles: target.to_string(),
precursor_smiles: vec!["CC".to_string(), "not-a-valid-smiles(((".to_string()],
sources: vec![],
source_template_count: 0,
best_upstream_score: None,
best_upstream_rank: 0,
min_base_step_cost: 0.0,
max_template_frequency: None,
mean_template_frequency: None,
features: CandidateFeatures::default(),
reranker_score: None,
};
let templates_by_id: HashMap<String, &RetroRule> = HashMap::new();
let features = extract_features(&candidate, &target_mol, &templates_by_id, None);
let num_i = feature_index("num_precursors").unwrap();
assert!(!features.missing[num_i]);
assert_eq!(features.values[num_i], 2.0);
for name in [
"target_heavy_atom_count",
"precursor_heavy_atom_count_sum",
"precursor_heavy_atom_count_max",
"heavy_atom_retention_ratio",
"net_charge_balanced",
"no_heavy_atom_gain",
] {
let i = feature_index(name).unwrap();
assert!(
features.missing[i],
"{name} must be missing on reparse failure"
);
}
}
#[test]
fn extract_features_is_deterministic_across_two_calls() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let candidate = candidate_for(target, &rules, ProposalMode::Exhaustive);
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let a = extract_features(&candidate, &target_mol, &templates_by_id, None);
let b = extract_features(&candidate, &target_mol, &templates_by_id, None);
assert_eq!(a.values, b.values);
assert_eq!(a.missing, b.missing);
}
}