diff --git a/core/Cargo.lock b/core/Cargo.lock index 50f446513..a96ab6be1 100644 --- a/core/Cargo.lock +++ b/core/Cargo.lock @@ -17,6 +17,15 @@ dependencies = [ "libc", ] +[[package]] +name = "arbitrary" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c3d036a3c4ab069c7b410a2ce876bd74808d2d0888a82667669f8e783a898bf1" +dependencies = [ + "derive_arbitrary", +] + [[package]] name = "argon2" version = "0.5.3" @@ -282,6 +291,12 @@ dependencies = [ "cfg-if", ] +[[package]] +name = "crossbeam-utils" +version = "0.8.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61803da095bee82a81bb1a452ecc25d3b2f1416d1897eb86430c6159ef717c17" + [[package]] name = "crypto-bigint" version = "0.5.5" @@ -340,6 +355,17 @@ version = "0.5.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7cd812cc2bc1d69d4764bd80df88b4317eaef9e773c75226407d9bc0876b211c" +[[package]] +name = "derive_arbitrary" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e567bd82dcff979e4b03460c307b3cdc9e96fde3d73bed1496d2bc75d9dd62a" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "digest" version = "0.10.7" @@ -1602,6 +1628,7 @@ dependencies = [ "solstone-core-entity-matching", "solstone-core-journal-io", "unicode-normalization", + "zip", ] [[package]] @@ -2567,12 +2594,41 @@ version = "1.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" +[[package]] +name = "zip" +version = "2.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fabe6324e908f85a1c52063ce7aa26b68dcb7eb6dbc83a2d148403c9bc3eba50" +dependencies = [ + "arbitrary", + "crc32fast", + "crossbeam-utils", + "displaydoc", + "flate2", + "indexmap", + "memchr", + "thiserror 2.0.19", + "zopfli", +] + [[package]] name = "zmij" version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" +[[package]] +name = "zopfli" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f05cd8797d63865425ff89b5c4a48804f35ba0ce8d125800027ad6017d2b5249" +dependencies = [ + "bumpalo", + "crc32fast", + "log", + "simd-adler32", +] + [[package]] name = "zune-core" version = "0.5.1" diff --git a/core/crates/solstone-core-entity/Cargo.toml b/core/crates/solstone-core-entity/Cargo.toml index 1aa0529d9..3641de457 100644 --- a/core/crates/solstone-core-entity/Cargo.toml +++ b/core/crates/solstone-core-entity/Cargo.toml @@ -21,6 +21,7 @@ sha2 = "0.10.9" solstone-core-journal-io.workspace = true unicode-normalization.workspace = true solstone-core-entity-matching.workspace = true +zip = { version = "=2.4.2", default-features = false, features = ["deflate"] } [dev-dependencies] diff --git a/core/crates/solstone-core-entity/src/store/mod.rs b/core/crates/solstone-core-entity/src/store/mod.rs index 49af4f3c9..10b904e18 100644 --- a/core/crates/solstone-core-entity/src/store/mod.rs +++ b/core/crates/solstone-core-entity/src/store/mod.rs @@ -11,6 +11,8 @@ mod map; mod paths; mod reconcile; mod repair; +#[allow(dead_code)] +pub(crate) mod voiceprints; mod write; pub use ambiguity::{ diff --git a/core/crates/solstone-core-entity/src/store/voiceprints.rs b/core/crates/solstone-core-entity/src/store/voiceprints.rs new file mode 100644 index 000000000..361cdcfda --- /dev/null +++ b/core/crates/solstone-core-entity/src/store/voiceprints.rs @@ -0,0 +1,416 @@ +// SPDX-License-Identifier: AGPL-3.0-only +// Copyright (c) 2026 sol pbc + +use std::collections::BTreeSet; +use std::error::Error; +use std::fmt; +use std::io::{Cursor, Read, Write}; + +use zip::write::SimpleFileOptions; +use zip::{CompressionMethod, ZipArchive, ZipWriter}; + +const EMBEDDING_WIDTH: usize = 256; +const EMBEDDINGS_MEMBER: &str = "embeddings.npy"; +const METADATA_MEMBER: &str = "metadata.npy"; + +#[derive(Debug, Clone, PartialEq)] +pub(crate) struct VoiceprintArchive { + pub embeddings: Vec, + pub rows: usize, + pub metadata: Vec, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum VoiceprintNpzError { + Archive(String), + Invalid(String), +} + +impl fmt::Display for VoiceprintNpzError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self { + Self::Archive(message) | Self::Invalid(message) => formatter.write_str(message), + } + } +} + +impl Error for VoiceprintNpzError {} + +pub(crate) fn read_voiceprints_npz(bytes: &[u8]) -> Result { + let mut archive = ZipArchive::new(Cursor::new(bytes)).map_err(|error| { + VoiceprintNpzError::Archive(format!("invalid voiceprint archive: {error}")) + })?; + let names = (0..archive.len()) + .map(|index| { + archive + .by_index(index) + .map(|file| file.name().to_owned()) + .map_err(|error| { + VoiceprintNpzError::Archive(format!("invalid voiceprint archive: {error}")) + }) + }) + .collect::, _>>()?; + let expected = BTreeSet::from([EMBEDDINGS_MEMBER.to_owned(), METADATA_MEMBER.to_owned()]); + if names != expected || archive.len() != expected.len() { + return Err(VoiceprintNpzError::Invalid( + "voiceprint archive must contain exactly embeddings.npy and metadata.npy".to_owned(), + )); + } + let embeddings = read_member(&mut archive, EMBEDDINGS_MEMBER)?; + let metadata = read_member(&mut archive, METADATA_MEMBER)?; + let (embedding_rows, values) = parse_embeddings(&embeddings)?; + let (metadata_rows, metadata) = parse_metadata(&metadata)?; + if embedding_rows != metadata_rows { + return Err(VoiceprintNpzError::Invalid( + "voiceprint embedding and metadata row counts differ".to_owned(), + )); + } + Ok(VoiceprintArchive { + embeddings: values, + rows: embedding_rows, + metadata, + }) +} + +pub(crate) fn write_voiceprints_npz( + embeddings: &[f32], + metadata: &[String], +) -> Result, VoiceprintNpzError> { + let expected_embeddings = metadata.len().checked_mul(EMBEDDING_WIDTH).ok_or_else(|| { + VoiceprintNpzError::Invalid("voiceprint row count is too large".to_owned()) + })?; + if embeddings.len() != expected_embeddings { + return Err(VoiceprintNpzError::Invalid(format!( + "voiceprint embeddings length {} does not match {} rows", + embeddings.len(), + metadata.len() + ))); + } + validate_metadata(metadata)?; + let embeddings = write_embeddings_npy(embeddings, metadata.len()); + let metadata = write_metadata_npy(metadata)?; + let cursor = Cursor::new(Vec::new()); + let mut writer = ZipWriter::new(cursor); + let options = SimpleFileOptions::default().compression_method(CompressionMethod::Deflated); + writer + .start_file(EMBEDDINGS_MEMBER, options) + .map_err(zip_error)?; + writer.write_all(&embeddings).map_err(io_error)?; + writer + .start_file(METADATA_MEMBER, options) + .map_err(zip_error)?; + writer.write_all(&metadata).map_err(io_error)?; + writer + .finish() + .map_err(zip_error) + .map(|cursor| cursor.into_inner()) +} + +fn read_member( + archive: &mut ZipArchive>, + name: &str, +) -> Result, VoiceprintNpzError> { + let mut member = archive.by_name(name).map_err(zip_error)?; + let mut bytes = Vec::new(); + member.read_to_end(&mut bytes).map_err(io_error)?; + Ok(bytes) +} + +fn parse_embeddings(bytes: &[u8]) -> Result<(usize, Vec), VoiceprintNpzError> { + let (header, payload) = parse_npy(bytes)?; + if header.descr != " Result<(usize, Vec), VoiceprintNpzError> { + let (header, payload) = parse_npy(bytes)?; + if header.fortran_order || header.shape.len() != 1 { + return Err(VoiceprintNpzError::Invalid( + "metadata.npy must be a C-order one-dimensional unicode array".to_owned(), + )); + } + let width = header + .descr + .strip_prefix("() + .map_err(|_| { + VoiceprintNpzError::Invalid("metadata.npy has an invalid unicode dtype".to_owned()) + })?; + let rows = header.shape[0]; + if width == 0 { + if rows == 0 && payload.is_empty() { + return Ok((0, Vec::new())); + } + return Err(VoiceprintNpzError::Invalid( + "metadata.npy has an invalid zero-width unicode dtype".to_owned(), + )); + } + let row_width = width.checked_mul(4).ok_or_else(|| { + VoiceprintNpzError::Invalid("metadata.npy unicode dtype is too large".to_owned()) + })?; + let expected = rows + .checked_mul(row_width) + .ok_or_else(|| VoiceprintNpzError::Invalid("metadata.npy shape is too large".to_owned()))?; + if payload.len() != expected { + return Err(VoiceprintNpzError::Invalid( + "metadata.npy payload length does not match its shape".to_owned(), + )); + } + let mut values = Vec::with_capacity(rows); + for row in payload.chunks_exact(row_width) { + let value = row + .chunks_exact(4) + .map(|bytes| u32::from_le_bytes(bytes.try_into().expect("exact chunk length"))) + .take_while(|codepoint| *codepoint != 0) + .map(|codepoint| { + char::from_u32(codepoint).ok_or_else(|| { + VoiceprintNpzError::Invalid( + "metadata.npy contains an invalid unicode code point".to_owned(), + ) + }) + }) + .collect::>()?; + values.push(value); + } + validate_metadata(&values)?; + Ok((rows, values)) +} + +fn write_embeddings_npy(values: &[f32], rows: usize) -> Vec { + let mut payload = Vec::with_capacity(values.len() * 4); + for value in values { + payload.extend_from_slice(&value.to_le_bytes()); + } + write_npy(" Result, VoiceprintNpzError> { + let width = values + .iter() + .map(|value| value.chars().count()) + .max() + .unwrap_or(0); + let mut payload = Vec::new(); + for value in values { + for character in value.chars() { + payload.extend_from_slice(&(character as u32).to_le_bytes()); + } + for _ in value.chars().count()..width { + payload.extend_from_slice(&0_u32.to_le_bytes()); + } + } + Ok(write_npy( + &format!(" Vec { + let mut header = format!("{{'descr': '{descr}', 'fortran_order': False, 'shape': {shape}, }}"); + let padding = (64 - ((10 + header.len() + 1) % 64)) % 64; + header.push_str(&" ".repeat(padding)); + header.push('\n'); + let mut bytes = Vec::with_capacity(10 + header.len() + payload.len()); + bytes.extend_from_slice(b"\x93NUMPY"); + bytes.extend_from_slice(&[1, 0]); + bytes.extend_from_slice(&(header.len() as u16).to_le_bytes()); + bytes.extend_from_slice(header.as_bytes()); + bytes.extend_from_slice(payload); + bytes +} + +struct NpyHeader { + descr: String, + fortran_order: bool, + shape: Vec, +} + +fn parse_npy(bytes: &[u8]) -> Result<(NpyHeader, &[u8]), VoiceprintNpzError> { + if bytes.len() < 10 || &bytes[..6] != b"\x93NUMPY" { + return Err(VoiceprintNpzError::Invalid( + "invalid NPY magic bytes".to_owned(), + )); + } + let version = (bytes[6], bytes[7]); + let (header_start, header_len): (usize, usize) = match version { + (1, 0) => (10, u16::from_le_bytes([bytes[8], bytes[9]]) as usize), + (2, 0) | (3, 0) if bytes.len() >= 12 => ( + 12, + u32::from_le_bytes([bytes[8], bytes[9], bytes[10], bytes[11]]) as usize, + ), + _ => { + return Err(VoiceprintNpzError::Invalid( + "unsupported NPY version".to_owned(), + )); + } + }; + let header_end = header_start + .checked_add(header_len) + .ok_or_else(|| VoiceprintNpzError::Invalid("NPY header is too large".to_owned()))?; + let header = bytes + .get(header_start..header_end) + .ok_or_else(|| VoiceprintNpzError::Invalid("truncated NPY header".to_owned()))?; + let header = std::str::from_utf8(header) + .map_err(|_| VoiceprintNpzError::Invalid("NPY header is not UTF-8".to_owned()))?; + let descr = header_string_value(header, "descr")?; + let fortran_order = match header_value(header, "fortran_order")? { + "False" => false, + "True" => true, + _ => { + return Err(VoiceprintNpzError::Invalid( + "NPY header has an invalid fortran_order value".to_owned(), + )); + } + }; + let shape_text = header_value(header, "shape")?; + let shape = shape_text + .strip_prefix('(') + .and_then(|value| value.strip_suffix(')')) + .ok_or_else(|| VoiceprintNpzError::Invalid("NPY header has an invalid shape".to_owned()))? + .split(',') + .filter_map(|part| { + let part = part.trim(); + (!part.is_empty()).then_some(part) + }) + .map(|part| { + part.parse::().map_err(|_| { + VoiceprintNpzError::Invalid("NPY header has an invalid shape dimension".to_owned()) + }) + }) + .collect::, _>>()?; + if shape.is_empty() { + return Err(VoiceprintNpzError::Invalid( + "NPY header must describe an array".to_owned(), + )); + } + Ok(( + NpyHeader { + descr, + fortran_order, + shape, + }, + &bytes[header_end..], + )) +} + +fn header_string_value(header: &str, key: &str) -> Result { + let value = header_value(header, key)?; + let value = value + .strip_prefix('\'') + .and_then(|value| value.strip_suffix('\'')) + .ok_or_else(|| VoiceprintNpzError::Invalid(format!("NPY header has an invalid {key}")))?; + Ok(value.to_owned()) +} + +fn header_value<'a>(header: &'a str, key: &str) -> Result<&'a str, VoiceprintNpzError> { + let prefix = format!("'{key}':"); + let value = header + .split(&prefix) + .nth(1) + .ok_or_else(|| VoiceprintNpzError::Invalid(format!("NPY header is missing {key}")))? + .trim_start(); + if value.starts_with('(') { + let end = value.find(')').ok_or_else(|| { + VoiceprintNpzError::Invalid(format!("NPY header has an invalid {key}")) + })?; + return Ok(&value[..=end]); + } + Ok(value + .split(',') + .next() + .ok_or_else(|| VoiceprintNpzError::Invalid(format!("NPY header has an invalid {key}")))? + .trim()) +} + +fn validate_metadata(values: &[String]) -> Result<(), VoiceprintNpzError> { + for value in values { + serde_json::from_str::(value).map_err(|error| { + VoiceprintNpzError::Invalid(format!("voiceprint metadata must be JSON: {error}")) + })?; + } + Ok(()) +} + +fn zip_error(error: zip::result::ZipError) -> VoiceprintNpzError { + VoiceprintNpzError::Archive(format!("voiceprint archive error: {error}")) +} + +fn io_error(error: std::io::Error) -> VoiceprintNpzError { + VoiceprintNpzError::Archive(format!("voiceprint archive I/O error: {error}")) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn round_trips_embeddings_and_unicode_json_metadata() { + let embeddings = (0..EMBEDDING_WIDTH * 2) + .map(|index| index as f32 / 10.0) + .collect::>(); + let metadata = vec![ + r#"{"day":"20260101","source":"screen"}"#.to_owned(), + r#"{"day":"20260102","label":"José"}"#.to_owned(), + ]; + + let bytes = write_voiceprints_npz(&embeddings, &metadata).unwrap(); + let actual = read_voiceprints_npz(&bytes).unwrap(); + + assert_eq!(actual.embeddings, embeddings); + assert_eq!(actual.rows, 2); + assert_eq!(actual.metadata, metadata); + } + + #[test] + fn rejects_object_metadata_dtype() { + let bytes = write_npy("|O", "(1,)", &[0; 8]); + assert!(matches!( + parse_metadata(&bytes), + Err(VoiceprintNpzError::Invalid(message)) if message.contains("never pickle") + )); + } + + #[test] + fn round_trips_empty_voiceprints() { + let bytes = write_voiceprints_npz(&[], &[]).unwrap(); + assert_eq!( + read_voiceprints_npz(&bytes).unwrap(), + VoiceprintArchive { + embeddings: Vec::new(), + rows: 0, + metadata: Vec::new(), + } + ); + } +}