use std::collections::{BTreeMap, HashMap};
use std::path::Path;
use rusqlite::params_from_iter;
use super::records::{SampleRecord, StrataSummary, Stratum};
use super::{open_eval_db, EvalError, Result};
pub(crate) fn resolve_merges(sample: &[SampleRecord], db: Option<&Path>) -> Result<Vec<bool>> {
let unknown: Vec<&str> = sample
.iter()
.filter(|r| r.is_merge.is_none())
.map(|r| r.sha.as_str())
.collect();
if unknown.is_empty() {
return Ok(sample.iter().map(|r| r.is_merge == Some(true)).collect());
}
let Some(db) = db else {
return Err(EvalError::Invalid(format!(
"{} sample rows carry no merge flag (a sample written before merges were \
excluded); pass --db with a copy of the tga database so merges can be \
resolved by SHA and excluded",
unknown.len()
)));
};
let stored = lookup(db, &unknown)?;
let flags: Vec<Option<bool>> = sample
.iter()
.map(|r| r.is_merge.or_else(|| stored.get(r.sha.as_str()).copied()))
.collect();
let missing: Vec<&str> = sample
.iter()
.zip(&flags)
.filter(|(_, f)| f.is_none())
.map(|(r, _)| r.sha.as_str())
.collect();
if let Some(first) = missing.first() {
return Err(EvalError::Invalid(format!(
"merge status unknown for {} sample rows: their SHAs are not in {} (first: {first}); \
pass the database the sample was drawn from",
missing.len(),
db.display()
)));
}
Ok(flags.into_iter().map(|f| f == Some(true)).collect())
}
pub(crate) fn scale_out_merges(
strata: &mut StrataSummary,
sample: &[SampleRecord],
merges: &[bool],
) -> u64 {
let mut total = 0;
let mut rows: BTreeMap<Stratum, (u64, u64)> = BTreeMap::new();
for (r, &m) in sample.iter().zip(merges) {
let e = rows.entry(r.stratum).or_default();
e.0 += 1;
e.1 += u64::from(m);
}
for (stratum, (n, m)) in rows {
let Some(counts) = strata.strata.get_mut(stratum.as_str()) else {
continue;
};
if m == 0 {
continue;
}
let removed = (counts.population as f64 * m as f64 / n as f64).round() as u64;
let removed = removed.min(counts.population);
counts.population -= removed;
strata.population = strata.population.saturating_sub(removed);
strata.merges_excluded += removed;
total += removed;
}
total
}
pub(crate) fn count_window_merges(db: &Path, strata: &StrataSummary) -> Result<u64> {
let bound = |s: &str| {
super::population::parse_ts(s).ok_or_else(|| {
EvalError::Invalid(format!("strata.json window bound {s:?} is not a timestamp"))
})
};
let (start, end) = (bound(&strata.window_start)?, bound(&strata.window_end)?);
let conn = open_eval_db(db)?;
let mut stmt = conn.prepare("SELECT timestamp FROM commits WHERE is_merge = 1")?;
let rows = stmt.query_map([], |r| r.get::<_, String>(0))?;
let mut n = 0;
for ts in rows {
if super::population::parse_ts(&ts?).is_some_and(|t| t >= start && t <= end) {
n += 1;
}
}
Ok(n)
}
fn lookup(db: &Path, shas: &[&str]) -> Result<HashMap<String, bool>> {
let conn = open_eval_db(db)?;
let mut out = HashMap::new();
for chunk in shas.chunks(500) {
let placeholders = vec!["?"; chunk.len()].join(",");
let sql = format!(
"SELECT sha, MAX(is_merge) FROM commits WHERE sha IN ({placeholders}) GROUP BY sha"
);
let mut stmt = conn.prepare(&sql)?;
let rows = stmt.query_map(params_from_iter(chunk.iter()), |r| {
Ok((r.get::<_, String>(0)?, r.get::<_, i64>(1)? != 0))
})?;
for row in rows {
let (sha, is_merge) = row?;
out.insert(sha, is_merge);
}
}
Ok(out)
}