use crate::LuaType;
use crate::Result;
use crate::State;
use crate::error::{ArgError, ErrorKind};
use std::f64::consts::PI;
pub(crate) fn open_math(state: &mut State) -> Result<()> {
state.new_table_with_capacity(24)?;
macro_rules! add_fn {
($name:expr, $func:expr) => {
#[cfg(feature = "snapshot")]
state
.set_table_str_key_named_rust_fn(-1, $name, concat!("math.", $name), $func)
.expect("math library registration cannot fail");
#[cfg(not(feature = "snapshot"))]
state
.set_table_str_key_rust_fn(-1, $name, $func)
.expect("math library registration cannot fail");
};
}
add_fn!("sin", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.sin())?;
Ok(1)
});
add_fn!("cos", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.cos())?;
Ok(1)
});
add_fn!("atan2", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
state.check_type(2, LuaType::Number)?;
let y = state.to_number(1)?;
let x = state.to_number(2)?;
state.set_top(0)?;
state.push_number(y.atan2(x))?;
Ok(1)
});
add_fn!("sqrt", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.sqrt())?;
Ok(1)
});
add_fn!("abs", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.abs())?;
Ok(1)
});
add_fn!("min", |state| {
state.consume_cost(1)?;
let n = state.get_top() as isize;
if n == 0 {
let e = ArgError {
arg_number: 1,
func_name: Some("min".to_string()),
expected: Some(LuaType::Number),
received: None,
};
return Err(state.error(ErrorKind::ArgError(e)));
}
state.check_type(1, LuaType::Number)?;
let mut result = state.to_number(1)?;
for i in 2..=n {
state.check_type(i, LuaType::Number)?;
let v = state.to_number(i)?;
if v < result {
result = v;
}
}
state.set_top(0)?;
state.push_number(result)?;
Ok(1)
});
add_fn!("max", |state| {
state.consume_cost(1)?;
let n = state.get_top() as isize;
if n == 0 {
let e = ArgError {
arg_number: 1,
func_name: Some("max".to_string()),
expected: Some(LuaType::Number),
received: None,
};
return Err(state.error(ErrorKind::ArgError(e)));
}
state.check_type(1, LuaType::Number)?;
let mut result = state.to_number(1)?;
for i in 2..=n {
state.check_type(i, LuaType::Number)?;
let v = state.to_number(i)?;
if v > result {
result = v;
}
}
state.set_top(0)?;
state.push_number(result)?;
Ok(1)
});
add_fn!("floor", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.floor())?;
Ok(1)
});
add_fn!("ceil", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.ceil())?;
Ok(1)
});
add_fn!("random", |state| {
state.consume_cost(1)?;
let num_args = state.get_top();
let result = match num_args {
0 => {
state.rng.random_f64()
}
1 => {
state.check_type(1, LuaType::Number)?;
let n = state.to_number(1)? as i64;
if n < 1 {
return Err(state.error(ErrorKind::RuntimeError(
"bad argument #1 to 'random' (interval is empty)".to_string(),
)));
}
state.rng.random_range_i64(1, n) as f64
}
2 => {
state.check_type(1, LuaType::Number)?;
state.check_type(2, LuaType::Number)?;
let m = state.to_number(1)? as i64;
let n = state.to_number(2)? as i64;
if m > n {
return Err(state.error(ErrorKind::RuntimeError(
"bad argument #1 to 'random' (interval is empty)".to_string(),
)));
}
state.rng.random_range_i64(m, n) as f64
}
_ => {
return Err(state.error(ErrorKind::RuntimeError(
"wrong number of arguments to 'random'".to_string(),
)));
}
};
state.set_top(0)?;
state.push_number(result)?;
Ok(1)
});
add_fn!("tan", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.tan())?;
Ok(1)
});
add_fn!("acos", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.acos())?;
Ok(1)
});
add_fn!("asin", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.asin())?;
Ok(1)
});
add_fn!("atan", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let y = state.to_number(1)?;
let result = if state.get_top() >= 2 {
state.check_type(2, LuaType::Number)?;
let x = state.to_number(2)?;
y.atan2(x)
} else {
y.atan()
};
state.set_top(0)?;
state.push_number(result)?;
Ok(1)
});
add_fn!("deg", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.to_degrees())?;
Ok(1)
});
add_fn!("rad", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.to_radians())?;
Ok(1)
});
add_fn!("exp", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
state.set_top(0)?;
state.push_number(x.exp())?;
Ok(1)
});
add_fn!("log", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
let result = if state.get_top() >= 2 {
state.check_type(2, LuaType::Number)?;
let base = state.to_number(2)?;
#[allow(clippy::float_cmp)]
let is_base_ten = base == 10.0;
if is_base_ten { x.log10() } else { x.log(base) }
} else {
x.ln()
};
state.set_top(0)?;
state.push_number(result)?;
Ok(1)
});
add_fn!("fmod", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
state.check_type(2, LuaType::Number)?;
let x = state.to_number(1)?;
let y = state.to_number(2)?;
state.set_top(0)?;
state.push_number(x % y)?;
Ok(1)
});
add_fn!("modf", |state| {
state.consume_cost(1)?;
state.check_type(1, LuaType::Number)?;
let x = state.to_number(1)?;
let integral = x.trunc();
let fractional = if matches!(x.partial_cmp(&integral), Some(std::cmp::Ordering::Equal)) {
0.0
} else {
x - integral
};
state.set_top(0)?;
state.push_number(integral)?;
state.push_number(fractional)?;
Ok(2)
});
state.set_table_str_key_number(-1, "pi", PI)?;
state.set_table_str_key_number(-1, "huge", f64::INFINITY)?;
state.set_global("math");
Ok(())
}