Skip to main content

readcon_db/
cooked_soa.rs

1//! Optional **cooked SoA** payload: derived binary numerics beside authoritative CON text.
2//!
3//! RCSO is **non-authoritative**. CON text in `frames` is the sole authority.
4//!
5//! # Why CON text is still required (RCSO is **not** fully equivalent)
6//!
7//! RCSO stores only POD numerics (positions and optional forces/velocities). It does
8//! **not** carry element symbols, masses, cell/angles, constraint masks, JSON metadata,
9//! section labels, or exact on-disk CON bytes. Therefore it cannot replace
10//! `frames` for xxHash3 dedup, join/split fidelity, `reindex`, formula/symbol indexes,
11//! or CON export. The improvement is skipping **CON parse on numeric hot paths** when a
12//! valid cooked blob exists—not omitting the text tier from storage.
13//!
14//! Layout (little-endian, version 1):
15//! - magic `RCSO` (4 bytes)
16//! - version `u32` (=1)
17//! - natoms `u32`
18//! - flags `u32` (bit0 = forces block present, bit1 = velocities block present)
19//! - pos_dtype `u8` (0 = f64 row-major N×3); pad 3 bytes
20//! - reserved `u32` (=0)
21//! - positions: `natoms * 3` f64 LE
22//! - forces (if flag): `natoms * 3` f64 LE
23//! - velocities (if flag): `natoms * 3` f64 LE
24//!
25//! Never used for xxHash dedup, join-split fidelity, or secondary index rebuild.
26//! Rebuild by re-parsing CON text → [`CookedSoa::encode_frame`].
27
28use readcon_core::types::ConFrame;
29
30use crate::error::{Error, Result};
31
32pub const COOKED_MAGIC: &[u8; 4] = b"RCSO";
33pub const COOKED_VERSION: u32 = 1;
34pub const DTYPE_F64: u8 = 0;
35
36pub const FLAG_FORCES: u32 = 1 << 0;
37pub const FLAG_VELOCITIES: u32 = 1 << 1;
38
39const HEADER_LEN: usize = 4 + 4 + 4 + 4 + 1 + 3 + 4; // 24
40
41/// Decoded cooked numerics (always f64 for v1).
42#[derive(Clone, Debug, PartialEq)]
43pub struct CookedSoa {
44    pub natoms: u32,
45    pub positions: Vec<[f64; 3]>,
46    pub forces: Option<Vec<[f64; 3]>>,
47    pub velocities: Option<Vec<[f64; 3]>>,
48}
49
50impl CookedSoa {
51    pub fn encode_frame(frame: &ConFrame) -> Result<Vec<u8>> {
52        let n = frame.atom_data.len();
53        if n > u32::MAX as usize {
54            return Err(Error::Message("too many atoms for cooked SoA".into()));
55        }
56        let natoms = n as u32;
57        let mut flags = 0u32;
58        let has_f = frame.atom_data.iter().any(|a| a.force.is_some());
59        let has_v = frame.atom_data.iter().any(|a| a.velocity.is_some());
60        if has_f {
61            flags |= FLAG_FORCES;
62        }
63        if has_v {
64            flags |= FLAG_VELOCITIES;
65        }
66
67        let mut out =
68            Vec::with_capacity(HEADER_LEN + n * 3 * 8 * (1 + has_f as usize + has_v as usize));
69        out.extend_from_slice(COOKED_MAGIC);
70        out.extend_from_slice(&COOKED_VERSION.to_le_bytes());
71        out.extend_from_slice(&natoms.to_le_bytes());
72        out.extend_from_slice(&flags.to_le_bytes());
73        out.push(DTYPE_F64);
74        out.extend_from_slice(&[0u8; 3]);
75        out.extend_from_slice(&0u32.to_le_bytes());
76
77        for a in &frame.atom_data {
78            for c in [a.x, a.y, a.z] {
79                out.extend_from_slice(&c.to_le_bytes());
80            }
81        }
82        if has_f {
83            for a in &frame.atom_data {
84                let f = a.force.unwrap_or([0.0; 3]);
85                for c in f {
86                    out.extend_from_slice(&c.to_le_bytes());
87                }
88            }
89        }
90        if has_v {
91            for a in &frame.atom_data {
92                let v = a.velocity.unwrap_or([0.0; 3]);
93                for c in v {
94                    out.extend_from_slice(&c.to_le_bytes());
95                }
96            }
97        }
98        Ok(out)
99    }
100
101    pub fn decode(bytes: &[u8]) -> Result<Self> {
102        if bytes.len() < HEADER_LEN {
103            return Err(Error::Message("cooked SoA truncated header".into()));
104        }
105        if &bytes[0..4] != COOKED_MAGIC {
106            return Err(Error::Message("cooked SoA bad magic".into()));
107        }
108        let version = u32::from_le_bytes(bytes[4..8].try_into().unwrap());
109        if version != COOKED_VERSION {
110            return Err(Error::Message(format!(
111                "cooked SoA unsupported version {version}"
112            )));
113        }
114        let natoms = u32::from_le_bytes(bytes[8..12].try_into().unwrap());
115        let flags = u32::from_le_bytes(bytes[12..16].try_into().unwrap());
116        let dtype = bytes[16];
117        if dtype != DTYPE_F64 {
118            return Err(Error::Message(format!(
119                "cooked SoA unsupported dtype {dtype}"
120            )));
121        }
122        let n = natoms as usize;
123        let block = n
124            .checked_mul(3)
125            .ok_or_else(|| Error::Message("overflow".into()))?;
126        let block_bytes = block
127            .checked_mul(8)
128            .ok_or_else(|| Error::Message("overflow".into()))?;
129        let mut need = HEADER_LEN + block_bytes;
130        let has_f = flags & FLAG_FORCES != 0;
131        let has_v = flags & FLAG_VELOCITIES != 0;
132        if has_f {
133            need = need
134                .checked_add(block_bytes)
135                .ok_or_else(|| Error::Message("overflow".into()))?;
136        }
137        if has_v {
138            need = need
139                .checked_add(block_bytes)
140                .ok_or_else(|| Error::Message("overflow".into()))?;
141        }
142        if bytes.len() < need {
143            return Err(Error::Message("cooked SoA truncated body".into()));
144        }
145
146        let mut off = HEADER_LEN;
147        let positions = read_vec3_block(&bytes[off..off + block_bytes], n)?;
148        off += block_bytes;
149        let forces = if has_f {
150            let f = read_vec3_block(&bytes[off..off + block_bytes], n)?;
151            off += block_bytes;
152            Some(f)
153        } else {
154            None
155        };
156        let velocities = if has_v {
157            let v = read_vec3_block(&bytes[off..off + block_bytes], n)?;
158            Some(v)
159        } else {
160            None
161        };
162        let _ = off;
163        Ok(Self {
164            natoms,
165            positions,
166            forces,
167            velocities,
168        })
169    }
170
171    /// Prefer cooked bytes; return None if missing/invalid (caller parses CON).
172    pub fn try_decode(bytes: &[u8]) -> Option<Self> {
173        Self::decode(bytes).ok()
174    }
175}
176
177/// Length-prefixed RCSO blobs for one collective on the caller comm.
178/// Grain: many frames per Bcast (ADIOS BP5 `MinDeferredSize` is 4 MiB).
179pub const BATCH_MAGIC: &[u8; 4] = b"RCSB";
180pub const BATCH_VERSION: u32 = 1;
181
182pub fn encode_batch(blobs: &[Vec<u8>]) -> Result<Vec<u8>> {
183    if blobs.len() > u32::MAX as usize {
184        return Err(Error::Message("too many frames in pack batch".into()));
185    }
186    let mut out = Vec::new();
187    out.extend_from_slice(BATCH_MAGIC);
188    out.extend_from_slice(&BATCH_VERSION.to_le_bytes());
189    out.extend_from_slice(&(blobs.len() as u32).to_le_bytes());
190    for b in blobs {
191        if b.len() > u32::MAX as usize {
192            return Err(Error::Message("RCSO blob exceeds u32 length".into()));
193        }
194        out.extend_from_slice(&(b.len() as u32).to_le_bytes());
195        out.extend_from_slice(b);
196    }
197    Ok(out)
198}
199
200pub fn decode_batch(bytes: &[u8]) -> Result<Vec<Vec<u8>>> {
201    if bytes.len() < 12 {
202        return Err(Error::Message("RCSB truncated header".into()));
203    }
204    if &bytes[0..4] != BATCH_MAGIC {
205        return Err(Error::Message("RCSB bad magic".into()));
206    }
207    let version = u32::from_le_bytes(bytes[4..8].try_into().unwrap());
208    if version != BATCH_VERSION {
209        return Err(Error::Message(format!(
210            "RCSB unsupported version {version}"
211        )));
212    }
213    let n = u32::from_le_bytes(bytes[8..12].try_into().unwrap()) as usize;
214    let mut off = 12usize;
215    let mut out = Vec::with_capacity(n);
216    for _ in 0..n {
217        if off + 4 > bytes.len() {
218            return Err(Error::Message("RCSB truncated length".into()));
219        }
220        let ln = u32::from_le_bytes(bytes[off..off + 4].try_into().unwrap()) as usize;
221        off += 4;
222        if off + ln > bytes.len() {
223            return Err(Error::Message("RCSB truncated blob".into()));
224        }
225        out.push(bytes[off..off + ln].to_vec());
226        off += ln;
227    }
228    Ok(out)
229}
230
231fn read_vec3_block(bytes: &[u8], n: usize) -> Result<Vec<[f64; 3]>> {
232    let mut out = Vec::with_capacity(n);
233    let mut i = 0;
234    for _ in 0..n {
235        let mut row = [0.0f64; 3];
236        for c in 0..3 {
237            let start = i;
238            let end = start + 8;
239            if end > bytes.len() {
240                return Err(Error::Message("cooked SoA short block".into()));
241            }
242            row[c] = f64::from_le_bytes(bytes[start..end].try_into().unwrap());
243            i = end;
244        }
245        out.push(row);
246    }
247    Ok(out)
248}
249
250#[cfg(test)]
251mod tests {
252    use super::*;
253    use readcon_core::iterators::ConFrameIterator;
254    use std::path::PathBuf;
255
256    fn fixture(name: &str) -> String {
257        let p = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
258            .join("resources/test")
259            .join(name);
260        std::fs::read_to_string(p).unwrap()
261    }
262
263    #[test]
264    fn encode_decode_positions_fixture() {
265        let text = fixture("tiny_cuh2.con");
266        let fr = ConFrameIterator::new(&text).next().unwrap().unwrap();
267        let bytes = CookedSoa::encode_frame(&fr).unwrap();
268        assert!(bytes.len() > HEADER_LEN);
269        let cooked = CookedSoa::decode(&bytes).unwrap();
270        assert_eq!(cooked.natoms as usize, fr.atom_data.len());
271        for (i, a) in fr.atom_data.iter().enumerate() {
272            assert_eq!(cooked.positions[i], [a.x, a.y, a.z]);
273        }
274        assert!(cooked.forces.is_none());
275    }
276
277    #[test]
278    fn encode_decode_forces_fixture() {
279        let text = fixture("tiny_cuh2_forces.con");
280        let fr = ConFrameIterator::new(&text).next().unwrap().unwrap();
281        let bytes = CookedSoa::encode_frame(&fr).unwrap();
282        let cooked = CookedSoa::decode(&bytes).unwrap();
283        assert!(cooked.forces.is_some());
284        let forces = cooked.forces.as_ref().unwrap();
285        for (i, a) in fr.atom_data.iter().enumerate() {
286            assert_eq!(cooked.positions[i], [a.x, a.y, a.z]);
287            if let Some(f) = a.force {
288                assert_eq!(forces[i], f);
289            }
290        }
291    }
292
293    #[test]
294    fn encode_decode_velocities_fixture() {
295        let text = fixture("tiny_cuh2_vel_forces.con");
296        let fr = ConFrameIterator::new(&text).next().unwrap().unwrap();
297        let bytes = CookedSoa::encode_frame(&fr).unwrap();
298        let cooked = CookedSoa::decode(&bytes).unwrap();
299        assert!(cooked.forces.is_some());
300        assert!(cooked.velocities.is_some());
301        let forces = cooked.forces.as_ref().unwrap();
302        let vels = cooked.velocities.as_ref().unwrap();
303        for (i, a) in fr.atom_data.iter().enumerate() {
304            assert_eq!(cooked.positions[i], [a.x, a.y, a.z]);
305            if let Some(f) = a.force {
306                assert_eq!(forces[i], f);
307            }
308            if let Some(v) = a.velocity {
309                assert_eq!(vels[i], v);
310            }
311        }
312        assert!((vels[0][0] - 0.001234).abs() < 1e-12);
313        assert!((forces[0][0] - 0.123456).abs() < 1e-12);
314    }
315
316    #[test]
317    fn bad_magic_rejected() {
318        assert!(CookedSoa::decode(b"XXXX").is_err());
319        assert!(CookedSoa::try_decode(b"XXXX").is_none());
320    }
321
322    #[test]
323    fn rcsb_batch_roundtrip() {
324        let a = vec![1u8, 2, 3];
325        let b = vec![4u8, 5];
326        let enc = encode_batch(&[a.clone(), b.clone()]).unwrap();
327        let dec = decode_batch(&enc).unwrap();
328        assert_eq!(dec, vec![a, b]);
329        assert!(decode_batch(&enc[..8]).is_err());
330    }
331}