use rust_decimal::{Decimal, RoundingStrategy};
use std::collections::{HashMap, HashSet};
use uuid::Uuid;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ConversionRounding {
Exact,
Places {
dp: u32,
strategy: RoundingStrategy,
},
}
impl ConversionRounding {
pub fn apply(self, d: Decimal) -> Decimal {
match self {
ConversionRounding::Exact => d,
ConversionRounding::Places { dp, strategy } => d.round_dp_with_strategy(dp, strategy),
}
}
}
#[derive(Debug, Clone, PartialEq, sqlx::FromRow)]
pub struct UomChainNode {
pub id: Uuid,
pub code: String,
pub relative_uom_id: Option<Uuid>,
pub relative_factor: Option<Decimal>,
pub factor: Decimal,
}
#[derive(Debug, Clone, PartialEq)]
pub struct UomChain {
nodes: Vec<UomChainNode>,
}
impl UomChain {
pub fn from_rows(
leaf_id: Uuid,
rows: Vec<UomChainNode>,
) -> Result<Self, UomConversionError> {
let by_id: HashMap<Uuid, UomChainNode> =
rows.into_iter().map(|n| (n.id, n)).collect();
let leaf = by_id.get(&leaf_id).ok_or_else(|| UomConversionError::DanglingRelative {
code: leaf_id.to_string(),
relative_uom_id: leaf_id,
})?;
let mut nodes = Vec::with_capacity(by_id.len());
let mut seen: HashSet<Uuid> = HashSet::new();
let mut cursor = leaf;
loop {
if !seen.insert(cursor.id) {
return Err(UomConversionError::CycleDetected { code: cursor.code.clone() });
}
let next = match cursor.relative_uom_id {
None => {
nodes.push(cursor.clone());
break;
}
Some(parent_id) => {
nodes.push(cursor.clone());
match by_id.get(&parent_id) {
Some(parent) => parent,
None => {
return Err(UomConversionError::DanglingRelative {
code: cursor.code.clone(),
relative_uom_id: parent_id,
})
}
}
}
};
cursor = next;
}
Ok(Self { nodes })
}
pub fn leaf(&self) -> &UomChainNode {
&self.nodes[0]
}
pub fn root(&self) -> Result<&UomChainNode, UomConversionError> {
let last = self.nodes.last().expect("a chain is never empty");
if last.relative_uom_id.is_some() {
return Err(UomConversionError::CycleDetected { code: last.code.clone() });
}
Ok(last)
}
pub fn derived_root_factor(&self) -> Result<Decimal, UomConversionError> {
self.root()?;
let mut f = Decimal::ONE;
for n in &self.nodes {
f *= n.relative_factor.unwrap_or(Decimal::ONE);
}
Ok(f)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum UomConversionError {
CrossTree {
from_code: String,
from_root_code: String,
to_code: String,
to_root_code: String,
},
CycleDetected { code: String },
DanglingRelative { code: String, relative_uom_id: Uuid },
}
impl UomConversionError {
pub fn code(&self) -> &'static str {
match self {
UomConversionError::CrossTree { .. } => "cross_tree_conversion",
UomConversionError::CycleDetected { .. } => "uom_tree_cycle",
UomConversionError::DanglingRelative { .. } => "dangling_relative_uom",
}
}
}
impl std::fmt::Display for UomConversionError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
UomConversionError::CrossTree {
from_code,
from_root_code,
to_code,
to_root_code,
} => write!(
f,
"cannot convert {from_code} -> {to_code}: the units live in different trees \
({from_code} is in the tree rooted at {from_root_code}, {to_code} is in the \
tree rooted at {to_root_code}) — there is no common reference unit"
),
UomConversionError::CycleDetected { code } => {
write!(f, "unit tree cycle detected at or above {code}: the parent links never reach a root")
}
UomConversionError::DanglingRelative { code, relative_uom_id } => {
write!(f, "unit {code} links at relative_uom_id {relative_uom_id}, which is not visible in this scope")
}
}
}
}
impl std::error::Error for UomConversionError {}
pub fn convert_quantity(
qty: Decimal,
from: &UomChain,
to: &UomChain,
rounding: ConversionRounding,
) -> Result<Decimal, UomConversionError> {
let from_root = from.root()?;
let to_root = to.root()?;
if from_root.id != to_root.id {
return Err(UomConversionError::CrossTree {
from_code: from.leaf().code.clone(),
from_root_code: from_root.code.clone(),
to_code: to.leaf().code.clone(),
to_root_code: to_root.code.clone(),
});
}
if from.leaf().id == to.leaf().id {
return Ok(rounding.apply(qty));
}
Ok(rounding.apply(qty * from.leaf().factor / to.leaf().factor))
}
#[cfg(test)]
mod tests {
use super::*;
fn node(id: Uuid, code: &str, parent: Option<Uuid>, rf: Option<Decimal>, factor: Decimal) -> UomChainNode {
UomChainNode { id, code: code.into(), relative_uom_id: parent, relative_factor: rf, factor }
}
#[test]
fn same_unit_converts_identically_under_exact() {
let id = Uuid::new_v4();
let chain = UomChain::from_rows(id, vec![node(id, "PCS", None, None, Decimal::ONE)]).unwrap();
let q = Decimal::new(1234, 2);
assert_eq!(
convert_quantity(q, &chain, &chain, ConversionRounding::Exact).unwrap(),
q
);
}
#[test]
fn chain_assembles_from_unordered_rows_leaf_first_root_last() {
let unit = Uuid::new_v4();
let pack = Uuid::new_v4();
let box_ = Uuid::new_v4();
let rows = vec![
node(pack, "PACK", Some(unit), Some(Decimal::from(10)), Decimal::from(10)),
node(box_, "BOX", Some(pack), Some(Decimal::from(12)), Decimal::from(120)),
node(unit, "UNIT", None, None, Decimal::ONE),
];
let chain = UomChain::from_rows(box_, rows).unwrap();
assert_eq!(chain.leaf().code, "BOX");
assert_eq!(chain.root().unwrap().code, "UNIT");
assert_eq!(chain.nodes.len(), 3);
assert_eq!(chain.derived_root_factor().unwrap(), Decimal::from(120));
}
#[test]
fn cycle_in_stored_links_fails_loudly() {
let a = Uuid::new_v4();
let b = Uuid::new_v4();
let rows = vec![
node(a, "A", Some(b), Some(Decimal::ONE), Decimal::ONE),
node(b, "B", Some(a), Some(Decimal::ONE), Decimal::ONE),
];
let err = UomChain::from_rows(a, rows).unwrap_err();
assert!(matches!(err, UomConversionError::CycleDetected { .. }));
}
#[test]
fn dangling_parent_link_fails_loudly() {
let a = Uuid::new_v4();
let ghost = Uuid::new_v4();
let rows = vec![node(a, "A", Some(ghost), Some(Decimal::ONE), Decimal::ONE)];
let err = UomChain::from_rows(a, rows).unwrap_err();
assert!(matches!(err, UomConversionError::DanglingRelative { .. }));
}
}