use alloc::collections::BTreeMap;
use alloc::sync::Arc;
use alloc::vec::Vec;
use serde::de::Error as _;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use crate::symbolic::SymbolicExpr;
#[derive(Debug, Serialize, Deserialize)]
enum FlatNode<A> {
Leaf(A),
Add {
x: usize,
y: usize,
degree_multiple: usize,
},
Sub {
x: usize,
y: usize,
degree_multiple: usize,
},
Neg {
x: usize,
degree_multiple: usize,
},
Mul {
x: usize,
y: usize,
degree_multiple: usize,
},
}
fn flatten_into<'a, A>(
node: &'a SymbolicExpr<A>,
nodes: &mut Vec<FlatNode<&'a A>>,
seen: &mut BTreeMap<*const SymbolicExpr<A>, usize>,
) -> usize {
let key: *const SymbolicExpr<A> = node;
if let Some(&idx) = seen.get(&key) {
return idx;
}
let flat = match node {
SymbolicExpr::Leaf(a) => FlatNode::Leaf(a),
SymbolicExpr::Add {
x,
y,
degree_multiple,
} => {
let x = flatten_into(x, nodes, seen);
let y = flatten_into(y, nodes, seen);
FlatNode::Add {
x,
y,
degree_multiple: *degree_multiple,
}
}
SymbolicExpr::Sub {
x,
y,
degree_multiple,
} => {
let x = flatten_into(x, nodes, seen);
let y = flatten_into(y, nodes, seen);
FlatNode::Sub {
x,
y,
degree_multiple: *degree_multiple,
}
}
SymbolicExpr::Neg { x, degree_multiple } => {
let x = flatten_into(x, nodes, seen);
FlatNode::Neg {
x,
degree_multiple: *degree_multiple,
}
}
SymbolicExpr::Mul {
x,
y,
degree_multiple,
} => {
let x = flatten_into(x, nodes, seen);
let y = flatten_into(y, nodes, seen);
FlatNode::Mul {
x,
y,
degree_multiple: *degree_multiple,
}
}
};
let idx = nodes.len();
nodes.push(flat);
seen.insert(key, idx);
idx
}
fn unflatten<A>(nodes: Vec<FlatNode<A>>) -> Option<SymbolicExpr<A>> {
let mut built: Vec<Arc<SymbolicExpr<A>>> = Vec::with_capacity(nodes.len());
for (i, flat) in nodes.into_iter().enumerate() {
let expr = match flat {
FlatNode::Leaf(a) => SymbolicExpr::Leaf(a),
FlatNode::Add {
x,
y,
degree_multiple,
} => {
if x >= i || y >= i {
return None;
}
SymbolicExpr::Add {
x: built[x].clone(),
y: built[y].clone(),
degree_multiple,
}
}
FlatNode::Sub {
x,
y,
degree_multiple,
} => {
if x >= i || y >= i {
return None;
}
SymbolicExpr::Sub {
x: built[x].clone(),
y: built[y].clone(),
degree_multiple,
}
}
FlatNode::Neg { x, degree_multiple } => {
if x >= i {
return None;
}
SymbolicExpr::Neg {
x: built[x].clone(),
degree_multiple,
}
}
FlatNode::Mul {
x,
y,
degree_multiple,
} => {
if x >= i || y >= i {
return None;
}
SymbolicExpr::Mul {
x: built[x].clone(),
y: built[y].clone(),
degree_multiple,
}
}
};
built.push(Arc::new(expr));
}
Arc::try_unwrap(built.pop()?).ok()
}
impl<A: Serialize> Serialize for SymbolicExpr<A> {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
let mut nodes: Vec<FlatNode<&A>> = Vec::new();
let mut seen: BTreeMap<*const Self, usize> = BTreeMap::new();
flatten_into(self, &mut nodes, &mut seen);
nodes.serialize(serializer)
}
}
impl<'de, A: Deserialize<'de>> Deserialize<'de> for SymbolicExpr<A> {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let nodes = Vec::<FlatNode<A>>::deserialize(deserializer)?;
unflatten(nodes).ok_or_else(|| D::Error::custom("invalid flattened symbolic expression"))
}
}
#[cfg(test)]
mod tests {
use alloc::format;
use alloc::sync::Arc;
use p3_baby_bear::BabyBear;
use crate::symbolic::expression::BaseLeaf;
use crate::symbolic::variable::BaseEntry;
use crate::symbolic::{SymbolicExpr, SymbolicExpression, SymbolicVariable};
type F = BabyBear;
fn var(index: usize) -> SymbolicExpression<F> {
SymbolicExpression::from(SymbolicVariable::new(BaseEntry::Main { offset: 0 }, index))
}
#[test]
fn json_round_trip_preserves_structure() {
let expr = var(0) * var(1) - var(0) + SymbolicExpression::from(F::new(7));
let json = serde_json::to_string(&expr).unwrap();
let decoded: SymbolicExpression<F> = serde_json::from_str(&json).unwrap();
assert_eq!(serde_json::to_string(&decoded).unwrap(), json);
assert_eq!(format!("{decoded:?}"), format!("{expr:?}"));
}
#[test]
fn shared_subtree_is_emitted_once() {
let shared = Arc::new(var(0));
let expr = SymbolicExpr::Add {
x: shared.clone(),
y: shared,
degree_multiple: 1,
};
let value: serde_json::Value = serde_json::to_value(&expr).unwrap();
assert_eq!(value.as_array().unwrap().len(), 2);
let decoded: SymbolicExpression<F> = serde_json::from_value(value).unwrap();
assert_eq!(format!("{decoded:?}"), format!("{expr:?}"));
}
#[test]
fn rejects_forward_reference() {
let json = r#"[{"Add":{"x":1,"y":2,"degree_multiple":1}}]"#;
let decoded: Result<SymbolicExpression<F>, _> = serde_json::from_str(json);
assert!(decoded.is_err());
}
#[test]
fn rejects_empty_arena() {
let decoded: Result<SymbolicExpression<F>, _> = serde_json::from_str("[]");
assert!(decoded.is_err());
}
#[test]
fn leaf_round_trips() {
let expr = SymbolicExpression::<F>::Leaf(BaseLeaf::Constant(F::new(42)));
let json = serde_json::to_string(&expr).unwrap();
let decoded: SymbolicExpression<F> = serde_json::from_str(&json).unwrap();
assert_eq!(format!("{decoded:?}"), format!("{expr:?}"));
}
}