use rust_decimal::Decimal;
use rust_decimal_macros::dec;
use serde::{Deserialize, Serialize};
use crate::error::{A2AError, A2AResult};
pub const DECIMALS_DP: u32 = 6;
pub const DECIMALS_MULTIPLIER: u64 = 1_000_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum SplitType {
Percentage,
Fixed,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Recipient {
pub address: String,
pub percent: Option<Decimal>,
pub amount: Option<Decimal>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SplitShare {
pub address: String,
pub amount: Decimal,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SplitResult {
pub shares: Vec<SplitShare>,
pub platform_fee: Decimal,
pub total_distributed: Decimal,
}
pub fn calculate_percentage_split(
total: Decimal,
platform_fee_pct: Decimal,
recipients: &[Recipient],
) -> A2AResult<SplitResult> {
if total <= Decimal::ZERO {
return Err(A2AError::validation("total must be greater than 0"));
}
if recipients.len() < 2 {
return Err(A2AError::validation(
"recipients must have at least 2 entries",
));
}
for (i, r) in recipients.iter().enumerate() {
if r.percent.is_none() {
return Err(A2AError::validation(format!(
"recipients[{i}].percent is required for percentage splits"
)));
}
}
let percent_sum: Decimal = recipients
.iter()
.map(|r| r.percent.unwrap_or_default())
.sum();
let tolerance = dec!(0.001);
if (percent_sum - dec!(100)).abs() > tolerance {
return Err(A2AError::PercentageSumMismatch {
actual: percent_sum,
});
}
let fee = (total * platform_fee_pct / dec!(100)).round_dp(DECIMALS_DP);
let remaining = total - fee;
let mut allocated = Decimal::ZERO;
let mut shares = Vec::with_capacity(recipients.len());
for (i, r) in recipients.iter().enumerate() {
let pct = r.percent.unwrap_or_default();
let share = if i == recipients.len() - 1 {
remaining - allocated
} else {
(remaining * pct / dec!(100)).round_dp(DECIMALS_DP)
};
allocated += share;
shares.push(SplitShare {
address: r.address.clone(),
amount: share,
});
}
Ok(SplitResult {
shares,
platform_fee: fee,
total_distributed: allocated + fee,
})
}
pub fn calculate_fixed_split(
total: Decimal,
platform_fee_pct: Decimal,
recipients: &[Recipient],
) -> A2AResult<SplitResult> {
if total <= Decimal::ZERO {
return Err(A2AError::validation("total must be greater than 0"));
}
if recipients.len() < 2 {
return Err(A2AError::validation(
"recipients must have at least 2 entries",
));
}
for (i, r) in recipients.iter().enumerate() {
if r.amount.is_none() {
return Err(A2AError::validation(format!(
"recipients[{i}].amount is required for fixed splits"
)));
}
}
let fee = (total * platform_fee_pct / dec!(100)).round_dp(DECIMALS_DP);
let remaining = total - fee;
let fixed_sum: Decimal = recipients
.iter()
.map(|r| r.amount.unwrap_or_default().round_dp(DECIMALS_DP))
.sum();
let one_unit = Decimal::new(1, DECIMALS_DP);
if (fixed_sum - remaining).abs() > one_unit {
return Err(A2AError::FixedSumMismatch {
expected: remaining,
actual: fixed_sum,
});
}
let mut shares = Vec::with_capacity(recipients.len());
let mut distributed = Decimal::ZERO;
for r in recipients {
let amount = r.amount.unwrap_or_default().round_dp(DECIMALS_DP);
distributed += amount;
shares.push(SplitShare {
address: r.address.clone(),
amount,
});
}
Ok(SplitResult {
shares,
platform_fee: fee,
total_distributed: distributed + fee,
})
}
#[cfg(test)]
mod tests {
use super::*;
use rust_decimal_macros::dec;
fn make_pct_recipients(pcts: &[(& str, Decimal)]) -> Vec<Recipient> {
pcts.iter()
.map(|(addr, pct)| Recipient {
address: (*addr).to_string(),
percent: Some(*pct),
amount: None,
})
.collect()
}
fn make_fixed_recipients(amounts: &[(&str, Decimal)]) -> Vec<Recipient> {
amounts
.iter()
.map(|(addr, amt)| Recipient {
address: (*addr).to_string(),
percent: None,
amount: Some(*amt),
})
.collect()
}
#[test]
fn percentage_split_basic_two_way() {
let r = make_pct_recipients(&[("0xA", dec!(50)), ("0xB", dec!(50))]);
let result = calculate_percentage_split(dec!(100), dec!(0), &r).unwrap();
assert_eq!(result.shares.len(), 2);
assert_eq!(result.shares[0].amount, dec!(50));
assert_eq!(result.shares[1].amount, dec!(50));
assert_eq!(result.platform_fee, dec!(0));
assert_eq!(result.total_distributed, dec!(100));
}
#[test]
fn percentage_split_three_way_with_platform_fee() {
let r = make_pct_recipients(&[
("0xAlice", dec!(50)),
("0xBob", dec!(30)),
("0xCharlie", dec!(20)),
]);
let result = calculate_percentage_split(dec!(100), dec!(2.5), &r).unwrap();
assert_eq!(result.platform_fee, dec!(2.500000));
assert_eq!(result.total_distributed, dec!(100));
assert_eq!(result.shares.len(), 3);
assert_eq!(result.shares[0].amount, dec!(48.750000));
assert_eq!(result.shares[1].amount, dec!(29.250000));
assert_eq!(result.shares[2].amount, dec!(19.500000));
}
#[test]
fn percentage_split_rounding_drift_prevention() {
let r = make_pct_recipients(&[
("0xA", dec!(33.333)),
("0xB", dec!(33.333)),
("0xC", dec!(33.334)),
]);
let result = calculate_percentage_split(dec!(100), dec!(0), &r).unwrap();
assert_eq!(result.total_distributed, dec!(100));
}
#[test]
fn percentage_split_large_amount() {
let r = make_pct_recipients(&[("0xA", dec!(70)), ("0xB", dec!(30))]);
let result =
calculate_percentage_split(dec!(1000000), dec!(1), &r).unwrap();
assert_eq!(result.platform_fee, dec!(10000.000000));
assert_eq!(result.total_distributed, dec!(1000000));
}
#[test]
fn percentage_split_tiny_amount() {
let r = make_pct_recipients(&[("0xA", dec!(60)), ("0xB", dec!(40))]);
let result = calculate_percentage_split(dec!(0.000001), dec!(0), &r).unwrap();
assert_eq!(result.total_distributed, dec!(0.000001));
}
#[test]
fn percentage_split_rejects_single_recipient() {
let r = make_pct_recipients(&[("0xA", dec!(100))]);
let err = calculate_percentage_split(dec!(100), dec!(0), &r).unwrap_err();
assert!(matches!(err, A2AError::Validation(_)));
}
#[test]
fn percentage_split_rejects_zero_total() {
let r = make_pct_recipients(&[("0xA", dec!(50)), ("0xB", dec!(50))]);
let err = calculate_percentage_split(dec!(0), dec!(0), &r).unwrap_err();
assert!(matches!(err, A2AError::Validation(_)));
}
#[test]
fn percentage_split_rejects_negative_total() {
let r = make_pct_recipients(&[("0xA", dec!(50)), ("0xB", dec!(50))]);
let err = calculate_percentage_split(dec!(-10), dec!(0), &r).unwrap_err();
assert!(matches!(err, A2AError::Validation(_)));
}
#[test]
fn percentage_split_rejects_wrong_sum() {
let r = make_pct_recipients(&[("0xA", dec!(50)), ("0xB", dec!(40))]);
let err = calculate_percentage_split(dec!(100), dec!(0), &r).unwrap_err();
assert!(matches!(err, A2AError::PercentageSumMismatch { .. }));
}
#[test]
fn percentage_split_rejects_missing_percent() {
let r = vec![
Recipient { address: "0xA".into(), percent: Some(dec!(50)), amount: None },
Recipient { address: "0xB".into(), percent: None, amount: None },
];
let err = calculate_percentage_split(dec!(100), dec!(0), &r).unwrap_err();
assert!(matches!(err, A2AError::Validation(_)));
}
#[test]
fn percentage_split_many_recipients() {
let r: Vec<_> = (0..10)
.map(|i| Recipient {
address: format!("0x{i:02}"),
percent: Some(dec!(10)),
amount: None,
})
.collect();
let result = calculate_percentage_split(dec!(1000), dec!(5), &r).unwrap();
assert_eq!(result.total_distributed, dec!(1000));
assert_eq!(result.shares.len(), 10);
}
#[test]
fn percentage_split_uneven_three_way() {
let r = make_pct_recipients(&[
("0xAlice", dec!(50)),
("0xBob", dec!(30)),
("0xCharlie", dec!(20)),
]);
let result = calculate_percentage_split(dec!(100), dec!(2.5), &r).unwrap();
assert_eq!(result.platform_fee, dec!(2.500000));
assert_eq!(result.shares[0].amount, dec!(48.750000));
assert_eq!(result.shares[1].amount, dec!(29.250000));
assert_eq!(result.shares[2].amount, dec!(19.500000));
assert_eq!(result.total_distributed, dec!(100));
}
#[test]
fn fixed_split_basic() {
let r = make_fixed_recipients(&[("0xA", dec!(60)), ("0xB", dec!(40))]);
let result = calculate_fixed_split(dec!(100), dec!(0), &r).unwrap();
assert_eq!(result.shares.len(), 2);
assert_eq!(result.shares[0].amount, dec!(60));
assert_eq!(result.shares[1].amount, dec!(40));
assert_eq!(result.total_distributed, dec!(100));
}
#[test]
fn fixed_split_with_platform_fee() {
let r = make_fixed_recipients(&[("0xA", dec!(55)), ("0xB", dec!(40))]);
let result = calculate_fixed_split(dec!(100), dec!(5), &r).unwrap();
assert_eq!(result.platform_fee, dec!(5.000000));
assert_eq!(result.total_distributed, dec!(100));
}
#[test]
fn fixed_split_rejects_wrong_sum() {
let r = make_fixed_recipients(&[("0xA", dec!(50)), ("0xB", dec!(40))]);
let err = calculate_fixed_split(dec!(100), dec!(0), &r).unwrap_err();
assert!(matches!(err, A2AError::FixedSumMismatch { .. }));
}
#[test]
fn fixed_split_rejects_single_recipient() {
let r = make_fixed_recipients(&[("0xA", dec!(100))]);
let err = calculate_fixed_split(dec!(100), dec!(0), &r).unwrap_err();
assert!(matches!(err, A2AError::Validation(_)));
}
#[test]
fn fixed_split_rejects_missing_amount() {
let r = vec![
Recipient { address: "0xA".into(), percent: None, amount: Some(dec!(50)) },
Recipient { address: "0xB".into(), percent: None, amount: None },
];
let err = calculate_fixed_split(dec!(100), dec!(0), &r).unwrap_err();
assert!(matches!(err, A2AError::Validation(_)));
}
#[test]
fn fixed_split_allows_one_unit_tolerance() {
let r = make_fixed_recipients(&[
("0xA", dec!(59.999999)),
("0xB", dec!(40.000000)),
]);
let result = calculate_fixed_split(dec!(100), dec!(0), &r);
assert!(result.is_ok());
}
}