use rust_decimal_macros::dec;
use tan::{
context::Context,
error::Error,
expr::{annotate_type, Expr},
util::{
args::{unpack_float_arg, unpack_int_arg},
module_util::require_module,
},
};
pub fn add_int(args: &[Expr]) -> Result<Expr, Error> {
let mut xs = Vec::new();
for arg in args {
let Some(n) = arg.as_int() else {
return Err(Error::invalid_arguments(
&format!("{arg} is not an Int"),
arg.range(),
));
};
xs.push(n);
}
let sum = add_int_impl(xs);
Ok(Expr::Int(sum))
}
fn add_int_impl(xs: Vec<i64>) -> i64 {
xs.iter().sum()
}
pub fn add_float(args: &[Expr]) -> Result<Expr, Error> {
let mut sum = 0.0;
for arg in args {
let Some(n) = arg.as_float() else {
return Err(Error::invalid_arguments(
&format!("{arg} is not a Float"),
arg.range(),
));
};
sum += n;
}
Ok(Expr::Float(sum))
}
pub fn add_dec(args: &[Expr]) -> Result<Expr, Error> {
let mut sum = dec!(0.0);
for arg in args {
let Some(n) = arg.as_decimal() else {
return Err(Error::invalid_arguments(
&format!("{arg} is not a Dec"),
arg.range(),
));
};
sum += n;
}
Ok(Expr::Dec(sum))
}
pub fn sub_int(args: &[Expr]) -> Result<Expr, Error> {
let mut diff = unpack_int_arg(args, 0, "n")?;
for arg in args.iter().skip(1) {
let Some(n) = arg.as_int() else {
return Err(Error::invalid_arguments(
&format!("{arg} is not an Int"),
arg.range(),
));
};
diff -= n;
}
Ok(Expr::Int(diff))
}
pub fn neg_int(args: &[Expr]) -> Result<Expr, Error> {
let n = unpack_int_arg(args, 0, "n")?;
Ok(Expr::Int(-n))
}
pub fn sub_float(args: &[Expr]) -> Result<Expr, Error> {
let [a, b] = args else {
return Err(Error::invalid_arguments(
"- requires at least two arguments",
None,
));
};
let Some(a) = a.as_float() else {
return Err(Error::invalid_arguments(
&format!("{a} is not a Float"),
a.range(),
));
};
let Some(b) = b.as_float() else {
return Err(Error::invalid_arguments(
&format!("{b} is not a Float"),
b.range(),
));
};
Ok(Expr::Float(a - b))
}
pub fn neg_float(args: &[Expr]) -> Result<Expr, Error> {
let n = unpack_float_arg(args, 0, "n")?;
Ok(Expr::Float(-n))
}
pub fn mul_int(args: &[Expr]) -> Result<Expr, Error> {
let mut product = 1;
for arg in args {
let Some(n) = arg.as_int() else {
return Err(Error::invalid_arguments(
&format!("{arg} is not an Int"),
arg.range(),
));
};
product *= n;
}
Ok(Expr::Int(product))
}
pub fn mul_float(args: &[Expr]) -> Result<Expr, Error> {
let mut product = 1.0;
for arg in args {
let Some(n) = arg.as_float() else {
return Err(Error::invalid_arguments(
&format!("{arg} is not a Float"),
arg.range(),
));
};
product *= n;
}
Ok(Expr::Float(product))
}
pub fn div_int(args: &[Expr]) -> Result<Expr, Error> {
let mut quotient = unpack_int_arg(args, 0, "n")?;
for arg in args.iter().skip(1) {
let Some(n) = arg.as_int() else {
return Err(Error::invalid_arguments(
&format!("{arg} is not an Int"),
arg.range(),
));
};
quotient /= n;
}
Ok(Expr::Int(quotient))
}
pub fn div_float(args: &[Expr]) -> Result<Expr, Error> {
let mut quotient = f64::NAN;
for arg in args {
let Some(n) = arg.as_float() else {
return Err(Error::invalid_arguments(
&format!("{arg} is not a Float"),
arg.range(),
));
};
if quotient.is_nan() {
quotient = n;
} else {
quotient /= n;
}
}
Ok(Expr::Float(quotient))
}
pub fn mod_int(args: &[Expr]) -> Result<Expr, Error> {
let x = unpack_int_arg(args, 0, "x")?;
let y = unpack_int_arg(args, 1, "y")?;
Ok(Expr::Int(x % y))
}
pub fn mod_float(args: &[Expr]) -> Result<Expr, Error> {
let x = unpack_float_arg(args, 0, "x")?;
let y = unpack_float_arg(args, 1, "y")?;
Ok(Expr::Float(x % y))
}
pub fn int_compare(args: &[Expr]) -> Result<Expr, Error> {
let [a, b] = args else {
return Err(Error::invalid_arguments(
"requires at least two arguments",
None,
));
};
let Some(a) = a.as_int() else {
return Err(Error::invalid_arguments(
&format!("{a} is not an Int"),
a.range(),
));
};
let Some(b) = b.as_int() else {
return Err(Error::invalid_arguments(
&format!("{b} is not an Int"),
b.range(),
));
};
let ordering = match a.cmp(&b) {
std::cmp::Ordering::Less => -1,
std::cmp::Ordering::Equal => 0,
std::cmp::Ordering::Greater => 1,
};
Ok(Expr::Int(ordering))
}
pub fn setup_lib_arithmetic(context: &mut Context) {
let module = require_module("prelude", context);
module.insert_invocable("+", annotate_type(Expr::foreign_func(&add_int), "Int"));
module.insert_invocable(
"+$$Int$$Int",
annotate_type(Expr::foreign_func(&add_int), "Int"),
);
module.insert_invocable(
"+$$Int$$Int$$Int",
annotate_type(Expr::foreign_func(&add_int), "Int"),
);
module.insert_invocable(
"+$$Int$$Int$$Int$$Int",
annotate_type(Expr::foreign_func(&add_int), "Int"),
);
module.insert_invocable(
"+$$Float$$Float",
annotate_type(Expr::foreign_func(&add_float), "Float"),
);
module.insert_invocable(
"+$$Float$$Float$$Float",
annotate_type(Expr::foreign_func(&add_float), "Float"),
);
module.insert_invocable(
"+$$Float$$Float$$Float$$Float",
annotate_type(Expr::foreign_func(&add_float), "Float"),
);
module.insert_invocable(
"+$$Dec$$Dec",
annotate_type(Expr::foreign_func(&add_dec), "Dec"),
);
module.insert_invocable("-", Expr::foreign_func(&sub_int));
module.insert_invocable("-$$Int", annotate_type(Expr::foreign_func(&neg_int), "Int"));
module.insert_invocable(
"-$$Int$$Int",
annotate_type(Expr::foreign_func(&sub_int), "Int"),
);
module.insert_invocable(
"-$$Int$$Int$$Int",
annotate_type(Expr::foreign_func(&sub_int), "Int"),
);
module.insert_invocable(
"-$$Int$$Int$$Int$$Int",
annotate_type(Expr::foreign_func(&sub_int), "Int"),
);
module.insert_invocable(
"-$$Float",
annotate_type(Expr::foreign_func(&neg_float), "Float"),
);
module.insert_invocable(
"-$$Float$$Float",
annotate_type(Expr::foreign_func(&sub_float), "Float"),
);
module.insert_invocable("*", Expr::foreign_func(&mul_int));
module.insert_invocable(
"*$$Int$$Int",
annotate_type(Expr::foreign_func(&mul_int), "Int"),
);
module.insert_invocable(
"*$$Int$$Int$$Int",
annotate_type(Expr::foreign_func(&mul_int), "Int"),
);
module.insert_invocable(
"*$$Float$$Float",
annotate_type(Expr::foreign_func(&mul_float), "Float"),
);
module.insert_invocable(
"*$$Float$$Float$$Float",
annotate_type(Expr::foreign_func(&mul_float), "Float"),
);
module.insert_invocable("/", annotate_type(Expr::foreign_func(&div_float), "Float"));
module.insert_invocable(
"/$$Int$$Int",
annotate_type(Expr::foreign_func(&div_int), "Int"),
);
module.insert_invocable(
"/$$Float$$Float",
annotate_type(Expr::foreign_func(&div_float), "Float"),
);
module.insert_invocable(
"/$$Float$$Float$$Float",
annotate_type(Expr::foreign_func(&div_float), "Float"),
);
module.insert_invocable("%", annotate_type(Expr::foreign_func(&mod_int), "Int"));
module.insert_invocable(
"%$$Int$$Int",
annotate_type(Expr::foreign_func(&mod_int), "Int"),
);
module.insert_invocable(
"%$$Float$$Float",
annotate_type(Expr::foreign_func(&mod_float), "Float"),
);
}