use std::path::Path;
use super::super::error::OutputError;
use super::codec::{
deserialize_stage_basis, deserialize_stage_cuts, deserialize_stage_states,
read_sorted_bin_files, serialize_stage_basis, serialize_stage_cuts, serialize_stage_states,
};
use super::records::{
OwnedPolicyBasisRecord, PolicyBasisRecord, PolicyCheckpoint, PolicyCheckpointMetadata,
StageCutsPayload, StageCutsReadResult, StageStatesPayload, StageStatesReadResult,
};
pub fn write_policy_checkpoint(
path: &Path,
stage_cuts: &[StageCutsPayload<'_>],
stage_bases: &[PolicyBasisRecord<'_>],
metadata: &PolicyCheckpointMetadata,
stage_states: &[StageStatesPayload<'_>],
) -> Result<(), OutputError> {
let cuts_dir = path.join("cuts");
std::fs::create_dir_all(&cuts_dir).map_err(|e| OutputError::io(&cuts_dir, e))?;
let basis_dir = path.join("basis");
std::fs::create_dir_all(&basis_dir).map_err(|e| OutputError::io(&basis_dir, e))?;
for payload in stage_cuts {
let filename = format!("stage_{:03}.bin", payload.stage_id);
let file_path = cuts_dir.join(&filename);
let buf = serialize_stage_cuts(
payload.stage_id,
payload.state_dimension,
payload.capacity,
payload.warm_start_count,
payload.cuts,
payload.active_cut_indices,
payload.populated_count,
);
std::fs::write(&file_path, &buf).map_err(|e| OutputError::io(&file_path, e))?;
}
for record in stage_bases {
let filename = format!("stage_{:03}.bin", record.stage_id);
let file_path = basis_dir.join(&filename);
let buf = serialize_stage_basis(record);
std::fs::write(&file_path, &buf).map_err(|e| OutputError::io(&file_path, e))?;
}
if !stage_states.is_empty() {
let states_dir = path.join("states");
std::fs::create_dir_all(&states_dir).map_err(|e| OutputError::io(&states_dir, e))?;
for payload in stage_states {
let filename = format!("stage_{:03}.bin", payload.stage_id);
let file_path = states_dir.join(&filename);
let buf = serialize_stage_states(payload);
std::fs::write(&file_path, &buf).map_err(|e| OutputError::io(&file_path, e))?;
}
}
let json = serde_json::to_string_pretty(metadata)
.map_err(|e| OutputError::serialization("policy_metadata", e.to_string()))?;
let meta_path = path.join("metadata.json");
std::fs::write(&meta_path, json.as_bytes()).map_err(|e| OutputError::io(&meta_path, e))?;
Ok(())
}
pub fn read_policy_checkpoint(path: &Path) -> Result<PolicyCheckpoint, OutputError> {
let meta_path = path.join("metadata.json");
let meta_bytes = std::fs::read(&meta_path).map_err(|e| OutputError::io(&meta_path, e))?;
let metadata: PolicyCheckpointMetadata = serde_json::from_slice(&meta_bytes)
.map_err(|e| OutputError::serialization("policy_metadata", e.to_string()))?;
let cuts_dir = path.join("cuts");
let mut stage_cuts: Vec<StageCutsReadResult> =
read_sorted_bin_files(&cuts_dir, "stage_cuts", deserialize_stage_cuts)?;
stage_cuts.sort_by_key(|r| r.stage_id);
let basis_dir = path.join("basis");
let mut stage_bases: Vec<OwnedPolicyBasisRecord> =
read_sorted_bin_files(&basis_dir, "stage_basis", deserialize_stage_basis)?;
stage_bases.sort_by_key(|r| r.stage_id);
let states_dir = path.join("states");
let stage_states: Vec<StageStatesReadResult> = if states_dir.is_dir() {
let mut ss = read_sorted_bin_files(&states_dir, "stage_states", deserialize_stage_states)?;
ss.sort_by_key(|r| r.stage_id);
ss
} else {
Vec::new()
};
Ok(PolicyCheckpoint {
metadata,
stage_cuts,
stage_bases,
stage_states,
})
}