use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use serde_json::{Map as JsonMap, Value as JsonValue};
use crate::MolRsError;
use crate::store::block::Column;
use crate::store::frame::Frame;
use crate::types::F;
pub type SchemaValue = JsonValue;
#[derive(Debug, Clone, Default)]
pub struct Trajectory {
pub frames: Vec<Frame>,
pub step: Option<Vec<i64>>,
pub time: Option<Vec<F>>,
}
impl Trajectory {
pub fn new() -> Self {
Self::default()
}
pub fn from_frames(frames: Vec<Frame>) -> Self {
Self {
frames,
step: None,
time: None,
}
}
pub fn len(&self) -> usize {
self.frames.len()
}
pub fn is_empty(&self) -> bool {
self.frames.is_empty()
}
pub fn validate(&self) -> Result<(), MolRsError> {
let n = self.frames.len();
if let Some(step) = &self.step
&& step.len() != n
{
return Err(MolRsError::validation(format!(
"trajectory.step length mismatch: expected {}, got {}",
n,
step.len()
)));
}
if let Some(time) = &self.time
&& time.len() != n
{
return Err(MolRsError::validation(format!(
"trajectory.time length mismatch: expected {}, got {}",
n,
time.len()
)));
}
Ok(())
}
}
#[derive(Debug, Clone, Copy, Serialize, Deserialize, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum ObservableKind {
Scalar,
Vector,
}
#[derive(Debug, Clone)]
pub enum ObservableData {
Column(Column),
}
#[derive(Debug, Clone)]
pub struct ObservableRecord {
pub name: String,
pub kind: ObservableKind,
pub description: String,
pub time_dependent: bool,
pub unit: Option<String>,
pub axes: Vec<String>,
pub sampling: Option<String>,
pub domain: Option<String>,
pub target: Option<String>,
pub extra: JsonMap<String, JsonValue>,
pub data: ObservableData,
}
impl ObservableRecord {
pub fn scalar(name: impl Into<String>, data: Column) -> Self {
Self {
name: name.into(),
kind: ObservableKind::Scalar,
description: String::new(),
time_dependent: false,
unit: None,
axes: Vec::new(),
sampling: None,
domain: None,
target: None,
extra: JsonMap::new(),
data: ObservableData::Column(data),
}
}
pub fn vector(name: impl Into<String>, data: Column) -> Self {
Self {
name: name.into(),
kind: ObservableKind::Vector,
description: String::new(),
time_dependent: false,
unit: None,
axes: Vec::new(),
sampling: None,
domain: None,
target: None,
extra: JsonMap::new(),
data: ObservableData::Column(data),
}
}
pub fn validate(&self) -> Result<(), MolRsError> {
match (&self.kind, &self.data) {
(ObservableKind::Scalar | ObservableKind::Vector, ObservableData::Column(_)) => Ok(()),
}
}
}
#[derive(Debug, Clone)]
pub struct MolRec {
pub meta: SchemaValue,
pub frame: Frame,
pub trajectory: Option<Trajectory>,
pub observables: BTreeMap<String, ObservableRecord>,
pub method: SchemaValue,
pub parameters: SchemaValue,
}
impl Default for MolRec {
fn default() -> Self {
Self::new(Frame::new())
}
}
impl MolRec {
pub fn new(frame: Frame) -> Self {
Self {
meta: empty_object(),
frame,
trajectory: None,
observables: BTreeMap::new(),
method: empty_object(),
parameters: empty_object(),
}
}
pub fn from_frames(frame: Frame, frames: Vec<Frame>) -> Self {
let mut rec = Self::new(frame);
let trajectory = Trajectory::from_frames(frames);
rec.trajectory = Some(trajectory);
rec
}
pub fn from_trajectory(trajectory: Trajectory) -> Result<Self, MolRsError> {
trajectory.validate()?;
let Some(frame) = trajectory.frames.first().cloned() else {
return Err(MolRsError::validation(
"cannot build MolRec from an empty trajectory",
));
};
let mut rec = Self::new(frame);
rec.trajectory = Some(trajectory);
Ok(rec)
}
pub fn count_frames(&self) -> usize {
match &self.trajectory {
Some(traj) if !traj.frames.is_empty() => traj.frames.len(),
_ => 1,
}
}
pub fn frame_at(&self, index: usize) -> Option<Frame> {
match &self.trajectory {
Some(traj) if !traj.frames.is_empty() => traj.frames.get(index).cloned(),
_ if index == 0 => Some(self.frame.clone()),
_ => None,
}
}
pub fn set_frame(&mut self, frame: Frame) {
self.frame = frame;
}
pub fn add_frame(&mut self, frame: Frame) {
match &mut self.trajectory {
Some(traj) => traj.frames.push(frame),
None => {
self.trajectory = Some(Trajectory::from_frames(vec![frame]));
}
}
}
pub fn set_trajectory(&mut self, trajectory: Option<Trajectory>) {
self.trajectory = trajectory;
}
pub fn add_observable(&mut self, observable: ObservableRecord) -> Option<ObservableRecord> {
self.observables.insert(observable.name.clone(), observable)
}
pub fn get_observable(&self, name: &str) -> Option<&ObservableRecord> {
self.observables.get(name)
}
pub fn remove_observable(&mut self, name: &str) -> Option<ObservableRecord> {
self.observables.remove(name)
}
}
fn empty_object() -> JsonValue {
JsonValue::Object(JsonMap::new())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn static_molrec_counts_one_frame() {
let rec = MolRec::new(Frame::new());
assert_eq!(rec.count_frames(), 1);
assert!(rec.frame_at(0).is_some());
assert!(rec.frame_at(1).is_none());
}
#[test]
fn from_trajectory_uses_first_frame_as_canonical() {
let mut traj = Trajectory::new();
traj.frames.push(Frame::new());
traj.frames.push(Frame::new());
let rec = MolRec::from_trajectory(traj).unwrap();
assert_eq!(rec.count_frames(), 2);
}
}