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.len() == metadata.len()
107 && self
108 .metadata
109 .iter()
110 .zip(metadata)
111 .all(|((ka, a), (kb, b))| ka == kb && metadata_value_eq(a, b))
112 }
113}
114
115fn metadata_value_eq(a: &serde_json::Value, b: &serde_json::Value) -> bool {
118 use serde_json::Value;
119 match (a, b) {
120 (Value::Number(a), Value::Number(b)) if a.is_f64() && b.is_f64() => {
121 a.as_f64().map(f64::to_bits) == b.as_f64().map(f64::to_bits)
122 }
123 (Value::Array(a), Value::Array(b)) => {
124 a.len() == b.len() && a.iter().zip(b).all(|(a, b)| metadata_value_eq(a, b))
125 }
126 (Value::Object(a), Value::Object(b)) => {
127 a.len() == b.len()
128 && a.iter()
129 .all(|(key, a)| b.get(key).is_some_and(|b| metadata_value_eq(a, b)))
130 }
131 _ => a == b,
132 }
133}
134
135impl<W: Write> ConFrameWriter<W> {
137 pub fn new(writer: W) -> Self {
143 Self::with_float_format(writer, FloatFormat::default())
144 }
145
146 pub fn with_precision(writer: W, precision: usize) -> Self {
153 Self::with_float_format(writer, FloatFormat::DecimalPlaces(precision))
154 }
155
156 pub fn with_float_format(writer: W, float_format: FloatFormat) -> Self {
161 Self {
162 writer: BufWriter::new(writer),
163 float_format,
164 canonical: false,
165 metadata_cache: None,
166 scratch: Vec::with_capacity(16 * 1024),
167 }
168 }
169
170 pub fn canonical(mut self, on: bool) -> Self {
175 self.set_canonical(on);
176 self
177 }
178
179 pub fn set_canonical(&mut self, on: bool) {
181 self.canonical = on;
182 if on {
183 self.metadata_cache = None;
184 }
185 }
186
187 pub fn is_canonical(&self) -> bool {
189 self.canonical
190 }
191
192 fn refresh_metadata_cache(&mut self, frame: &ConFrame) {
193 let spec_version = frame.header.spec_version;
194 let has_vel = frame.has_velocities();
195 let has_frc = frame.has_forces();
196 let has_eng = frame.has_energies();
197 let has_chg = frame.has_charges();
198 let has_spn = frame.has_spins();
199 let has_mm = frame.has_magmoms();
200 let has_dsp = frame.has_displacements();
201 let has_spr = frame.has_spreads();
202
203 let cache_hit = !self.canonical
204 && self.metadata_cache.as_ref().is_some_and(|c| {
205 c.matches(
206 spec_version,
207 has_vel,
208 has_frc,
209 has_eng,
210 has_chg,
211 has_spn,
212 has_mm,
213 has_dsp,
214 has_spr,
215 &frame.header.metadata,
216 )
217 });
218 if cache_hit {
219 return;
220 }
221
222 let mut meta_obj = serde_json::Map::new();
223 meta_obj.insert(meta::CON_SPEC_VERSION.into(), json!(spec_version));
224 let mut sections = Vec::new();
225 if has_vel {
226 sections.push(json!(SECTION_VELOCITIES));
227 }
228 if has_frc {
229 sections.push(json!(SECTION_FORCES));
230 }
231 if has_eng {
232 sections.push(json!(SECTION_ENERGIES));
233 }
234 if has_chg {
235 sections.push(json!(SECTION_CHARGES));
236 }
237 if has_spn {
238 sections.push(json!(SECTION_SPINS));
239 }
240 if has_mm {
241 sections.push(json!(SECTION_MAGMOMS));
242 }
243 if has_dsp {
244 sections.push(json!(SECTION_DISPLACEMENTS));
245 }
246 if has_spr {
247 sections.push(json!(SECTION_SPREADS));
248 }
249 let validate = frame
250 .header
251 .metadata
252 .get(meta::VALIDATE)
253 .and_then(|value| value.as_bool())
254 .unwrap_or(false);
255 if !sections.is_empty() || validate {
256 meta_obj.insert(meta::SECTIONS.into(), json!(sections));
257 }
258 for (k, v) in &frame.header.metadata {
259 if k == meta::CON_SPEC_VERSION || k == meta::SECTIONS {
260 continue;
261 }
262 meta_obj.insert(k.clone(), v.clone());
263 }
264 if let Some(u) = meta_obj.get(meta::UNITS).cloned() {
265 if let Ok(c) = crate::units::canonicalize_units_object(&u) {
266 meta_obj.insert(meta::UNITS.into(), c);
267 }
268 }
269 if spec_version >= 3 {
270 let need_default = match meta_obj.get(meta::UNITS) {
271 None => true,
272 Some(u) => crate::units::validate_v3_units_metadata(u).is_err(),
273 };
274 if need_default {
275 meta_obj.insert(meta::UNITS.into(), crate::units::default_v3_units_json());
276 }
277 }
278 let serialized = serde_json::Value::Object(meta_obj).to_string();
279 self.metadata_cache = Some(MetadataCacheEntry {
280 spec_version,
281 has_velocities: has_vel,
282 has_forces: has_frc,
283 has_energies: has_eng,
284 has_charges: has_chg,
285 has_spins: has_spn,
286 has_magmoms: has_mm,
287 has_displacements: has_dsp,
288 has_spreads: has_spr,
289 metadata: frame.header.metadata.clone(),
290 serialized,
291 });
292 }
293
294 pub fn write_frame(&mut self, frame: &ConFrame) -> io::Result<()> {
296 let prec = self.float_format;
297 self.refresh_metadata_cache(frame);
298 let meta_line = self
299 .metadata_cache
300 .as_ref()
301 .expect("metadata_cache populated above")
302 .serialized
303 .clone();
304 self.scratch.clear();
305 {
306 let buf = &mut self.scratch;
307
308 let _ = writeln!(buf, "{}", frame.header.prebox_header.user);
310 let _ = writeln!(buf, "{meta_line}");
311 push_f64(buf, frame.header.boxl[0], prec);
312 buf.push(b' ');
313 push_f64(buf, frame.header.boxl[1], prec);
314 buf.push(b' ');
315 push_f64(buf, frame.header.boxl[2], prec);
316 buf.push(b'\n');
317 push_f64(buf, frame.header.angles[0], prec);
318 buf.push(b' ');
319 push_f64(buf, frame.header.angles[1], prec);
320 buf.push(b' ');
321 push_f64(buf, frame.header.angles[2], prec);
322 buf.push(b'\n');
323 let _ = writeln!(buf, "{}", frame.header.postbox_header[0]);
324 let _ = writeln!(buf, "{}", frame.header.postbox_header[1]);
325 let _ = writeln!(buf, "{}", frame.header.natm_types);
326
327 for (i, n) in frame.header.natms_per_type.iter().enumerate() {
328 if i > 0 {
329 buf.push(b' ');
330 }
331 push_u64(buf, *n as u64);
332 }
333 buf.push(b'\n');
334
335 for (i, m) in frame.header.masses_per_type.iter().enumerate() {
336 if i > 0 {
337 buf.push(b' ');
338 }
339 push_f64(buf, *m, prec);
340 }
341 buf.push(b'\n');
342
343 let mut atom_idx_offset = 0;
345 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate() {
346 let symbol = &frame.atom_data[atom_idx_offset].symbol;
347 let _ = writeln!(buf, "{symbol}");
348 let _ = writeln!(buf, "Coordinates of Component {}", type_idx + 1);
349
350 for i in 0..num_atoms_in_type {
351 let atom = &frame.atom_data[atom_idx_offset + i];
352 push_xyz_line(
353 buf,
354 atom.x,
355 atom.y,
356 atom.z,
357 prec,
358 encode_fixed_bitmask(atom.fixed),
359 atom.atom_id,
360 );
361 }
362 atom_idx_offset += num_atoms_in_type;
363 }
364
365 if frame.has_velocities() {
367 buf.push(b'\n');
368
369 let mut vel_idx_offset = 0;
370 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
371 {
372 let symbol = &frame.atom_data[vel_idx_offset].symbol;
373 let _ = writeln!(buf, "{symbol}");
374 let _ = writeln!(buf, "Velocities of Component {}", type_idx + 1);
375
376 for i in 0..num_atoms_in_type {
377 let atom = &frame.atom_data[vel_idx_offset + i];
378 let [vx, vy, vz] = atom.velocity.unwrap_or([0.0; 3]);
379 push_xyz_line(
380 buf,
381 vx,
382 vy,
383 vz,
384 prec,
385 encode_fixed_bitmask(atom.fixed),
386 atom.atom_id,
387 );
388 }
389 vel_idx_offset += num_atoms_in_type;
390 }
391 }
392
393 if frame.has_forces() {
395 buf.push(b'\n');
396
397 let mut force_idx_offset = 0;
398 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
399 {
400 let symbol = &frame.atom_data[force_idx_offset].symbol;
401 let _ = writeln!(buf, "{symbol}");
402 let _ = writeln!(buf, "Forces of Component {}", type_idx + 1);
403
404 for i in 0..num_atoms_in_type {
405 let atom = &frame.atom_data[force_idx_offset + i];
406 let [fx, fy, fz] = atom.force.unwrap_or([0.0; 3]);
407 push_xyz_line(
408 buf,
409 fx,
410 fy,
411 fz,
412 prec,
413 encode_fixed_bitmask(atom.fixed),
414 atom.atom_id,
415 );
416 }
417 force_idx_offset += num_atoms_in_type;
418 }
419 }
420
421 if frame.has_energies() {
423 buf.push(b'\n');
424
425 let mut energy_idx_offset = 0;
426 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
427 {
428 let symbol = &frame.atom_data[energy_idx_offset].symbol;
429 let _ = writeln!(buf, "{symbol}");
430 let _ = writeln!(buf, "Energies of Component {}", type_idx + 1);
431
432 for i in 0..num_atoms_in_type {
433 let atom = &frame.atom_data[energy_idx_offset + i];
434 let e = atom.energy.unwrap_or(0.0);
435 push_scalar_line(
436 buf,
437 e,
438 prec,
439 encode_fixed_bitmask(atom.fixed),
440 atom.atom_id,
441 );
442 }
443 energy_idx_offset += num_atoms_in_type;
444 }
445 }
446
447 if frame.has_charges() {
448 buf.push(b'\n');
449 let mut off = 0;
450 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
451 {
452 let symbol = &frame.atom_data[off].symbol;
453 let _ = writeln!(buf, "{symbol}");
454 let _ = writeln!(buf, "Charges of Component {}", type_idx + 1);
455 for i in 0..num_atoms_in_type {
456 let atom = &frame.atom_data[off + i];
457 let q = atom.charge.unwrap_or(0.0);
458 push_scalar_line(
459 buf,
460 q,
461 prec,
462 encode_fixed_bitmask(atom.fixed),
463 atom.atom_id,
464 );
465 }
466 off += num_atoms_in_type;
467 }
468 }
469
470 if frame.has_spins() {
471 buf.push(b'\n');
472 let mut off = 0;
473 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
474 {
475 let symbol = &frame.atom_data[off].symbol;
476 let _ = writeln!(buf, "{symbol}");
477 let _ = writeln!(buf, "Spins of Component {}", type_idx + 1);
478 for i in 0..num_atoms_in_type {
479 let atom = &frame.atom_data[off + i];
480 let s = atom.spin.unwrap_or(0.0);
481 push_scalar_line(
482 buf,
483 s,
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_magmoms() {
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, "Magmoms of Component {}", type_idx + 1);
501 for i in 0..num_atoms_in_type {
502 let atom = &frame.atom_data[off + i];
503 let [mx, my, mz] = atom.magmom.unwrap_or([0.0; 3]);
504 push_xyz_line(
505 buf,
506 mx,
507 my,
508 mz,
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_displacements() {
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, "Displacements of Component {}", type_idx + 1);
526 for i in 0..num_atoms_in_type {
527 let atom = &frame.atom_data[off + i];
528 let [dx, dy, dz] = atom.displacement.unwrap_or([0.0; 3]);
529 push_xyz_line(
530 buf,
531 dx,
532 dy,
533 dz,
534 prec,
535 encode_fixed_bitmask(atom.fixed),
536 atom.atom_id,
537 );
538 }
539 off += num_atoms_in_type;
540 }
541 }
542
543 if frame.has_spreads() {
544 buf.push(b'\n');
545 let mut off = 0;
546 for (type_idx, &num_atoms_in_type) in frame.header.natms_per_type.iter().enumerate()
547 {
548 let symbol = &frame.atom_data[off].symbol;
549 let _ = writeln!(buf, "{symbol}");
550 let _ = writeln!(buf, "Spreads of Component {}", type_idx + 1);
551 for i in 0..num_atoms_in_type {
552 let atom = &frame.atom_data[off + i];
553 let [sx, sy, sz] = atom.spread.unwrap_or([0.0; 3]);
554 if !crate::types::spread_row_ok([sx, sy, sz]) {
555 return Err(std::io::Error::new(
556 std::io::ErrorKind::InvalidInput,
557 format!(
558 "spreads: atom {} has a negative or non-finite spread",
559 off + i
560 ),
561 ));
562 }
563 push_xyz_line(
564 buf,
565 sx,
566 sy,
567 sz,
568 prec,
569 encode_fixed_bitmask(atom.fixed),
570 atom.atom_id,
571 );
572 }
573 off += num_atoms_in_type;
574 }
575 }
576 }
577
578 self.writer.write_all(&self.scratch)
579 }
580
581 pub fn flush(&mut self) -> io::Result<()> {
583 self.writer.flush()
584 }
585
586 pub fn extend<'a>(&mut self, frames: impl Iterator<Item = &'a ConFrame>) -> io::Result<()> {
590 for frame in frames {
591 self.write_frame(frame)?;
592 }
593 Ok(())
594 }
595}
596
597fn push_u64(buf: &mut Vec<u8>, n: u64) {
598 push_u128(buf, u128::from(n));
599}
600
601fn push_u128(buf: &mut Vec<u8>, mut n: u128) {
602 let mut tmp = [0u8; 40];
603 let mut i = 40;
604 if n == 0 {
605 buf.push(b'0');
606 return;
607 }
608 while n > 0 {
609 i -= 1;
610 tmp[i] = b'0' + (n % 10) as u8;
611 n /= 10;
612 }
613 buf.extend_from_slice(&tmp[i..]);
614}
615
616fn push_f64(buf: &mut Vec<u8>, value: f64, format: FloatFormat) {
617 match format {
618 FloatFormat::DecimalPlaces(precision) => push_f64_prec(buf, value, precision),
619 FloatFormat::RoundTrip => {
620 let _ = write!(buf, "{value:e}");
621 }
622 }
623}
624
625fn push_f64_prec(buf: &mut Vec<u8>, v: f64, prec: usize) {
628 if !v.is_finite() || prec != 6 {
631 let _ = write!(buf, "{v:.prec$}");
632 return;
633 }
634 if v.is_sign_negative() {
635 buf.push(b'-');
636 }
637 let ax = v.abs();
638 if prec == 0 {
639 push_u128(buf, ax.round() as u128);
640 return;
641 }
642 let scale = 10u128.pow(prec as u32);
643 let n = (ax * scale as f64).round() as u128;
644 let int_part = n / scale;
645 let frac = n % scale;
646 push_u128(buf, int_part);
647 buf.push(b'.');
648 let mut tmp = [b'0'; 20];
649 let mut x = frac;
650 let mut i = prec;
651 while i > 0 {
652 i -= 1;
653 tmp[i] = b'0' + (x % 10) as u8;
654 x /= 10;
655 }
656 buf.extend_from_slice(&tmp[..prec]);
657}
658
659fn push_xyz_line(
660 buf: &mut Vec<u8>,
661 x: f64,
662 y: f64,
663 z: f64,
664 prec: FloatFormat,
665 fixed: u8,
666 atom_id: u64,
667) {
668 push_f64(buf, x, prec);
669 buf.push(b' ');
670 push_f64(buf, y, prec);
671 buf.push(b' ');
672 push_f64(buf, z, prec);
673 buf.push(b' ');
674 push_u64(buf, u64::from(fixed));
675 buf.push(b' ');
676 push_u64(buf, atom_id);
677 buf.push(b'\n');
678}
679
680fn push_scalar_line(buf: &mut Vec<u8>, v: f64, prec: FloatFormat, fixed: u8, atom_id: u64) {
681 push_f64(buf, v, prec);
682 buf.push(b' ');
683 push_u64(buf, u64::from(fixed));
684 buf.push(b' ');
685 push_u64(buf, atom_id);
686 buf.push(b'\n');
687}
688
689impl ConFrameWriter<File> {
691 pub fn from_path<P: AsRef<Path>>(path: P) -> io::Result<Self> {
695 let file = File::create(path)?;
696 Ok(Self::new(file))
697 }
698
699 pub fn from_path_with_precision<P: AsRef<Path>>(path: P, precision: usize) -> io::Result<Self> {
701 let file = File::create(path)?;
702 Ok(Self::with_precision(file, precision))
703 }
704}
705
706impl ConFrameWriter<flate2::write::GzEncoder<File>> {
708 pub fn from_path_gzip<P: AsRef<Path>>(path: P) -> io::Result<Self> {
710 let encoder = crate::compression::gzip_writer(path.as_ref())?;
711 Ok(Self::new(encoder))
712 }
713
714 pub fn from_path_gzip_with_precision<P: AsRef<Path>>(
716 path: P,
717 precision: usize,
718 ) -> io::Result<Self> {
719 let encoder = crate::compression::gzip_writer(path.as_ref())?;
720 Ok(Self::with_precision(encoder, precision))
721 }
722}
723
724#[cfg(feature = "zstd")]
727impl ConFrameWriter<zstd::stream::write::AutoFinishEncoder<'static, File>> {
728 pub fn from_path_zstd<P: AsRef<Path>>(path: P) -> io::Result<Self> {
730 let encoder = crate::compression::zstd_writer(path.as_ref())?;
731 Ok(Self::new(encoder))
732 }
733
734 pub fn from_path_zstd_with_precision<P: AsRef<Path>>(
736 path: P,
737 precision: usize,
738 ) -> io::Result<Self> {
739 let encoder = crate::compression::zstd_writer(path.as_ref())?;
740 Ok(Self::with_precision(encoder, precision))
741 }
742}
743
744#[cfg(test)]
745mod float_format_tests {
746 use super::push_f64_prec;
747
748 fn formatted(v: f64, prec: usize) -> String {
749 let mut buf = Vec::new();
750 push_f64_prec(&mut buf, v, prec);
751 String::from_utf8(buf).expect("utf8")
752 }
753
754 #[test]
755 fn matches_std_fixed_precision() {
756 let vals = [
757 0.0,
758 -0.0,
759 1.0,
760 -1.0,
761 0.9045,
762 6.975_299_999_999_995,
763 63.546,
764 1.008,
765 15.3456,
766 90.0,
767 218.0,
768 1e-6,
769 -1.23456789,
770 10.0,
771 ];
772 for prec in [0usize, 6] {
773 for v in vals {
774 let got = formatted(v, prec);
775 let exp = format!("{v:.prec$}");
776 assert_eq!(got, exp, "v={v:?} prec={prec}");
777 }
778 }
779 }
780}