1use 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; #[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 pub fn try_decode(bytes: &[u8]) -> Option<Self> {
173 Self::decode(bytes).ok()
174 }
175}
176
177pub 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}