use std::collections::HashMap;
use std::io::Write;
use serde::Serialize;
use sha2::{Digest, Sha256};
use crate::candidate::{
CandidatePool, FEATURE_NAMES_V1, FEATURE_SCHEMA_VERSION, ProposalMode, UpstreamScoreStatus,
extract_features, feature_schema_hash,
};
use crate::chem_env::{ChemEnv, Molecule, RetroRule};
#[derive(Debug, Clone, PartialEq, Serialize)]
pub struct ProposalModeSummary {
pub mode: &'static str,
pub top_k: Option<usize>,
pub rules_offset: Option<usize>,
pub scorer_identity: Option<String>,
pub scorer_model_sha256: Option<String>,
pub scorer_status: Option<UpstreamScoreStatus>,
}
impl ProposalModeSummary {
pub fn from_mode(mode: &ProposalMode) -> Self {
match mode {
ProposalMode::Exhaustive => Self {
mode: "exhaustive",
top_k: None,
rules_offset: None,
scorer_identity: None,
scorer_model_sha256: None,
scorer_status: None,
},
ProposalMode::BondIndexed { top_k } => Self {
mode: "bond_indexed",
top_k: Some(*top_k),
rules_offset: None,
scorer_identity: None,
scorer_model_sha256: None,
scorer_status: None,
},
ProposalMode::ScorerConditioned { input, top_k } => Self {
mode: "scorer_conditioned",
top_k: Some(*top_k),
rules_offset: Some(input.rules_offset),
scorer_identity: Some(input.scorer_identity.clone()),
scorer_model_sha256: Some(input.scorer_model_sha256.clone()),
scorer_status: Some(input.status),
},
}
}
}
#[derive(Debug, Clone, Serialize)]
pub struct SourceRow {
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, Serialize)]
pub struct CandidateRow {
pub group_id: String,
pub target_id: String,
pub target_smiles: String,
pub candidate_id: String,
pub precursor_smiles: Vec<String>,
pub source_template_count: usize,
pub best_upstream_rank: usize,
pub sources: Vec<SourceRow>,
pub feature_schema_version: u32,
pub feature_values: Vec<f32>,
pub feature_missing: Vec<bool>,
}
pub fn candidate_rows_for_pool(
pool: &CandidatePool,
target_mol: &Molecule,
templates_by_id: &HashMap<String, &RetroRule>,
stock: Option<&ChemEnv>,
) -> Vec<CandidateRow> {
let mut candidates: Vec<&crate::candidate::ReactionCandidate> =
pool.candidates.iter().collect();
candidates.sort_by(|a, b| a.candidate_id.cmp(&b.candidate_id));
candidates
.into_iter()
.map(|c| {
let features = extract_features(c, target_mol, templates_by_id, stock);
let sources = c
.sources
.iter()
.map(|s| SourceRow {
template_id: s.template_id.clone(),
rule_name: s.rule_name.clone(),
original_rank: s.original_rank,
upstream_score: s.upstream_score,
upstream_score_status: s.upstream_score_status,
template_log_frequency_raw: s.template_log_frequency_raw,
base_step_cost: s.base_step_cost,
})
.collect();
CandidateRow {
group_id: pool.group_id.clone(),
target_id: pool.target_id.clone(),
target_smiles: pool.target_smiles.clone(),
candidate_id: c.candidate_id.clone(),
precursor_smiles: c.precursor_smiles.clone(),
source_template_count: c.source_template_count,
best_upstream_rank: c.best_upstream_rank,
sources,
feature_schema_version: FEATURE_SCHEMA_VERSION,
feature_values: features.values,
feature_missing: features.missing,
}
})
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ProposalStatus {
Ok,
TargetParseFailed,
TargetIdMismatch,
}
#[derive(Debug, Clone, Serialize)]
pub struct TargetPoolRecord {
pub group_id: String,
pub target_id: String,
pub target_smiles: String,
pub candidate_count: usize,
pub proposal_status: ProposalStatus,
}
pub fn target_pool_record_for_pool(pool: &CandidatePool) -> TargetPoolRecord {
TargetPoolRecord {
group_id: pool.group_id.clone(),
target_id: pool.target_id.clone(),
target_smiles: pool.target_smiles.clone(),
candidate_count: pool.candidates.len(),
proposal_status: ProposalStatus::Ok,
}
}
pub fn target_pool_record_for_failure(
group_id: &str,
requested_target_smiles: &str,
) -> TargetPoolRecord {
TargetPoolRecord {
group_id: group_id.to_string(),
target_id: requested_target_smiles.to_string(),
target_smiles: requested_target_smiles.to_string(),
candidate_count: 0,
proposal_status: ProposalStatus::TargetParseFailed,
}
}
pub fn target_pool_record_for_target_id_mismatch(
group_id: &str,
requested_target_id: &str,
) -> TargetPoolRecord {
TargetPoolRecord {
group_id: group_id.to_string(),
target_id: requested_target_id.to_string(),
target_smiles: requested_target_id.to_string(),
candidate_count: 0,
proposal_status: ProposalStatus::TargetIdMismatch,
}
}
fn validate_group_index(records: &[TargetPoolRecord]) -> anyhow::Result<()> {
let mut seen = std::collections::HashSet::new();
for record in records {
if !seen.insert(record.group_id.as_str()) {
anyhow::bail!(
"duplicate group_id {:?} in the target/group index",
record.group_id
);
}
}
Ok(())
}
fn validate_candidate_rows(rows: &[CandidateRow]) -> anyhow::Result<()> {
let mut seen_candidate_ids_by_group: HashMap<&str, std::collections::HashSet<&str>> =
HashMap::new();
let mut target_id_by_group: HashMap<&str, &str> = HashMap::new();
let mut target_smiles_by_group: HashMap<&str, &str> = HashMap::new();
for row in rows {
if !seen_candidate_ids_by_group
.entry(row.group_id.as_str())
.or_default()
.insert(row.candidate_id.as_str())
{
anyhow::bail!(
"duplicate candidate_id {:?} within group_id {:?}",
row.candidate_id,
row.group_id
);
}
match target_id_by_group.get(row.group_id.as_str()) {
None => {
target_id_by_group.insert(row.group_id.as_str(), row.target_id.as_str());
}
Some(&existing) if existing != row.target_id => {
anyhow::bail!(
"group_id {:?} has rows with inconsistent target_id ({:?} vs {:?})",
row.group_id,
existing,
row.target_id
);
}
_ => {}
}
match target_smiles_by_group.get(row.group_id.as_str()) {
None => {
target_smiles_by_group.insert(row.group_id.as_str(), row.target_smiles.as_str());
}
Some(&existing) if existing != row.target_smiles => {
anyhow::bail!(
"group_id {:?} has rows with inconsistent target_smiles ({:?} vs {:?})",
row.group_id,
existing,
row.target_smiles
);
}
_ => {}
}
if row.precursor_smiles.is_empty() {
anyhow::bail!(
"candidate {:?} (group_id {:?}) has an empty precursor_smiles list",
row.candidate_id,
row.group_id
);
}
if row.sources.is_empty() {
anyhow::bail!(
"candidate {:?} (group_id {:?}) has an empty sources list",
row.candidate_id,
row.group_id
);
}
if row.feature_values.len() != FEATURE_NAMES_V1.len()
|| row.feature_missing.len() != FEATURE_NAMES_V1.len()
{
anyhow::bail!(
"candidate {:?} has feature_values.len()={} feature_missing.len()={}, \
expected {} (FEATURE_NAMES_V1.len())",
row.candidate_id,
row.feature_values.len(),
row.feature_missing.len(),
FEATURE_NAMES_V1.len()
);
}
for (i, (&value, &missing)) in row
.feature_values
.iter()
.zip(&row.feature_missing)
.enumerate()
{
if !missing && !value.is_finite() {
anyhow::bail!(
"candidate {:?} feature[{}] ({:?}) is non-finite ({}) but not marked missing",
row.candidate_id,
i,
FEATURE_NAMES_V1.get(i).copied().unwrap_or("?"),
value
);
}
}
}
Ok(())
}
fn validate_rows_consistent_with_group_index(
rows: &[CandidateRow],
target_pool_records: &[TargetPoolRecord],
) -> anyhow::Result<()> {
let index: HashMap<&str, &TargetPoolRecord> = target_pool_records
.iter()
.map(|r| (r.group_id.as_str(), r))
.collect();
let mut row_count_by_group: HashMap<&str, usize> = HashMap::new();
for row in rows {
match index.get(row.group_id.as_str()) {
None => anyhow::bail!(
"candidate row's group_id {:?} has no entry in the target/group index",
row.group_id
),
Some(record) => {
if record.target_id != row.target_id {
anyhow::bail!(
"group_id {:?}: candidate row target_id {:?} does not match \
group index target_id {:?}",
row.group_id,
row.target_id,
record.target_id
);
}
if record.target_smiles != row.target_smiles {
anyhow::bail!(
"group_id {:?}: candidate row target_smiles {:?} does not match \
group index target_smiles {:?}",
row.group_id,
row.target_smiles,
record.target_smiles
);
}
}
}
*row_count_by_group.entry(row.group_id.as_str()).or_insert(0) += 1;
}
for record in target_pool_records {
let actual = row_count_by_group
.get(record.group_id.as_str())
.copied()
.unwrap_or(0);
if actual != record.candidate_count {
anyhow::bail!(
"group_id {:?}: group index claims candidate_count={}, but {} candidate \
row(s) were actually found",
record.group_id,
record.candidate_count,
actual
);
}
}
Ok(())
}
pub fn write_target_pool_jsonl<W: Write>(
records: &[TargetPoolRecord],
mut writer: W,
) -> anyhow::Result<String> {
validate_group_index(records)?;
let mut hasher = Sha256::new();
for record in records {
let line = serde_json::to_vec(record)?;
hasher.update(&line);
hasher.update(b"\n");
writer.write_all(&line)?;
writer.write_all(b"\n")?;
}
Ok(format!("sha256:{}", crate::sha256_hex(hasher.finalize())))
}
pub fn write_jsonl<W: Write>(rows: &[CandidateRow], mut writer: W) -> anyhow::Result<String> {
validate_candidate_rows(rows)?;
let mut hasher = Sha256::new();
for row in rows {
let line = serde_json::to_vec(row)?;
hasher.update(&line);
hasher.update(b"\n");
writer.write_all(&line)?;
writer.write_all(b"\n")?;
}
Ok(format!("sha256:{}", crate::sha256_hex(hasher.finalize())))
}
pub fn rules_content_hash(rules: &[RetroRule]) -> String {
let mut sorted: Vec<&RetroRule> = rules.iter().collect();
sorted.sort_by(|a, b| a.template_id.cmp(&b.template_id));
let mut hasher = Sha256::new();
hasher.update(b"renkin-retrospect-rules-v2\0");
hasher.update((sorted.len() as u64).to_be_bytes());
for rule in sorted {
for field in [
rule.template_id.as_str(),
rule.name.as_str(),
rule.smirks.as_str(),
] {
hasher.update((field.len() as u64).to_be_bytes());
hasher.update(field.as_bytes());
}
hasher.update(rule.weight.to_bits().to_be_bytes());
hasher.update(rule.required_elements.to_be_bytes());
}
format!("sha256:{}", crate::sha256_hex(hasher.finalize()))
}
#[derive(Debug, Clone, Serialize)]
pub struct PoolProvenance {
pub renkin_git_commit: String,
pub cargo_lock_sha256: String,
pub chematic_version: String,
pub target_input_sha256: String,
pub stock_source: Option<String>,
pub embedded_fallback_used: bool,
pub export_config: serde_json::Value,
}
impl Default for PoolProvenance {
fn default() -> Self {
Self {
renkin_git_commit: String::new(),
cargo_lock_sha256: String::new(),
chematic_version: String::new(),
target_input_sha256: String::new(),
stock_source: None,
embedded_fallback_used: false,
export_config: serde_json::Value::Null,
}
}
}
#[derive(Debug, Clone, Serialize)]
pub struct PoolManifest {
pub manifest_schema_version: u32,
pub feature_schema_version: u32,
pub feature_names: Vec<&'static str>,
pub feature_schema_hash: String,
pub proposal_mode: ProposalModeSummary,
pub rules_content_hash: String,
pub rules_count: usize,
pub stock_identity: Option<String>,
pub stock_compound_count: Option<usize>,
pub stock_content_sha256: Option<String>,
pub target_count: usize,
pub group_count: usize,
pub candidate_count: usize,
pub candidate_jsonl_sha256: String,
pub target_group_index_sha256: String,
pub provenance: PoolProvenance,
}
pub const MANIFEST_SCHEMA_VERSION: u32 = 2;
#[allow(clippy::too_many_arguments)]
pub fn build_manifest(
rows: &[CandidateRow],
candidate_jsonl_sha256: &str,
target_pool_records: &[TargetPoolRecord],
target_group_index_sha256: &str,
rules: &[RetroRule],
mode: &ProposalMode,
stock: Option<(&str, &ChemEnv)>,
provenance: PoolProvenance,
) -> anyhow::Result<PoolManifest> {
validate_group_index(target_pool_records)?;
validate_rows_consistent_with_group_index(rows, target_pool_records)?;
let target_ids: std::collections::HashSet<&str> = target_pool_records
.iter()
.map(|r| r.target_id.as_str())
.collect();
Ok(PoolManifest {
manifest_schema_version: MANIFEST_SCHEMA_VERSION,
feature_schema_version: FEATURE_SCHEMA_VERSION,
feature_names: FEATURE_NAMES_V1.to_vec(),
feature_schema_hash: feature_schema_hash(),
proposal_mode: ProposalModeSummary::from_mode(mode),
rules_content_hash: rules_content_hash(rules),
rules_count: rules.len(),
stock_identity: stock.map(|(id, _)| id.to_string()),
stock_compound_count: stock.map(|(_, env)| env.bb_count()),
stock_content_sha256: stock.map(|(_, env)| env.content_sha256()),
target_count: target_ids.len(),
group_count: target_pool_records.len(),
candidate_count: rows.len(),
candidate_jsonl_sha256: candidate_jsonl_sha256.to_string(),
target_group_index_sha256: target_group_index_sha256.to_string(),
provenance,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::candidate::{ProposalConfig, index_rules_by_template_id, propose_one_step};
use crate::chem_env::{default_rules, mol_from_smiles};
#[test]
fn candidate_rows_carry_group_id_from_pool() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let pool =
propose_one_step("rxn-example-42", target, &rules, &ProposalConfig::default()).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let rows = candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
assert!(!rows.is_empty());
for row in &rows {
assert_eq!(row.group_id, "rxn-example-42");
assert_eq!(row.target_id, pool.target_id);
}
}
#[test]
fn target_pool_record_for_pool_reports_real_candidate_count() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let pool =
propose_one_step("rxn-example-1", target, &rules, &ProposalConfig::default()).unwrap();
let record = target_pool_record_for_pool(&pool);
assert_eq!(record.group_id, "rxn-example-1");
assert_eq!(record.target_id, pool.target_id);
assert_eq!(record.candidate_count, pool.candidates.len());
assert_eq!(record.proposal_status, ProposalStatus::Ok);
}
#[test]
fn target_pool_record_for_pool_still_emits_a_record_with_zero_candidates() {
let rules = vec![RetroRule {
name: "unreachable".to_string(),
template_id: "rule:unreachable".to_string(),
smirks: "[Xe:1]>>[Xe:1]".to_string(),
weight: 1.0,
required_elements: 0,
}];
let pool =
propose_one_step("rxn-example-2", "CCO", &rules, &ProposalConfig::default()).unwrap();
assert_eq!(pool.candidates.len(), 0);
let record = target_pool_record_for_pool(&pool);
assert_eq!(record.candidate_count, 0);
assert_eq!(record.proposal_status, ProposalStatus::Ok);
}
#[test]
fn target_pool_record_for_failure_is_distinguishable_from_a_real_zero_candidate_outcome() {
let record = target_pool_record_for_failure("rxn-example-3", "not-a-valid-smiles(((");
assert_eq!(record.group_id, "rxn-example-3");
assert_eq!(record.candidate_count, 0);
assert_eq!(record.proposal_status, ProposalStatus::TargetParseFailed);
assert_ne!(
record.proposal_status,
ProposalStatus::Ok,
"a parse failure must never be reported as a real zero-candidate outcome"
);
}
#[test]
fn target_pool_record_for_target_id_mismatch_keeps_the_requested_id_and_zero_candidates() {
let record = target_pool_record_for_target_id_mismatch("rxn-example-4", "requested-id");
assert_eq!(record.group_id, "rxn-example-4");
assert_eq!(record.target_id, "requested-id");
assert_eq!(record.target_smiles, "requested-id");
assert_eq!(record.candidate_count, 0);
assert_eq!(record.proposal_status, ProposalStatus::TargetIdMismatch);
assert_ne!(
record.proposal_status,
ProposalStatus::Ok,
"a target_id mismatch must never be reported as a real zero-candidate outcome"
);
}
#[test]
fn write_target_pool_jsonl_is_one_valid_json_object_per_line() {
let records = vec![
TargetPoolRecord {
group_id: "g1".to_string(),
target_id: "t1".to_string(),
target_smiles: "t1".to_string(),
candidate_count: 3,
proposal_status: ProposalStatus::Ok,
},
target_pool_record_for_failure("g2", "bad-smiles"),
];
let mut buf: Vec<u8> = Vec::new();
write_target_pool_jsonl(&records, &mut buf).unwrap();
let text = String::from_utf8(buf).unwrap();
let lines: Vec<&str> = text.lines().collect();
assert_eq!(lines.len(), 2);
for line in &lines {
let parsed: serde_json::Value =
serde_json::from_str(line).expect("each line must be valid JSON");
assert!(parsed.is_object());
}
let second: serde_json::Value = serde_json::from_str(lines[1]).unwrap();
assert_eq!(second["proposal_status"], "target_parse_failed");
}
#[test]
fn candidate_rows_are_sorted_by_candidate_id() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let pool = propose_one_step("group:1", target, &rules, &ProposalConfig::default()).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let rows = candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
assert!(!rows.is_empty());
let ids: Vec<&str> = rows.iter().map(|r| r.candidate_id.as_str()).collect();
let mut sorted_ids = ids.clone();
sorted_ids.sort_unstable();
assert_eq!(ids, sorted_ids);
}
#[test]
fn candidate_rows_carry_feature_schema_version_and_matching_lengths() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let pool = propose_one_step("group:1", target, &rules, &ProposalConfig::default()).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let rows = candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
for row in &rows {
assert_eq!(row.feature_schema_version, FEATURE_SCHEMA_VERSION);
assert_eq!(row.feature_values.len(), FEATURE_NAMES_V1.len());
assert_eq!(row.feature_missing.len(), FEATURE_NAMES_V1.len());
}
}
#[test]
fn candidate_rows_export_full_source_provenance_including_scorer_status() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let pool = propose_one_step("group:1", target, &rules, &ProposalConfig::default()).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let rows = candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
assert!(!rows.is_empty());
for row in &rows {
assert_eq!(row.sources.len(), row.source_template_count);
for source in &row.sources {
assert!(!source.template_id.is_empty());
assert!(!source.rule_name.is_empty());
assert_eq!(
source.upstream_score_status,
UpstreamScoreStatus::NotApplicable
);
}
}
}
#[test]
fn write_jsonl_is_one_valid_json_object_per_line() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let pool = propose_one_step("group:1", target, &rules, &ProposalConfig::default()).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let rows = candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
let mut buf: Vec<u8> = Vec::new();
write_jsonl(&rows, &mut buf).unwrap();
let text = String::from_utf8(buf).unwrap();
let lines: Vec<&str> = text.lines().collect();
assert_eq!(lines.len(), rows.len());
for line in &lines {
let parsed: serde_json::Value =
serde_json::from_str(line).expect("each line must be valid JSON");
assert!(parsed.is_object());
}
}
#[test]
fn two_export_runs_produce_byte_identical_jsonl() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let mut outputs = Vec::new();
for _ in 0..2 {
let pool =
propose_one_step("group:1", target, &rules, &ProposalConfig::default()).unwrap();
let rows = candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
let mut buf: Vec<u8> = Vec::new();
write_jsonl(&rows, &mut buf).unwrap();
outputs.push(buf);
}
assert_eq!(outputs[0], outputs[1]);
}
#[test]
fn rules_content_hash_differs_on_smirks_change_stable_otherwise() {
let rules_a = default_rules();
let mut rules_b = default_rules();
rules_b[0].smirks = format!("{}X", rules_b[0].smirks);
assert_ne!(rules_content_hash(&rules_a), rules_content_hash(&rules_b));
assert_eq!(
rules_content_hash(&rules_a),
rules_content_hash(&default_rules())
);
}
#[test]
fn rules_content_hash_is_order_independent() {
let mut rules_a = default_rules();
let mut rules_b = rules_a.clone();
rules_b.reverse();
rules_a.sort_by(|a, b| a.name.cmp(&b.name));
rules_b.sort_by(|a, b| a.name.cmp(&b.name));
assert_eq!(rules_content_hash(&rules_a), rules_content_hash(&rules_b));
}
fn export_pool_with_group_index(
pool: &CandidatePool,
target_mol: &Molecule,
templates_by_id: &HashMap<String, &RetroRule>,
) -> (Vec<CandidateRow>, String, Vec<TargetPoolRecord>, String) {
let rows = candidate_rows_for_pool(pool, target_mol, templates_by_id, None);
let records = vec![target_pool_record_for_pool(pool)];
let mut pool_buf: Vec<u8> = Vec::new();
let candidate_hash = write_jsonl(&rows, &mut pool_buf).unwrap();
let mut group_buf: Vec<u8> = Vec::new();
let group_hash = write_target_pool_jsonl(&records, &mut group_buf).unwrap();
(rows, candidate_hash, records, group_hash)
}
#[test]
fn manifest_records_mode_rules_and_stock_identity() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let pool = propose_one_step("group:1", target, &rules, &ProposalConfig::default()).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let (rows, candidate_hash, records, group_hash) =
export_pool_with_group_index(&pool, &target_mol, &templates_by_id);
let manifest = build_manifest(
&rows,
&candidate_hash,
&records,
&group_hash,
&rules,
&ProposalMode::Exhaustive,
None,
PoolProvenance::default(),
)
.unwrap();
assert_eq!(manifest.manifest_schema_version, MANIFEST_SCHEMA_VERSION);
assert_eq!(manifest.feature_schema_version, FEATURE_SCHEMA_VERSION);
assert_eq!(manifest.feature_names.len(), FEATURE_NAMES_V1.len());
assert_eq!(manifest.feature_schema_hash, feature_schema_hash());
assert_eq!(
manifest.proposal_mode,
ProposalModeSummary {
mode: "exhaustive",
top_k: None,
rules_offset: None,
scorer_identity: None,
scorer_model_sha256: None,
scorer_status: None,
}
);
assert_eq!(manifest.rules_count, rules.len());
assert_eq!(manifest.rules_content_hash, rules_content_hash(&rules));
assert!(manifest.stock_identity.is_none());
assert!(manifest.stock_compound_count.is_none());
assert!(manifest.stock_content_sha256.is_none());
assert_eq!(manifest.target_count, 1);
assert_eq!(manifest.group_count, 1);
assert_eq!(manifest.candidate_count, rows.len());
assert_eq!(manifest.candidate_jsonl_sha256, candidate_hash);
assert_eq!(manifest.target_group_index_sha256, group_hash);
let stock_a = ChemEnv::in_memory(&["CCO", "CC(=O)O"]);
let manifest_with_stock = build_manifest(
&rows,
&candidate_hash,
&records,
&group_hash,
&rules,
&ProposalMode::BondIndexed { top_k: 5 },
Some(("in_memory:test", &stock_a)),
PoolProvenance::default(),
)
.unwrap();
assert_eq!(
manifest_with_stock.stock_identity,
Some("in_memory:test".to_string())
);
assert_eq!(manifest_with_stock.stock_compound_count, Some(2));
assert_eq!(
manifest_with_stock.stock_content_sha256,
Some(stock_a.content_sha256())
);
assert_eq!(
manifest_with_stock.proposal_mode,
ProposalModeSummary {
mode: "bond_indexed",
top_k: Some(5),
rules_offset: None,
scorer_identity: None,
scorer_model_sha256: None,
scorer_status: None,
}
);
let stock_b = ChemEnv::in_memory(&["CCN"]);
let manifest_with_swapped_stock = build_manifest(
&rows,
&candidate_hash,
&records,
&group_hash,
&rules,
&ProposalMode::BondIndexed { top_k: 5 },
Some(("in_memory:test", &stock_b)),
PoolProvenance::default(),
)
.unwrap();
assert_ne!(
manifest_with_stock.stock_content_sha256,
manifest_with_swapped_stock.stock_content_sha256,
"a swapped stock under an unchanged label must still change the content hash"
);
}
#[test]
fn manifest_serializes_to_json() {
let rules = default_rules();
let manifest = build_manifest(
&[],
"sha256:empty",
&[],
"sha256:empty",
&rules,
&ProposalMode::Exhaustive,
None,
PoolProvenance::default(),
)
.unwrap();
let json = serde_json::to_string(&manifest).unwrap();
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
assert!(parsed["feature_names"].is_array());
assert_eq!(parsed["proposal_mode"]["mode"], "exhaustive");
}
#[test]
fn manifest_records_scorer_conditioned_provenance() {
let mut rules = default_rules();
let n_handcrafted = rules.len();
rules.push(RetroRule {
name: "extracted_0".to_string(),
template_id: "smirks-sha256:fake0".to_string(),
smirks: "[C:1][C:2]>>[C:1].[C:2]".to_string(),
weight: 3.0,
required_elements: 0,
});
let mode = ProposalMode::ScorerConditioned {
input: crate::candidate::ScorerConditionedInput {
scores: vec![(n_handcrafted, 0.9, 0)],
status: UpstreamScoreStatus::Available,
rules_offset: n_handcrafted,
scorer_identity: "test-scorer-v1".to_string(),
scorer_model_sha256: "sha256:modelbytes".to_string(),
},
top_k: 1,
};
let manifest = build_manifest(
&[],
"sha256:empty",
&[],
"sha256:empty",
&rules,
&mode,
None,
PoolProvenance::default(),
)
.unwrap();
assert_eq!(manifest.proposal_mode.mode, "scorer_conditioned");
assert_eq!(manifest.proposal_mode.rules_offset, Some(n_handcrafted));
assert_eq!(
manifest.proposal_mode.scorer_identity,
Some("test-scorer-v1".to_string())
);
assert_eq!(
manifest.proposal_mode.scorer_model_sha256,
Some("sha256:modelbytes".to_string())
);
assert_eq!(
manifest.proposal_mode.scorer_status,
Some(UpstreamScoreStatus::Available)
);
}
#[test]
fn manifest_records_embedded_fallback_used_provenance() {
let rules = default_rules();
let provenance = PoolProvenance {
embedded_fallback_used: true,
..PoolProvenance::default()
};
let manifest = build_manifest(
&[],
"sha256:empty",
&[],
"sha256:empty",
&rules,
&ProposalMode::Exhaustive,
None,
provenance,
)
.unwrap();
assert!(manifest.provenance.embedded_fallback_used);
let json = serde_json::to_string(&manifest).unwrap();
let parsed: serde_json::Value = serde_json::from_str(&json).unwrap();
assert_eq!(parsed["provenance"]["embedded_fallback_used"], true);
}
#[test]
fn build_manifest_rejects_group_id_missing_from_group_index() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let pool = propose_one_step("group:1", target, &rules, &ProposalConfig::default()).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let rows = candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
assert!(!rows.is_empty());
let result = build_manifest(
&rows,
"sha256:whatever",
&[],
"sha256:whatever",
&rules,
&ProposalMode::Exhaustive,
None,
PoolProvenance::default(),
);
assert!(
result.is_err(),
"a group_id present in rows but absent from the group index must be a hard error"
);
}
#[test]
fn build_manifest_rejects_duplicate_group_id_in_group_index() {
let rules = default_rules();
let records = vec![
TargetPoolRecord {
group_id: "g1".to_string(),
target_id: "t1".to_string(),
target_smiles: "t1".to_string(),
candidate_count: 0,
proposal_status: ProposalStatus::Ok,
},
TargetPoolRecord {
group_id: "g1".to_string(),
target_id: "t2".to_string(),
target_smiles: "t2".to_string(),
candidate_count: 0,
proposal_status: ProposalStatus::Ok,
},
];
let result = build_manifest(
&[],
"sha256:empty",
&records,
"sha256:whatever",
&rules,
&ProposalMode::Exhaustive,
None,
PoolProvenance::default(),
);
assert!(
result.is_err(),
"a duplicate group_id in the group index must be a hard error"
);
}
#[test]
fn build_manifest_rejects_candidate_count_mismatch() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let pool = propose_one_step("group:1", target, &rules, &ProposalConfig::default()).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let rows = candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
assert!(
rows.len() > 1,
"fixture must have more than one candidate for this check to bite"
);
let mut record = target_pool_record_for_pool(&pool);
record.candidate_count += 1; let result = build_manifest(
&rows,
"sha256:whatever",
&[record],
"sha256:whatever",
&rules,
&ProposalMode::Exhaustive,
None,
PoolProvenance::default(),
);
assert!(
result.is_err(),
"a group index candidate_count that disagrees with the actual row count must be rejected"
);
}
#[test]
fn write_jsonl_rejects_duplicate_candidate_id_within_one_group() {
let base = CandidateRow {
group_id: "g1".to_string(),
target_id: "t1".to_string(),
target_smiles: "t1".to_string(),
candidate_id: "same-id".to_string(),
precursor_smiles: vec!["CCO".to_string()],
source_template_count: 1,
best_upstream_rank: 0,
sources: vec![SourceRow {
template_id: "rule:x".to_string(),
rule_name: "x".to_string(),
original_rank: 0,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
template_log_frequency_raw: None,
base_step_cost: 1.0,
}],
feature_schema_version: FEATURE_SCHEMA_VERSION,
feature_values: vec![0.0; FEATURE_NAMES_V1.len()],
feature_missing: vec![true; FEATURE_NAMES_V1.len()],
};
let rows = vec![base.clone(), base];
let mut buf: Vec<u8> = Vec::new();
assert!(write_jsonl(&rows, &mut buf).is_err());
}
#[test]
fn write_jsonl_allows_same_candidate_id_across_different_groups() {
let mut row_a = CandidateRow {
group_id: "g1".to_string(),
target_id: "t1".to_string(),
target_smiles: "t1".to_string(),
candidate_id: "same-id".to_string(),
precursor_smiles: vec!["CCO".to_string()],
source_template_count: 1,
best_upstream_rank: 0,
sources: vec![SourceRow {
template_id: "rule:x".to_string(),
rule_name: "x".to_string(),
original_rank: 0,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
template_log_frequency_raw: None,
base_step_cost: 1.0,
}],
feature_schema_version: FEATURE_SCHEMA_VERSION,
feature_values: vec![0.0; FEATURE_NAMES_V1.len()],
feature_missing: vec![true; FEATURE_NAMES_V1.len()],
};
let mut row_b = row_a.clone();
row_b.group_id = "g2".to_string();
row_a.candidate_id = "same-id".to_string();
row_b.candidate_id = "same-id".to_string();
let mut buf: Vec<u8> = Vec::new();
assert!(write_jsonl(&[row_a, row_b], &mut buf).is_ok());
}
#[test]
fn write_jsonl_rejects_malformed_feature_vector_lengths() {
let mut row = CandidateRow {
group_id: "g1".to_string(),
target_id: "t1".to_string(),
target_smiles: "t1".to_string(),
candidate_id: "id-1".to_string(),
precursor_smiles: vec!["CCO".to_string()],
source_template_count: 1,
best_upstream_rank: 0,
sources: vec![SourceRow {
template_id: "rule:x".to_string(),
rule_name: "x".to_string(),
original_rank: 0,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
template_log_frequency_raw: None,
base_step_cost: 1.0,
}],
feature_schema_version: FEATURE_SCHEMA_VERSION,
feature_values: vec![0.0; FEATURE_NAMES_V1.len() - 1],
feature_missing: vec![true; FEATURE_NAMES_V1.len()],
};
let mut buf: Vec<u8> = Vec::new();
assert!(write_jsonl(std::slice::from_ref(&row), &mut buf).is_err());
row.feature_values = vec![0.0; FEATURE_NAMES_V1.len()];
row.feature_missing = vec![false; FEATURE_NAMES_V1.len()];
row.feature_values[0] = f32::NAN;
let mut buf2: Vec<u8> = Vec::new();
assert!(
write_jsonl(&[row], &mut buf2).is_err(),
"a non-finite value not marked missing must be rejected"
);
}
#[test]
fn write_jsonl_rejects_empty_precursor_or_source_lists() {
let mut row = CandidateRow {
group_id: "g1".to_string(),
target_id: "t1".to_string(),
target_smiles: "t1".to_string(),
candidate_id: "id-1".to_string(),
precursor_smiles: vec![],
source_template_count: 1,
best_upstream_rank: 0,
sources: vec![SourceRow {
template_id: "rule:x".to_string(),
rule_name: "x".to_string(),
original_rank: 0,
upstream_score: None,
upstream_score_status: UpstreamScoreStatus::NotApplicable,
template_log_frequency_raw: None,
base_step_cost: 1.0,
}],
feature_schema_version: FEATURE_SCHEMA_VERSION,
feature_values: vec![0.0; FEATURE_NAMES_V1.len()],
feature_missing: vec![true; FEATURE_NAMES_V1.len()],
};
let mut buf: Vec<u8> = Vec::new();
assert!(write_jsonl(std::slice::from_ref(&row), &mut buf).is_err());
row.precursor_smiles = vec!["CCO".to_string()];
row.sources = vec![];
let mut buf2: Vec<u8> = Vec::new();
assert!(write_jsonl(&[row], &mut buf2).is_err());
}
#[test]
fn write_functions_return_the_digest_of_exactly_what_they_wrote() {
let rules = default_rules();
let target = "CC(=O)c1ccccc1";
let target_mol = mol_from_smiles(target).unwrap();
let pool = propose_one_step("group:1", target, &rules, &ProposalConfig::default()).unwrap();
let templates_by_id = index_rules_by_template_id(&rules).unwrap();
let rows = candidate_rows_for_pool(&pool, &target_mol, &templates_by_id, None);
let mut buf: Vec<u8> = Vec::new();
let returned_hash = write_jsonl(&rows, &mut buf).unwrap();
let mut hasher = Sha256::new();
hasher.update(&buf);
let recomputed = format!("sha256:{}", crate::sha256_hex(hasher.finalize()));
assert_eq!(returned_hash, recomputed);
}
}