#![forbid(unsafe_code)]
use renkin::candidate::{
CandidateProposalContext, ProposalConfig, propose_phase_nanos, reset_propose_phase_nanos,
};
use renkin::chem_env::{load_rules_from_file, mol_from_smiles};
use renkin::pool_export::{
PoolProvenance, build_manifest, candidate_rows_for_pool, target_pool_record_for_failure,
target_pool_record_for_pool, target_pool_record_for_target_id_mismatch, write_jsonl,
write_target_pool_jsonl,
};
use serde::Deserialize;
use sha2::{Digest, Sha256};
use std::fs::File;
use std::io::{BufRead, BufReader, BufWriter};
use std::time::Instant;
#[derive(Debug, Deserialize)]
struct GroupInput {
group_id: String,
target_id: String,
}
fn arg_value(flag: &str, default: &str) -> String {
let args: Vec<String> = std::env::args().collect();
args.iter()
.position(|a| a == flag)
.and_then(|i| args.get(i + 1))
.cloned()
.unwrap_or_else(|| default.to_string())
}
fn arg_opt(flag: &str) -> Option<String> {
let args: Vec<String> = std::env::args().collect();
args.iter()
.position(|a| a == flag)
.and_then(|i| args.get(i + 1))
.cloned()
}
fn sha256_of_file(path: &str) -> String {
let bytes = std::fs::read(path).unwrap_or_else(|e| panic!("read {path}: {e}"));
format!("sha256:{}", renkin::sha256_hex(Sha256::digest(&bytes)))
}
fn chematic_version_from_lockfile(lockfile: &str) -> String {
let mut lines = lockfile.lines().peekable();
while let Some(line) = lines.next() {
if line.trim() == "[[package]]" {
let mut name = None;
let mut version = None;
while let Some(&next) = lines.peek() {
if next.trim() == "[[package]]" || next.trim().is_empty() {
break;
}
let next = lines.next().unwrap();
if let Some(v) = next
.strip_prefix("name = \"")
.and_then(|s| s.strip_suffix('"'))
{
name = Some(v.to_string());
} else if let Some(v) = next
.strip_prefix("version = \"")
.and_then(|s| s.strip_suffix('"'))
{
version = Some(v.to_string());
}
}
if name.as_deref() == Some("chematic") {
return version.unwrap_or_default();
}
}
}
String::new()
}
fn percentile(sorted: &[usize], pct: f64) -> usize {
if sorted.is_empty() {
return 0;
}
let idx = ((sorted.len() - 1) as f64 * pct).round() as usize;
sorted[idx.min(sorted.len() - 1)]
}
fn percentile_f64(sorted: &[f64], pct: f64) -> f64 {
if sorted.is_empty() {
return 0.0;
}
let idx = ((sorted.len() - 1) as f64 * pct).round() as usize;
sorted[idx.min(sorted.len() - 1)]
}
fn main() {
let start = Instant::now();
let groups_path = arg_value("--groups", "data/reranker_groups_uspto50k_test.jsonl");
let templates_path = arg_value("--templates", "data/templates_extracted_500.smi");
let pool_output = arg_value("--pool-output", "data/pool_gen_output.jsonl");
let groups_output = arg_value("--groups-output", "data/pool_gen_groups.jsonl");
let manifest_output = arg_value("--manifest-output", "data/pool_gen_manifest.json");
let limit: Option<usize> =
arg_opt("--limit").map(|s| s.parse().expect("--limit must be an integer"));
let group_inputs: Vec<GroupInput> = {
let file = File::open(&groups_path).unwrap_or_else(|e| panic!("open {groups_path}: {e}"));
let mut rows: Vec<GroupInput> = BufReader::new(file)
.lines()
.map(|l| l.unwrap())
.filter(|l| !l.trim().is_empty())
.map(|l| serde_json::from_str(&l).unwrap_or_else(|e| panic!("parse {l:?}: {e}")))
.collect();
if let Some(n) = limit {
rows.truncate(n);
}
rows
};
let n_targets_requested = group_inputs.len();
let rules = load_rules_from_file(&templates_path);
eprintln!("loaded {} rules from {templates_path}", rules.len());
eprintln!("processing {n_targets_requested} group(s) from {groups_path}");
let ctx = CandidateProposalContext::new(&rules, false);
let config = ProposalConfig::default();
let templates_by_id = renkin::candidate::index_rules_by_template_id(&rules)
.expect("rules must have consistent template_id -> rule mapping");
reset_propose_phase_nanos();
let mut candidate_rows = Vec::new();
let mut group_records = Vec::new();
let mut candidate_counts: Vec<usize> = Vec::new();
let mut n_parse_failed = 0usize;
let mut n_zero_candidate = 0usize;
let mut n_target_id_mismatch = 0usize;
let mut per_target_seconds: Vec<f64> = Vec::new();
for (i, g) in group_inputs.iter().enumerate() {
if i % 500 == 0 && i > 0 {
eprintln!(" {i}/{n_targets_requested}...");
}
match mol_from_smiles(&g.target_id) {
Err(_) => {
n_parse_failed += 1;
group_records.push(target_pool_record_for_failure(&g.group_id, &g.target_id));
}
Ok(target_mol) => {
let t_target = Instant::now();
let result = ctx.propose_one_step(&g.group_id, &g.target_id, &config);
per_target_seconds.push(t_target.elapsed().as_secs_f64());
match result {
Err(e) => {
n_parse_failed += 1;
eprintln!(" {}: propose_one_step error: {e}", g.group_id);
group_records
.push(target_pool_record_for_failure(&g.group_id, &g.target_id));
}
Ok(pool) if pool.target_id != g.target_id => {
n_target_id_mismatch += 1;
eprintln!(
" {}: target_id mismatch -- requested {:?}, propose_one_step derived {:?}",
g.group_id, g.target_id, pool.target_id
);
group_records.push(target_pool_record_for_target_id_mismatch(
&g.group_id,
&g.target_id,
));
}
Ok(pool) => {
if pool.candidates.is_empty() {
n_zero_candidate += 1;
}
candidate_counts.push(pool.candidates.len());
group_records.push(target_pool_record_for_pool(&pool));
let rows =
candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
candidate_rows.extend(rows);
}
}
}
}
}
candidate_rows.sort_by(|a, b| {
(a.group_id.as_str(), a.candidate_id.as_str())
.cmp(&(b.group_id.as_str(), b.candidate_id.as_str()))
});
let pool_file =
File::create(&pool_output).unwrap_or_else(|e| panic!("create {pool_output}: {e}"));
let candidate_jsonl_sha256 =
write_jsonl(&candidate_rows, BufWriter::new(pool_file)).expect("write pool jsonl");
let groups_file =
File::create(&groups_output).unwrap_or_else(|e| panic!("create {groups_output}: {e}"));
let target_group_index_sha256 =
write_target_pool_jsonl(&group_records, BufWriter::new(groups_file))
.expect("write group index jsonl");
let renkin_git_commit = std::process::Command::new("git")
.args(["rev-parse", "HEAD"])
.output()
.ok()
.and_then(|o| String::from_utf8(o.stdout).ok())
.map(|s| s.trim().to_string())
.unwrap_or_default();
let provenance = PoolProvenance {
renkin_git_commit,
cargo_lock_sha256: sha256_of_file("Cargo.lock"),
chematic_version: chematic_version_from_lockfile(
&std::fs::read_to_string("Cargo.lock").unwrap_or_default(),
),
target_input_sha256: sha256_of_file(&groups_path),
stock_source: None,
embedded_fallback_used: false,
export_config: serde_json::json!({
"groups_path": groups_path,
"templates_path": templates_path,
"limit": limit,
"proposal_mode": "exhaustive",
}),
};
let manifest = build_manifest(
&candidate_rows,
&candidate_jsonl_sha256,
&group_records,
&target_group_index_sha256,
&rules,
&config.mode,
None,
provenance,
)
.expect("build manifest");
let manifest_json = serde_json::to_string_pretty(&manifest).expect("serialize manifest");
std::fs::write(&manifest_output, &manifest_json)
.unwrap_or_else(|e| panic!("write {manifest_output}: {e}"));
candidate_counts.sort_unstable();
per_target_seconds.sort_by(|a, b| a.partial_cmp(b).unwrap());
let elapsed = start.elapsed();
let phase_nanos = propose_phase_nanos();
let feasibility_summary = serde_json::json!({
"n_groups_requested": n_targets_requested,
"n_groups_parse_failed": n_parse_failed,
"n_groups_target_id_mismatch": n_target_id_mismatch,
"n_groups_zero_candidates": n_zero_candidate,
"n_candidate_rows": candidate_rows.len(),
"candidates_per_group_p50": percentile(&candidate_counts, 0.50),
"candidates_per_group_p90": percentile(&candidate_counts, 0.90),
"candidates_per_group_p95": percentile(&candidate_counts, 0.95),
"candidates_per_group_max": candidate_counts.last().copied().unwrap_or(0),
"wall_clock_seconds": elapsed.as_secs_f64(),
"pool_output": pool_output,
"groups_output": groups_output,
"manifest_output": manifest_output,
"candidate_jsonl_sha256": candidate_jsonl_sha256,
"target_group_index_sha256": target_group_index_sha256,
"proposal_mode": "exhaustive",
"propose_phase_seconds": {
"select": phase_nanos.select as f64 / 1e9,
"raw_propose": phase_nanos.raw_propose as f64 / 1e9,
"merge": phase_nanos.merge as f64 / 1e9,
},
"proposal_seconds_per_target_p50": percentile_f64(&per_target_seconds, 0.50),
"proposal_seconds_per_target_p95": percentile_f64(&per_target_seconds, 0.95),
"proposal_seconds_per_target_max": per_target_seconds.last().copied().unwrap_or(0.0),
});
eprintln!(
"{}",
serde_json::to_string_pretty(&feasibility_summary).unwrap()
);
}