Skip to main content

readcon_core/
rcso.rs

1//! RCSO cooked SoA bytes for a caller-side `MPI_Bcast`.
2//!
3//! Layout matches readcon-db `cooked_soa` v1 (little-endian):
4//! magic `RCSO`, version, natoms, flags, dtype, pad, reserved,
5//! then positions, optional forces, optional velocities.
6//!
7//! This crate never calls `MPI_Init` or names WORLD. Rank 0 encodes;
8//! the caller broadcasts; workers decode.
9
10use crate::error::ParseError;
11use crate::types::ConFrame;
12
13/// Four-byte magic. Identical to readcon-db.
14pub const RCSO_MAGIC: &[u8; 4] = b"RCSO";
15/// Layout version 1.
16pub const RCSO_VERSION: u32 = 1;
17/// Positions/forces/velocities stored as f64.
18pub const RCSO_DTYPE_F64: u8 = 0;
19/// Bit 0: forces block present.
20pub const RCSO_FLAG_FORCES: u32 = 1 << 0;
21/// Bit 1: velocities block present.
22pub const RCSO_FLAG_VELOCITIES: u32 = 1 << 1;
23
24const HEADER_LEN: usize = 24;
25
26/// Length-prefixed RCSO blobs for one collective (ADIOS BP5 grain: many frames).
27pub const RCSB_MAGIC: &[u8; 4] = b"RCSB";
28/// Batch envelope version 1.
29pub const RCSB_VERSION: u32 = 1;
30
31/// Decoded cooked numerics.
32#[derive(Clone, Debug, PartialEq)]
33pub struct Rcso {
34    pub natoms: u32,
35    pub positions: Vec<[f64; 3]>,
36    pub forces: Option<Vec<[f64; 3]>>,
37    pub velocities: Option<Vec<[f64; 3]>>,
38}
39
40fn err(msg: impl Into<String>) -> ParseError {
41    ParseError::ValidationError(msg.into())
42}
43
44impl Rcso {
45    /// Pack one parsed frame. CON text stays the authority.
46    pub fn encode_frame(frame: &ConFrame) -> Result<Vec<u8>, ParseError> {
47        let n = frame.atom_data.len();
48        if n > u32::MAX as usize {
49            return Err(err("too many atoms for RCSO"));
50        }
51        let natoms = n as u32;
52        let has_f = frame.atom_data.iter().any(|a| a.force.is_some());
53        let has_v = frame.atom_data.iter().any(|a| a.velocity.is_some());
54        let mut flags = 0u32;
55        if has_f {
56            flags |= RCSO_FLAG_FORCES;
57        }
58        if has_v {
59            flags |= RCSO_FLAG_VELOCITIES;
60        }
61
62        let mut out =
63            Vec::with_capacity(HEADER_LEN + n * 3 * 8 * (1 + has_f as usize + has_v as usize));
64        out.extend_from_slice(RCSO_MAGIC);
65        out.extend_from_slice(&RCSO_VERSION.to_le_bytes());
66        out.extend_from_slice(&natoms.to_le_bytes());
67        out.extend_from_slice(&flags.to_le_bytes());
68        out.push(RCSO_DTYPE_F64);
69        out.extend_from_slice(&[0u8; 3]);
70        out.extend_from_slice(&0u32.to_le_bytes());
71
72        for a in &frame.atom_data {
73            for c in [a.x, a.y, a.z] {
74                out.extend_from_slice(&c.to_le_bytes());
75            }
76        }
77        if has_f {
78            for a in &frame.atom_data {
79                let f = a.force.unwrap_or([0.0; 3]);
80                for c in f {
81                    out.extend_from_slice(&c.to_le_bytes());
82                }
83            }
84        }
85        if has_v {
86            for a in &frame.atom_data {
87                let v = a.velocity.unwrap_or([0.0; 3]);
88                for c in v {
89                    out.extend_from_slice(&c.to_le_bytes());
90                }
91            }
92        }
93        Ok(out)
94    }
95
96    /// Decode a v1 RCSO blob.
97    pub fn decode(bytes: &[u8]) -> Result<Self, ParseError> {
98        if bytes.len() < HEADER_LEN {
99            return Err(err("RCSO truncated header"));
100        }
101        if &bytes[0..4] != RCSO_MAGIC {
102            return Err(err("RCSO bad magic"));
103        }
104        let version = u32::from_le_bytes(bytes[4..8].try_into().unwrap());
105        if version != RCSO_VERSION {
106            return Err(err(format!("RCSO unsupported version {version}")));
107        }
108        let natoms = u32::from_le_bytes(bytes[8..12].try_into().unwrap());
109        let flags = u32::from_le_bytes(bytes[12..16].try_into().unwrap());
110        let dtype = bytes[16];
111        if dtype != RCSO_DTYPE_F64 {
112            return Err(err(format!("RCSO unsupported dtype {dtype}")));
113        }
114        let n = natoms as usize;
115        let block = n.checked_mul(3).ok_or_else(|| err("RCSO overflow"))?;
116        let block_bytes = block.checked_mul(8).ok_or_else(|| err("RCSO overflow"))?;
117        let mut need = HEADER_LEN + block_bytes;
118        let has_f = flags & RCSO_FLAG_FORCES != 0;
119        let has_v = flags & RCSO_FLAG_VELOCITIES != 0;
120        if has_f {
121            need = need
122                .checked_add(block_bytes)
123                .ok_or_else(|| err("RCSO overflow"))?;
124        }
125        if has_v {
126            need = need
127                .checked_add(block_bytes)
128                .ok_or_else(|| err("RCSO overflow"))?;
129        }
130        if bytes.len() < need {
131            return Err(err("RCSO truncated body"));
132        }
133
134        let mut off = HEADER_LEN;
135        let positions = read_vec3_block(&bytes[off..off + block_bytes], n)?;
136        off += block_bytes;
137        let forces = if has_f {
138            let f = read_vec3_block(&bytes[off..off + block_bytes], n)?;
139            off += block_bytes;
140            Some(f)
141        } else {
142            None
143        };
144        let velocities = if has_v {
145            Some(read_vec3_block(&bytes[off..off + block_bytes], n)?)
146        } else {
147            None
148        };
149        Ok(Self {
150            natoms,
151            positions,
152            forces,
153            velocities,
154        })
155    }
156}
157
158fn read_vec3_block(bytes: &[u8], n: usize) -> Result<Vec<[f64; 3]>, ParseError> {
159    let mut out = Vec::with_capacity(n);
160    for i in 0..n {
161        let base = i * 24;
162        let x = f64::from_le_bytes(bytes[base..base + 8].try_into().unwrap());
163        let y = f64::from_le_bytes(bytes[base + 8..base + 16].try_into().unwrap());
164        let z = f64::from_le_bytes(bytes[base + 16..base + 24].try_into().unwrap());
165        out.push([x, y, z]);
166    }
167    Ok(out)
168}
169
170/// Pack many RCSO blobs into one RCSB envelope (one Bcast).
171pub fn encode_batch(blobs: &[Vec<u8>]) -> Result<Vec<u8>, ParseError> {
172    if blobs.len() > u32::MAX as usize {
173        return Err(err("too many frames in RCSB batch"));
174    }
175    let mut out = Vec::new();
176    out.extend_from_slice(RCSB_MAGIC);
177    out.extend_from_slice(&RCSB_VERSION.to_le_bytes());
178    out.extend_from_slice(&(blobs.len() as u32).to_le_bytes());
179    for b in blobs {
180        if b.len() > u32::MAX as usize {
181            return Err(err("RCSO blob exceeds u32 length"));
182        }
183        out.extend_from_slice(&(b.len() as u32).to_le_bytes());
184        out.extend_from_slice(b);
185    }
186    Ok(out)
187}
188
189/// Split an RCSB envelope into RCSO blobs.
190pub fn decode_batch(bytes: &[u8]) -> Result<Vec<Vec<u8>>, ParseError> {
191    if bytes.len() < 12 {
192        return Err(err("RCSB truncated header"));
193    }
194    if &bytes[0..4] != RCSB_MAGIC {
195        return Err(err("RCSB bad magic"));
196    }
197    let version = u32::from_le_bytes(bytes[4..8].try_into().unwrap());
198    if version != RCSB_VERSION {
199        return Err(err(format!("RCSB unsupported version {version}")));
200    }
201    let n = u32::from_le_bytes(bytes[8..12].try_into().unwrap()) as usize;
202    let mut off = 12usize;
203    let mut out = Vec::with_capacity(n);
204    for _ in 0..n {
205        if off + 4 > bytes.len() {
206            return Err(err("RCSB truncated length"));
207        }
208        let ln = u32::from_le_bytes(bytes[off..off + 4].try_into().unwrap()) as usize;
209        off += 4;
210        if off + ln > bytes.len() {
211            return Err(err("RCSB truncated blob"));
212        }
213        out.push(bytes[off..off + ln].to_vec());
214        off += ln;
215    }
216    Ok(out)
217}
218
219#[cfg(test)]
220mod tests {
221    use super::*;
222    use crate::iterators::ConFrameIterator;
223
224    fn first_frame(text: &str) -> ConFrame {
225        ConFrameIterator::new(text)
226            .next()
227            .expect("frame")
228            .expect("parse")
229    }
230
231    #[test]
232    fn v2_minimal_roundtrip_positions() {
233        let text = include_str!("../resources/conformance/valid/v2_minimal.con");
234        let frame = first_frame(text);
235        let blob = Rcso::encode_frame(&frame).unwrap();
236        assert_eq!(&blob[0..4], b"RCSO");
237        assert_eq!(u32::from_le_bytes(blob[4..8].try_into().unwrap()), 1);
238        let got = Rcso::decode(&blob).unwrap();
239        assert_eq!(got.natoms as usize, frame.atom_data.len());
240        assert!(got.forces.is_none());
241        assert!(got.velocities.is_none());
242        for (a, p) in frame.atom_data.iter().zip(&got.positions) {
243            assert_eq!([a.x, a.y, a.z], *p);
244        }
245    }
246
247    #[test]
248    fn forces_fixture_keeps_force_block() {
249        let text = include_str!("../resources/test/tiny_cuh2_forces.con");
250        let frame = first_frame(text);
251        assert!(frame.atom_data.iter().any(|a| a.force.is_some()));
252        let blob = Rcso::encode_frame(&frame).unwrap();
253        let flags = u32::from_le_bytes(blob[12..16].try_into().unwrap());
254        assert_ne!(flags & RCSO_FLAG_FORCES, 0);
255        let got = Rcso::decode(&blob).unwrap();
256        let forces = got.forces.expect("forces");
257        for (a, f) in frame.atom_data.iter().zip(&forces) {
258            assert_eq!(a.force.unwrap_or([0.0; 3]), *f);
259        }
260    }
261
262    #[test]
263    fn rejects_bad_magic() {
264        assert!(Rcso::decode(b"NOPE").is_err());
265    }
266
267    #[test]
268    fn rejects_truncated_and_bad_fields() {
269        assert!(Rcso::decode(&[0u8; 8]).is_err());
270        let frame = first_frame(include_str!(
271            "../resources/conformance/valid/v2_minimal.con"
272        ));
273        let good = Rcso::encode_frame(&frame).unwrap();
274        let mut ver = good.clone();
275        ver[4..8].copy_from_slice(&2u32.to_le_bytes());
276        assert!(Rcso::decode(&ver).is_err());
277        let mut dt = good.clone();
278        dt[16] = 1;
279        assert!(Rcso::decode(&dt).is_err());
280        assert!(Rcso::decode(&good[..20]).is_err());
281        assert!(decode_batch(b"xxxx").is_err());
282        assert!(decode_batch(b"RCSB\x02\x00\x00\x00\x00\x00\x00\x00").is_err());
283    }
284
285    #[test]
286    fn velocities_fixture_keeps_velocity_block() {
287        let text = include_str!("../resources/test/tiny_cuh2.convel");
288        let frame = first_frame(text);
289        assert!(frame.atom_data.iter().any(|a| a.velocity.is_some()));
290        let blob = Rcso::encode_frame(&frame).unwrap();
291        let flags = u32::from_le_bytes(blob[12..16].try_into().unwrap());
292        assert_ne!(flags & RCSO_FLAG_VELOCITIES, 0);
293        let got = Rcso::decode(&blob).unwrap();
294        let vels = got.velocities.expect("velocities");
295        for (a, v) in frame.atom_data.iter().zip(&vels) {
296            assert_eq!(a.velocity.unwrap_or([0.0; 3]), *v);
297        }
298    }
299
300    #[test]
301    fn rcsb_batch_holds_two_blobs() {
302        let text = include_str!("../resources/conformance/valid/v2_minimal.con");
303        let frame = first_frame(text);
304        let a = Rcso::encode_frame(&frame).unwrap();
305        let b = a.clone();
306        let batch = encode_batch(&[a.clone(), b.clone()]).unwrap();
307        assert_eq!(&batch[0..4], b"RCSB");
308        let parts = decode_batch(&batch).unwrap();
309        assert_eq!(parts.len(), 2);
310        assert_eq!(parts[0], a);
311        assert_eq!(parts[1], b);
312    }
313}