Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 26 additions & 2 deletions src/controllers/apotheosis.rs
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,12 @@ use std::fs::{self, File};
use std::io::{Read, Write};
use std::path::{Path, PathBuf};

/// Last path segment of the type name, e.g.
/// "apotheosis2::datalayer::algorithms::TlshDistance" -> "TlshDistance".
fn distance_name<T>() -> &'static str {
std::any::type_name::<T>().rsplit("::").next().unwrap()
}

#[derive(serde::Serialize, serde::Deserialize)]
#[serde(bound(
serialize = "R: ApotheosisRecord + serde::Serialize, D: DistanceAlgorithm<R::MetricId> + serde::Serialize, R::MetricId: serde::Serialize",
Expand Down Expand Up @@ -137,6 +143,9 @@ where
file.write_all(&(M0 as u32).to_le_bytes())?;
file.write_all(&(EF as u32).to_le_bytes())?;
file.write_all(&[HEURISTIC as u8])?;
let name = distance_name::<D>().as_bytes();
file.write_all(&[u8::try_from(name.len())?])?;
file.write_all(name)?;

// Then write the model data
bincode::serialize_into(file, self)?;
Expand All @@ -149,7 +158,7 @@ where
{
let mut file = File::open(path)?;

// Read and verify header (17 bytes)
// Read and verify header (17 bytes + distance type name)
let mut header = [0u8; 17];
file.read_exact(&mut header)?;

Expand All @@ -169,6 +178,22 @@ where
).into());
}

// Explicit check of distance type (does not rely on bincode)
let mut name_len = [0u8; 1];
file.read_exact(&mut name_len)?;
let mut name = vec![0u8; name_len[0] as usize];
file.read_exact(&mut name)?;
let file_distance = String::from_utf8_lossy(&name);

let expected = distance_name::<D>();
if file_distance != expected {
return Err(format!(
"Distance type mismatch. File was built with {} but is being loaded as {}",
file_distance, expected
)
.into());
}

let decoded = bincode::deserialize_from(file)?;
Ok(decoded)
}
Expand All @@ -179,7 +204,6 @@ where
///
/// # Parameters
/// * `path` - Base filename for output (e.g., "model" creates "model_layer0.gexf", "model_layer1.gexf", etc.)

pub fn draw<P: AsRef<Path>>(&self, path: P) {
let base_path = path.as_ref();

Expand Down
2 changes: 1 addition & 1 deletion src/datalayer/algorithms.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ impl DistanceAlgorithm<u32> for NormalDistance {
pub struct TlshDistance;
impl DistanceAlgorithm<TlshDefault> for TlshDistance {
fn calculate_distance(&self, a: &TlshDefault, b: &TlshDefault) -> u32 {
let diff = a.diff(&b, true);
let diff = a.diff(b, true);
diff as u32
}
}
Loading