use std::collections::{BTreeMap, HashMap};
use std::fs;
use std::path::{Path, PathBuf};
use serde::{Deserialize, Serialize};
use super::population::{load_commits, CommitRow};
use super::records::{SampleRecord, Stratum};
use super::sample::{write_json, write_jsonl};
use super::score::read_sample;
use super::subsample::{create_private_dir, create_private_file};
use super::verdict::{resolve_verdicts, CarryPolicy};
use super::{io_err, open_eval_db, EvalError, Result};
use crate::classify::ClassificationPipeline;
use crate::core::config::Config;
#[derive(Debug, Clone)]
pub struct RepredictParams {
pub sample: PathBuf,
pub db: PathBuf,
pub config: Config,
pub config_path: PathBuf,
pub out: PathBuf,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct HashedFile {
pub path: String,
pub blake3: String,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
#[non_exhaustive]
pub struct Provenance {
pub tga_version: String,
pub config: HashedFile,
pub rules_files: Vec<HashedFile>,
pub source_sample: HashedFile,
pub db: String,
pub rows: u64,
pub changed: u64,
pub abstentions: u64,
pub carried: BTreeMap<String, u64>,
pub superseded: u64,
}
#[derive(Debug, Clone, Default)]
#[non_exhaustive]
pub struct RepredictSummary {
pub provenance: Provenance,
pub files: Vec<PathBuf>,
}
pub fn provenance_path(out: &Path) -> PathBuf {
out.with_extension("provenance.json")
}
fn hashed(path: &Path) -> Result<HashedFile> {
let bytes = fs::read(path).map_err(io_err(path))?;
let shown = fs::canonicalize(path).unwrap_or_else(|_| path.to_path_buf());
Ok(HashedFile {
path: shown.display().to_string(),
blake3: blake3::hash(&bytes).to_hex().to_string(),
})
}
pub fn run_repredict(params: &RepredictParams) -> Result<RepredictSummary> {
let sample = read_sample(¶ms.sample)?;
if sample.is_empty() {
return Err(EvalError::Invalid(format!(
"{} holds no rows",
params.sample.display()
)));
}
let prov_path = provenance_path(¶ms.out);
for f in [¶ms.out, &prov_path] {
if f.exists() {
return Err(EvalError::Invalid(format!(
"{} already exists; pick a new --out",
f.display()
)));
}
}
let conn = open_eval_db(¶ms.db)?;
let all = load_commits(&conn)?;
let by_key: HashMap<(&str, &str), &CommitRow> = all
.iter()
.map(|c| ((c.sha.as_str(), c.repo.as_str()), c))
.collect();
let found: Vec<Option<&CommitRow>> = sample
.iter()
.map(|r| by_key.get(&(r.sha.as_str(), r.repo.as_str())).copied())
.collect();
let missing: Vec<&SampleRecord> = sample
.iter()
.zip(&found)
.filter(|(_, c)| c.is_none())
.map(|(r, _)| r)
.collect();
if let Some(first) = missing.first() {
return Err(EvalError::Invalid(format!(
"{} sample rows are not in {} (first: {} in {}); pass the database the sample \
was drawn from",
missing.len(),
params.db.display(),
first.sha,
first.repo
)));
}
let commits: Vec<&CommitRow> = found.into_iter().flatten().collect();
let engine = ClassificationPipeline::new(params.config.clone()).build_rule_engine()?;
let policy = CarryPolicy::from_config(¶ms.config)?;
let (resolved, _drifted) = resolve_verdicts(&engine, &policy, &commits);
let (mut changed, mut abstentions, mut superseded) = (0u64, 0u64, 0u64);
let mut carried: BTreeMap<String, u64> = BTreeMap::new();
let records: Vec<SampleRecord> = sample
.iter()
.zip(&commits)
.zip(&resolved)
.map(|((r, c), v)| {
changed += u64::from(r.predicted_category != v.category);
abstentions +=
u64::from(Stratum::classify(v.tier, &v.category, v.confidence) == Stratum::Unknown);
if v.carried {
*carried.entry(v.tier.as_str().to_string()).or_default() += 1;
}
superseded += u64::from(v.superseded);
SampleRecord {
method: v.tier.as_str().to_string(),
rule_id: v.rule_id.clone(),
predicted_category: v.category.clone(),
confidence: v.confidence,
is_merge: r.is_merge.or(Some(c.is_merge)),
..r.clone()
}
})
.collect();
let rules_files = params
.config
.classification
.as_ref()
.map(|c| {
c.rules_files
.iter()
.map(|p| hashed(p))
.collect::<Result<Vec<_>>>()
})
.transpose()?
.unwrap_or_default();
let provenance = Provenance {
tga_version: env!("CARGO_PKG_VERSION").to_string(),
config: hashed(¶ms.config_path)?,
rules_files,
source_sample: hashed(¶ms.sample)?,
db: params.db.display().to_string(),
rows: records.len() as u64,
changed,
abstentions,
carried,
superseded,
};
if let Some(dir) = params.out.parent().filter(|d| !d.as_os_str().is_empty()) {
create_private_dir(dir)?;
}
create_private_file(¶ms.out)?;
let mut created = vec![params.out.as_path()];
let written = create_private_file(&prov_path).and_then(|()| {
created.push(prov_path.as_path());
write_jsonl(¶ms.out, &records)?;
write_json(&prov_path, &provenance)
});
if let Err(e) = written {
for f in created {
let _ = fs::remove_file(f);
}
return Err(e);
}
Ok(RepredictSummary {
provenance,
files: vec![params.out.clone(), prov_path],
})
}