use super::PER_M_DENOMINATOR;
use super::error::{
CoreError, AMOUNT_EXCEEDS_MAX_U64, ARITHMETIC_OVERFLOW, INVALID_SWAP_FEE_PER_M,
};
#[cfg(feature = "wasm")]
use equanetwork_macros::wasm_expose;
#[cfg_attr(feature = "wasm", wasm_expose)]
pub fn fee_from_pre_fee_amount(pre_fee_amount: u64, swap_fee_per_m: u32) -> Result<u64, CoreError> {
if swap_fee_per_m > PER_M_DENOMINATOR as u32 {
Err(INVALID_SWAP_FEE_PER_M)
} else if swap_fee_per_m == 0 || pre_fee_amount == 0 {
Ok(0)
} else {
let numerator = <u128>::from(pre_fee_amount)
.checked_mul(swap_fee_per_m as u128)
.ok_or(ARITHMETIC_OVERFLOW)?;
let fee_amount: u64 = numerator
.div_ceil(PER_M_DENOMINATOR as u128)
.try_into()
.map_err(|_| AMOUNT_EXCEEDS_MAX_U64)?;
Ok(fee_amount)
}
}
#[cfg_attr(feature = "wasm", wasm_expose)]
pub fn fee_from_post_fee_amount(
post_fee_amount: u64,
swap_fee_per_m: u32,
) -> Result<u64, CoreError> {
if swap_fee_per_m > PER_M_DENOMINATOR as u32 {
Err(INVALID_SWAP_FEE_PER_M)
} else if swap_fee_per_m == 0 || post_fee_amount == 0 {
Ok(0)
} else if swap_fee_per_m == PER_M_DENOMINATOR as u32 {
Ok(u64::MAX)
} else {
let numerator = <u128>::from(post_fee_amount)
.checked_mul(PER_M_DENOMINATOR as u128)
.ok_or(ARITHMETIC_OVERFLOW)?;
let denominator = <u128>::from(PER_M_DENOMINATOR as u32) - <u128>::from(swap_fee_per_m);
let pre_fee_amount = numerator.div_ceil(denominator);
let fee_amount: u64 = pre_fee_amount
.checked_sub(post_fee_amount.into())
.ok_or(ARITHMETIC_OVERFLOW)?
.try_into()
.map_err(|_| AMOUNT_EXCEEDS_MAX_U64)?;
Ok(fee_amount)
}
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
#[rstest]
#[case(1000, 10_000, 10)]
#[case(1000, 0, 0)]
#[case(1000, 1_000_000, 1000)]
#[case(0, 10_000, 0)]
#[case(9, 10_000, 1)]
fn test_fee_from_pre_fee_amount(
#[case] pre_fee_amount: u64,
#[case] swap_fee_per_m: u32,
#[case] expected_fee_amount: u64,
) {
let fee_amount = fee_from_pre_fee_amount(pre_fee_amount, swap_fee_per_m).unwrap();
assert_eq!(fee_amount, expected_fee_amount);
}
#[rstest]
#[case(990, 10_000, 10)]
#[case(1000, 0, 0)]
#[case(1000, 1_000_000, u64::MAX)]
#[case(0, 10_000, 0)]
#[case(9, 10_000, 1)]
fn test_fee_from_post_fee_amount(
#[case] post_fee_amount: u64,
#[case] swap_fee_per_m: u32,
#[case] expected_fee_amount: u64,
) {
let fee_amount = fee_from_post_fee_amount(post_fee_amount, swap_fee_per_m).unwrap();
assert_eq!(fee_amount, expected_fee_amount);
}
}