1use crate::types::{
2 ConFrame, SECTION_CHARGES, SECTION_DISPLACEMENTS, SECTION_ENERGIES, SECTION_FORCES,
3 SECTION_MAGMOMS, SECTION_SPINS, SECTION_SPREADS, SECTION_VELOCITIES, encode_fixed_bitmask,
4 meta,
5};
6use serde_json::json;
7use std::fs::File;
8use std::io::{self, BufWriter, Write};
9use std::path::Path;
10
11const DEFAULT_FLOAT_PRECISION: usize = 6;
13
14#[derive(Clone, Copy, Debug, PartialEq, Eq)]
16pub enum FloatFormat {
17 DecimalPlaces(usize),
19 RoundTrip,
22}
23
24impl Default for FloatFormat {
25 fn default() -> Self {
26 Self::DecimalPlaces(DEFAULT_FLOAT_PRECISION)
27 }
28}
29
30pub struct ConFrameWriter<W: Write> {
45 writer: BufWriter<W>,
46 float_format: FloatFormat,
47 canonical: bool,
52 metadata_cache: Option<MetadataCacheEntry>,
59 scratch: Vec<u8>,
62}
63
64#[derive(Debug)]
65struct MetadataCacheEntry {
66 spec_version: u32,
70 has_velocities: bool,
71 has_forces: bool,
72 has_energies: bool,
73 has_charges: bool,
74 has_spins: bool,
75 has_magmoms: bool,
76 has_displacements: bool,
77 has_spreads: bool,
78 metadata: std::collections::BTreeMap<String, serde_json::Value>,
79 serialized: String,
81}
82
83impl MetadataCacheEntry {
84 fn matches(
85 &self,
86 spec_version: u32,
87 has_velocities: bool,
88 has_forces: bool,
89 has_energies: bool,
90 has_charges: bool,
91 has_spins: bool,
92 has_magmoms: bool,
93 has_displacements: bool,
94 has_spreads: bool,
95 metadata: &std::collections::BTreeMap<String, serde_json::Value>,
96 ) -> bool {
97 self.spec_version == spec_version
98 && self.has_velocities == has_velocities
99 && self.has_forces == has_forces
100 && self.has_energies == has_energies
101 && self.has_charges == has_charges
102 && self.has_spins == has_spins
103 && self.has_magmoms == has_magmoms
104 && self.has_displacements == has_displacements
105 && self.has_spreads == has_spreads
106 && &self.metadata == metadata
107 }
108}
109
110impl<W: Write> ConFrameWriter<W> {
112 pub fn new(writer: W) -> Self {
118 Self::with_float_format(writer, FloatFormat::default())
119 }
120
121 pub fn with_precision(writer: W, precision: usize) -> Self {
128 Self::with_float_format(writer, FloatFormat::DecimalPlaces(precision))
129 }
130
131 pub fn with_float_format(writer: W, float_format: FloatFormat) -> Self {
136 Self {
137 writer: BufWriter::new(writer),
138 float_format,
139 canonical: false,
140 metadata_cache: None,
141 scratch: Vec::with_capacity(16 * 1024),
142 }
143 }
144
145 pub fn canonical(mut self, on: bool) -> Self {
150 self.set_canonical(on);
151 self
152 }
153
154 pub fn set_canonical(&mut self, on: bool) {
156 self.canonical = on;
157 if on {
158 self.metadata_cache = None;
159 }
160 }
161
162 pub fn is_canonical(&self) -> bool {
164 self.canonical
165 }
166
167 fn refresh_metadata_cache(&mut self, frame: &ConFrame) {
168 let spec_version = frame.header.spec_version;
169 let has_vel = frame.has_velocities();
170 let has_frc = frame.has_forces();
171 let has_eng = frame.has_energies();
172 let has_chg = frame.has_charges();
173 let has_spn = frame.has_spins();
174 let has_mm = frame.has_magmoms();
175 let has_dsp = frame.has_displacements();
176 let has_spr = frame.has_spreads();
177
178 let cache_hit = !self.canonical
179 && self.metadata_cache.as_ref().is_some_and(|c| {
180 c.matches(
181 spec_version,
182 has_vel,
183 has_frc,
184 has_eng,
185 has_chg,
186 has_spn,
187 has_mm,
188 has_dsp,
189 has_spr,
190 &frame.header.metadata,
191 )
192 });
193 if cache_hit {
194 return;
195 }
196
197 let mut meta_obj = serde_json::Map::new();
198 meta_obj.insert(meta::CON_SPEC_VERSION.into(), json!(spec_version));
199 let mut sections = Vec::new();
200 if has_vel {
201 sections.push(json!(SECTION_VELOCITIES));
202 }
203 if has_frc {
204 sections.push(json!(SECTION_FORCES));
205 }
206 if has_eng {
207 sections.push(json!(SECTION_ENERGIES));
208 }
209 if has_chg {
210 sections.push(json!(SECTION_CHARGES));
211 }
212 if has_spn {
213 sections.push(json!(SECTION_SPINS));
214 }
215 if has_mm {
216 sections.push(json!(SECTION_MAGMOMS));
217 }
218 if has_dsp {
219 sections.push(json!(SECTION_DISPLACEMENTS));
220 }
221 if has_spr {
222 sections.push(json!(SECTION_SPREADS));
223 }
224 let validate = frame
225 .header
226 .metadata
227 .get(meta::VALIDATE)
228 .and_then(|value| value.as_bool())
229 .unwrap_or(false);
230 if !sections.is_empty() || validate {
231 meta_obj.insert(meta::SECTIONS.into(), json!(sections));
232 }
233 for (k, v) in &frame.header.metadata {
234 if k == meta::CON_SPEC_VERSION || k == meta::SECTIONS {
235 continue;
236 }
237 meta_obj.insert(k.clone(), v.clone());
238 }
239 if let Some(u) = meta_obj.get(meta::UNITS).cloned() {
240 if let Ok(c) = crate::units::canonicalize_units_object(&u) {
241 meta_obj.insert(meta::UNITS.into(), c);
242 }
243 }
244 if spec_version >= 3 {
245 let need_default = match meta_obj.get(meta::UNITS) {
246 None => true,
247 Some(u) => crate::units::validate_v3_units_metadata(u).is_err(),
248 };
249 if need_default {
250 meta_obj.insert(meta::UNITS.into(), crate::units::default_v3_units_json());
251 }
252 }
253 let serialized = serde_json::Value::Object(meta_obj).to_string();
254 self.metadata_cache = Some(MetadataCacheEntry {
255 spec_version,
256 has_velocities: has_vel,
257 has_forces: has_frc,
258 has_energies: has_eng,
259 has_charges: has_chg,
260 has_spins: has_spn,
261 has_magmoms: has_mm,
262 has_displacements: has_dsp,
263 has_spreads: has_spr,
264 metadata: frame.header.metadata.clone(),
265 serialized,
266 });
267 }
268
269 pub fn write_frame(&mut self, frame: &ConFrame) -> io::Result<()> {
271 let prec = self.float_format;
272 self.refresh_metadata_cache(frame);
273 let meta_line = self
274 .metadata_cache
275 .as_ref()
276 .expect("metadata_cache populated above")
277 .serialized
278 .clone();
279 self.scratch.clear();
280 {
281 let buf = &mut self.scratch;
282
283 let _ = writeln!(buf, "{}", frame.header.prebox_header.user);
285 let _ = writeln!(buf, "{meta_line}");
286 push_f64(buf, frame.header.boxl[0], prec);
287 buf.push(b' ');
288 push_f64(buf, frame.header.boxl[1], prec);
289 buf.push(b' ');
290 push_f64(buf, frame.header.boxl[2], prec);
291 buf.push(b'\n');
292 push_f64(buf, frame.header.angles[0], prec);
293 buf.push(b' ');
294 push_f64(buf, frame.header.angles[1], prec);
295 buf.push(b' ');
296 push_f64(buf, frame.header.angles[2], prec);
297 buf.push(b'\n');
298 let _ = writeln!(buf, "{}", frame.header.postbox_header[0]);
299 let _ = writeln!(buf, "{}", frame.header.postbox_header[1]);
300 let _ = writeln!(buf, "{}", frame.header.natm_types);
301
302 for (i, n) in frame.header.natms_per_type.iter().enumerate() {
303 if i > 0 {
304 buf.push(b' ');
305 }
306 push_u64(buf, *n as u64);
307 }
308 buf.push(b'\n');
309
310 for (i, m) in frame.header.masses_per_type.iter().enumerate() {
311 if i > 0 {
312 buf.push(b' ');
313 }
314 push_f64(buf, *m, prec);
315 }
316 buf.push(b'\n');
317
318 let mut atom_idx_offset = 0;
320 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate() {
321 let symbol = &frame.atom_data[atom_idx_offset].symbol;
322 let _ = writeln!(buf, "{symbol}");
323 let _ = writeln!(buf, "Coordinates of Component {}", type_idx + 1);
324
325 for i in 0..num_atoms_in_type {
326 let atom = &frame.atom_data[atom_idx_offset + i];
327 push_xyz_line(
328 buf,
329 atom.x,
330 atom.y,
331 atom.z,
332 prec,
333 encode_fixed_bitmask(atom.fixed),
334 atom.atom_id,
335 );
336 }
337 atom_idx_offset += num_atoms_in_type;
338 }
339
340 if frame.has_velocities() {
342 buf.push(b'\n');
343
344 let mut vel_idx_offset = 0;
345 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
346 {
347 let symbol = &frame.atom_data[vel_idx_offset].symbol;
348 let _ = writeln!(buf, "{symbol}");
349 let _ = writeln!(buf, "Velocities of Component {}", type_idx + 1);
350
351 for i in 0..num_atoms_in_type {
352 let atom = &frame.atom_data[vel_idx_offset + i];
353 let [vx, vy, vz] = atom.velocity.unwrap_or([0.0; 3]);
354 push_xyz_line(
355 buf,
356 vx,
357 vy,
358 vz,
359 prec,
360 encode_fixed_bitmask(atom.fixed),
361 atom.atom_id,
362 );
363 }
364 vel_idx_offset += num_atoms_in_type;
365 }
366 }
367
368 if frame.has_forces() {
370 buf.push(b'\n');
371
372 let mut force_idx_offset = 0;
373 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
374 {
375 let symbol = &frame.atom_data[force_idx_offset].symbol;
376 let _ = writeln!(buf, "{symbol}");
377 let _ = writeln!(buf, "Forces of Component {}", type_idx + 1);
378
379 for i in 0..num_atoms_in_type {
380 let atom = &frame.atom_data[force_idx_offset + i];
381 let [fx, fy, fz] = atom.force.unwrap_or([0.0; 3]);
382 push_xyz_line(
383 buf,
384 fx,
385 fy,
386 fz,
387 prec,
388 encode_fixed_bitmask(atom.fixed),
389 atom.atom_id,
390 );
391 }
392 force_idx_offset += num_atoms_in_type;
393 }
394 }
395
396 if frame.has_energies() {
398 buf.push(b'\n');
399
400 let mut energy_idx_offset = 0;
401 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
402 {
403 let symbol = &frame.atom_data[energy_idx_offset].symbol;
404 let _ = writeln!(buf, "{symbol}");
405 let _ = writeln!(buf, "Energies of Component {}", type_idx + 1);
406
407 for i in 0..num_atoms_in_type {
408 let atom = &frame.atom_data[energy_idx_offset + i];
409 let e = atom.energy.unwrap_or(0.0);
410 push_scalar_line(
411 buf,
412 e,
413 prec,
414 encode_fixed_bitmask(atom.fixed),
415 atom.atom_id,
416 );
417 }
418 energy_idx_offset += num_atoms_in_type;
419 }
420 }
421
422 if frame.has_charges() {
423 buf.push(b'\n');
424 let mut off = 0;
425 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
426 {
427 let symbol = &frame.atom_data[off].symbol;
428 let _ = writeln!(buf, "{symbol}");
429 let _ = writeln!(buf, "Charges of Component {}", type_idx + 1);
430 for i in 0..num_atoms_in_type {
431 let atom = &frame.atom_data[off + i];
432 let q = atom.charge.unwrap_or(0.0);
433 push_scalar_line(
434 buf,
435 q,
436 prec,
437 encode_fixed_bitmask(atom.fixed),
438 atom.atom_id,
439 );
440 }
441 off += num_atoms_in_type;
442 }
443 }
444
445 if frame.has_spins() {
446 buf.push(b'\n');
447 let mut off = 0;
448 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
449 {
450 let symbol = &frame.atom_data[off].symbol;
451 let _ = writeln!(buf, "{symbol}");
452 let _ = writeln!(buf, "Spins of Component {}", type_idx + 1);
453 for i in 0..num_atoms_in_type {
454 let atom = &frame.atom_data[off + i];
455 let s = atom.spin.unwrap_or(0.0);
456 push_scalar_line(
457 buf,
458 s,
459 prec,
460 encode_fixed_bitmask(atom.fixed),
461 atom.atom_id,
462 );
463 }
464 off += num_atoms_in_type;
465 }
466 }
467
468 if frame.has_magmoms() {
469 buf.push(b'\n');
470 let mut off = 0;
471 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
472 {
473 let symbol = &frame.atom_data[off].symbol;
474 let _ = writeln!(buf, "{symbol}");
475 let _ = writeln!(buf, "Magmoms of Component {}", type_idx + 1);
476 for i in 0..num_atoms_in_type {
477 let atom = &frame.atom_data[off + i];
478 let [mx, my, mz] = atom.magmom.unwrap_or([0.0; 3]);
479 push_xyz_line(
480 buf,
481 mx,
482 my,
483 mz,
484 prec,
485 encode_fixed_bitmask(atom.fixed),
486 atom.atom_id,
487 );
488 }
489 off += num_atoms_in_type;
490 }
491 }
492
493 if frame.has_displacements() {
494 buf.push(b'\n');
495 let mut off = 0;
496 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
497 {
498 let symbol = &frame.atom_data[off].symbol;
499 let _ = writeln!(buf, "{symbol}");
500 let _ = writeln!(buf, "Displacements of Component {}", type_idx + 1);
501 for i in 0..num_atoms_in_type {
502 let atom = &frame.atom_data[off + i];
503 let [dx, dy, dz] = atom.displacement.unwrap_or([0.0; 3]);
504 push_xyz_line(
505 buf,
506 dx,
507 dy,
508 dz,
509 prec,
510 encode_fixed_bitmask(atom.fixed),
511 atom.atom_id,
512 );
513 }
514 off += num_atoms_in_type;
515 }
516 }
517
518 if frame.has_spreads() {
519 buf.push(b'\n');
520 let mut off = 0;
521 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
522 {
523 let symbol = &frame.atom_data[off].symbol;
524 let _ = writeln!(buf, "{symbol}");
525 let _ = writeln!(buf, "Spreads of Component {}", type_idx + 1);
526 for i in 0..num_atoms_in_type {
527 let atom = &frame.atom_data[off + i];
528 let [sx, sy, sz] = atom.spread.unwrap_or([0.0; 3]);
529 if !crate::types::spread_row_ok([sx, sy, sz]) {
530 return Err(std::io::Error::new(
531 std::io::ErrorKind::InvalidInput,
532 format!(
533 "spreads: atom {} has a negative or non-finite spread",
534 off + i
535 ),
536 ));
537 }
538 push_xyz_line(
539 buf,
540 sx,
541 sy,
542 sz,
543 prec,
544 encode_fixed_bitmask(atom.fixed),
545 atom.atom_id,
546 );
547 }
548 off += num_atoms_in_type;
549 }
550 }
551 }
552
553 self.writer.write_all(&self.scratch)
554 }
555
556 pub fn extend<'a>(&mut self, frames: impl Iterator<Item = &'a ConFrame>) -> io::Result<()> {
560 for frame in frames {
561 self.write_frame(frame)?;
562 }
563 Ok(())
564 }
565}
566
567fn push_u64(buf: &mut Vec<u8>, n: u64) {
568 push_u128(buf, u128::from(n));
569}
570
571fn push_u128(buf: &mut Vec<u8>, mut n: u128) {
572 let mut tmp = [0u8; 40];
573 let mut i = 40;
574 if n == 0 {
575 buf.push(b'0');
576 return;
577 }
578 while n > 0 {
579 i -= 1;
580 tmp[i] = b'0' + (n % 10) as u8;
581 n /= 10;
582 }
583 buf.extend_from_slice(&tmp[i..]);
584}
585
586fn push_f64(buf: &mut Vec<u8>, value: f64, format: FloatFormat) {
587 match format {
588 FloatFormat::DecimalPlaces(precision) => push_f64_prec(buf, value, precision),
589 FloatFormat::RoundTrip => {
590 let _ = write!(buf, "{value:e}");
591 }
592 }
593}
594
595fn push_f64_prec(buf: &mut Vec<u8>, v: f64, prec: usize) {
598 if !v.is_finite() || prec != 6 {
601 let _ = write!(buf, "{v:.prec$}");
602 return;
603 }
604 if v.is_sign_negative() {
605 buf.push(b'-');
606 }
607 let ax = v.abs();
608 if prec == 0 {
609 push_u128(buf, ax.round() as u128);
610 return;
611 }
612 let scale = 10u128.pow(prec as u32);
613 let n = (ax * scale as f64).round() as u128;
614 let int_part = n / scale;
615 let frac = n % scale;
616 push_u128(buf, int_part);
617 buf.push(b'.');
618 let mut tmp = [b'0'; 20];
619 let mut x = frac;
620 let mut i = prec;
621 while i > 0 {
622 i -= 1;
623 tmp[i] = b'0' + (x % 10) as u8;
624 x /= 10;
625 }
626 buf.extend_from_slice(&tmp[..prec]);
627}
628
629fn push_xyz_line(
630 buf: &mut Vec<u8>,
631 x: f64,
632 y: f64,
633 z: f64,
634 prec: FloatFormat,
635 fixed: u8,
636 atom_id: u64,
637) {
638 push_f64(buf, x, prec);
639 buf.push(b' ');
640 push_f64(buf, y, prec);
641 buf.push(b' ');
642 push_f64(buf, z, prec);
643 buf.push(b' ');
644 push_u64(buf, u64::from(fixed));
645 buf.push(b' ');
646 push_u64(buf, atom_id);
647 buf.push(b'\n');
648}
649
650fn push_scalar_line(buf: &mut Vec<u8>, v: f64, prec: FloatFormat, fixed: u8, atom_id: u64) {
651 push_f64(buf, v, prec);
652 buf.push(b' ');
653 push_u64(buf, u64::from(fixed));
654 buf.push(b' ');
655 push_u64(buf, atom_id);
656 buf.push(b'\n');
657}
658
659impl ConFrameWriter<File> {
661 pub fn from_path<P: AsRef<Path>>(path: P) -> io::Result<Self> {
665 let file = File::create(path)?;
666 Ok(Self::new(file))
667 }
668
669 pub fn from_path_with_precision<P: AsRef<Path>>(path: P, precision: usize) -> io::Result<Self> {
671 let file = File::create(path)?;
672 Ok(Self::with_precision(file, precision))
673 }
674}
675
676impl ConFrameWriter<flate2::write::GzEncoder<File>> {
678 pub fn from_path_gzip<P: AsRef<Path>>(path: P) -> io::Result<Self> {
680 let encoder = crate::compression::gzip_writer(path.as_ref())?;
681 Ok(Self::new(encoder))
682 }
683
684 pub fn from_path_gzip_with_precision<P: AsRef<Path>>(
686 path: P,
687 precision: usize,
688 ) -> io::Result<Self> {
689 let encoder = crate::compression::gzip_writer(path.as_ref())?;
690 Ok(Self::with_precision(encoder, precision))
691 }
692}
693
694#[cfg(feature = "zstd")]
697impl ConFrameWriter<zstd::stream::write::AutoFinishEncoder<'static, File>> {
698 pub fn from_path_zstd<P: AsRef<Path>>(path: P) -> io::Result<Self> {
700 let encoder = crate::compression::zstd_writer(path.as_ref())?;
701 Ok(Self::new(encoder))
702 }
703
704 pub fn from_path_zstd_with_precision<P: AsRef<Path>>(
706 path: P,
707 precision: usize,
708 ) -> io::Result<Self> {
709 let encoder = crate::compression::zstd_writer(path.as_ref())?;
710 Ok(Self::with_precision(encoder, precision))
711 }
712}
713
714#[cfg(test)]
715mod float_format_tests {
716 use super::push_f64_prec;
717
718 fn formatted(v: f64, prec: usize) -> String {
719 let mut buf = Vec::new();
720 push_f64_prec(&mut buf, v, prec);
721 String::from_utf8(buf).expect("utf8")
722 }
723
724 #[test]
725 fn matches_std_fixed_precision() {
726 let vals = [
727 0.0,
728 -0.0,
729 1.0,
730 -1.0,
731 0.9045,
732 6.975_299_999_999_995,
733 63.546,
734 1.008,
735 15.3456,
736 90.0,
737 218.0,
738 1e-6,
739 -1.23456789,
740 10.0,
741 ];
742 for prec in [0usize, 6] {
743 for v in vals {
744 let got = formatted(v, prec);
745 let exp = format!("{v:.prec$}");
746 assert_eq!(got, exp, "v={v:?} prec={prec}");
747 }
748 }
749 }
750}