use crate::arithmetic::{fp_div_i, fp_mul, fp_mul_i, fp_sqrt};
use crate::constants::*;
use crate::error::SolMathError;
use crate::exp_coeffs::{
EXP2_PHASE_Q62, EXP_LN2_RESIDUAL_Q96, EXP_PHASES, EXP_PHASE_BITS, EXP_POLY_GUARD,
EXP_RAW_TO_Q63_FRAC_Q28, EXP_RAW_TO_Q63_HI, EXP_REMEZ_Q22, EXP_STEP_Q63,
};
use crate::expm1_lut::{
EXPM1_INV_LN2_Q56, EXPM1_LUT_SEGMENTS, EXPM1_LUT_STEP, EXPM1_LUT_STEP_SHIFT,
EXPM1_MID_EXP_RAW_Q22, EXPM1_RAW_TO_Q43_G31, EXPM1_R_MIN,
};
use crate::hp::pow_fixed_hp;
use crate::ln2_lut::{K_LN2_MAX, K_LN2_MIN, K_LN2_RAW};
use crate::ln_lut::{
LN_LUT_HALF_STEP, LN_LUT_MID_LOG, LN_LUT_SEGMENTS, LN_LUT_STEP, LN_Q42_RECIP_G32,
};
use crate::overflow::checked_mul_div_u;
#[inline(always)]
fn round_shift_signed(value: i128, shift: u32) -> i128 {
let half = 1i128 << (shift - 1);
if value >= 0 {
(value + half) >> shift
} else {
-((-value + half) >> shift)
}
}
#[inline(always)]
fn round_shift_i64(value: i64, shift: u32) -> i64 {
let half = 1i64 << (shift - 1);
if value >= 0 {
(value + half) >> shift
} else {
-((-value + half) >> shift)
}
}
#[inline(always)]
fn mul_q42(a: i64, b: i64) -> i64 {
round_shift_i64(a * b, 42)
}
#[inline(always)]
fn mul_q43(a: i64, b: i64) -> i64 {
round_shift_i64(a * b, 43)
}
#[inline(always)]
fn mul_q63_i64(a: i64, b: i64) -> i64 {
round_shift_signed(a as i128 * b as i128, 63) as i64
}
#[inline]
fn ln_mantissa_lut(m: u128, k: i32) -> i128 {
debug_assert!(m >= SCALE && m < 2 * SCALE);
debug_assert!((K_LN2_MIN..=K_LN2_MAX).contains(&k));
if m == SCALE {
return K_LN2_RAW[(k - K_LN2_MIN) as usize] as i128;
}
let offset = m - SCALE;
let j = ((offset as u64) / (LN_LUT_STEP as u64)) as usize;
debug_assert!(j < LN_LUT_SEGMENTS);
let midpoint = SCALE + j as u128 * LN_LUT_STEP + LN_LUT_HALF_STEP;
let d = m as i64 - midpoint as i64;
let q = round_shift_i64(d * LN_Q42_RECIP_G32[j], 32);
let q2 = mul_q42(q, q);
let q3 = mul_q42(q2, q);
let local_q42 = q - q2 / 2 + q3 / 3;
let local_raw = round_shift_signed(local_q42 as i128 * SCALE_I, 42);
let k_log = K_LN2_RAW[(k - K_LN2_MIN) as usize];
LN_LUT_MID_LOG[j] as i128 + local_raw + k_log as i128
}
#[cold]
#[inline(never)]
fn normalize_ln_fallback(value: u128) -> (u128, i32) {
let bit_length = 128 - value.leading_zeros() as i32;
let mut k = bit_length - 40;
let shift = |exponent: i32| {
if exponent >= 0 {
value >> exponent as u32
} else {
value << (-exponent) as u32
}
};
let mut m = shift(k);
if m < SCALE {
k -= 1;
m = shift(k);
} else if m >= 2 * SCALE {
k += 1;
m = shift(k);
}
(m, k)
}
#[inline]
fn normalize_ln(value: u128) -> (u128, i32) {
if value >= SCALE && value < 2 * SCALE {
(value, 0)
} else if value >= SCALE / 2 && value < SCALE {
(value * 2, -1)
} else {
normalize_ln_fallback(value)
}
}
pub fn ln_fixed_i(x: u128) -> Result<i128, SolMathError> {
if x == 0 {
return Err(SolMathError::DomainError);
}
if x.abs_diff(SCALE) < 1_000_000 {
return Ok(x as i128 - SCALE_I);
}
let (m, k) = normalize_ln(x);
Ok(ln_mantissa_lut(m, k))
}
pub fn ln_1p_fixed(x: i128) -> Result<i128, SolMathError> {
if x <= -SCALE_I {
return Err(SolMathError::DomainError);
}
if x == SCALE_I {
return Ok(LN2_I);
}
if x.unsigned_abs() < 1_000_000 {
return Ok(x);
}
let one_plus_x = if x < 0 {
(SCALE_I + x) as u128
} else {
SCALE.checked_add(x as u128).ok_or(SolMathError::Overflow)?
};
let (m, k) = normalize_ln(one_plus_x);
Ok(ln_mantissa_lut(m, k))
}
pub fn exp_fixed_i(x: i128) -> Result<i128, SolMathError> {
let max_x = 40 * SCALE_I;
if x <= -max_x {
return Ok(0);
}
if x >= max_x {
return Err(SolMathError::Overflow);
}
if (-1_000_000..1_000_000).contains(&x) {
return Ok(SCALE_I + x);
}
let x64 = x as i64;
let octave_estimate = round_shift_i64(x64 * EXPM1_INV_LN2_Q56, 56) as i32;
let raw_residual = x64 - octave_estimate as i64 * LN2_I as i64;
let scaled_residual = raw_residual * EXP_RAW_TO_Q63_HI
+ round_shift_i64(raw_residual * EXP_RAW_TO_Q63_FRAC_Q28, 28);
let octave_residual_q63 = round_shift_i64(scaled_residual, 1)
- round_shift_i64(octave_estimate as i64 * EXP_LN2_RESIDUAL_Q96, 33);
let mut subcell = round_shift_i64(
raw_residual * EXPM1_INV_LN2_Q56,
(56 - EXP_PHASE_BITS) as u32,
) as i32;
let mut r_q63 = octave_residual_q63 - subcell as i64 * EXP_STEP_Q63;
let half_step_q63 = (EXP_STEP_Q63 + 1) / 2;
if r_q63 > half_step_q63 {
subcell += 1;
r_q63 -= EXP_STEP_Q63;
} else if r_q63 < -half_step_q63 {
subcell -= 1;
r_q63 += EXP_STEP_Q63;
}
debug_assert!(r_q63.abs() <= half_step_q63 + 1);
let poly = mul_q63_i64(EXP_REMEZ_Q22[0], r_q63) + EXP_REMEZ_Q22[1];
let poly = mul_q63_i64(poly, r_q63) + EXP_REMEZ_Q22[2];
let poly = mul_q63_i64(poly, r_q63) + EXP_REMEZ_Q22[3];
let poly = mul_q63_i64(poly, r_q63) + EXP_REMEZ_Q22[4];
let poly = mul_q63_i64(poly, r_q63) + EXP_REMEZ_Q22[5];
let cell = octave_estimate * EXP_PHASES as i32 + subcell;
let octave = cell >> EXP_PHASE_BITS;
let phase = (cell & (EXP_PHASES as i32 - 1)) as usize;
let (guarded, guard) = if phase == 0 {
(poly as i128, EXP_POLY_GUARD)
} else {
(
poly as i128 * EXP2_PHASE_Q62[phase] as i128,
EXP_POLY_GUARD + 62,
)
};
let shift = guard - octave;
if shift >= 128 {
Ok(0)
} else if shift > 0 {
Ok(round_shift_signed(guarded, shift as u32))
} else if shift == 0 {
Ok(guarded)
} else {
guarded
.checked_shl((-shift) as u32)
.ok_or(SolMathError::Overflow)
}
}
pub fn pow_fixed(base: u128, exponent: u128) -> Result<u128, SolMathError> {
if base == 0 && exponent == 0 {
return Err(SolMathError::DomainError); }
if exponent == 0 {
return Ok(SCALE); }
if base == 0 {
return Ok(0); }
if exponent == SCALE {
return Ok(base); }
if base == SCALE {
return Ok(SCALE); }
if exponent == 2 * SCALE {
return fp_mul(base, base);
}
if exponent == SCALE / 2 {
return fp_sqrt(base);
}
if exponent % SCALE == 0 {
let n = exponent / SCALE;
if n >= 1 && n <= 20 {
return pow_fixed_hp(base, exponent);
}
}
if base.abs_diff(SCALE) < 1_000_000 && exponent > SCALE {
return pow_fixed_hp(base, exponent);
}
let ln_base = ln_fixed_i(base)?; if ln_base == 0 && base != SCALE {
return pow_fixed_hp(base, exponent);
}
let exp_i = match i128::try_from(exponent) {
Ok(v) => v,
Err(_) if base < SCALE => return Ok(0),
Err(_) => return Err(SolMathError::Overflow),
};
let product = match fp_mul_i(exp_i, ln_base) {
Ok(v) => v,
Err(SolMathError::Overflow) if base < SCALE => return Ok(0),
Err(e) => return Err(e),
}; let result = exp_fixed_i(product)?;
Ok(if result <= 0 { 0 } else { result as u128 })
}
#[cfg(test)]
mod adversarial_power_tests {
use super::*;
#[test]
fn huge_positive_exponents_preserve_direction() {
assert_eq!(pow_fixed(SCALE / 2, 1u128 << 127), Ok(0));
assert_eq!(
pow_fixed(2 * SCALE, 1u128 << 127),
Err(SolMathError::Overflow)
);
assert_eq!(pow_fixed(SCALE / 10, i128::MAX as u128), Ok(0));
}
#[test]
fn near_one_large_exponent_uses_high_precision_log() {
let exponent = 1u128 << 80;
assert_eq!(
pow_fixed(SCALE + 1, exponent),
pow_fixed_hp(SCALE + 1, exponent)
);
assert_eq!(
pow_fixed(SCALE + 1, 1u128 << 86),
Err(SolMathError::Overflow)
);
}
}
pub fn pow_int(base: u128, n: u128) -> Result<u128, SolMathError> {
match n {
0 => Ok(SCALE),
1 => Ok(base),
_ if base == 0 => Ok(0),
_ if base == SCALE => Ok(SCALE),
2 => fp_mul(base, base),
3 => Ok(fp_mul(fp_mul(base, base)?, base)?),
4 => {
let x2 = fp_mul(base, base)?;
fp_mul(x2, x2)
}
_ => {
let ln_base = ln_fixed_i(base)?;
let n_i = match i128::try_from(n) {
Ok(value) => value,
Err(_) if base < SCALE => return Ok(0),
Err(_) => return Err(SolMathError::Overflow),
};
let total = match n_i.checked_mul(ln_base) {
Some(v) => v,
None => {
if ln_base > 0 {
return Err(SolMathError::Overflow);
} else {
return Ok(0);
}
}
};
if total.unsigned_abs() < (39 * SCALE_I) as u128 {
let exponent = n.checked_mul(SCALE).ok_or(SolMathError::Overflow)?;
pow_fixed_hp(base, exponent)
} else {
let half = pow_int(base, n / 2)?;
let mut result =
checked_mul_div_u(half, half, SCALE).ok_or(SolMathError::Overflow)?;
if n % 2 == 1 {
result =
checked_mul_div_u(result, base, SCALE).ok_or(SolMathError::Overflow)?;
}
Ok(result)
}
}
}
}
pub fn pow_fixed_i(base: i128, exponent: i128) -> Result<i128, SolMathError> {
if base == 0 && exponent == 0 {
return Err(SolMathError::DomainError); }
if exponent == 0 {
return Ok(SCALE_I); }
if base == 0 {
if exponent < 0 {
return Err(SolMathError::Overflow); }
return Ok(0); }
if exponent == SCALE_I {
return Ok(base); }
if base == SCALE_I {
return Ok(SCALE_I); }
if exponent < 0 {
let positive_exponent = exponent.checked_neg().ok_or(SolMathError::Overflow)?;
let pos_result = pow_fixed_i(base, positive_exponent)?;
if pos_result == 0 {
return Err(SolMathError::Overflow); }
return fp_div_i(SCALE_I, pos_result);
}
if base < 0 {
if exponent % SCALE_I != 0 {
return Err(SolMathError::DomainError); }
let n = exponent / SCALE_I;
let abs_base = base.unsigned_abs();
let abs_result = pow_fixed(abs_base, exponent as u128)?;
if abs_result > i128::MAX as u128 {
return Err(SolMathError::Overflow);
}
Ok(if n % 2 == 0 {
abs_result as i128
} else {
-(abs_result as i128)
})
} else {
let result = pow_fixed(base as u128, exponent as u128)?;
if result > i128::MAX as u128 {
return Err(SolMathError::Overflow);
}
Ok(result as i128)
}
}
pub fn expm1_fixed(x: i128) -> Result<i128, SolMathError> {
let limit = 40 * SCALE_I;
if x <= -limit {
return Ok(-SCALE_I);
}
if x >= limit {
return Err(SolMathError::Overflow);
}
if x.unsigned_abs() < 1_000_000 {
return Ok(x);
}
let x64 = x as i64; let mut k = round_shift_i64(x64 * EXPM1_INV_LN2_Q56, 56) as i32;
debug_assert!((K_LN2_MIN..=64).contains(&k));
let mut r = x64 - K_LN2_RAW[(k - K_LN2_MIN) as usize];
const HALF_LN2_RAW: i64 = 346_573_590_280;
if r > HALF_LN2_RAW {
k += 1;
r = x64 - K_LN2_RAW[(k - K_LN2_MIN) as usize];
} else if r < -HALF_LN2_RAW {
k -= 1;
r = x64 - K_LN2_RAW[(k - K_LN2_MIN) as usize];
}
let offset = (r - EXPM1_R_MIN) as u64;
debug_assert!(r >= EXPM1_R_MIN);
let j = ((offset >> EXPM1_LUT_STEP_SHIFT) as usize).min(EXPM1_LUT_SEGMENTS - 1);
let midpoint = EXPM1_R_MIN + j as i64 * EXPM1_LUT_STEP + EXPM1_LUT_STEP / 2;
let delta = r - midpoint;
let q = round_shift_i64(delta * EXPM1_RAW_TO_Q43_G31, 31);
let q2 = mul_q43(q, q);
let q3 = mul_q43(q2, q);
let local_q43 = (1i64 << 43) + q + q2 / 2 + q3 / 6;
let exp_r_raw_q22 =
round_shift_signed(EXPM1_MID_EXP_RAW_Q22[j] as i128 * local_q43 as i128, 43);
let shift = 22 - k;
let exp_x = if shift > 0 {
round_shift_signed(exp_r_raw_q22, shift as u32)
} else if shift == 0 {
exp_r_raw_q22
} else {
exp_r_raw_q22
.checked_shl((-shift) as u32)
.ok_or(SolMathError::Overflow)?
};
Ok(exp_x - SCALE_I)
}
#[cfg(test)]
mod boundary_tests {
use super::*;
#[test]
fn pow_fixed_rejects_unsigned_exponent_that_cannot_be_signed() {
assert_eq!(
pow_fixed(2 * SCALE, i128::MAX as u128 + 1),
Err(SolMathError::Overflow)
);
}
#[test]
fn pow_int_handles_zero_and_extreme_exponents_consistently() {
assert_eq!(pow_int(0, 5), Ok(0));
assert_eq!(pow_int(SCALE, u128::MAX), Ok(SCALE));
assert_eq!(pow_int(SCALE / 2, i128::MAX as u128 + 1), Ok(0));
assert_eq!(
pow_int(2 * SCALE, i128::MAX as u128 + 1),
Err(SolMathError::Overflow)
);
}
#[test]
fn pow_fixed_i_rejects_unnegatable_min_exponent() {
assert_eq!(
pow_fixed_i(2 * SCALE_I, i128::MIN),
Err(SolMathError::Overflow)
);
}
#[test]
fn ln_1p_has_explicit_domain_and_exact_special_values() {
assert_eq!(ln_1p_fixed(-SCALE_I), Err(SolMathError::DomainError));
assert_eq!(ln_1p_fixed(i128::MIN), Err(SolMathError::DomainError));
assert_eq!(ln_1p_fixed(0), Ok(0));
assert_eq!(ln_1p_fixed(SCALE_I), ln_fixed_i(2 * SCALE));
assert!(ln_1p_fixed(i128::MAX).is_ok());
}
#[test]
fn ln_1p_preserves_raw_increments_near_zero() {
assert_eq!(ln_1p_fixed(1), Ok(1));
assert_eq!(ln_1p_fixed(-1), Ok(-1));
assert_eq!(ln_1p_fixed(2), Ok(2));
assert_eq!(ln_1p_fixed(-2), Ok(-2));
}
#[test]
fn ln_1p_and_expm1_round_trip_financial_rates() {
for x in [
-900_000_000_000,
-500_000_000_000,
-10_000_000_000,
-1_000_000,
1_000_000,
10_000_000_000,
500_000_000_000,
5 * SCALE_I,
] {
let recovered = expm1_fixed(ln_1p_fixed(x).unwrap()).unwrap();
assert!((recovered - x).abs() <= 12, "x={x}, recovered={recovered}");
}
}
#[test]
fn expm1_has_explicit_limits_and_preserves_raw_increments() {
let limit = 40 * SCALE_I;
assert_eq!(expm1_fixed(i128::MIN), Ok(-SCALE_I));
assert_eq!(expm1_fixed(-limit), Ok(-SCALE_I));
assert_eq!(expm1_fixed(limit), Err(SolMathError::Overflow));
assert_eq!(expm1_fixed(i128::MAX), Err(SolMathError::Overflow));
assert_eq!(expm1_fixed(-1), Ok(-1));
assert_eq!(expm1_fixed(0), Ok(0));
assert_eq!(expm1_fixed(1), Ok(1));
}
#[test]
fn expm1_matches_known_ordinary_values() {
assert!(expm1_fixed(SCALE_I).unwrap().abs_diff(1_718_281_828_459) <= 3);
assert!(expm1_fixed(-SCALE_I).unwrap().abs_diff(-632_120_558_829) <= 1);
}
#[test]
fn expm1_is_monotone_across_every_lut_boundary_and_exponent() {
for k in -58..=58 {
let k_ln2 = K_LN2_RAW[(k - K_LN2_MIN) as usize] as i128;
for j in 1..EXPM1_LUT_SEGMENTS {
let boundary = EXPM1_R_MIN as i128 + j as i128 * EXPM1_LUT_STEP as i128;
let x_left = k_ln2 + boundary - 1;
let x_right = k_ln2 + boundary;
if x_left <= -40 * SCALE_I || x_right >= 40 * SCALE_I {
continue;
}
let left = expm1_fixed(x_left).unwrap();
let right = expm1_fixed(x_right).unwrap();
assert!(left <= right, "k={k}, j={j}: {left} > {right}");
}
}
}
}