use pyo3::exceptions::PyValueError;
use pyo3::prelude::*;
use crate::bridge;
use crate::chem_env::{
ChemEnv, default_rules, elem_symbols_to_mask, load_rules_from_file, mol_from_smiles,
};
use crate::search::{SearchConfig, diagnose, find_routes};
#[pyfunction]
#[pyo3(name = "find_routes", signature = (target, depth=5, max_routes=5, beam_width=0, building_blocks=None, avoid_elements="", require_elements="", verbose=false, bb_prices_path=None, templates_path=None, template_metadata_path=None, reranker_model_path=None, reranker_freq_table_path=None, top_templates=None, search_mode="standard", coverage_templates_path=None, coverage_timeout_seconds=None, search_diagnostics=false, spectator_bond_policy="off"))]
#[allow(clippy::too_many_arguments)]
pub fn find_routes_py(
target: &str,
depth: u32,
max_routes: usize,
beam_width: usize,
building_blocks: Option<Vec<String>>,
avoid_elements: &str,
require_elements: &str,
verbose: bool,
bb_prices_path: Option<&str>,
templates_path: Option<&str>,
template_metadata_path: Option<&str>,
reranker_model_path: Option<&str>,
reranker_freq_table_path: Option<&str>,
top_templates: Option<usize>,
search_mode: &str,
coverage_templates_path: Option<&str>,
coverage_timeout_seconds: Option<u64>,
search_diagnostics: bool,
spectator_bond_policy: &str,
) -> PyResult<String> {
if search_mode != "standard" && search_mode != "coverage" {
return Err(PyValueError::new_err(format!(
"invalid search_mode {search_mode:?} (expected \"standard\" or \"coverage\")"
)));
}
let spectator_bond_policy = match spectator_bond_policy {
"off" => crate::spectator_bond::SpectatorBondPolicy::Off,
"diagnostics_only" => crate::spectator_bond::SpectatorBondPolicy::DiagnosticsOnly,
"gated" => crate::spectator_bond::SpectatorBondPolicy::Gated,
other => {
return Err(PyValueError::new_err(format!(
"invalid spectator_bond_policy {other:?} (expected \"off\", \"diagnostics_only\", \
or \"gated\")"
)));
}
};
if search_mode == "standard" {
if coverage_templates_path.is_some() {
return Err(PyValueError::new_err(
"coverage_templates_path requires search_mode=\"coverage\"",
));
}
if coverage_timeout_seconds.is_some() {
return Err(PyValueError::new_err(
"coverage_timeout_seconds requires search_mode=\"coverage\"",
));
}
}
if search_mode == "coverage" && coverage_timeout_seconds == Some(0) {
return Err(PyValueError::new_err(
"coverage_timeout_seconds must be a positive integer (got 0)",
));
}
let env = match building_blocks {
Some(ref bbs) => {
let refs: Vec<&str> = bbs.iter().map(|s| s.as_str()).collect();
ChemEnv::in_memory(&refs)
}
None => ChemEnv::load("data/building_blocks.smi")
.unwrap_or_else(|_| ChemEnv::in_memory(crate::DEFAULT_BUILDING_BLOCKS)),
};
let mut rules = default_rules();
if let Some(path) = templates_path {
let mut extra = load_rules_from_file(path);
if let Some(k) = top_templates {
extra = crate::chem_env::top_templates_by_weight(extra, k);
}
rules.extend(extra);
}
let template_metadata = template_metadata_path
.map(crate::evidence::load_template_metadata)
.transpose()
.map_err(|e| PyValueError::new_err(e.to_string()))?;
if let Some(ref tm) = template_metadata {
let known_ids: std::collections::HashSet<&str> =
rules.iter().map(|r| r.template_id.as_str()).collect();
crate::evidence::warn_unknown_templates(tm, &known_ids);
}
let bb_price_map = bb_prices_path.map(|path| {
std::fs::read_to_string(path)
.ok()
.map(|content| {
content
.lines()
.filter(|l| !l.is_empty() && !l.starts_with('#'))
.filter_map(|l| {
let (smiles, price) = l.split_once(',')?;
let price: f64 = price.trim().parse().ok()?;
Some((smiles.trim().to_string(), price))
})
.collect::<std::collections::HashMap<String, f64>>()
})
.unwrap_or_default()
});
let reranker: Option<std::sync::Arc<dyn crate::candidate::CandidateReranker>> =
match (reranker_model_path, reranker_freq_table_path) {
(Some(model_path), Some(freq_path)) => {
match crate::reranker::RuntimeReranker::from_paths(model_path, freq_path) {
Ok(r) => Some(std::sync::Arc::new(r)),
Err(e) => {
eprintln!(
"warning: failed to load reranker_model_path/reranker_freq_table_path \
({e:#}); falling back to legacy ordering for this run"
);
None
}
}
}
(None, None) => None,
_ => {
eprintln!(
"warning: reranker_model_path and reranker_freq_table_path must both be \
given; falling back to legacy ordering for this run"
);
None
}
};
let config = SearchConfig {
max_depth: depth,
max_routes,
beam_width,
forbidden_elements: elem_symbols_to_mask(avoid_elements),
required_element_present: elem_symbols_to_mask(require_elements),
verbose,
bb_price_map,
template_metadata: template_metadata.map(|tm| tm.templates),
reranker,
spectator_bond_policy,
..Default::default()
};
struct CoverageModeMeta {
selected_stage: &'static str,
stage2_invoked: bool,
stage1_timeout: bool,
stage2_timeout: bool,
stage1_elapsed_ms: f64,
stage2_elapsed_ms: Option<f64>,
total_elapsed_ms: f64,
reranker_failures_summed: u64,
}
let (routes, stats, coverage_meta) = if search_mode == "coverage" {
let coverage_path = coverage_templates_path.ok_or_else(|| {
PyValueError::new_err("search_mode=\"coverage\" requires coverage_templates_path")
})?;
crate::coverage_mode::validate_coverage_mode_config(&config)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
let coverage_rules = crate::coverage_mode::load_coverage_rules(coverage_path)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
let coverage_timeout = coverage_timeout_seconds.map(std::time::Duration::from_secs);
let result = crate::coverage_mode::run_coverage_mode(
target,
&env,
&rules,
&config,
&coverage_rules,
coverage_timeout,
)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
let meta = CoverageModeMeta {
selected_stage: match result.selected_stage {
crate::coverage_mode::SelectedStage::Stage1 => "stage1",
crate::coverage_mode::SelectedStage::Stage2 => "stage2",
},
stage2_invoked: result.stage2_invoked,
stage1_timeout: result.stage1_timeout,
stage2_timeout: result.stage2_timeout,
stage1_elapsed_ms: result.stage1_elapsed_ms,
stage2_elapsed_ms: result.stage2_elapsed_ms,
total_elapsed_ms: result.total_elapsed_ms,
reranker_failures_summed: result.reranker_failures,
};
(result.routes, result.stats, Some(meta))
} else {
let (routes, stats) = find_routes(target, &env, &rules, &config)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
(routes, stats, None)
};
let reranker_failures_for_output = coverage_meta
.as_ref()
.map(|m| m.reranker_failures_summed)
.unwrap_or(stats.reranker_failures);
let mut output = if routes.is_empty() {
let (causes, suggestions) = diagnose(&stats, depth);
serde_json::json!({
"target": target,
"routes_found": 0,
"routes": [],
"diagnostics": {
"nodes_expanded": stats.nodes_expanded,
"max_depth_reached": stats.max_depth_reached,
"beam_limit_hit": stats.beam_limit_hit,
"matched_templates": stats.matched_templates,
"stock_hits": stats.stock_hits,
"likely_causes": causes,
"suggestions": suggestions,
}
})
} else {
let joint_success_probability = 1.0
- routes
.iter()
.map(|r| 1.0 - r.success_probability)
.product::<f64>();
serde_json::json!({
"target": target,
"routes_found": routes.len(),
"routes": routes,
"joint_success_probability": joint_success_probability,
})
};
if search_diagnostics {
output["search_diagnostics"] = serde_json::to_value(&stats.crowd_out)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
}
if config.reranker.is_some() {
output["reranker_failures"] = serde_json::Value::from(reranker_failures_for_output);
}
if let Some(ref m) = coverage_meta {
output["search_mode"] = serde_json::Value::from("coverage");
output["selected_stage"] = serde_json::Value::from(m.selected_stage);
output["stage2_invoked"] = serde_json::Value::from(m.stage2_invoked);
output["stage1_timeout"] = serde_json::Value::from(m.stage1_timeout);
output["stage2_timeout"] = serde_json::Value::from(m.stage2_timeout);
output["stage1_elapsed_ms"] = serde_json::Value::from(m.stage1_elapsed_ms);
output["stage2_elapsed_ms"] = serde_json::to_value(m.stage2_elapsed_ms)
.map_err(|e| PyValueError::new_err(e.to_string()))?;
output["total_elapsed_ms"] = serde_json::Value::from(m.total_elapsed_ms);
}
serde_json::to_string(&output).map_err(|e| PyValueError::new_err(e.to_string()))
}
fn py_reverse_smirks(s: &str) -> Option<String> {
let (lhs, rhs) = s.split_once(">>")?;
Some(format!("{rhs}>>{lhs}"))
}
fn py_is_valid_smiles(s: &str) -> bool {
let has_aromatic = s
.bytes()
.any(|b| matches!(b, b'c' | b'n' | b'o' | b's' | b'p'));
!has_aromatic || s.bytes().any(|b| b.is_ascii_digit())
}
fn py_predict_forward_core(
reactants: &[&str],
rules: &[crate::chem_env::RetroRule],
max_results: usize,
) -> Result<Vec<serde_json::Value>, String> {
use chematic::rxn::run_reactants;
use chematic::smiles::canonical_smiles as canon;
let mols: Vec<_> = reactants
.iter()
.filter_map(|s| mol_from_smiles(s).ok())
.collect();
if mols.len() != reactants.len() {
return Err("one or more reactant SMILES failed to parse".into());
}
let mol_refs: Vec<_> = mols.iter().collect();
let mut preds: Vec<serde_json::Value> = rules
.iter()
.filter(|r| !r.smirks.is_empty())
.filter_map(|rule| {
let fwd = py_reverse_smirks(&rule.smirks)?;
let outcomes = run_reactants(&fwd, &mol_refs).ok()?;
if outcomes.is_empty() { return None; }
let products: Vec<String> = outcomes
.into_iter()
.flat_map(|ms| ms.iter().map(canon).collect::<Vec<_>>())
.filter(|s| py_is_valid_smiles(s))
.collect();
if products.is_empty() { return None; }
Some(serde_json::json!({ "template": rule.name, "products": products, "weight": rule.weight }))
})
.collect();
preds.sort_unstable_by(|a, b| {
b["weight"]
.as_f64()
.unwrap_or(0.0)
.partial_cmp(&a["weight"].as_f64().unwrap_or(0.0))
.unwrap_or(std::cmp::Ordering::Equal)
});
preds.truncate(max_results);
Ok(preds)
}
#[pyfunction]
#[pyo3(name = "predict_forward", signature = (reactants, templates_path=None, max_results=5))]
pub fn predict_forward_py(
reactants: Vec<String>,
templates_path: Option<&str>,
max_results: usize,
) -> PyResult<String> {
let mut rules = default_rules();
if let Some(path) = templates_path {
rules.extend(load_rules_from_file(path));
}
let refs: Vec<&str> = reactants.iter().map(|s| s.as_str()).collect();
let preds =
py_predict_forward_core(&refs, &rules, max_results).map_err(PyValueError::new_err)?;
serde_json::to_string(&preds).map_err(|e| PyValueError::new_err(e.to_string()))
}
#[pyfunction]
#[pyo3(name = "validate_forward", signature = (route_json, templates_path=None, max_results=5))]
pub fn validate_forward_py(
route_json: &str,
templates_path: Option<&str>,
max_results: usize,
) -> PyResult<String> {
use chematic::smiles::canonical_smiles as canon;
let v: serde_json::Value = serde_json::from_str(route_json)
.map_err(|e| PyValueError::new_err(format!("invalid JSON: {e}")))?;
let steps = v["steps"]
.as_array()
.ok_or_else(|| PyValueError::new_err("route JSON must have a 'steps' array"))?;
let mut rules = default_rules();
if let Some(path) = templates_path {
rules.extend(load_rules_from_file(path));
}
let mut results: Vec<serde_json::Value> = Vec::new();
for (idx, step) in steps.iter().enumerate() {
let target = step["target"].as_str().unwrap_or("");
let prec_refs: Vec<&str> = step["precursors"]
.as_array()
.map(|a| a.iter().filter_map(|v| v.as_str()).collect())
.unwrap_or_default();
let preds = py_predict_forward_core(&prec_refs, &rules, max_results)
.map_err(PyValueError::new_err)?;
let target_canon = mol_from_smiles(target)
.ok()
.map(|m| canon(&m))
.unwrap_or_else(|| target.to_string());
let verified = preds.iter().any(|p| {
p["products"]
.as_array()
.map(|a| a.iter().any(|v| v.as_str() == Some(&target_canon)))
.unwrap_or(false)
});
results.push(serde_json::json!({
"step_index": idx, "target": target, "verified": verified, "top_predictions": preds
}));
}
serde_json::to_string(&results).map_err(|e| PyValueError::new_err(e.to_string()))
}
#[pyfunction]
#[pyo3(name = "audit_route", signature = (content, format="auto", stock_text="", policy="standard"))]
pub fn audit_route_py(
content: &str,
format: &str,
stock_text: &str,
policy: &str,
) -> PyResult<String> {
let policy: bridge::AuditPolicy = policy.parse().map_err(PyValueError::new_err)?;
let stock = (!stock_text.trim().is_empty()).then(|| bridge::parse_stock_text(stock_text));
let rules = default_rules();
let report = bridge::build_audit_route_report_with_policy(
content,
format,
stock.as_ref(),
&rules,
policy,
)
.map_err(|e| PyValueError::new_err(format!("{e:#}")))?;
serde_json::to_string(&report).map_err(|e| PyValueError::new_err(e.to_string()))
}
#[pymodule]
pub fn renkin(_py: Python, m: &Bound<'_, PyModule>) -> PyResult<()> {
m.add_function(wrap_pyfunction!(find_routes_py, m)?)?;
m.add_function(wrap_pyfunction!(predict_forward_py, m)?)?;
m.add_function(wrap_pyfunction!(validate_forward_py, m)?)?;
m.add_function(wrap_pyfunction!(audit_route_py, m)?)?;
m.add("__version__", env!("CARGO_PKG_VERSION"))?;
Ok(())
}