1use crate::error::ParseError;
11use crate::types::ConFrame;
12
13pub const RCSO_MAGIC: &[u8; 4] = b"RCSO";
15pub const RCSO_VERSION: u32 = 1;
17pub const RCSO_DTYPE_F64: u8 = 0;
19pub const RCSO_FLAG_FORCES: u32 = 1 << 0;
21pub const RCSO_FLAG_VELOCITIES: u32 = 1 << 1;
23
24const HEADER_LEN: usize = 24;
25
26pub const RCSB_MAGIC: &[u8; 4] = b"RCSB";
28pub const RCSB_VERSION: u32 = 1;
30
31#[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 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 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
170pub 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
189pub 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}