1use crate::crc32::crc32;
9use crate::error::{AprFormatError, Result};
10use crate::types::{Compression, Header, Metadata, ModelInfo, ModelType, SaveOptions, HEADER_SIZE};
11use serde::{de::DeserializeOwned, Serialize};
12use std::fs::File;
13use std::io::{BufReader, BufWriter, Read, Write};
14use std::path::Path;
15
16pub const MMAP_THRESHOLD: u64 = 1024 * 1024;
21
22#[allow(clippy::unnecessary_wraps)]
27fn compress_payload(data: &[u8], compression: Compression) -> Result<(Vec<u8>, Compression)> {
28 match compression {
29 Compression::None => Ok((data.to_vec(), Compression::None)),
30 #[cfg(feature = "compression")]
31 Compression::ZstdDefault => {
32 let compressed = zstd::encode_all(std::io::Cursor::new(data), 3).map_err(|e| {
33 AprFormatError::Serialization(format!("Zstd compression failed: {e}"))
34 })?;
35 Ok((compressed, Compression::ZstdDefault))
36 }
37 #[cfg(feature = "compression")]
38 Compression::ZstdMax => {
39 let compressed = zstd::encode_all(std::io::Cursor::new(data), 19).map_err(|e| {
40 AprFormatError::Serialization(format!("Zstd compression failed: {e}"))
41 })?;
42 Ok((compressed, Compression::ZstdMax))
43 }
44 #[cfg(not(feature = "compression"))]
45 Compression::ZstdDefault | Compression::ZstdMax => Ok((data.to_vec(), Compression::None)),
46 #[cfg(feature = "compression")]
47 Compression::Lz4 => {
48 let compressed = lz4_flex::compress_prepend_size(data);
49 Ok((compressed, Compression::Lz4))
50 }
51 #[cfg(not(feature = "compression"))]
52 Compression::Lz4 => Ok((data.to_vec(), Compression::None)),
53 }
54}
55
56fn decompress_payload(data: &[u8], compression: Compression) -> Result<Vec<u8>> {
58 match compression {
59 Compression::None => Ok(data.to_vec()),
60 #[cfg(feature = "compression")]
61 Compression::ZstdDefault | Compression::ZstdMax => {
62 zstd::decode_all(std::io::Cursor::new(data)).map_err(|e| {
63 AprFormatError::Serialization(format!("Zstd decompression failed: {e}"))
64 })
65 }
66 #[cfg(not(feature = "compression"))]
67 Compression::ZstdDefault | Compression::ZstdMax => Err(AprFormatError::FormatError {
68 message: "Zstd compression not supported (enable `compression` feature)".to_string(),
69 }),
70 #[cfg(feature = "compression")]
71 Compression::Lz4 => lz4_flex::decompress_size_prepended(data)
72 .map_err(|e| AprFormatError::Serialization(format!("LZ4 decompression failed: {e}"))),
73 #[cfg(not(feature = "compression"))]
74 Compression::Lz4 => Err(AprFormatError::FormatError {
75 message: "LZ4 compression not supported (enable `compression` feature)".to_string(),
76 }),
77 }
78}
79
80#[allow(clippy::needless_pass_by_value)]
85pub fn save<M: Serialize>(
86 model: &M,
87 model_type: ModelType,
88 path: impl AsRef<Path>,
89 options: SaveOptions,
90) -> Result<()> {
91 let path = path.as_ref();
92
93 if options.quality_score == Some(0) {
95 return Err(AprFormatError::ValidationError {
96 message: "Jidoka: Refusing to save model with quality_score=0. \
97 Fix validation errors or use score=None to skip validation."
98 .to_string(),
99 });
100 }
101
102 let payload_uncompressed = bincode::serialize(model)
103 .map_err(|e| AprFormatError::Serialization(format!("Failed to serialize model: {e}")))?;
104
105 let (payload_compressed, compression) =
106 compress_payload(&payload_uncompressed, options.compression)?;
107
108 let metadata_bytes = rmp_serde::to_vec_named(&options.metadata)
109 .map_err(|e| AprFormatError::Serialization(format!("Failed to serialize metadata: {e}")))?;
110
111 let mut header = Header::new(model_type);
112 header.compression = compression;
113 header.metadata_size = metadata_bytes.len() as u32;
114 header.payload_size = payload_compressed.len() as u32;
115 header.uncompressed_size = payload_uncompressed.len() as u32;
116
117 if options.metadata.license.is_some() {
118 header.flags = header.flags.with_licensed();
119 }
120 header.quality_score = options.quality_score.unwrap_or(0);
121
122 let mut content = Vec::new();
123 content.extend_from_slice(&header.to_bytes());
124 content.extend_from_slice(&metadata_bytes);
125 content.extend_from_slice(&payload_compressed);
126
127 let checksum = crc32(&content);
128 content.extend_from_slice(&checksum.to_le_bytes());
129
130 let file = File::create(path)?;
131 let mut writer = BufWriter::new(file);
132 writer.write_all(&content)?;
133 writer.flush()?;
134 Ok(())
135}
136
137pub fn load<M: DeserializeOwned>(path: impl AsRef<Path>, expected_type: ModelType) -> Result<M> {
142 let file = File::open(path.as_ref())?;
143 let mut reader = BufReader::new(file);
144 let mut content = Vec::new();
145 reader.read_to_end(&mut content)?;
146 load_from_bytes(&content, expected_type)
147}
148
149pub fn load_from_bytes<M: DeserializeOwned>(data: &[u8], expected_type: ModelType) -> Result<M> {
157 if data.len() < HEADER_SIZE + 4 {
158 return Err(AprFormatError::FormatError {
159 message: format!("Data too small: {} bytes", data.len()),
160 });
161 }
162
163 let stored_checksum = u32::from_le_bytes([
165 data[data.len() - 4],
166 data[data.len() - 3],
167 data[data.len() - 2],
168 data[data.len() - 1],
169 ]);
170 let computed_checksum = crc32(&data[..data.len() - 4]);
171 if stored_checksum != computed_checksum {
172 return Err(AprFormatError::ChecksumMismatch {
173 expected: stored_checksum,
174 actual: computed_checksum,
175 });
176 }
177
178 let header = Header::from_bytes(&data[..HEADER_SIZE])?;
179 if header.model_type != expected_type {
180 return Err(AprFormatError::FormatError {
181 message: format!(
182 "Model type mismatch: data contains {:?}, expected {:?}",
183 header.model_type, expected_type
184 ),
185 });
186 }
187
188 let metadata_end = HEADER_SIZE + header.metadata_size as usize;
189 let payload_end = metadata_end + header.payload_size as usize;
190 if payload_end > data.len() - 4 {
191 return Err(AprFormatError::InvalidOffset);
192 }
193
194 let payload_compressed = &data[metadata_end..payload_end];
195 let payload_uncompressed = decompress_payload(payload_compressed, header.compression)?;
196
197 bincode::deserialize(&payload_uncompressed)
198 .map_err(|e| AprFormatError::Serialization(format!("Failed to deserialize model: {e}")))
199}
200
201fn model_info_from(header: &Header, metadata: Metadata) -> ModelInfo {
203 ModelInfo {
204 model_type: header.model_type,
205 format_version: header.version,
206 metadata,
207 payload_size: header.payload_size as usize,
208 uncompressed_size: header.uncompressed_size as usize,
209 encrypted: header.flags.is_encrypted(),
210 signed: header.flags.is_signed(),
211 streaming: header.flags.is_streaming(),
212 licensed: header.flags.is_licensed(),
213 trueno_native: header.flags.is_trueno_native(),
214 quantized: header.flags.is_quantized(),
215 has_model_card: header.flags.has_model_card(),
216 }
217}
218
219pub fn inspect_bytes(data: &[u8]) -> Result<ModelInfo> {
228 if data.len() < HEADER_SIZE {
229 return Err(AprFormatError::FormatError {
230 message: format!("Data too small: {} bytes", data.len()),
231 });
232 }
233 let header = Header::from_bytes(&data[..HEADER_SIZE])?;
234 let metadata_end = HEADER_SIZE + header.metadata_size as usize;
235 if metadata_end > data.len() {
236 return Err(AprFormatError::FormatError {
237 message: "Metadata extends beyond data boundary".to_string(),
238 });
239 }
240 let metadata_bytes = &data[HEADER_SIZE..metadata_end];
241 let metadata: Metadata = rmp_serde::from_slice(metadata_bytes)
242 .map_err(|e| AprFormatError::Serialization(format!("Failed to parse metadata: {e}")))?;
243 Ok(model_info_from(&header, metadata))
244}
245
246pub fn inspect(path: impl AsRef<Path>) -> Result<ModelInfo> {
251 let path = path.as_ref();
252 let file = File::open(path)?;
253 let mut reader = BufReader::new(file);
254
255 let mut header_bytes = [0u8; HEADER_SIZE];
256 reader.read_exact(&mut header_bytes)?;
257 let header = Header::from_bytes(&header_bytes)?;
258
259 let mut metadata_bytes = vec![0u8; header.metadata_size as usize];
260 reader.read_exact(&mut metadata_bytes)?;
261 let metadata: Metadata = rmp_serde::from_slice(&metadata_bytes)
262 .map_err(|e| AprFormatError::Serialization(format!("Failed to parse metadata: {e}")))?;
263
264 Ok(model_info_from(&header, metadata))
265}
266
267#[cfg(feature = "mmap")]
280pub fn load_mmap<M: DeserializeOwned>(
281 path: impl AsRef<Path>,
282 expected_type: ModelType,
283) -> Result<M> {
284 let file = File::open(path.as_ref())?;
285 let mmap = unsafe { memmap2::Mmap::map(&file)? };
288 load_from_bytes(&mmap, expected_type)
289}
290
291#[cfg(not(feature = "mmap"))]
300pub fn load_mmap<M: DeserializeOwned>(
301 path: impl AsRef<Path>,
302 expected_type: ModelType,
303) -> Result<M> {
304 load(path, expected_type)
305}
306
307pub fn load_auto<M: DeserializeOwned>(
316 path: impl AsRef<Path>,
317 expected_type: ModelType,
318) -> Result<M> {
319 let metadata = std::fs::metadata(path.as_ref())?;
320 if metadata.len() > MMAP_THRESHOLD {
321 load_mmap(path, expected_type)
322 } else {
323 load(path, expected_type)
324 }
325}
326
327#[cfg(test)]
328mod tests {
329 use super::*;
330 use serde::{Deserialize, Serialize};
331
332 #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
333 struct TestModel {
334 name: String,
335 values: Vec<f32>,
336 }
337
338 #[test]
339 fn test_save_load_roundtrip() {
340 let model = TestModel {
341 name: "test_model".to_string(),
342 values: vec![1.0, 2.0, 3.0, 4.0],
343 };
344 let dir = tempfile::tempdir().expect("create temp dir");
345 let path = dir.path().join("test.apr");
346 save(
347 &model,
348 ModelType::LinearRegression,
349 &path,
350 SaveOptions::default(),
351 )
352 .expect("save");
353 let loaded: TestModel = load(&path, ModelType::LinearRegression).expect("load");
354 assert_eq!(model, loaded);
355 }
356
357 #[test]
358 fn test_save_rejects_quality_score_zero() {
359 let model = TestModel {
360 name: "bad".to_string(),
361 values: vec![],
362 };
363 let dir = tempfile::tempdir().expect("create temp dir");
364 let path = dir.path().join("nope.apr");
365 let options = SaveOptions {
366 quality_score: Some(0),
367 ..Default::default()
368 };
369 assert!(save(&model, ModelType::LinearRegression, &path, options).is_err());
370 }
371
372 #[test]
373 fn test_load_wrong_model_type() {
374 let model = TestModel {
375 name: "t".to_string(),
376 values: vec![1.0],
377 };
378 let dir = tempfile::tempdir().expect("create temp dir");
379 let path = dir.path().join("t.apr");
380 save(
381 &model,
382 ModelType::LinearRegression,
383 &path,
384 SaveOptions::default(),
385 )
386 .expect("save");
387 let result: Result<TestModel> = load(&path, ModelType::KMeans);
388 assert!(result.is_err());
389 }
390
391 #[test]
392 fn test_load_from_bytes_corrupted_checksum() {
393 let model = TestModel {
394 name: "c".to_string(),
395 values: vec![1.0],
396 };
397 let dir = tempfile::tempdir().expect("create temp dir");
398 let path = dir.path().join("c.apr");
399 save(
400 &model,
401 ModelType::LinearRegression,
402 &path,
403 SaveOptions::default(),
404 )
405 .expect("save");
406 let mut data = std::fs::read(&path).expect("read");
407 data[HEADER_SIZE + 2] ^= 0xFF;
408 let result: Result<TestModel> = load_from_bytes(&data, ModelType::LinearRegression);
409 assert!(matches!(
410 result,
411 Err(AprFormatError::ChecksumMismatch { .. })
412 ));
413 }
414
415 #[test]
416 fn test_save_with_metadata_license_sets_flag() {
417 use crate::types::{LicenseInfo, LicenseTier};
418 let model = TestModel {
419 name: "l".to_string(),
420 values: vec![1.0],
421 };
422 let dir = tempfile::tempdir().expect("create temp dir");
423 let path = dir.path().join("l.apr");
424 let metadata = Metadata {
425 license: Some(LicenseInfo {
426 uuid: "u".to_string(),
427 hash: "h".to_string(),
428 expiry: None,
429 seats: None,
430 licensee: Some("X".to_string()),
431 tier: LicenseTier::Enterprise,
432 }),
433 ..Metadata::default()
434 };
435 let options = SaveOptions {
436 metadata,
437 compression: Compression::None,
438 quality_score: None,
439 };
440 save(&model, ModelType::LinearRegression, &path, options).expect("save");
441 let data = std::fs::read(&path).expect("read");
442 let header = Header::from_bytes(&data[..HEADER_SIZE]).expect("hdr");
443 assert!(header.flags.is_licensed());
444 }
445}