use crate::common::constants::{FOUR_WAD, MAX_POW_RELATIVE_ERROR, TWO_WAD, WAD};
use crate::common::errors::PoolError;
use crate::common::log_exp_math;
use alloy_primitives::U256;
pub fn mul_up_fixed(a: &U256, b: &U256) -> Result<U256, PoolError> {
let product = a.checked_mul(*b).ok_or(PoolError::MathOverflow)?;
if product.is_zero() {
return Ok(U256::ZERO);
}
let result = (product - U256::ONE) / WAD + U256::ONE;
Ok(result)
}
pub fn div_up_fixed(a: &U256, b: &U256) -> Result<U256, PoolError> {
let result = mul_div_up_fixed(a, &WAD, b)?;
Ok(result)
}
pub fn mul_down_fixed(a: &U256, b: &U256) -> Result<U256, PoolError> {
let product = a.checked_mul(*b).ok_or(PoolError::MathOverflow)?;
let result = product / WAD;
Ok(result)
}
pub fn div_down_fixed(a: &U256, b: &U256) -> Result<U256, PoolError> {
if a.is_zero() {
return Ok(U256::ZERO);
}
if b.is_zero() {
return Err(PoolError::MathOverflow);
}
let a_inflated = a.checked_mul(WAD).ok_or(PoolError::MathOverflow)?;
let result = a_inflated / b;
Ok(result)
}
pub fn div_up(a: &U256, b: &U256) -> Result<U256, PoolError> {
if b.is_zero() {
return Ok(U256::ZERO);
}
let result = U256::ONE + (a - U256::ONE) / b;
Ok(result)
}
pub fn mul_div_up_fixed(a: &U256, b: &U256, c: &U256) -> Result<U256, PoolError> {
let product = a.checked_mul(*b).ok_or(PoolError::MathOverflow)?;
if product.is_zero() {
return Ok(U256::ZERO);
}
let result = (product - U256::ONE) / c + U256::ONE;
Ok(result)
}
pub fn pow_down_fixed(base: &U256, exponent: &U256) -> Result<U256, PoolError> {
pow_down_fixed_with_version(base, exponent, 0)
}
pub fn pow_down_fixed_with_version(
base: &U256,
exponent: &U256,
version: u32,
) -> Result<U256, PoolError> {
if *exponent == WAD && version != 1 {
return Ok(*base);
}
if *exponent == TWO_WAD && version != 1 {
return mul_up_fixed(base, base);
}
if *exponent == FOUR_WAD && version != 1 {
let square = mul_up_fixed(base, base)?;
return mul_up_fixed(&square, &square);
}
let raw = log_exp_math::pow(base, exponent)?;
let max_error = mul_up_fixed(&raw, &MAX_POW_RELATIVE_ERROR)? + U256::ONE;
if raw < max_error {
return Ok(U256::ZERO);
}
Ok(raw - max_error)
}
pub fn pow_up_fixed(base: &U256, exponent: &U256) -> Result<U256, PoolError> {
pow_up_fixed_with_version(base, exponent, 0)
}
pub fn pow_up_fixed_with_version(
base: &U256,
exponent: &U256,
version: u32,
) -> Result<U256, PoolError> {
if *exponent == WAD && version != 1 {
return Ok(*base);
}
if *exponent == TWO_WAD && version != 1 {
return mul_up_fixed(base, base);
}
if *exponent == FOUR_WAD && version != 1 {
let square = mul_up_fixed(base, base)?;
return mul_up_fixed(&square, &square);
}
let raw = log_exp_math::pow(base, exponent)?;
let max_error = mul_up_fixed(&raw, &MAX_POW_RELATIVE_ERROR)? + U256::ONE;
Ok(raw + max_error)
}
pub fn complement_fixed(x: &U256) -> Result<U256, PoolError> {
if *x < WAD {
Ok(WAD - x)
} else {
Ok(U256::ZERO)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn mul_down_fixed_overflow_returns_err() {
let a = U256::MAX;
let b = U256::from(2u64);
assert!(matches!(
mul_down_fixed(&a, &b),
Err(PoolError::MathOverflow)
));
}
#[test]
fn mul_up_fixed_overflow_returns_err() {
let a = U256::MAX;
let b = U256::from(2u64);
assert!(matches!(mul_up_fixed(&a, &b), Err(PoolError::MathOverflow)));
}
#[test]
fn mul_div_up_fixed_overflow_returns_err() {
let a = U256::MAX;
let b = U256::from(2u64);
assert!(matches!(
mul_div_up_fixed(&a, &b, &WAD),
Err(PoolError::MathOverflow)
));
}
#[test]
fn div_down_fixed_inflated_overflow_returns_err() {
let a = U256::MAX;
assert!(matches!(
div_down_fixed(&a, &U256::from(1u64)),
Err(PoolError::MathOverflow)
));
}
#[test]
fn mul_down_fixed_normal_ok() {
let result = mul_down_fixed(&WAD, &WAD).unwrap();
assert_eq!(result, WAD);
}
}