use std::collections::HashMap;
use std::fs;
use std::io::{BufWriter, Write};
use std::path::Path;
pub fn load(path: &Path, expected_depth: u32) -> HashMap<String, i32> {
let content = match fs::read_to_string(path) {
Ok(c) => c,
Err(e) => {
eprintln!("teacher cache: cannot read {:?}: {e}", path);
return HashMap::new();
}
};
let mut map = HashMap::new();
let mut skipped = 0usize;
let mut depth_mismatch = 0usize;
for line in content.lines() {
let line = line.trim();
if line.is_empty() {
continue;
}
let Ok(val) = serde_json::from_str::<serde_json::Value>(line) else {
skipped += 1;
continue;
};
let Some(sfen) = val.get("sfen").and_then(|v| v.as_str()) else {
skipped += 1;
continue;
};
let Some(cp) = val.get("score_cp").and_then(|v| v.as_i64()) else {
skipped += 1;
continue;
};
match val.get("label_depth").and_then(|v| v.as_u64()) {
Some(d) if d as u32 == expected_depth => {}
Some(_) => {
depth_mismatch += 1;
continue;
}
None => {
skipped += 1;
continue;
}
}
map.insert(sfen.to_string(), cp as i32);
}
if skipped > 0 {
eprintln!("teacher cache: {skipped} lines skipped (unparseable)");
}
if depth_mismatch > 0 {
eprintln!(
"teacher cache: {depth_mismatch} entries skipped (label_depth != {expected_depth})"
);
}
eprintln!(
"teacher cache: {} entries loaded from {:?}",
map.len(),
path
);
map
}
pub fn write(path: &Path, entries: &HashMap<String, i32>, label_depth: u32) -> std::io::Result<()> {
let tmp_path = path.with_extension("jsonl.tmp");
{
let f = fs::File::create(&tmp_path)?;
let mut w = BufWriter::new(f);
for (sfen, &cp) in entries {
writeln!(
w,
r#"{{"sfen":{},"label_depth":{},"score_cp":{}}}"#,
json_string(sfen),
label_depth,
cp
)?;
}
w.flush()?;
}
fs::rename(&tmp_path, path)?;
Ok(())
}
fn json_string(s: &str) -> String {
format!("\"{}\"", s.replace('\\', "\\\\").replace('"', "\\\""))
}
#[cfg(test)]
mod tests {
use super::*;
use tempfile::NamedTempFile;
const SFEN_A: &str = "lnsgkgsnl/1r5b1/ppppppppp/9/9/9/PPPPPPPPP/1B5R1/LNSGKGSNL b - 1";
const SFEN_B: &str = "lnsgkgsnl/1r5b1/ppppppppp/9/9/2P6/PP1PPPPPP/1B5R1/LNSGKGSNL w - 2";
#[test]
fn roundtrip() {
let f = NamedTempFile::new().unwrap();
let mut expected = HashMap::new();
expected.insert(SFEN_A.to_string(), 48i32);
expected.insert(SFEN_B.to_string(), -120i32);
write(f.path(), &expected, 4).unwrap();
let loaded = load(f.path(), 4);
assert_eq!(loaded, expected);
}
#[test]
fn broken_lines_skipped() {
let mut f = NamedTempFile::new().unwrap();
writeln!(f, "not json").unwrap();
writeln!(f, r#"{{"sfen":"{SFEN_A}","label_depth":4,"score_cp":100}}"#).unwrap();
let loaded = load(f.path(), 4);
assert_eq!(loaded.len(), 1);
assert_eq!(loaded[SFEN_A], 100);
}
#[test]
fn missing_score_cp_skipped() {
let mut f = NamedTempFile::new().unwrap();
writeln!(f, r#"{{"sfen":"{SFEN_A}","label_depth":4}}"#).unwrap();
let loaded = load(f.path(), 4);
assert!(loaded.is_empty());
}
#[test]
fn truncated_trailing_line_is_skipped_not_fatal() {
let mut f = NamedTempFile::new().unwrap();
writeln!(f, r#"{{"sfen":"{SFEN_A}","label_depth":4,"score_cp":100}}"#).unwrap();
write!(f, r#"{{"sfen":"{SFEN_B}","label_depth":4,"sco"#).unwrap(); let loaded = load(f.path(), 4);
assert_eq!(loaded.len(), 1);
assert_eq!(loaded[SFEN_A], 100);
}
#[test]
fn wrong_depth_entries_are_filtered_out_and_reported() {
let mut f = NamedTempFile::new().unwrap();
writeln!(f, r#"{{"sfen":"{SFEN_A}","label_depth":1,"score_cp":999}}"#).unwrap();
writeln!(f, r#"{{"sfen":"{SFEN_B}","label_depth":4,"score_cp":100}}"#).unwrap();
let loaded = load(f.path(), 4);
assert_eq!(loaded.len(), 1);
assert_eq!(loaded[SFEN_B], 100);
assert!(
!loaded.contains_key(SFEN_A),
"depth-1 entry must not be usable as a depth-4 cache hit"
);
}
#[test]
fn duplicate_key_resolves_to_last_occurrence_in_file() {
let mut f = NamedTempFile::new().unwrap();
writeln!(f, r#"{{"sfen":"{SFEN_A}","label_depth":4,"score_cp":100}}"#).unwrap();
writeln!(f, r#"{{"sfen":"{SFEN_A}","label_depth":4,"score_cp":250}}"#).unwrap();
let loaded = load(f.path(), 4);
assert_eq!(loaded[SFEN_A], 250);
}
#[test]
fn write_is_atomic_no_tmp_file_left_behind_on_success() {
let f = NamedTempFile::new().unwrap();
let mut entries = HashMap::new();
entries.insert(SFEN_A.to_string(), 48i32);
write(f.path(), &entries, 4).unwrap();
let tmp_path = f.path().with_extension("jsonl.tmp");
assert!(
!tmp_path.exists(),
"the intermediate .tmp file must be renamed away, not left behind"
);
assert_eq!(load(f.path(), 4), entries);
}
}