equanetwork-math 0.1.0

The Equa Network program math library
Documentation
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);
    }
}