Skip to main content

apr_format/
core_io.rs

1//! v1 (`APRN`) core I/O — spike slice for issue #2231 Stage 1.
2//!
3//! Moved from `aprender-core/src/format/core_io.rs`, rewired to the sovereign
4//! [`crate::error::AprFormatError`] and the deduplicated [`crate::crc32::crc32`].
5//! Demonstrates the byte-identical save/load path with no dependency on
6//! `aprender-core`.
7
8use 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
16/// Threshold for switching to mmap loading (1MB).
17///
18/// Files larger than this use memory-mapped I/O (when the `mmap` feature is on);
19/// smaller files use standard read-to-heap (lower overhead for small data).
20pub const MMAP_THRESHOLD: u64 = 1024 * 1024;
21
22/// Compress payload based on algorithm (spec §3.3).
23///
24/// Without the `compression` feature, compressed variants fall back to `None`
25/// (mirrors the legacy `aprender-core` behavior when `format-compression` is off).
26#[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
56/// Decompress payload based on algorithm (spec §3.3).
57fn 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/// Save a model to `.apr` (v1 `APRN`) format.
81///
82/// # Errors
83/// Returns an error on I/O failure, serialization error, or a refused quality gate.
84#[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    // APR-POKA-001: Jidoka gate — refuse to write if validation explicitly failed.
94    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
137/// Load a model from a `.apr` (v1 `APRN`) file.
138///
139/// # Errors
140/// Returns an error on I/O failure, format error, checksum failure, or type mismatch.
141pub 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
149/// Load a model from a byte slice (spec §1.1 — single-binary deployment).
150///
151/// Enables the `include_bytes!()` pattern for embedding models directly in
152/// executables.
153///
154/// # Errors
155/// Returns an error on format error, type mismatch, or checksum failure.
156pub 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    // Verify checksum (Jidoka: stop the line on corruption).
164    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
201/// Build a [`ModelInfo`] from a parsed header + decoded metadata.
202fn 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
219/// Inspect model data without loading the payload (spec §1.1).
220///
221/// Useful for validating embedded models or checking metadata without
222/// deserializing the full model.
223///
224/// # Errors
225/// Returns an error on a too-small buffer, a bad header, or metadata that
226/// extends past the data boundary.
227pub 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
246/// Inspect a model file without loading the payload.
247///
248/// # Errors
249/// Returns an error on I/O failure or a format error.
250pub 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/// Load a model using memory-mapped I/O (zero-copy where possible).
268///
269/// Maps the file directly into the address space (via `memmap2`) and parses
270/// from the mapped slice, avoiding a read-to-heap copy. Falls back to standard
271/// [`load`] when the `mmap` feature is disabled, preserving the same API.
272///
273/// # Safety
274/// Uses OS-level memory mapping; the file must not be modified while loaded.
275///
276/// # Errors
277/// Returns an error on file-not-found, a format error, a type mismatch, or a
278/// checksum failure.
279#[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    // SAFETY: standard memmap2 usage; the caller upholds the no-concurrent-write
286    // contract documented above (same precondition as the pre-extraction core).
287    let mmap = unsafe { memmap2::Mmap::map(&file)? };
288    load_from_bytes(&mmap, expected_type)
289}
290
291/// Load a model using memory-mapped I/O — `mmap`-feature-disabled fallback.
292///
293/// Without the `mmap` feature this delegates to the standard heap-backed
294/// [`load`], keeping the public API identical.
295///
296/// # Errors
297/// Returns an error on file-not-found, a format error, a type mismatch, or a
298/// checksum failure.
299#[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
307/// Load a model with automatic strategy selection based on file size.
308///
309/// Files larger than [`MMAP_THRESHOLD`] use [`load_mmap`]; smaller files use
310/// [`load`] (lower overhead for small files).
311///
312/// # Errors
313/// Returns an error on file-not-found, a format error, a type mismatch, or a
314/// checksum failure.
315pub 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}