diff --git a/Cargo.lock b/Cargo.lock index 1de1627..2ad6567 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2494,6 +2494,15 @@ version = "1.9.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ba39f3699c378cd8970968dcbff9c43159ea4cfbd88d43c00b22f2ef10a435d2" +[[package]] +name = "relative-path" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bca40a312222d8ba74837cb474edef44b37f561da5f773981007a10bbaa992b0" +dependencies = [ + "serde", +] + [[package]] name = "reqwest" version = "0.13.4" @@ -2579,7 +2588,7 @@ dependencies = [ "proc-macro2", "quote", "regex", - "relative-path", + "relative-path 1.9.3", "rustc_version", "syn 2.0.119", "unicode-ident", @@ -2629,6 +2638,7 @@ dependencies = [ "rand_chacha 0.10.0", "rand_distr", "rayon", + "relative-path 2.0.1", "reqwest", "serde", "serde_json", diff --git a/Cargo.toml b/Cargo.toml index 8e8da1e..927d4e9 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -14,6 +14,7 @@ periodic_table = "0.5" proptest = "1" rand = "0.10" rayon = "1" +relative-path = { version = "2", features = ["serde"] } serde = { version = "1", features = ["derive"] } serde_json = "1.0.*" sysinfo = { version = "0.39.6", default-features = false, features = ["system"] } diff --git a/crates/rustiq-core/Cargo.toml b/crates/rustiq-core/Cargo.toml index 703bbac..3d156b3 100644 --- a/crates/rustiq-core/Cargo.toml +++ b/crates/rustiq-core/Cargo.toml @@ -27,6 +27,7 @@ npyz = "0.9.1" physical_constants = "0.5" periodic_table.workspace = true rayon.workspace = true +relative-path.workspace = true tokio = { workspace = true, features = ["fs", "io-util"], optional = true } serde_json.workspace = true serde.workspace = true diff --git a/crates/rustiq-core/src/persistence/data.rs b/crates/rustiq-core/src/persistence/data.rs new file mode 100644 index 0000000..86dccc8 --- /dev/null +++ b/crates/rustiq-core/src/persistence/data.rs @@ -0,0 +1,8 @@ +mod ao_eri_artifact; +mod artifact; +mod rustiq_data; + +pub use ao_eri_artifact::AoEriArtifact; +pub use artifact::Artifact; +pub use rustiq_data::RustiQData; +pub(crate) use rustiq_data::AO_ERI_ARTIFACT; diff --git a/crates/rustiq-core/src/persistence/data/ao_eri_artifact.rs b/crates/rustiq-core/src/persistence/data/ao_eri_artifact.rs new file mode 100644 index 0000000..3e93b56 --- /dev/null +++ b/crates/rustiq-core/src/persistence/data/ao_eri_artifact.rs @@ -0,0 +1,29 @@ +use crate::{eri::CompactEri, persistence::ArtifactError}; + +use super::{ + artifact::{private, Artifact}, + rustiq_data::{RustiQData, AO_ERI_ARTIFACT}, +}; + +/// AO electron-repulsion integrals in compact storage order. +pub struct AoEriArtifact; + +impl private::Sealed for AoEriArtifact {} + +impl Artifact for AoEriArtifact { + type Value = CompactEri; + + fn get(data: &mut RustiQData) -> Result, ArtifactError> { + if data.ao_eri.is_some() { + return Ok(data.ao_eri.as_ref()); + } + if !data.manifest.artifacts.contains_key(AO_ERI_ARTIFACT) { + return Ok(None); + } + data.read_eri().map(Some) + } + + fn set(data: &mut RustiQData, value: CompactEri) -> Result<(), ArtifactError> { + data.set_eri(value) + } +} diff --git a/crates/rustiq-core/src/persistence/data/artifact.rs b/crates/rustiq-core/src/persistence/data/artifact.rs new file mode 100644 index 0000000..dd1df18 --- /dev/null +++ b/crates/rustiq-core/src/persistence/data/artifact.rs @@ -0,0 +1,16 @@ +use super::rustiq_data::RustiQData; +use crate::persistence::ArtifactError; + +pub(crate) mod private { + pub trait Sealed {} +} + +/// A known scientific artifact. Only RustiQ's declared artifact markers implement this trait. +pub trait Artifact: private::Sealed { + type Value; + + #[doc(hidden)] + fn get(data: &mut RustiQData) -> Result, ArtifactError>; + #[doc(hidden)] + fn set(data: &mut RustiQData, value: Self::Value) -> Result<(), ArtifactError>; +} diff --git a/crates/rustiq-core/src/persistence/data/rustiq_data.rs b/crates/rustiq-core/src/persistence/data/rustiq_data.rs new file mode 100644 index 0000000..6ab7648 --- /dev/null +++ b/crates/rustiq-core/src/persistence/data/rustiq_data.rs @@ -0,0 +1,419 @@ +use std::collections::{BTreeMap, HashSet}; + +use relative_path::RelativePath; + +use crate::eri::CompactEri; + +use super::super::{ + read_compact_eri, validate_storage_path, write_compact_eri, AoEriAttributes, + ArtifactAttributes, ArtifactError, ArtifactManifest, Manifest, ManifestError, ManifestKind, + PersistenceReadError, PersistenceWriteError, Producer, ScientificIdentityManifest, Storage, + StorageError, AO_ERI_COMPUTATION_VERSION, AO_ERI_PATH, COMPACT_ERI_REPRESENTATION, FORMAT_NAME, + FORMAT_VERSION, MANIFEST_PATH, +}; +use super::artifact::Artifact; + +pub(crate) const AO_ERI_ARTIFACT: &str = "ao_eri"; +const MAX_MANIFEST_BYTES: u64 = 1024 * 1024; + +/// Known scientific artifacts and their manifest, with values loaded on demand. +#[derive(Debug)] +pub struct RustiQData { + pub(super) manifest: Manifest, + source: Option, + pub(super) ao_eri: Option, + basis_functions: Option, +} + +impl RustiQData { + /// Gets a known scientific artifact, loading and caching it on first access. + pub fn get(&mut self) -> Result, ArtifactError> { + A::get(self) + } + + /// Sets a known scientific artifact using its statically selected value type. + pub fn set(&mut self, value: A::Value) -> Result<(), ArtifactError> { + A::set(self, value) + } + + pub(crate) fn new_with_identity( + identity: super::super::ScientificIdentity, + basis_functions: usize, + kind: ManifestKind, + ) -> Self { + Self { + manifest: Manifest { + format: FORMAT_NAME.to_owned(), + format_version: FORMAT_VERSION, + kind, + producer: Producer { + name: "RustiQ".to_owned(), + version: env!("CARGO_PKG_VERSION").to_owned(), + }, + scientific_identity: ScientificIdentityManifest { + version: identity.version, + digest: identity.digest, + }, + artifacts: BTreeMap::new(), + }, + source: None, + ao_eri: None, + basis_functions: Some(basis_functions), + } + } + + pub(crate) fn read_from(mut source: Storage) -> Result { + let manifest: Manifest = + source.read_json(RelativePath::new(MANIFEST_PATH), MAX_MANIFEST_BYTES)?; + if manifest.format != FORMAT_NAME || manifest.format_version != FORMAT_VERSION { + return Err(ManifestError::UnsupportedFormat.into()); + } + + let mut paths: HashSet = HashSet::new(); + for artifact in manifest.artifacts.values() { + validate_artifact_path(&artifact.path)?; + let key = artifact.path.as_str().to_lowercase(); + if paths.iter().any(|other| paths_conflict(other, &key)) { + return Err( + ManifestError::ConflictingArtifactPath(artifact.path.to_string()).into(), + ); + } + paths.insert(key); + } + + let basis_functions = manifest + .artifacts + .get(AO_ERI_ARTIFACT) + .and_then(|artifact| match &artifact.attributes { + ArtifactAttributes::AoEri(attributes) => Some(attributes.basis_functions), + ArtifactAttributes::Unknown(_) => None, + }); + + Ok(Self { + manifest, + source: Some(source), + ao_eri: None, + basis_functions, + }) + } + + pub(crate) fn manifest(&self) -> &Manifest { + &self.manifest + } + + /// Replaces the AO ERI artifact after checking its compact length. + pub fn set_eri(&mut self, eri: CompactEri) -> Result<(), ArtifactError> { + let basis_functions = self.basis_functions.ok_or(ArtifactError::Missing)?; + validate_eri_len(&eri, basis_functions)?; + self.ao_eri = Some(eri); + Ok(()) + } + + /// Validates and decodes the AO ERI on first access, then reuses the object. + pub fn read_eri(&mut self) -> Result<&CompactEri, ArtifactError> { + if self.ao_eri.is_none() { + let artifact = self + .manifest + .artifacts + .get(AO_ERI_ARTIFACT) + .ok_or(ArtifactError::Missing)?; + let ArtifactAttributes::AoEri(attributes) = &artifact.attributes else { + return Err(ArtifactError::InvalidMetadata( + "AO ERI attributes are missing".into(), + )); + }; + if artifact.path.as_str() != AO_ERI_PATH + || artifact.representation != COMPACT_ERI_REPRESENTATION + || attributes.computation_version != AO_ERI_COMPUTATION_VERSION + { + return Err(ArtifactError::UnsupportedRepresentation( + artifact.representation.clone(), + )); + } + + let source = self.source.as_mut().ok_or(ArtifactError::Missing)?; + let metadata = source.artifact_metadata(&artifact.path)?; + if metadata.size != artifact.size || metadata.digest != artifact.digest { + return Err(ArtifactError::IntegrityMismatch(artifact.path.to_string())); + } + + self.ao_eri = Some( + source.with_artifact::<_, ArtifactError, _>(&artifact.path, |reader| { + read_compact_eri(reader, attributes.basis_functions).map_err(Into::into) + })?, + ); + } + + Ok(self + .ao_eri + .as_ref() + .expect("ERI was loaded or already present")) + } + + pub(crate) fn take_eri(&mut self) -> Option { + self.ao_eri.take() + } + + pub(crate) fn write_to(&mut self, destination: Storage) -> Result<(), PersistenceWriteError> { + let eri = self.ao_eri.take(); + let result = self.write_inner(destination, eri.as_ref()); + self.ao_eri = eri; + result + } + + pub(crate) fn write_with_eri( + &mut self, + destination: Storage, + eri: &CompactEri, + ) -> Result<(), PersistenceWriteError> { + self.write_inner(destination, Some(eri)) + } + + fn write_inner( + &mut self, + mut destination: Storage, + eri: Option<&CompactEri>, + ) -> Result<(), PersistenceWriteError> { + let mut manifest = self.manifest.clone(); + + for (name, artifact) in &self.manifest.artifacts { + if name == AO_ERI_ARTIFACT && eri.is_some() { + continue; + } + + let source = self.source.as_mut().ok_or_else(|| { + ArtifactError::InvalidMetadata(format!("artifact {name} has no source")) + })?; + validate_artifact_path(&artifact.path)?; + + let metadata = + source.with_artifact::<_, PersistenceWriteError, _>(&artifact.path, |input| { + destination.write_artifact::( + &artifact.path, + |output| { + std::io::copy(input, output) + .map_err(StorageError::from) + .map_err(PersistenceWriteError::from)?; + Ok(()) + }, + ) + })?; + + if metadata.size != artifact.size || metadata.digest != artifact.digest { + return Err(ArtifactError::IntegrityMismatch(artifact.path.to_string()).into()); + } + } + + if let Some(eri) = eri { + let basis_functions = self.basis_functions.ok_or(ArtifactError::Missing)?; + validate_eri_len(eri, basis_functions)?; + let path = RelativePath::new(AO_ERI_PATH); + let metadata = destination + .write_artifact::(path, |writer| { + write_compact_eri(writer, eri).map_err(Into::into) + })?; + + manifest.artifacts.insert( + AO_ERI_ARTIFACT.to_owned(), + ArtifactManifest { + path: AO_ERI_PATH.into(), + size: metadata.size, + representation: COMPACT_ERI_REPRESENTATION.to_owned(), + digest: metadata.digest, + attributes: ArtifactAttributes::AoEri(AoEriAttributes { + basis_functions, + computation_version: AO_ERI_COMPUTATION_VERSION, + }), + }, + ); + } + + destination.write_json(RelativePath::new(MANIFEST_PATH), &manifest)?; + destination.finish()?; + Ok(()) + } +} + +fn validate_eri_len(eri: &CompactEri, basis_functions: usize) -> Result<(), ArtifactError> { + let expected = CompactEri::checked_storage_len(basis_functions).ok_or_else(|| { + ArtifactError::InvalidMetadata("basis-function count overflows compact ERI storage".into()) + })?; + if eri.len() != expected { + return Err(ArtifactError::InvalidValueCount { + basis_functions, + expected, + actual: eri.len(), + }); + } + Ok(()) +} + +fn validate_artifact_path(path: &RelativePath) -> Result<(), ArtifactError> { + validate_storage_path(path).map_err(|_| ArtifactError::InvalidPath(path.to_string()))?; + if path.as_str().split('/').next() == Some(MANIFEST_PATH) { + return Err(ArtifactError::InvalidPath(path.to_string())); + } + Ok(()) +} + +fn paths_conflict(left: &str, right: &str) -> bool { + left == right + || left + .strip_prefix(right) + .is_some_and(|suffix| suffix.starts_with('/')) + || right + .strip_prefix(left) + .is_some_and(|suffix| suffix.starts_with('/')) +} + +#[cfg(test)] +mod tests { + use std::fs; + + use super::*; + use crate::persistence::{ + data::AoEriArtifact, sha256, Sha256Digest, SCIENTIFIC_IDENTITY_VERSION, + }; + + fn new_data() -> RustiQData { + RustiQData::new_with_identity( + crate::persistence::ScientificIdentity { + version: SCIENTIFIC_IDENTITY_VERSION, + digest: Sha256Digest::from([7; 32]), + }, + 2, + ManifestKind::Unknown("test-data".to_owned()), + ) + } + + #[test] + fn eri_is_loaded_only_on_request_and_then_cached_as_an_object() { + let root = tempfile::tempdir().unwrap(); + let entry = root.path().join("entry"); + fs::create_dir(&entry).unwrap(); + let mut data = new_data(); + assert!(data.get::().unwrap().is_none()); + data.set::(CompactEri::Zeroed(2)).unwrap(); + data.write_to(Storage::folder(&entry)).unwrap(); + + let mut restored = RustiQData::read_from(Storage::folder(&entry)).unwrap(); + assert!(restored.ao_eri.is_none()); + let first = restored.get::().unwrap().unwrap() as *const CompactEri; + fs::remove_file(entry.join(AO_ERI_PATH)).unwrap(); + let second = restored.get::().unwrap().unwrap() as *const CompactEri; + assert_eq!(first, second); + assert_eq!(restored.read_eri().unwrap().len(), 6); + } + + #[test] + fn unloaded_unknown_artifact_is_copied_without_decoding() { + let root = tempfile::tempdir().unwrap(); + let first = root.path().join("first"); + let second = root.path().join("second"); + fs::create_dir(&first).unwrap(); + fs::create_dir(&second).unwrap(); + + let mut data = new_data(); + data.set_eri(CompactEri::Zeroed(2)).unwrap(); + data.write_to(Storage::folder(&first)).unwrap(); + + let unknown = b"opaque future data"; + let unknown_path = first.join("arrays/post-hf/future.npy"); + fs::create_dir_all(unknown_path.parent().unwrap()).unwrap(); + fs::write(&unknown_path, unknown).unwrap(); + + let manifest_path = first.join(MANIFEST_PATH); + let mut manifest: Manifest = + serde_json::from_slice(&fs::read(&manifest_path).unwrap()).unwrap(); + manifest.artifacts.insert( + "future".into(), + ArtifactManifest { + path: "arrays/post-hf/future.npy".into(), + size: unknown.len() as u64, + representation: "rustiq-future-v1".into(), + digest: sha256(unknown), + attributes: ArtifactAttributes::Unknown(BTreeMap::new()), + }, + ); + fs::write(&manifest_path, serde_json::to_vec(&manifest).unwrap()).unwrap(); + + let mut restored = RustiQData::read_from(Storage::folder(&first)).unwrap(); + restored.write_to(Storage::folder(&second)).unwrap(); + assert_eq!( + fs::read(second.join("arrays/post-hf/future.npy")).unwrap(), + unknown + ); + assert_eq!( + RustiQData::read_from(Storage::folder(&second)) + .unwrap() + .manifest, + manifest + ); + } + + #[test] + fn artifact_paths_are_portable_and_relative() { + assert!(validate_artifact_path(RelativePath::new("arrays/integrals/ao-eri.npy")).is_ok()); + + for path in [ + "", + "/arrays/integrals/ao-eri.npy", + "arrays\\integrals\\ao-eri.npy", + "C:/arrays/ao-eri.npy", + "arrays/../ao-eri.npy", + "arrays/./ao-eri.npy", + "arrays//ao-eri.npy", + "arrays/ao-eri.npy/", + "manifest.json", + "manifest.json/child", + ] { + assert!( + validate_artifact_path(RelativePath::new(path)).is_err(), + "{path} must be rejected" + ); + } + } + + #[test] + fn rejects_corrupt_payload_and_conflicting_paths() { + let root = tempfile::tempdir().unwrap(); + let first = root.path().join("first"); + let second = root.path().join("second"); + fs::create_dir(&first).unwrap(); + fs::create_dir(&second).unwrap(); + + let mut data = new_data(); + data.set_eri(CompactEri::Zeroed(2)).unwrap(); + data.write_to(Storage::folder(&first)).unwrap(); + fs::write(first.join(AO_ERI_PATH), b"bad").unwrap(); + + let mut restored = RustiQData::read_from(Storage::folder(&first)).unwrap(); + assert!(matches!( + restored.write_to(Storage::folder(&second)), + Err(PersistenceWriteError::Artifact( + ArtifactError::IntegrityMismatch(_) + )) + )); + + let manifest_path = first.join(MANIFEST_PATH); + let mut manifest: Manifest = + serde_json::from_slice(&fs::read(&manifest_path).unwrap()).unwrap(); + manifest.artifacts.insert( + "future".into(), + ArtifactManifest { + path: "arrays/integrals".into(), + size: 0, + representation: "future-v1".into(), + digest: sha256(b""), + attributes: ArtifactAttributes::Unknown(BTreeMap::new()), + }, + ); + fs::write(&manifest_path, serde_json::to_vec(&manifest).unwrap()).unwrap(); + assert!(matches!( + RustiQData::read_from(Storage::folder(&first)), + Err(PersistenceReadError::Manifest( + ManifestError::ConflictingArtifactPath(_) + )) + )); + } +} diff --git a/crates/rustiq-core/src/persistence/eri_cache.rs b/crates/rustiq-core/src/persistence/eri_cache.rs index c60f70a..771018c 100644 --- a/crates/rustiq-core/src/persistence/eri_cache.rs +++ b/crates/rustiq-core/src/persistence/eri_cache.rs @@ -1,6 +1,6 @@ use std::{ fs::{self, File}, - io::{self, BufReader, BufWriter, Seek, Write}, + io::{self, Seek}, path::{Path, PathBuf}, }; @@ -12,16 +12,13 @@ use crate::{ }; use super::{ - ao_eri_identity, read_compact_eri, sha256_reader, validate_compact_eri_header, AoEriAttributes, - ArtifactAttributes, ArtifactManifest, Manifest, Producer, ScientificIdentity, - ScientificIdentityManifest, AO_ERI_COMPUTATION_VERSION, AO_ERI_PATH, - COMPACT_ERI_REPRESENTATION, FORMAT_NAME, FORMAT_VERSION, MANIFEST_PATH, - SCIENTIFIC_IDENTITY_VERSION, + ao_eri_identity, sha256_reader, validate_compact_eri_header, AoEriAttributes, + ArtifactAttributes, ArtifactManifest, Manifest, ManifestKind, RustiQData, ScientificIdentity, + Storage, AO_ERI_COMPUTATION_VERSION, AO_ERI_PATH, COMPACT_ERI_REPRESENTATION, FORMAT_NAME, + FORMAT_VERSION, SCIENTIFIC_IDENTITY_VERSION, }; -const CACHE_KIND: &str = "integral-cache"; -const AO_ERI_ARTIFACT: &str = "ao_eri"; -const MAX_MANIFEST_BYTES: u64 = 1024 * 1024; +use super::data::AO_ERI_ARTIFACT; /// A directory-backed AO ERI cache entry available for management. #[derive(Clone, Debug, PartialEq, Eq)] @@ -312,14 +309,15 @@ impl EriCache { if !metadata.is_dir() || metadata.file_type().is_symlink() { return None; } - let manifest = read_manifest(&entry)?; + let mut data = RustiQData::read_from(Storage::folder(&entry)).ok()?; + let manifest = data.manifest(); let artifact = manifest.artifacts.get(AO_ERI_ARTIFACT)?; let attributes = ao_eri_attributes(artifact)?; - if !manifest_is_valid(&manifest, identity) || attributes.basis_functions != basis_functions - { + if !manifest_is_valid(manifest, identity) || attributes.basis_functions != basis_functions { return None; } - read_validated_payload(&entry, artifact, basis_functions) + data.read_eri().ok()?; + data.take_eri() } fn store_identity( @@ -340,52 +338,11 @@ impl EriCache { } fs::create_dir_all(parent)?; let temporary = Builder::new().prefix(".rustiq-eri-").tempdir_in(parent)?; - let payload_path = temporary.path().join(AO_ERI_PATH); - fs::create_dir_all(payload_path.parent().expect("AO ERI path has a parent"))?; - { - let mut writer = BufWriter::new(File::create(&payload_path)?); - super::write_compact_eri(&mut writer, eri) - .map_err(|e| io::Error::new(io::ErrorKind::InvalidData, e))?; - writer.flush()?; - writer.into_inner()?.sync_all()?; - } - let payload_metadata = fs::metadata(&payload_path)?; - let payload_digest = sha256_reader(BufReader::new(File::open(&payload_path)?))?; - let manifest = Manifest { - format: FORMAT_NAME.to_owned(), - format_version: FORMAT_VERSION, - kind: CACHE_KIND.to_owned(), - producer: Producer { - name: "RustiQ".to_owned(), - version: env!("CARGO_PKG_VERSION").to_owned(), - }, - scientific_identity: ScientificIdentityManifest { - version: identity.version, - digest: identity.digest, - }, - artifacts: [( - AO_ERI_ARTIFACT.to_owned(), - ArtifactManifest { - path: AO_ERI_PATH.to_owned(), - size: payload_metadata.len(), - representation: COMPACT_ERI_REPRESENTATION.to_owned(), - digest: payload_digest, - attributes: ArtifactAttributes::AoEri(AoEriAttributes { - basis_functions, - computation_version: AO_ERI_COMPUTATION_VERSION, - }), - }, - )] - .into_iter() - .collect(), - }; - { - let mut writer = BufWriter::new(File::create(temporary.path().join(MANIFEST_PATH))?); - serde_json::to_writer_pretty(&mut writer, &manifest).map_err(io::Error::other)?; - writer.write_all(b"\n")?; - writer.flush()?; - writer.into_inner()?.sync_all()?; - } + let mut entry_data = + RustiQData::new_with_identity(identity, basis_functions, ManifestKind::IntegralCache); + entry_data + .write_with_eri(Storage::folder(temporary.path()), eri) + .map_err(io::Error::other)?; let temporary_path = temporary.keep(); if fs::symlink_metadata(&final_entry).is_ok() { if self.load_identity(identity, basis_functions).is_some() { @@ -412,28 +369,22 @@ impl EriCache { } fn read_manifest(entry: &Path) -> Option { - let manifest_path = entry.join(MANIFEST_PATH); - fs::symlink_metadata(&manifest_path) + RustiQData::read_from(Storage::folder(entry)) .ok() - .filter(|metadata| { - metadata.is_file() - && !metadata.file_type().is_symlink() - && metadata.len() <= MAX_MANIFEST_BYTES - })?; - serde_json::from_reader(BufReader::new(File::open(manifest_path).ok()?)).ok() + .map(|data| data.manifest().clone()) } fn manifest_is_valid(manifest: &Manifest, identity: ScientificIdentity) -> bool { manifest.format == FORMAT_NAME && manifest.format_version == FORMAT_VERSION - && manifest.kind == CACHE_KIND + && manifest.kind == ManifestKind::IntegralCache && manifest.scientific_identity.version == identity.version && manifest.scientific_identity.digest == identity.digest && manifest .artifacts .get(AO_ERI_ARTIFACT) .is_some_and(|artifact| { - artifact.path == AO_ERI_PATH + artifact.path.as_str() == AO_ERI_PATH && artifact.representation == COMPACT_ERI_REPRESENTATION && ao_eri_attributes(artifact).is_some_and(|attributes| { attributes.computation_version == AO_ERI_COMPUTATION_VERSION @@ -448,28 +399,6 @@ fn ao_eri_attributes(artifact: &ArtifactManifest) -> Option<&AoEriAttributes> { } } -fn read_validated_payload( - entry: &Path, - artifact: &ArtifactManifest, - basis_functions: usize, -) -> Option { - CompactEri::checked_storage_len(basis_functions)?; - let payload_path = entry.join(AO_ERI_PATH); - let metadata = fs::symlink_metadata(&payload_path).ok()?; - if !metadata.is_file() || metadata.file_type().is_symlink() { - return None; - } - let mut file = File::open(payload_path).ok()?; - if file.metadata().ok()?.len() != artifact.size { - return None; - } - if sha256_reader(&mut file).ok()? != artifact.digest { - return None; - } - file.rewind().ok()?; - read_compact_eri(BufReader::new(file), basis_functions).ok() -} - fn validate_payload(entry: &Path, artifact: &ArtifactManifest, basis_functions: usize) -> bool { if CompactEri::checked_storage_len(basis_functions).is_none() { return false; @@ -525,6 +454,7 @@ pub(super) fn is_fingerprint(value: &str) -> bool { #[cfg(test)] mod tests { use super::*; + use crate::persistence::MANIFEST_PATH; use crate::{config::DEFAULT_ERI_SCHWARZ_THRESHOLD, test_utils::load_sto3g_basis}; fn input() -> (Molecule, Basis) { @@ -562,6 +492,43 @@ mod tests { eri.ordered_values() ); } + #[test] + fn stored_entry_is_explicitly_an_integral_cache_with_ao_eri() { + let temporary = tempfile::tempdir().unwrap(); + let cache = EriCache::new(temporary.path()); + let (molecule, basis) = input(); + let eri = CompactEri::Zeroed(basis.nbasis()); + + cache.store(&molecule, &basis, threshold(), &eri).unwrap(); + + let identity = ao_eri_identity(molecule.geometry(), &basis, threshold()); + let manifest = read_manifest(&cache.entry_path(identity)).unwrap(); + assert_eq!(manifest.kind, ManifestKind::IntegralCache); + assert!(manifest.artifacts.contains_key(AO_ERI_ARTIFACT)); + } + + #[test] + fn unknown_kind_is_never_accepted_as_a_cache_hit() { + let temporary = tempfile::tempdir().unwrap(); + let cache = EriCache::new(temporary.path()); + let (molecule, basis) = input(); + let eri = CompactEri::Zeroed(basis.nbasis()); + cache.store(&molecule, &basis, threshold(), &eri).unwrap(); + + let identity = ao_eri_identity(molecule.geometry(), &basis, threshold()); + let manifest_path = cache.entry_path(identity).join(MANIFEST_PATH); + let mut manifest: serde_json::Value = + serde_json::from_reader(File::open(&manifest_path).unwrap()).unwrap(); + manifest["kind"] = + serde_json::to_value(ManifestKind::Unknown("future-state".to_owned())).unwrap(); + fs::write(&manifest_path, serde_json::to_vec(&manifest).unwrap()).unwrap(); + + assert!(cache.load(&molecule, &basis, threshold()).is_none()); + let entries = cache.entries().unwrap(); + assert_eq!(entries.len(), 1); + assert!(!entries[0].verified); + } + #[test] fn corruption_is_a_cache_miss() { let temporary = tempfile::tempdir().unwrap(); diff --git a/crates/rustiq-core/src/persistence/error.rs b/crates/rustiq-core/src/persistence/error.rs new file mode 100644 index 0000000..02e115c --- /dev/null +++ b/crates/rustiq-core/src/persistence/error.rs @@ -0,0 +1,111 @@ +use std::io; + +use thiserror::Error; + +#[derive(Debug, Error)] +pub(crate) enum StorageError { + #[error("storage I/O failed: {0}")] + Io(#[from] io::Error), + #[error("unsafe storage path: {0}")] + InvalidPath(String), + #[error("storage entry has an unexpected type: {0}")] + UnexpectedEntryType(String), +} + +#[derive(Debug, Error)] +pub enum NpyError { + #[error("could not read NPY data: {0}")] + Read(#[source] io::Error), + #[error("could not write NPY data: {0}")] + Write(#[source] io::Error), + #[error("AO ERI NPY must be one-dimensional, found shape {0:?}")] + InvalidEriShape(Box<[u64]>), + #[error("matrix NPY must be two-dimensional, found shape {0:?}")] + InvalidMatrixShape(Box<[u64]>), + #[error("NPY has shape {actual:?}, expected {expected:?}")] + InvalidShape { + expected: Box<[u64]>, + actual: Box<[u64]>, + }, + #[error("NPY dtype is not a supported f64 representation: {0}")] + InvalidDtype(String), + #[error( + "AO ERI NPY has {actual} values, expected {expected} for {basis_functions} basis functions" + )] + InvalidValueCount { + basis_functions: usize, + expected: usize, + actual: usize, + }, + #[error("NPY dimensions exceed supported limits")] + DimensionOverflow, +} + +#[derive(Debug, Error)] +pub(crate) enum ManifestError { + #[error("manifest storage failed: {0}")] + Storage(#[from] StorageError), + #[error("could not decode or encode persistence manifest: {0}")] + Json(#[from] serde_json::Error), + #[error("manifest is larger than the supported limit")] + TooLarge, + #[error("unsupported persistence format or version")] + UnsupportedFormat, + #[error("artifact path conflicts with another artifact: {0}")] + ConflictingArtifactPath(String), +} + +#[derive(Debug, Error)] +pub enum ArtifactError { + #[error("AO ERI artifact is missing")] + Missing, + #[error("unsupported artifact representation: {0}")] + UnsupportedRepresentation(String), + #[error("invalid artifact metadata: {0}")] + InvalidMetadata(String), + #[error("artifact integrity check failed: {0}")] + IntegrityMismatch(String), + #[error("invalid artifact path: {0}")] + InvalidPath(String), + #[error( + "AO ERI payload has {actual} values, expected {expected} for {basis_functions} basis functions" + )] + InvalidValueCount { + basis_functions: usize, + expected: usize, + actual: usize, + }, + #[error("artifact access failed: {0}")] + AccessFailed(String), + #[error("artifact NPY data is invalid: {0}")] + Npy(#[from] NpyError), +} + +#[derive(Debug, Error)] +pub(crate) enum PersistenceReadError { + #[error(transparent)] + Manifest(#[from] ManifestError), + #[error(transparent)] + Artifact(#[from] ArtifactError), +} + +#[derive(Debug, Error)] +pub(crate) enum PersistenceWriteError { + #[error(transparent)] + Storage(#[from] StorageError), + #[error(transparent)] + Manifest(#[from] ManifestError), + #[error(transparent)] + Artifact(#[from] ArtifactError), + #[error(transparent)] + Npy(#[from] NpyError), +} + +impl From for ArtifactError { + fn from(error: StorageError) -> Self { + match error { + StorageError::InvalidPath(path) => Self::InvalidPath(path), + other => Self::AccessFailed(other.to_string()), + } + } +} diff --git a/crates/rustiq-core/src/persistence/manifest.rs b/crates/rustiq-core/src/persistence/manifest.rs index 61a89d9..4a0ca36 100644 --- a/crates/rustiq-core/src/persistence/manifest.rs +++ b/crates/rustiq-core/src/persistence/manifest.rs @@ -1,28 +1,66 @@ use std::collections::BTreeMap; +use relative_path::RelativePathBuf; use serde::{Deserialize, Deserializer, Serialize}; use serde_json::Value; use super::{Sha256Digest, COMPACT_ERI_REPRESENTATION}; +#[derive(Clone, Debug, PartialEq, Eq)] +pub(crate) enum ManifestKind { + IntegralCache, + Unknown(String), +} + +impl ManifestKind { + pub(crate) fn as_str(&self) -> &str { + match self { + Self::IntegralCache => "integral-cache", + Self::Unknown(value) => value, + } + } +} + +impl Serialize for ManifestKind { + fn serialize(&self, serializer: S) -> Result + where + S: serde::Serializer, + { + serializer.serialize_str(self.as_str()) + } +} + +impl<'de> Deserialize<'de> for ManifestKind { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + let value = String::deserialize(deserializer)?; + Ok(match value.as_str() { + "integral-cache" => Self::IntegralCache, + _ => Self::Unknown(value), + }) + } +} + #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct Manifest { +pub(crate) struct Manifest { pub format: String, pub format_version: u32, - pub kind: String, + pub kind: ManifestKind, pub producer: Producer, pub scientific_identity: ScientificIdentityManifest, pub artifacts: BTreeMap, } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct Producer { +pub(crate) struct Producer { pub name: String, pub version: String, } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] -pub struct ScientificIdentityManifest { +pub(crate) struct ScientificIdentityManifest { pub version: u32, pub digest: Sha256Digest, } @@ -33,21 +71,21 @@ pub struct ScientificIdentityManifest { /// retain their attributes so newer manifests remain inspectable by older readers. #[derive(Clone, Debug, PartialEq, Eq, Serialize)] #[serde(untagged)] -pub enum ArtifactAttributes { +pub(crate) enum ArtifactAttributes { AoEri(AoEriAttributes), Unknown(BTreeMap), } #[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)] #[serde(deny_unknown_fields)] -pub struct AoEriAttributes { +pub(crate) struct AoEriAttributes { pub basis_functions: usize, pub computation_version: u32, } #[derive(Clone, Debug, PartialEq, Eq, Serialize)] -pub struct ArtifactManifest { - pub path: String, +pub(crate) struct ArtifactManifest { + pub path: RelativePathBuf, pub size: u64, pub representation: String, pub digest: Sha256Digest, @@ -56,7 +94,7 @@ pub struct ArtifactManifest { #[derive(Deserialize)] struct RawArtifactManifest { - path: String, + path: RelativePathBuf, size: u64, representation: String, digest: Sha256Digest, @@ -97,6 +135,8 @@ fn decode_attributes( #[cfg(test)] mod tests { + use proptest::prelude::*; + use super::*; use crate::persistence::{ AO_ERI_COMPUTATION_VERSION, AO_ERI_PATH, FORMAT_NAME, FORMAT_VERSION, @@ -109,7 +149,7 @@ mod tests { artifacts.insert( "ao_eri".to_string(), ArtifactManifest { - path: AO_ERI_PATH.to_string(), + path: RelativePathBuf::from(AO_ERI_PATH), size: 176, representation: COMPACT_ERI_REPRESENTATION.to_string(), digest: Sha256Digest::from([0x22; 32]), @@ -123,7 +163,7 @@ mod tests { let manifest = Manifest { format: FORMAT_NAME.to_string(), format_version: FORMAT_VERSION, - kind: "integral-cache".to_string(), + kind: ManifestKind::IntegralCache, producer: Producer { name: "RustiQ".to_string(), version: "0.1.0".to_string(), @@ -169,12 +209,25 @@ mod tests { assert_eq!(serde_json::from_str::(&json).unwrap(), manifest); } + proptest! { + #[test] + fn unknown_manifest_kind_round_trips(value in "[A-Za-z0-9_-]{1,64}") { + prop_assume!(value != "integral-cache"); + + let kind = ManifestKind::Unknown(value.clone()); + let json = serde_json::to_string(&kind).unwrap(); + let restored: ManifestKind = serde_json::from_str(&json).unwrap(); + + prop_assert_eq!(restored, ManifestKind::Unknown(value)); + } + } + #[test] fn unknown_artifact_attributes_are_preserved() { let json = r#"{ "format": "rustiq-persistence", "format_version": 1, - "kind": "checkpoint", + "kind": "future-state", "producer": {"name": "RustiQ", "version": "0.2.0"}, "scientific_identity": { "version": 1, diff --git a/crates/rustiq-core/src/persistence/mod.rs b/crates/rustiq-core/src/persistence/mod.rs index a65860c..1a5bad4 100644 --- a/crates/rustiq-core/src/persistence/mod.rs +++ b/crates/rustiq-core/src/persistence/mod.rs @@ -6,50 +6,37 @@ mod cache_names; mod checksum; +mod data; mod eri_cache; +mod error; mod identity; mod manifest; mod npy; +mod storage; +pub use crate::eri::CompactEri; pub use checksum::{sha256, sha256_reader, verify_sha256, Sha256Digest, Sha256DigestParseError}; +pub use data::{AoEriArtifact, Artifact, RustiQData}; pub use eri_cache::{EriCache, EriCacheEntry}; +pub use error::{ArtifactError, NpyError}; +pub(crate) use error::{ManifestError, PersistenceReadError, PersistenceWriteError, StorageError}; #[allow(unused_imports)] pub(crate) use identity::{ao_eri_identity, ScientificIdentity, AO_ERI_COMPUTATION_VERSION}; -pub use manifest::{ - AoEriAttributes, ArtifactAttributes, ArtifactManifest, Manifest, Producer, +pub(crate) use manifest::{ + AoEriAttributes, ArtifactAttributes, ArtifactManifest, Manifest, ManifestKind, Producer, ScientificIdentityManifest, }; +pub(crate) use storage::Storage; #[allow(unused_imports)] pub(crate) use npy::{ read_compact_eri, read_dmatrix, validate_compact_eri_header, write_compact_eri, }; +pub(crate) use storage::validate_path as validate_storage_path; -pub const FORMAT_NAME: &str = "rustiq-persistence"; -pub const FORMAT_VERSION: u32 = 1; -pub const MANIFEST_PATH: &str = "manifest.json"; -pub const AO_ERI_PATH: &str = "arrays/integrals/ao-eri.npy"; -pub const COMPACT_ERI_REPRESENTATION: &str = "rustiq-compact-eri-v1"; -pub const SCIENTIFIC_IDENTITY_VERSION: u32 = 1; - -use thiserror::Error; - -#[derive(Debug, Error)] -pub enum PersistenceError { - #[error("could not read NPY data: {0}")] - NpyRead(#[source] std::io::Error), - #[error("could not write NPY data: {0}")] - NpyWrite(#[source] std::io::Error), - #[error("AO ERI NPY must be one-dimensional, found shape {0:?}")] - InvalidEriShape(Vec), - #[error("matrix NPY must be two-dimensional, found shape {0:?}")] - InvalidMatrixShape(Vec), - #[error("AO ERI payload has {actual} values, expected {expected} for {basis_functions} basis functions")] - InvalidValueCount { - basis_functions: usize, - expected: usize, - actual: usize, - }, - #[error("NPY dtype is not a supported f64 representation: {0}")] - InvalidDtype(String), -} +pub(crate) const FORMAT_NAME: &str = "rustiq-persistence"; +pub(crate) const FORMAT_VERSION: u32 = 1; +pub(crate) const MANIFEST_PATH: &str = "manifest.json"; +pub(crate) const AO_ERI_PATH: &str = "arrays/integrals/ao-eri.npy"; +pub(crate) const COMPACT_ERI_REPRESENTATION: &str = "rustiq-compact-eri-v1"; +pub(crate) const SCIENTIFIC_IDENTITY_VERSION: u32 = 1; diff --git a/crates/rustiq-core/src/persistence/npy.rs b/crates/rustiq-core/src/persistence/npy.rs index cf6cb5d..7001418 100644 --- a/crates/rustiq-core/src/persistence/npy.rs +++ b/crates/rustiq-core/src/persistence/npy.rs @@ -5,43 +5,135 @@ use npyz::{DType, NpyFile, TypeChar, WriterBuilder}; use crate::eri::CompactEri; -use super::PersistenceError; - -pub(crate) fn write_compact_eri( - writer: impl Write, - eri: &CompactEri, -) -> Result<(), PersistenceError> { - let shape = [eri.len() as u64]; - let mut writer = npyz::WriteOptions::new() - .default_dtype() - .shape(&shape) - .writer(writer) - .begin_nd() - .map_err(PersistenceError::NpyWrite)?; - writer - .extend(eri.ordered_values().iter().copied()) - .map_err(PersistenceError::NpyWrite)?; - writer.finish().map_err(PersistenceError::NpyWrite) +use super::NpyError; + +/// Conversion between scientific values and an NPY byte stream. +trait NpyConvert: Sized { + type Shape: Copy; + type NpyShape: AsRef<[u64]>; + + fn npy_shape(shape: Self::Shape) -> Result; + + fn decode_with_shape(npy: NpyFile, shape: Self::Shape) -> Result; + + /// Checks a parsed header before decoding any array values. + fn try_from_npy_with_shape( + npy: NpyFile, + shape: Self::Shape, + ) -> Result { + let expected = Self::npy_shape(shape)?; + if npy.shape() != expected.as_ref() { + return Err(NpyError::InvalidShape { + expected: expected.as_ref().to_vec().into_boxed_slice(), + actual: npy.shape().to_vec().into_boxed_slice(), + }); + } + Self::decode_with_shape(npy, shape) + } + + /// Parses a stream once, checks its shape, then decodes its values. + fn try_read_with_shape(reader: impl Read, shape: Self::Shape) -> Result { + Self::try_from_npy_with_shape(NpyFile::new(reader).map_err(NpyError::Read)?, shape) + } + + fn write_npy(&self, writer: impl Write) -> Result<(), NpyError>; +} + +impl NpyConvert for CompactEri { + type Shape = usize; + type NpyShape = [u64; 1]; + + fn npy_shape(basis_functions: usize) -> Result { + let length = + CompactEri::checked_storage_len(basis_functions).ok_or(NpyError::DimensionOverflow)?; + Ok([u64::try_from(length).map_err(|_| NpyError::DimensionOverflow)?]) + } + + fn decode_with_shape( + npy: NpyFile, + basis_functions: usize, + ) -> Result { + decode_compact_eri(npy, basis_functions) + } + + fn write_npy(&self, writer: impl Write) -> Result<(), NpyError> { + let shape = [self.len() as u64]; + let mut writer = npyz::WriteOptions::new() + .default_dtype() + .shape(&shape) + .writer(writer) + .begin_nd() + .map_err(NpyError::Write)?; + writer + .extend(self.ordered_values().iter().copied()) + .map_err(NpyError::Write)?; + writer.finish().map_err(NpyError::Write) + } +} + +impl NpyConvert for DMatrix { + type Shape = (usize, usize); + type NpyShape = [u64; 2]; + + fn npy_shape((rows, columns): (usize, usize)) -> Result { + let rows = u64::try_from(rows).map_err(|_| NpyError::DimensionOverflow)?; + let columns = u64::try_from(columns).map_err(|_| NpyError::DimensionOverflow)?; + Ok([rows, columns]) + } + + fn decode_with_shape( + npy: NpyFile, + _shape: (usize, usize), + ) -> Result { + decode_dmatrix(npy) + } + + fn write_npy(&self, writer: impl Write) -> Result<(), NpyError> { + let shape = [self.nrows() as u64, self.ncols() as u64]; + let mut writer = npyz::WriteOptions::new() + .default_dtype() + .order(npyz::Order::Fortran) + .shape(&shape) + .writer(writer) + .begin_nd() + .map_err(NpyError::Write)?; + writer + .extend(self.as_slice().iter().copied()) + .map_err(NpyError::Write)?; + writer.finish().map_err(NpyError::Write) + } +} + +pub(crate) fn write_compact_eri(writer: impl Write, eri: &CompactEri) -> Result<(), NpyError> { + eri.write_npy(writer) } pub(crate) fn read_compact_eri( reader: impl Read, basis_functions: usize, -) -> Result { - let expected = CompactEri::checked_storage_len(basis_functions).ok_or( - PersistenceError::InvalidValueCount { - basis_functions, - expected: 0, - actual: 0, - }, - )?; - let npy = NpyFile::new(reader).map_err(PersistenceError::NpyRead)?; +) -> Result { + CompactEri::try_read_with_shape(reader, basis_functions) +} + +fn decode_compact_eri( + npy: NpyFile, + basis_functions: usize, +) -> Result { if npy.shape().len() != 1 { - return Err(PersistenceError::InvalidEriShape(npy.shape().to_vec())); + return Err(NpyError::InvalidEriShape( + npy.shape().to_vec().into_boxed_slice(), + )); } - let actual = usize::try_from(npy.shape()[0]).unwrap_or(usize::MAX); + let actual = usize::try_from(npy.shape()[0]) + .map_err(|_| NpyError::InvalidEriShape(npy.shape().to_vec().into_boxed_slice()))?; + let expected = + CompactEri::checked_storage_len(basis_functions).ok_or(NpyError::InvalidValueCount { + basis_functions, + expected: 0, + actual, + })?; if actual != expected { - return Err(PersistenceError::InvalidValueCount { + return Err(NpyError::InvalidValueCount { basis_functions, expected, actual, @@ -49,14 +141,14 @@ pub(crate) fn read_compact_eri( } let values = npy.into_vec::().map_err(|error| { if error.kind() == std::io::ErrorKind::InvalidData { - PersistenceError::InvalidDtype(error.to_string()) + NpyError::InvalidDtype(error.to_string()) } else { - PersistenceError::NpyRead(error) + NpyError::Read(error) } })?; CompactEri::from_ordered_values(basis_functions, values).map_err(|error| match error { crate::eri::CompactEriBuildError::InvalidLength { expected, actual } => { - PersistenceError::InvalidValueCount { + NpyError::InvalidValueCount { basis_functions, expected, actual, @@ -104,23 +196,28 @@ pub(crate) fn validate_compact_eri_header( .is_some_and(|end| end == file_size) } -pub(crate) fn read_dmatrix(reader: impl Read) -> Result, PersistenceError> { - let npy = NpyFile::new(reader).map_err(PersistenceError::NpyRead)?; +pub(crate) fn read_dmatrix(reader: impl Read) -> Result, NpyError> { + decode_dmatrix(NpyFile::new(reader).map_err(NpyError::Read)?) +} + +fn decode_dmatrix(npy: NpyFile) -> Result, NpyError> { if npy.shape().len() != 2 { - return Err(PersistenceError::InvalidMatrixShape(npy.shape().to_vec())); + return Err(NpyError::InvalidMatrixShape( + npy.shape().to_vec().into_boxed_slice(), + )); } let rows = usize::try_from(npy.shape()[0]) - .map_err(|_| PersistenceError::InvalidMatrixShape(npy.shape().to_vec()))?; + .map_err(|_| NpyError::InvalidMatrixShape(npy.shape().to_vec().into_boxed_slice()))?; let columns = usize::try_from(npy.shape()[1]) - .map_err(|_| PersistenceError::InvalidMatrixShape(npy.shape().to_vec()))?; + .map_err(|_| NpyError::InvalidMatrixShape(npy.shape().to_vec().into_boxed_slice()))?; rows.checked_mul(columns) - .ok_or_else(|| PersistenceError::InvalidMatrixShape(npy.shape().to_vec()))?; + .ok_or_else(|| NpyError::InvalidMatrixShape(npy.shape().to_vec().into_boxed_slice()))?; let order = npy.order(); let values = npy.into_vec::().map_err(|error| { if error.kind() == std::io::ErrorKind::InvalidData { - PersistenceError::InvalidDtype(error.to_string()) + NpyError::InvalidDtype(error.to_string()) } else { - PersistenceError::NpyRead(error) + NpyError::Read(error) } })?; Ok(matrix_from_npy_values(rows, columns, order, values)) @@ -155,6 +252,16 @@ mod tests { write_compact_eri(&mut bytes, &source).unwrap(); let restored = read_compact_eri(bytes.as_slice(), basis_functions).unwrap(); assert_eq!(restored.ordered_values(), source.ordered_values()); + assert_eq!( + CompactEri::try_read_with_shape(bytes.as_slice(), basis_functions) + .unwrap() + .ordered_values(), + source.ordered_values() + ); + assert!(matches!( + CompactEri::try_read_with_shape(bytes.as_slice(), basis_functions - 1), + Err(NpyError::InvalidShape { .. }) + )); } #[test] @@ -183,6 +290,80 @@ mod tests { assert_eq!(matrix[(1, 2)], 6.0); } + #[test] + fn dmatrix_trait_writes_fortran_order_and_reads_both_orders() { + let matrix = DMatrix::from_row_slice(2, 3, &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]); + let mut bytes = Vec::new(); + matrix.write_npy(&mut bytes).unwrap(); + let npy = NpyFile::new(bytes.as_slice()).unwrap(); + assert_eq!(npy.order(), npyz::Order::Fortran); + assert_eq!(npy.shape(), &[2, 3]); + assert_eq!(npy.into_vec::().unwrap(), matrix.as_slice()); + assert_eq!(read_dmatrix(bytes.as_slice()).unwrap(), matrix); + assert_eq!( + DMatrix::::try_from_npy_with_shape( + NpyFile::new(bytes.as_slice()).unwrap(), + (2, 3) + ) + .unwrap(), + matrix + ); + assert_eq!( + DMatrix::::try_read_with_shape(bytes.as_slice(), (2, 3)).unwrap(), + matrix + ); + assert!(matches!( + DMatrix::::try_read_with_shape(bytes.as_slice(), (3, 2)), + Err(NpyError::InvalidShape { .. }) + )); + let truncated = &bytes[..bytes.len() - 8]; + assert!(matches!( + DMatrix::::try_read_with_shape(truncated, (3, 2)), + Err(NpyError::InvalidShape { .. }) + )); + assert!(matches!( + DMatrix::::try_read_with_shape(truncated, (2, 3)), + Err(NpyError::Read(_)) + )); + + let mut c_bytes = Vec::new(); + let mut writer = npyz::WriteOptions::new() + .default_dtype() + .shape(&[2, 3]) + .writer(&mut c_bytes) + .begin_nd() + .unwrap(); + writer.extend([1.0_f64, 2.0, 3.0, 4.0, 5.0, 6.0]).unwrap(); + writer.finish().unwrap(); + assert_eq!(read_dmatrix(c_bytes.as_slice()).unwrap(), matrix); + } + + #[test] + fn reads_python_numpy_matrix_fixtures_in_c_and_fortran_order() { + for hex in [ + include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/tests/data/persistence/dmatrix-python-c-v1.npy.hex" + )), + include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/tests/data/persistence/dmatrix-python-fortran-v1.npy.hex" + )), + ] { + let (pairs, remainder) = hex.trim().as_bytes().as_chunks::<2>(); + assert!(remainder.is_empty()); + let bytes: Vec = pairs + .iter() + .map(|pair| u8::from_str_radix(std::str::from_utf8(pair).unwrap(), 16).unwrap()) + .collect(); + + assert_eq!( + read_dmatrix(bytes.as_slice()).unwrap(), + DMatrix::from_row_slice(2, 3, &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0]) + ); + } + } + #[test] fn rejects_wrong_value_count() { let mut bytes = Vec::new(); @@ -197,7 +378,7 @@ mod tests { writer.finish().unwrap(); assert!(matches!( read_compact_eri(bytes.as_slice(), 2), - Err(PersistenceError::InvalidValueCount { .. }) + Err(NpyError::InvalidShape { .. }) )); } @@ -215,7 +396,7 @@ mod tests { writer.finish().unwrap(); assert!(matches!( read_compact_eri(bytes.as_slice(), 1), - Err(PersistenceError::InvalidDtype(_)) + Err(NpyError::InvalidDtype(_)) )); } @@ -227,7 +408,7 @@ mod tests { bytes.truncate(bytes.len() - 1); assert!(matches!( read_compact_eri(bytes.as_slice(), 3), - Err(PersistenceError::NpyRead(_)) + Err(NpyError::Read(_)) )); } diff --git a/crates/rustiq-core/src/persistence/storage.rs b/crates/rustiq-core/src/persistence/storage.rs new file mode 100644 index 0000000..7633beb --- /dev/null +++ b/crates/rustiq-core/src/persistence/storage.rs @@ -0,0 +1,342 @@ +use std::{ + fs::{self, File, OpenOptions}, + io::{self, BufReader, BufWriter, Read, Write}, + path::PathBuf, +}; + +use relative_path::RelativePath; +use serde::{de::DeserializeOwned, Serialize}; +use sha2::{Digest, Sha256}; + +use super::{sha256_reader, ManifestError, Sha256Digest, StorageError}; + +#[derive(Debug)] +pub(crate) enum Storage { + Folder(PathBuf), +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub(crate) struct ArtifactMetadata { + pub(crate) size: u64, + pub(crate) digest: Sha256Digest, +} + +struct DigestWriter { + inner: W, + hasher: Sha256, + size: u64, +} + +impl DigestWriter { + fn new(inner: W) -> Self { + Self { + inner, + hasher: Sha256::new(), + size: 0, + } + } + + fn finish(self) -> (W, ArtifactMetadata) { + ( + self.inner, + ArtifactMetadata { + size: self.size, + digest: self.hasher.finalize().into(), + }, + ) + } +} + +impl Write for DigestWriter { + fn write(&mut self, buffer: &[u8]) -> io::Result { + let written = self.inner.write(buffer)?; + self.hasher.update(&buffer[..written]); + let written_u64 = + u64::try_from(written).map_err(|_| io::Error::other("artifact size exceeds u64"))?; + self.size = self + .size + .checked_add(written_u64) + .ok_or_else(|| io::Error::other("artifact size exceeds u64"))?; + Ok(written) + } + + fn flush(&mut self) -> io::Result<()> { + self.inner.flush() + } +} + +impl Storage { + pub(crate) fn folder(root: impl Into) -> Self { + Self::Folder(root.into()) + } + + pub(crate) fn artifact_metadata( + &mut self, + path: &RelativePath, + ) -> Result { + match self { + Self::Folder(root) => FolderStorage { root: root.clone() }.artifact_metadata(path), + } + } + + pub(crate) fn with_artifact(&mut self, path: &RelativePath, read: F) -> Result + where + E: From, + F: FnOnce(&mut dyn Read) -> Result, + { + match self { + Self::Folder(root) => FolderStorage { root: root.clone() }.with_artifact(path, read), + } + } + + pub(crate) fn write_artifact( + &mut self, + path: &RelativePath, + write: F, + ) -> Result + where + E: From, + F: FnOnce(&mut dyn Write) -> Result<(), E>, + { + match self { + Self::Folder(root) => FolderStorage { root: root.clone() }.write_artifact(path, write), + } + } + + pub(crate) fn read_json( + &mut self, + path: &RelativePath, + max_size: u64, + ) -> Result { + let size = match self { + Self::Folder(root) => FolderStorage { root: root.clone() }.artifact_size(path)?, + }; + if size > max_size { + return Err(ManifestError::TooLarge); + } + + self.with_artifact(path, |reader| { + serde_json::from_reader(reader).map_err(ManifestError::from) + }) + } + + pub(crate) fn write_json( + &mut self, + path: &RelativePath, + value: &T, + ) -> Result<(), ManifestError> { + self.write_artifact::(path, |writer| { + serde_json::to_writer_pretty(&mut *writer, value)?; + writer.write_all(b"\n").map_err(StorageError::from)?; + Ok(()) + })?; + Ok(()) + } + + pub(crate) fn finish(self) -> Result<(), StorageError> { + match self { + Self::Folder(_) => Ok(()), + } + } +} + +#[derive(Debug)] +struct FolderStorage { + root: PathBuf, +} + +impl FolderStorage { + fn artifact_size(&self, path: &RelativePath) -> Result { + Ok(self.open_artifact(path)?.metadata()?.len()) + } + + fn artifact_metadata(&self, path: &RelativePath) -> Result { + let file = self.open_artifact(path)?; + let size = file.metadata()?.len(); + let digest = sha256_reader(BufReader::new(file))?; + Ok(ArtifactMetadata { size, digest }) + } + + fn with_artifact(&self, path: &RelativePath, read: F) -> Result + where + E: From, + F: FnOnce(&mut dyn Read) -> Result, + { + let mut reader = BufReader::new(self.open_artifact(path).map_err(E::from)?); + read(&mut reader) + } + + fn write_artifact(&mut self, path: &RelativePath, write: F) -> Result + where + E: From, + F: FnOnce(&mut dyn Write) -> Result<(), E>, + { + let file = self.create_artifact(path).map_err(E::from)?; + let writer = BufWriter::new(file); + let mut writer = DigestWriter::new(writer); + write(&mut writer)?; + writer + .flush() + .map_err(StorageError::from) + .map_err(E::from)?; + let (writer, metadata) = writer.finish(); + let file = writer + .into_inner() + .map_err(|error| StorageError::Io(error.into_error())) + .map_err(E::from)?; + file.sync_all() + .map_err(StorageError::from) + .map_err(E::from)?; + Ok(metadata) + } + + fn open_artifact(&self, path: &RelativePath) -> Result { + validate_path(path)?; + let components: Vec<_> = path.as_str().split('/').collect(); + let mut native = self.root.clone(); + + for (index, component) in components.iter().enumerate() { + native.push(component); + let metadata = fs::symlink_metadata(&native)?; + if metadata.file_type().is_symlink() { + return Err(StorageError::UnexpectedEntryType(path.to_string())); + } + + let is_last = index + 1 == components.len(); + if (!is_last && !metadata.is_dir()) || (is_last && !metadata.is_file()) { + return Err(StorageError::UnexpectedEntryType(path.to_string())); + } + } + + Ok(File::open(native)?) + } + + fn create_artifact(&self, path: &RelativePath) -> Result { + validate_path(path)?; + let root_metadata = fs::symlink_metadata(&self.root)?; + if !root_metadata.is_dir() || root_metadata.file_type().is_symlink() { + return Err(StorageError::UnexpectedEntryType( + self.root.display().to_string(), + )); + } + + let mut components = path.as_str().split('/').peekable(); + let mut parent = self.root.clone(); + while let Some(component) = components.next() { + if components.peek().is_none() { + parent.push(component); + return Ok(OpenOptions::new() + .write(true) + .create_new(true) + .open(parent)?); + } + + parent.push(component); + match fs::symlink_metadata(&parent) { + Ok(metadata) => { + if !metadata.is_dir() || metadata.file_type().is_symlink() { + return Err(StorageError::UnexpectedEntryType(path.to_string())); + } + } + Err(error) if error.kind() == io::ErrorKind::NotFound => { + fs::create_dir(&parent)?; + } + Err(error) => return Err(error.into()), + } + } + + Err(StorageError::InvalidPath(path.to_string())) + } +} + +pub(crate) fn validate_path(path: &RelativePath) -> Result<(), StorageError> { + let value = path.as_str(); + if value.is_empty() || value.starts_with('/') || value.contains('\\') { + return Err(StorageError::InvalidPath(path.to_string())); + } + + for component in value.split('/') { + if component.is_empty() + || matches!(component, "." | "..") + || component.ends_with(['.', ' ']) + || component.chars().any(|character| { + character < '\u{20}' || matches!(character, '<' | '>' | ':' | '"' | '|' | '?' | '*') + }) + || is_windows_reserved_name(component) + { + return Err(StorageError::InvalidPath(path.to_string())); + } + } + + Ok(()) +} + +fn is_windows_reserved_name(component: &str) -> bool { + let stem = component + .split('.') + .next() + .unwrap_or(component) + .to_ascii_uppercase(); + + matches!( + stem.as_str(), + "CON" | "PRN" | "AUX" | "NUL" | "CONIN$" | "CONOUT$" + ) || (stem.len() == 4 + && matches!(&stem[..3], "COM" | "LPT") + && matches!(stem.as_bytes()[3], b'1'..=b'9')) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn folder_storage_streams_artifacts_and_json() { + let root = tempfile::tempdir().unwrap(); + let mut storage = Storage::folder(root.path()); + let artifact_path = RelativePath::new("arrays/test.bin"); + + let metadata = storage + .write_artifact::(artifact_path, |writer| { + writer.write_all(b"payload")?; + Ok(()) + }) + .unwrap(); + assert_eq!(metadata.size, 7); + assert_eq!(metadata.digest, super::super::sha256(b"payload")); + + let payload = storage + .with_artifact::<_, StorageError, _>(artifact_path, |reader| { + let mut bytes = Vec::new(); + reader.read_to_end(&mut bytes)?; + Ok(bytes) + }) + .unwrap(); + assert_eq!(payload, b"payload"); + + let manifest_path = RelativePath::new("manifest.json"); + storage + .write_json(manifest_path, &serde_json::json!({"version": 1})) + .unwrap(); + let manifest: serde_json::Value = storage.read_json(manifest_path, 1024).unwrap(); + assert_eq!(manifest["version"], 1); + storage.finish().unwrap(); + } + + #[test] + fn rejects_non_portable_windows_paths() { + for path in [ + "CON", + "con.npy", + "arrays/NUL.bin", + "arrays/COM1.npy", + "arrays/LPT9.npy", + "arrays/trailing.", + "arrays/trailing ", + "arrays/bad?.npy", + "arrays/bad|name.npy", + ] { + assert!(validate_path(RelativePath::new(path)).is_err(), "{path}"); + } + } +} diff --git a/crates/rustiq-core/tests/data/persistence/dmatrix-python-c-v1.npy.hex b/crates/rustiq-core/tests/data/persistence/dmatrix-python-c-v1.npy.hex new file mode 100644 index 0000000..3491ea1 --- /dev/null +++ b/crates/rustiq-core/tests/data/persistence/dmatrix-python-c-v1.npy.hex @@ -0,0 +1 @@ +934e554d5059010076007b276465736372273a20273c6638272c2027666f727472616e5f6f72646572273a2046616c73652c20277368617065273a2028322c2033292c207d202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020200a000000000000f03f00000000000000400000000000000840000000000000104000000000000014400000000000001840 diff --git a/crates/rustiq-core/tests/data/persistence/dmatrix-python-fortran-v1.npy.hex b/crates/rustiq-core/tests/data/persistence/dmatrix-python-fortran-v1.npy.hex new file mode 100644 index 0000000..34e329d --- /dev/null +++ b/crates/rustiq-core/tests/data/persistence/dmatrix-python-fortran-v1.npy.hex @@ -0,0 +1 @@ +934e554d5059010076007b276465736372273a20273c6638272c2027666f727472616e5f6f72646572273a20547275652c20277368617065273a2028322c2033292c207d20202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020202020200a000000000000f03f00000000000010400000000000000040000000000000144000000000000008400000000000001840 diff --git a/docs/persistence-format-v1.md b/docs/persistence-format-v1.md index 444093d..f7e2f65 100644 --- a/docs/persistence-format-v1.md +++ b/docs/persistence-format-v1.md @@ -87,10 +87,40 @@ untouched. Names and fingerprints are validated rather than interpreted as paths ## Logical entries +The core Rust API exposes `persistence::RustiQData` as a typed scientific +facade. Physical storage selection is an internal persistence concern: the +directory-backed ERI cache resolves folder storage internally, and the future +portable `.rustiq` API will resolve ZIP/ZIP64 internally rather than exposing a +generic storage backend. Internal readers load the bounded manifest first and +open scientific artifacts on demand. `read_eri` validates and decodes the AO ERI +NPY on first access and keeps the resulting `CompactEri` for subsequent accesses. +Unknown representations can be copied internally as verified byte streams without +being decoded. + +The public scientific API is typed: `set_eri` and `read_eri` operate on +`CompactEri`. The generic `get::()` returns +`Result, ArtifactError>`, while +`set::(value)` accepts only `CompactEri`; only declared artifact +marker types are accepted. Artifact access reports `ArtifactError`, while NPY +format failures are represented by `NpyError`; storage, manifest, and persistence +orchestration errors remain internal. A future known artifact gets its own marker, +typed field, and accessors in `RustiQData`. NPY parsing and conversion remain +internal persistence details. Compact ERIs are decoded only with the +basis-function count from their typed manifest attributes; the NPY length is never +used to infer that scientific context. Matrix readers accept C and Fortran order +and matrix writers emit Fortran order. + - `manifest.json` is UTF-8 JSON and describes the format, producer, scientific identity and artifacts. - `arrays/integrals/ao-eri.npy` is the AO electron-repulsion integral artifact. +Artifact paths are portable relative UTF-8 paths with `/` as the only separator, +independent of the host operating system. Empty components, `.`, `..`, +backslashes, drive-like prefixes containing `:`, absolute paths, Windows-reserved +device names, Windows-invalid filename characters, trailing dots/spaces, and paths +conflicting with `manifest.json` are rejected before storage access. Logical paths +are also checked case-insensitively to avoid cross-platform collisions. + Every artifact records common envelope metadata: - logical path;