use std::sync::Arc;
use num_bigint::BigInt;
use num_rational::Ratio;
use smallvec::SmallVec;
use crate::api::context::Context;
use crate::api::expr::{BoolEx, Ex};
use crate::base::arena::{
Arena, FN_AIRYAI, FN_AIRYAIPRIME, FN_AIRYBI, FN_AIRYBIPRIME, FN_ASSOC_LAGUERRE,
FN_ASSOC_LEGENDRE, FN_BETAINC, FN_BETAINC_REGULARIZED, FN_CHI, FN_DIRICHLET_ETA, FN_ELLIPTIC_E,
FN_ELLIPTIC_F, FN_ELLIPTIC_K, FN_ELLIPTIC_PI, FN_ERFCINV, FN_ERFI, FN_ERFINV, FN_EXPINT,
FN_FRESNELC, FN_FRESNELS, FN_GEGENBAUER, FN_JACOBI, FN_LOWERGAMMA, FN_POLYLOG, FN_SHI,
FN_UPPERGAMMA,
};
use crate::base::node::{ExprId, ExprNode};
#[derive(Debug, Clone)]
pub struct ParseError {
pub message: String,
pub position: usize,
}
impl std::fmt::Display for ParseError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"parse error at position {}: {}",
self.position, self.message
)
}
}
impl std::error::Error for ParseError {}
pub fn parse(ctx: &Context, input: &str) -> Result<Ex, ParseError> {
let id = parse_with_mode(ctx, input, Mode::STRICT)?;
Ok(Ex::from_raw_parts(ctx.id, Arc::clone(&ctx.inner), id))
}
pub fn parse_bool(ctx: &Context, input: &str) -> Result<BoolEx, ParseError> {
let id = parse_with_mode(ctx, input, Mode::RELATIONS)?;
Ok(BoolEx::from_raw_parts(ctx.id, Arc::clone(&ctx.inner), id))
}
pub fn parse_implicit(ctx: &Context, input: &str) -> Result<Ex, ParseError> {
let id = parse_with_mode(ctx, input, Mode::IMPLICIT)?;
Ok(Ex::from_raw_parts(ctx.id, Arc::clone(&ctx.inner), id))
}
fn parse_with_mode(ctx: &Context, input: &str, mode: Mode) -> Result<ExprId, ParseError> {
let mut parser = Parser::new(input, mode);
ctx.with_arena_mut(|arena| {
let result = parser.parse_expr(arena, 0)?;
if parser.current != Token::Eof {
return Err(ParseError {
message: format!("unexpected token {:?} after expression", parser.current),
position: parser.lexer.pos,
});
}
if mode.relations && !is_bool_node(arena, result) {
return Err(ParseError {
message: format!(
"expected a relation or Boolean expression, got the numeric expression '{}'",
arena.display(result)
),
position: parser.lexer.pos,
});
}
Ok(result)
})
}
impl Context {
pub fn parse_bool(
&self,
input: &str,
) -> Result<crate::api::expr::BoolEx, crate::base::errors::SymplexError> {
parse_bool(self, input).map_err(|e| crate::base::errors::SymplexError::ComputationFailed {
operation: "parse_bool",
reason: e.to_string(),
})
}
pub fn parse_implicit(
&self,
input: &str,
) -> Result<crate::api::expr::Ex, crate::base::errors::SymplexError> {
parse_implicit(self, input).map_err(|e| {
crate::base::errors::SymplexError::ComputationFailed {
operation: "parse_implicit",
reason: e.to_string(),
}
})
}
}
#[derive(Clone, Copy)]
struct Mode {
relations: bool,
implicit_app: bool,
}
impl Mode {
const STRICT: Mode = Mode {
relations: false,
implicit_app: false,
};
const RELATIONS: Mode = Mode {
relations: true,
implicit_app: false,
};
const IMPLICIT: Mode = Mode {
relations: false,
implicit_app: true,
};
}
fn is_bool_node(arena: &Arena, id: ExprId) -> bool {
matches!(
arena.node(id),
ExprNode::BoolTrue
| ExprNode::BoolFalse
| ExprNode::Gt(_, _)
| ExprNode::Ge(_, _)
| ExprNode::Eq_(_, _)
| ExprNode::Ne(_, _)
| ExprNode::And(_)
| ExprNode::Or(_)
| ExprNode::Not(_)
)
}
const KNOWN_FUNCTIONS: &[&str] = &[
"sin",
"cos",
"tan",
"exp",
"ln",
"log",
"sqrt",
"cbrt",
"abs",
"asin",
"arcsin",
"acos",
"arccos",
"atan",
"arctan",
"sinh",
"cosh",
"tanh",
"asinh",
"arcsinh",
"acosh",
"arccosh",
"atanh",
"arctanh",
"sign",
"sgn",
"floor",
"ceil",
"ceiling",
"gamma",
"erf",
"erfc",
"heaviside",
"diracdelta",
"dirac_delta",
"lambertw",
"w",
"factorial",
"digamma",
"loggamma",
"cot",
"sec",
"csc",
"coth",
"sech",
"csch",
"acot",
"arccot",
"re",
"im",
"conjugate",
"conj",
"arg",
"si",
"ci",
"ei",
"li",
"zeta",
"erfi",
"erfinv",
"erfcinv",
"e1",
"shi",
"chi",
"fresnels",
"fresnelc",
"dirichlet_eta",
"airyai",
"airybi",
"airyaiprime",
"airybiprime",
"elliptic_k",
"elliptic_e",
"rootof",
"conditionset",
"integral",
"atan2",
"polygamma",
"kroneckerdelta",
"kronecker_delta",
"binomial",
"c",
"beta",
"b",
"besselj",
"bessely",
"besseli",
"besselk",
"expint",
"lowergamma",
"uppergamma",
"polylog",
"elliptic_f",
"elliptic_pi",
"limit",
"laplacetransform",
"inverselaplacetransform",
"residue",
"dsolve",
"gegenbauer",
"assoc_legendre",
"assoc_laguerre",
"series",
"jacobi",
"betainc",
"betainc_regularized",
"min",
"max",
"sum",
"product",
];
fn is_known_function(name_lower: &str) -> bool {
KNOWN_FUNCTIONS.contains(&name_lower)
}
fn is_implicit_unary_function(name_lower: &str) -> bool {
matches!(
name_lower,
"sin"
| "cos"
| "tan"
| "cot"
| "sec"
| "csc"
| "sinh"
| "cosh"
| "tanh"
| "coth"
| "sech"
| "csch"
| "asin"
| "acos"
| "atan"
| "acot"
| "arcsin"
| "arccos"
| "arctan"
| "arccot"
| "asinh"
| "acosh"
| "atanh"
| "arcsinh"
| "arccosh"
| "arctanh"
| "exp"
| "ln"
| "log"
| "sqrt"
| "cbrt"
| "abs"
| "floor"
| "ceil"
| "ceiling"
| "sign"
| "sgn"
| "gamma"
| "erf"
| "erfc"
| "factorial"
)
}
fn constant_of(arena: &Arena, name: &str) -> Option<ExprId> {
Some(match name {
"pi" | "Pi" | "PI" => arena.pi,
"e" | "E" => arena.e_const,
"I" | "i" => arena.i_unit,
"inf" | "oo" | "Inf" => arena.infinity,
"zoo" => arena.complex_infinity,
"nan" => arena.nan,
"EulerGamma" | "euler_gamma" => arena.euler_gamma,
"Catalan" => arena.catalan,
"GoldenRatio" | "golden_ratio" => arena.golden_ratio,
_ => return None,
})
}
#[derive(Debug, Clone, PartialEq)]
enum Token {
Int(BigInt),
Rational(Ratio<BigInt>),
Ident(String),
Plus,
Minus,
Star,
Slash,
Caret,
LParen,
RParen,
Comma,
Bang,
Eq,
DotDot,
Lt,
Le,
Gt,
Ge,
EqEq,
Ne,
Amp,
Pipe,
Tilde,
Eof,
}
struct Lexer<'a> {
input: &'a str,
pos: usize,
}
impl<'a> Lexer<'a> {
fn new(input: &'a str) -> Self {
Lexer { input, pos: 0 }
}
fn skip_whitespace(&mut self) {
while self.pos < self.input.len() && self.input.as_bytes()[self.pos].is_ascii_whitespace() {
self.pos += 1;
}
}
fn peek_is(&self, b: u8) -> bool {
self.input.as_bytes().get(self.pos) == Some(&b)
}
fn next_token(&mut self) -> Result<Token, ParseError> {
self.skip_whitespace();
if self.pos >= self.input.len() {
return Ok(Token::Eof);
}
let b = self.input.as_bytes()[self.pos];
match b {
b'+' => {
self.pos += 1;
Ok(Token::Plus)
}
b'-' => {
self.pos += 1;
Ok(Token::Minus)
}
b'*' => {
self.pos += 1;
if self.pos < self.input.len() && self.input.as_bytes()[self.pos] == b'*' {
self.pos += 1;
Ok(Token::Caret) } else {
Ok(Token::Star)
}
}
b'/' => {
self.pos += 1;
Ok(Token::Slash)
}
b'^' => {
self.pos += 1;
Ok(Token::Caret)
}
b'(' => {
self.pos += 1;
Ok(Token::LParen)
}
b')' => {
self.pos += 1;
Ok(Token::RParen)
}
b',' => {
self.pos += 1;
Ok(Token::Comma)
}
b'!' => {
self.pos += 1;
if self.peek_is(b'=') {
self.pos += 1;
Ok(Token::Ne)
} else {
Ok(Token::Bang)
}
}
b'=' => {
self.pos += 1;
if self.peek_is(b'=') {
self.pos += 1;
Ok(Token::EqEq)
} else {
Ok(Token::Eq)
}
}
b'<' => {
self.pos += 1;
if self.peek_is(b'=') {
self.pos += 1;
Ok(Token::Le)
} else {
Ok(Token::Lt)
}
}
b'>' => {
self.pos += 1;
if self.peek_is(b'=') {
self.pos += 1;
Ok(Token::Ge)
} else {
Ok(Token::Gt)
}
}
b'&' => {
self.pos += 1;
if self.peek_is(b'&') {
self.pos += 1;
}
Ok(Token::Amp)
}
b'|' => {
self.pos += 1;
if self.peek_is(b'|') {
self.pos += 1;
}
Ok(Token::Pipe)
}
b'~' => {
self.pos += 1;
Ok(Token::Tilde)
}
b'.' if self.pos + 1 < self.input.len()
&& self.input.as_bytes()[self.pos + 1] == b'.' =>
{
self.pos += 2;
Ok(Token::DotDot)
}
b'0'..=b'9' => {
let start = self.pos;
while self.pos < self.input.len()
&& self.input.as_bytes()[self.pos].is_ascii_digit()
{
self.pos += 1;
}
if self.pos < self.input.len()
&& self.input.as_bytes()[self.pos] == b'.'
&& self.pos + 1 < self.input.len()
&& self.input.as_bytes()[self.pos + 1].is_ascii_digit()
{
self.pos += 1; let frac_start = self.pos;
while self.pos < self.input.len()
&& self.input.as_bytes()[self.pos].is_ascii_digit()
{
self.pos += 1;
}
let decimal_places = self.pos - frac_start;
let full_str = self.input[start..self.pos].replace('.', "");
let numer = full_str.parse::<BigInt>().map_err(|e| ParseError {
message: format!(
"invalid number '{}': {}",
&self.input[start..self.pos],
e
),
position: start,
})?;
let mut denom = BigInt::from(1);
for _ in 0..decimal_places {
denom *= 10;
}
let ratio = Ratio::new(numer, denom);
return Ok(Token::Rational(ratio));
}
let s = &self.input[start..self.pos];
let n = s.parse::<BigInt>().map_err(|e| ParseError {
message: format!("invalid integer '{}': {}", s, e),
position: start,
})?;
Ok(Token::Int(n))
}
b'a'..=b'z' | b'A'..=b'Z' | b'_' => {
let start = self.pos;
while self.pos < self.input.len() {
let c = self.input.as_bytes()[self.pos];
if c.is_ascii_alphanumeric() || c == b'_' {
self.pos += 1;
} else {
break;
}
}
let s = self.input[start..self.pos].to_string();
Ok(Token::Ident(s))
}
_ => Err(ParseError {
message: format!("unexpected character '{}'", b as char),
position: self.pos,
}),
}
}
}
const BP_OR: (u8, u8) = (1, 2);
const BP_AND: (u8, u8) = (3, 4);
const BP_REL: (u8, u8) = (5, 6);
const BP_ADD: (u8, u8) = (7, 8);
const BP_MUL: (u8, u8) = (9, 10);
const BP_NEG: u8 = 11;
const BP_POW: (u8, u8) = (14, 13);
const BP_NOT: u8 = BP_REL.0;
#[derive(Clone, Copy)]
enum Infix {
Add,
Sub,
Mul,
Div,
Pow,
Rel(RelOp),
And,
Or,
}
#[derive(Clone, Copy)]
enum RelOp {
Lt,
Le,
Gt,
Ge,
Eq,
Ne,
}
impl RelOp {
fn of(token: &Token) -> Option<RelOp> {
Some(match token {
Token::Lt => RelOp::Lt,
Token::Le => RelOp::Le,
Token::Gt => RelOp::Gt,
Token::Ge => RelOp::Ge,
Token::EqEq => RelOp::Eq,
Token::Ne => RelOp::Ne,
_ => return None,
})
}
fn text(self) -> &'static str {
match self {
RelOp::Lt => "<",
RelOp::Le => "<=",
RelOp::Gt => ">",
RelOp::Ge => ">=",
RelOp::Eq => "==",
RelOp::Ne => "!=",
}
}
}
fn starts_operand(token: &Token) -> bool {
matches!(
token,
Token::Int(_) | Token::Rational(_) | Token::Ident(_) | Token::LParen
)
}
struct Parser<'a> {
lexer: Lexer<'a>,
current: Token,
depth: usize,
mode: Mode,
app_arg: bool,
}
impl<'a> Parser<'a> {
fn new(input: &'a str, mode: Mode) -> Self {
let mut lexer = Lexer::new(input);
let current = lexer.next_token().unwrap_or(Token::Eof);
Parser {
lexer,
current,
depth: 0,
mode,
app_arg: false,
}
}
fn error(&self, message: String) -> ParseError {
ParseError {
message,
position: self.lexer.pos,
}
}
fn outside_app_arg<T>(
&mut self,
f: impl FnOnce(&mut Self) -> Result<T, ParseError>,
) -> Result<T, ParseError> {
let saved = std::mem::replace(&mut self.app_arg, false);
let result = f(self);
self.app_arg = saved;
result
}
fn numeric_operands(
&self,
arena: &Arena,
op: &str,
lhs: ExprId,
rhs: ExprId,
) -> Result<(), ParseError> {
for id in [lhs, rhs] {
if is_bool_node(arena, id) {
return Err(self.error(format!(
"'{op}' needs numeric operands, but '{}' is a Boolean expression",
arena.display(id)
)));
}
}
Ok(())
}
fn boolean_operand(&self, arena: &Arena, op: &str, id: ExprId) -> Result<(), ParseError> {
if is_bool_node(arena, id) {
Ok(())
} else {
Err(self.error(format!(
"'{op}' needs Boolean operands (relations), but '{}' is numeric",
arena.display(id)
)))
}
}
fn combine(
&self,
arena: &mut Arena,
op: Infix,
lhs: ExprId,
rhs: ExprId,
) -> Result<ExprId, ParseError> {
match op {
Infix::Add => {
self.numeric_operands(arena, "+", lhs, rhs)?;
Ok(arena.add(&[lhs, rhs]))
}
Infix::Sub => {
self.numeric_operands(arena, "-", lhs, rhs)?;
Ok(arena.sub(lhs, rhs))
}
Infix::Mul => {
self.numeric_operands(arena, "*", lhs, rhs)?;
Ok(arena.mul(&[lhs, rhs]))
}
Infix::Div => {
self.numeric_operands(arena, "/", lhs, rhs)?;
Ok(arena.div(lhs, rhs))
}
Infix::Pow => {
self.numeric_operands(arena, "^", lhs, rhs)?;
Ok(arena.pow(lhs, rhs))
}
Infix::Rel(rel) => {
self.numeric_operands(arena, rel.text(), lhs, rhs)?;
if RelOp::of(&self.current).is_some() {
return Err(self.error(
"chained comparisons are not supported; write 'a < b & b < c'".into(),
));
}
Ok(match rel {
RelOp::Lt => arena.gt(rhs, lhs),
RelOp::Le => arena.ge(rhs, lhs),
RelOp::Gt => arena.gt(lhs, rhs),
RelOp::Ge => arena.ge(lhs, rhs),
RelOp::Eq => arena.eq_(lhs, rhs),
RelOp::Ne => arena.ne_(lhs, rhs),
})
}
Infix::And => {
self.boolean_operand(arena, "&", lhs)?;
self.boolean_operand(arena, "&", rhs)?;
let mut items: SmallVec<[ExprId; 4]> = match arena.node(lhs) {
ExprNode::And(children) => children.iter().copied().collect(),
_ => SmallVec::from_slice(&[lhs]),
};
items.push(rhs);
Ok(arena.and(&items))
}
Infix::Or => {
self.boolean_operand(arena, "|", lhs)?;
self.boolean_operand(arena, "|", rhs)?;
let mut items: SmallVec<[ExprId; 4]> = match arena.node(lhs) {
ExprNode::Or(children) => children.iter().copied().collect(),
_ => SmallVec::from_slice(&[lhs]),
};
items.push(rhs);
Ok(arena.or(&items))
}
}
}
fn advance(&mut self) -> Result<Token, ParseError> {
let old = std::mem::replace(&mut self.current, Token::Eof);
self.current = self.lexer.next_token()?;
Ok(old)
}
fn expect(&mut self, expected: &Token) -> Result<(), ParseError> {
if &self.current == expected {
self.advance()?;
Ok(())
} else {
Err(ParseError {
message: format!("expected {:?}, got {:?}", expected, self.current),
position: self.lexer.pos,
})
}
}
fn parse_expr(&mut self, arena: &mut Arena, min_bp: u8) -> Result<ExprId, ParseError> {
self.depth += 1;
if self.depth > 128 {
return Err(ParseError {
message: "expression nesting too deep (max 128 levels)".into(),
position: self.lexer.pos,
});
}
let mut lhs = self.parse_prefix(arena)?;
loop {
if self.current == Token::Bang {
self.advance()?;
lhs = arena.intern(ExprNode::Factorial(lhs));
continue;
}
let relations = self.mode.relations;
let (op, (l_bp, r_bp), implicit) = match &self.current {
Token::Plus => (Infix::Add, BP_ADD, false),
Token::Minus => (Infix::Sub, BP_ADD, false),
Token::Star => (Infix::Mul, BP_MUL, false),
Token::Slash => (Infix::Div, BP_MUL, false),
Token::Caret => (Infix::Pow, BP_POW, false),
Token::Lt | Token::Le | Token::Gt | Token::Ge | Token::EqEq | Token::Ne
if relations =>
{
match RelOp::of(&self.current) {
Some(rel) => (Infix::Rel(rel), BP_REL, false),
None => break,
}
}
Token::Amp if relations => (Infix::And, BP_AND, false),
Token::Pipe if relations => (Infix::Or, BP_OR, false),
Token::Ident(name) if relations && name == "and" => (Infix::And, BP_AND, false),
Token::Ident(name) if relations && name == "or" => (Infix::Or, BP_OR, false),
Token::Ident(name)
if self.app_arg && is_implicit_unary_function(&name.to_ascii_lowercase()) =>
{
break;
}
Token::Int(_) | Token::Rational(_) | Token::Ident(_) | Token::LParen => {
(Infix::Mul, BP_MUL, true)
}
_ => break,
};
if l_bp < min_bp {
break;
}
if !implicit {
self.advance()?;
}
let rhs = self.parse_expr(arena, r_bp)?;
lhs = self.combine(arena, op, lhs, rhs)?;
}
self.depth -= 1;
Ok(lhs)
}
fn parse_prefix(&mut self, arena: &mut Arena) -> Result<ExprId, ParseError> {
match self.current.clone() {
Token::Int(n) => {
self.advance()?;
Ok(arena.big_int(n))
}
Token::Rational(ratio) => {
self.advance()?;
let nid = arena.intern_num(ratio);
Ok(arena.intern(ExprNode::Num(nid)))
}
Token::Ident(name) => {
self.advance()?;
let name_lower = name.to_ascii_lowercase();
let is_call = !self.mode.implicit_app
|| (is_known_function(&name_lower) && name_lower.len() > 1);
if self.current == Token::LParen && is_call {
return self.outside_app_arg(|p| p.parse_function_call(arena, &name));
}
if self.mode.relations {
match name.as_str() {
"True" | "true" => return Ok(arena.bool_true()),
"False" | "false" => return Ok(arena.bool_false()),
"not" => return self.parse_not(arena),
_ => {}
}
}
if self.mode.implicit_app && is_implicit_unary_function(&name_lower) {
if !starts_operand(&self.current) {
return Err(self.error(format!(
"function '{name}' needs an argument (write '{name}(x)' or '{name} x')"
)));
}
let saved = std::mem::replace(&mut self.app_arg, true);
let arg = self.parse_expr(arena, BP_MUL.0);
self.app_arg = saved;
return self.call_1(arena, &name, &name_lower, arg?);
}
match constant_of(arena, &name) {
Some(c) => Ok(c),
None => Ok(arena.symbol(&name)),
}
}
Token::Minus => {
self.advance()?;
let operand = self.parse_expr(arena, BP_NEG)?;
if is_bool_node(arena, operand) {
return Err(self.error(format!(
"'-' needs a numeric operand, but '{}' is a Boolean expression",
arena.display(operand)
)));
}
Ok(arena.neg(operand))
}
Token::Tilde | Token::Bang if self.mode.relations => {
self.advance()?;
self.parse_not(arena)
}
Token::LParen => {
self.advance()?;
let inner = self.outside_app_arg(|p| p.parse_expr(arena, 0))?;
self.expect(&Token::RParen)?;
Ok(inner)
}
other => Err(ParseError {
message: format!("expected expression, got {:?}", other),
position: self.lexer.pos,
}),
}
}
fn parse_not(&mut self, arena: &mut Arena) -> Result<ExprId, ParseError> {
let operand = self.parse_expr(arena, BP_NOT)?;
self.boolean_operand(arena, "not", operand)?;
Ok(arena.not(operand))
}
const MAX_FN_ARGS: usize = 32;
fn parse_function_call(&mut self, arena: &mut Arena, name: &str) -> Result<ExprId, ParseError> {
self.expect(&Token::LParen)?;
if self.current == Token::RParen {
self.advance()?;
return Err(ParseError {
message: format!("function '{}' requires an argument", name),
position: self.lexer.pos,
});
}
let name_lower = name.to_ascii_lowercase();
let mut args: Vec<ExprId> = vec![self.parse_expr(arena, 0)?];
if matches!(name_lower.as_str(), "sum" | "product") && self.current == Token::Comma {
self.advance()?;
let var = self.parse_expr(arena, 0)?;
if self.current == Token::Eq {
self.advance()?;
let lo = self.parse_expr(arena, 0)?;
self.expect(&Token::DotDot)?;
let hi = self.parse_expr(arena, 0)?;
self.expect(&Token::RParen)?;
return self.make_sum_product(arena, name, &name_lower, args[0], var, lo, hi);
}
args.push(var);
}
while self.current == Token::Comma {
self.advance()?;
if args.len() >= Self::MAX_FN_ARGS {
return Err(ParseError {
message: format!(
"function '{}' has too many arguments (max {})",
name,
Self::MAX_FN_ARGS
),
position: self.lexer.pos,
});
}
args.push(self.parse_expr(arena, 0)?);
}
self.expect(&Token::RParen)?;
if self.mode.relations
&& let Some(result) = self.call_boolean(arena, name, &name_lower, &args)?
{
return Ok(result);
}
if matches!(name_lower.as_str(), "min" | "max") {
if args.len() < 2 {
return Err(ParseError {
message: format!("function '{}' requires at least 2 arguments", name),
position: self.lexer.pos,
});
}
let ids: SmallVec<[ExprId; 4]> = args.iter().copied().collect();
return Ok(arena.intern(if name_lower == "min" {
ExprNode::Min(ids)
} else {
ExprNode::Max(ids)
}));
}
match args.len() {
1 => self.call_1(arena, name, &name_lower, args[0]),
2 => self.call_2(arena, name, &name_lower, args[0], args[1]),
3 => self.call_3(arena, name, &name_lower, args[0], args[1], args[2]),
4 => self.call_4(arena, name, &name_lower, args[0], args[1], args[2], args[3]),
n => Err(ParseError {
message: format!(
"unknown {n}-argument function '{}'. Only min and max take more than 4 arguments",
name
),
position: self.lexer.pos,
}),
}
}
fn call_boolean(
&self,
arena: &mut Arena,
name: &str,
name_lower: &str,
args: &[ExprId],
) -> Result<Option<ExprId>, ParseError> {
let rel = match name_lower {
"eq" => Some(RelOp::Eq),
"ne" => Some(RelOp::Ne),
"lt" => Some(RelOp::Lt),
"le" => Some(RelOp::Le),
"gt" => Some(RelOp::Gt),
"ge" => Some(RelOp::Ge),
_ => None,
};
if let Some(rel) = rel {
let [lhs, rhs] = args else {
return Err(self.error(format!(
"'{name}' takes exactly 2 arguments, got {}",
args.len()
)));
};
self.numeric_operands(arena, rel.text(), *lhs, *rhs)?;
return Ok(Some(match rel {
RelOp::Lt => arena.gt(*rhs, *lhs),
RelOp::Le => arena.ge(*rhs, *lhs),
RelOp::Gt => arena.gt(*lhs, *rhs),
RelOp::Ge => arena.ge(*lhs, *rhs),
RelOp::Eq => arena.eq_(*lhs, *rhs),
RelOp::Ne => arena.ne_(*lhs, *rhs),
}));
}
match name_lower {
"and" | "or" => {
for &a in args {
self.boolean_operand(arena, name, a)?;
}
Ok(Some(if name_lower == "and" {
arena.and(args)
} else {
arena.or(args)
}))
}
"not" => {
let [a] = args else {
return Err(self.error(format!(
"'{name}' takes exactly 1 argument, got {}",
args.len()
)));
};
self.boolean_operand(arena, name, *a)?;
Ok(Some(arena.not(*a)))
}
_ => Ok(None),
}
}
#[allow(clippy::too_many_arguments)]
fn make_sum_product(
&self,
arena: &mut Arena,
name: &str,
name_lower: &str,
body: ExprId,
var: ExprId,
lo: ExprId,
hi: ExprId,
) -> Result<ExprId, ParseError> {
if !matches!(arena.node(var), ExprNode::Symbol(_)) {
return Err(ParseError {
message: format!(
"the index of '{}' must be a symbol, got '{}'",
name,
arena.display(var)
),
position: self.lexer.pos,
});
}
Ok(arena.intern(if name_lower == "sum" {
ExprNode::Sum(body, var, lo, hi)
} else {
ExprNode::Product_(body, var, lo, hi)
}))
}
#[allow(clippy::too_many_arguments)]
fn call_4(
&self,
arena: &mut Arena,
name: &str,
name_lower: &str,
arg: ExprId,
arg2: ExprId,
arg3: ExprId,
arg4: ExprId,
) -> Result<ExprId, ParseError> {
match name_lower {
"series" => Ok(arena.intern(ExprNode::Series(arg, arg2, arg3, arg4))),
"integral" => Ok(arena.definite_integral(arg, arg2, arg3, arg4)),
"sum" | "product" => {
self.make_sum_product(arena, name, name_lower, arg, arg2, arg3, arg4)
}
"jacobi" => Ok(apply_named(arena, FN_JACOBI, &[arg, arg2, arg3, arg4])),
"betainc" => Ok(apply_named(arena, FN_BETAINC, &[arg, arg2, arg3, arg4])),
"betainc_regularized" => Ok(apply_named(
arena,
FN_BETAINC_REGULARIZED,
&[arg, arg2, arg3, arg4],
)),
_ => Err(ParseError {
message: format!(
"unknown 4-argument function '{}'. Supported: Series, Sum, Product, Integral, \
jacobi, betainc, betainc_regularized",
name
),
position: self.lexer.pos,
}),
}
}
fn call_3(
&self,
arena: &mut Arena,
name: &str,
name_lower: &str,
arg: ExprId,
arg2: ExprId,
arg3: ExprId,
) -> Result<ExprId, ParseError> {
match name_lower {
"limit" => Ok(arena.intern(ExprNode::Limit(arg, arg2, arg3))),
"laplacetransform" => Ok(arena.intern(ExprNode::LaplaceTransform(arg, arg2, arg3))),
"inverselaplacetransform" => {
Ok(arena.intern(ExprNode::InverseLaplaceTransform(arg, arg2, arg3)))
}
"residue" => Ok(arena.intern(ExprNode::Residue(arg, arg2, arg3))),
"dsolve" => Ok(arena.intern(ExprNode::DSolve(arg, arg2, arg3))),
"gegenbauer" => Ok(apply_named(arena, FN_GEGENBAUER, &[arg, arg2, arg3])),
"assoc_legendre" => Ok(apply_named(arena, FN_ASSOC_LEGENDRE, &[arg, arg2, arg3])),
"assoc_laguerre" => Ok(apply_named(arena, FN_ASSOC_LAGUERRE, &[arg, arg2, arg3])),
_ => Err(ParseError {
message: format!(
"unknown 3-argument function '{}'. Supported: Limit, LaplaceTransform, \
InverseLaplaceTransform, Residue, DSolve, min, max, gegenbauer, \
assoc_legendre, assoc_laguerre",
name
),
position: self.lexer.pos,
}),
}
}
fn call_2(
&self,
arena: &mut Arena,
name: &str,
name_lower: &str,
arg: ExprId,
arg2: ExprId,
) -> Result<ExprId, ParseError> {
match name_lower {
"log" => {
let ln_x = arena.ln(arg);
let ln_base = arena.ln(arg2);
Ok(arena.div(ln_x, ln_base))
}
"rootof" => Ok(arena.intern(ExprNode::RootOf(arg, arg2))),
"conditionset" => Ok(arena.intern(ExprNode::ConditionSet(arg, arg2))),
"integral" => Ok(arena.intern(ExprNode::Integral(arg, arg2))),
"atan2" => Ok(arena.atan2(arg, arg2)),
"polygamma" => Ok(arena.polygamma(arg, arg2)),
"kroneckerdelta" | "kronecker_delta" => Ok(arena.kronecker_delta(arg, arg2)),
"binomial" | "c" => Ok(arena.binomial(arg, arg2)),
"beta" | "b" => Ok(arena.beta(arg, arg2)),
"besselj" => Ok(arena.besselj(arg, arg2)),
"bessely" => Ok(arena.bessely(arg, arg2)),
"besseli" => Ok(arena.besseli(arg, arg2)),
"besselk" => Ok(arena.besselk(arg, arg2)),
"expint" => Ok(apply_named(arena, FN_EXPINT, &[arg, arg2])),
"lowergamma" => Ok(apply_named(arena, FN_LOWERGAMMA, &[arg, arg2])),
"uppergamma" => Ok(apply_named(arena, FN_UPPERGAMMA, &[arg, arg2])),
"polylog" => Ok(apply_named(arena, FN_POLYLOG, &[arg, arg2])),
"elliptic_f" => Ok(apply_named(arena, FN_ELLIPTIC_F, &[arg, arg2])),
"elliptic_pi" => Ok(apply_named(arena, FN_ELLIPTIC_PI, &[arg, arg2])),
_ => Err(ParseError {
message: format!(
"unknown 2-argument function '{}'. Supported: log, atan2, polygamma, \
binomial, beta, besselj, bessely, besseli, besselk, expint, lowergamma, \
uppergamma, polylog, elliptic_f, elliptic_pi, min, max, KroneckerDelta, \
RootOf, ConditionSet, Integral",
name
),
position: self.lexer.pos,
}),
}
}
fn call_1(
&self,
arena: &mut Arena,
name: &str,
name_lower: &str,
arg: ExprId,
) -> Result<ExprId, ParseError> {
match name_lower {
"sin" => Ok(arena.sin(arg)),
"cos" => Ok(arena.cos(arg)),
"tan" => Ok(arena.tan(arg)),
"exp" => Ok(arena.exp(arg)),
"ln" | "log" => Ok(arena.ln(arg)),
"sqrt" => Ok(arena.sqrt(arg)),
"cbrt" => Ok(arena.cbrt(arg)),
"abs" => Ok(arena.abs(arg)),
"asin" | "arcsin" => Ok(arena.asin(arg)),
"acos" | "arccos" => Ok(arena.acos(arg)),
"atan" | "arctan" => Ok(arena.atan(arg)),
"sinh" => Ok(arena.sinh(arg)),
"cosh" => Ok(arena.cosh(arg)),
"tanh" => Ok(arena.tanh(arg)),
"asinh" | "arcsinh" => Ok(arena.asinh(arg)),
"acosh" | "arccosh" => Ok(arena.acosh(arg)),
"atanh" | "arctanh" => Ok(arena.atanh(arg)),
"sign" | "sgn" => Ok(arena.sign(arg)),
"floor" => Ok(arena.floor(arg)),
"ceil" | "ceiling" => Ok(arena.ceiling(arg)),
"gamma" => Ok(arena.intern(crate::base::node::ExprNode::Gamma(arg))),
"erf" => Ok(arena.intern(crate::base::node::ExprNode::Erf(arg))),
"erfc" => Ok(arena.intern(crate::base::node::ExprNode::Erfc(arg))),
"heaviside" => Ok(arena.intern(crate::base::node::ExprNode::Heaviside(arg))),
"diracdelta" | "dirac_delta" => {
Ok(arena.intern(crate::base::node::ExprNode::DiracDelta(arg)))
}
"lambertw" | "w" => Ok(arena.intern(crate::base::node::ExprNode::LambertW(arg))),
"factorial" => Ok(arena.intern(crate::base::node::ExprNode::Factorial(arg))),
"digamma" => Ok(arena.intern(crate::base::node::ExprNode::Digamma(arg))),
"loggamma" => Ok(arena.intern(crate::base::node::ExprNode::LogGamma(arg))),
"cot" => {
let c = arena.cos(arg);
let s = arena.sin(arg);
Ok(arena.div(c, s))
}
"sec" => {
let c = arena.cos(arg);
Ok(arena.div(arena.one, c))
}
"csc" => {
let s = arena.sin(arg);
Ok(arena.div(arena.one, s))
}
"coth" => {
let c = arena.cosh(arg);
let s = arena.sinh(arg);
Ok(arena.div(c, s))
}
"sech" => {
let c = arena.cosh(arg);
Ok(arena.div(arena.one, c))
}
"csch" => {
let s = arena.sinh(arg);
Ok(arena.div(arena.one, s))
}
"acot" | "arccot" => {
let inv = arena.div(arena.one, arg);
Ok(arena.atan(inv))
}
"re" => Ok(arena.re(arg)),
"im" => Ok(arena.im(arg)),
"conjugate" | "conj" => Ok(arena.conjugate(arg)),
"arg" => Ok(arena.arg(arg)),
"si" => Ok(arena.si(arg)),
"ci" => Ok(arena.ci(arg)),
"ei" => Ok(arena.ei(arg)),
"li" => Ok(arena.li(arg)),
"zeta" => Ok(arena.zeta(arg)),
"erfi" => Ok(apply_named(arena, FN_ERFI, &[arg])),
"erfinv" => Ok(apply_named(arena, FN_ERFINV, &[arg])),
"erfcinv" => Ok(apply_named(arena, FN_ERFCINV, &[arg])),
"e1" => Ok(apply_named(arena, FN_EXPINT, &[arena.one, arg])),
"shi" => Ok(apply_named(arena, FN_SHI, &[arg])),
"chi" => Ok(apply_named(arena, FN_CHI, &[arg])),
"fresnels" => Ok(apply_named(arena, FN_FRESNELS, &[arg])),
"fresnelc" => Ok(apply_named(arena, FN_FRESNELC, &[arg])),
"dirichlet_eta" => Ok(apply_named(arena, FN_DIRICHLET_ETA, &[arg])),
"airyai" => Ok(apply_named(arena, FN_AIRYAI, &[arg])),
"airybi" => Ok(apply_named(arena, FN_AIRYBI, &[arg])),
"airyaiprime" => Ok(apply_named(arena, FN_AIRYAIPRIME, &[arg])),
"airybiprime" => Ok(apply_named(arena, FN_AIRYBIPRIME, &[arg])),
"elliptic_k" => Ok(apply_named(arena, FN_ELLIPTIC_K, &[arg])),
"elliptic_e" => Ok(apply_named(arena, FN_ELLIPTIC_E, &[arg])),
_ => Err(ParseError {
message: format!(
"unknown function '{}'. Supported: sin, cos, tan, cot, sec, csc, exp, ln, log, \
sqrt, cbrt, abs, asin, acos, atan, acot, sinh, cosh, tanh, coth, sech, csch, \
asinh, acosh, atanh, sign, floor, ceil, gamma, erf, erfc, heaviside, \
diracdelta, lambertw, factorial, digamma, loggamma, re, im, conjugate, arg, \
Si, Ci, Ei, li, zeta, polygamma, binomial, beta, besselj, bessely, besseli, \
besselk, erfi, erfinv, erfcinv, E1, expint, Shi, Chi, fresnels, fresnelc, \
lowergamma, uppergamma, polylog, dirichlet_eta, airyai, airybi, \
airyaiprime, airybiprime, elliptic_k, elliptic_e, elliptic_f, elliptic_pi, \
gegenbauer, jacobi, assoc_legendre, assoc_laguerre, betainc, \
betainc_regularized, min, max, \
KroneckerDelta, Limit, RootOf, ConditionSet, LaplaceTransform, \
InverseLaplaceTransform, Residue, DSolve, Series, Sum, Product, Integral",
name
),
position: self.lexer.pos,
}),
}
}
}
fn apply_named(arena: &mut Arena, name: &str, args: &[ExprId]) -> ExprId {
let sid = arena.symbols.intern(name);
arena.intern(ExprNode::Apply(sid, args.iter().copied().collect()))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::api::context::Context;
fn parse_and_display(input: &str) -> String {
let ctx = Context::new();
let ex = parse(&ctx, input).unwrap();
format!("{ex}")
}
#[test]
fn parse_integer() {
assert_eq!(parse_and_display("42"), "42");
}
#[test]
fn parse_symbol() {
assert_eq!(parse_and_display("x"), "x");
}
#[test]
fn parse_addition() {
assert_eq!(parse_and_display("x + y"), "x + y");
}
#[test]
fn parse_polynomial() {
let s = parse_and_display("x^2 + 2*x + 1");
assert!(s.contains("x^2") && s.contains("2*x"), "got: {s}");
}
#[test]
fn parse_function_sin() {
assert_eq!(parse_and_display("sin(x)"), "sin(x)");
}
#[test]
fn parse_nested_functions() {
assert_eq!(parse_and_display("sin(cos(x))"), "sin(cos(x))");
}
#[test]
fn parse_negation() {
let s = parse_and_display("-x");
assert!(s.contains("x") && s.starts_with('-'), "got: {s}");
}
#[test]
fn parse_power_right_assoc() {
let s = parse_and_display("x^2^3");
assert!(s.contains("x"), "got: {s}");
}
#[test]
fn parse_constant_pi() {
assert_eq!(parse_and_display("pi"), "pi");
}
#[test]
fn parse_precedence() {
let s = parse_and_display("2 + 3*x");
assert!(s.contains("3*x"), "got: {s}");
}
#[test]
fn parse_parens() {
let s = parse_and_display("(x + 1)^2");
assert!(
s.contains("(1 + x)^2") || s.contains("(x + 1)^2"),
"got: {s}"
);
}
#[test]
fn parse_division() {
let s = parse_and_display("x / y");
assert!(
s.contains("1/y") || s.contains("x*1/y") || s.contains("x/y") || s.contains("y^(-1)"),
"got: {s}"
);
}
#[test]
fn parse_empty_string_error() {
let ctx = Context::new();
assert!(parse(&ctx, "").is_err());
}
#[test]
fn parse_unknown_function_error() {
let ctx = Context::new();
assert!(parse(&ctx, "foo(x)").is_err());
}
#[test]
fn parse_unary_minus_precedence() {
let s = parse_and_display("-x^2");
assert!(s.contains("x^2"), "got: {s}");
}
#[test]
fn parse_subtraction() {
let s = parse_and_display("x - y");
assert!(s.contains("x") && s.contains("y"), "got: {s}");
}
#[test]
fn parse_multiple_operations() {
let s = parse_and_display("2*x + 3*y - z");
assert!(
s.contains("2*x") && s.contains("3*y") && s.contains("z"),
"got: {s}"
);
}
#[test]
fn parse_constant_e() {
assert_eq!(parse_and_display("e"), "E");
}
#[test]
fn parse_exp_function() {
let s = parse_and_display("exp(x)");
assert!(s.contains("x"), "got: {s}");
}
#[test]
fn parse_sqrt_function() {
let s = parse_and_display("sqrt(x)");
assert!(s.contains("x"), "got: {s}");
}
#[test]
fn parse_complex_nested() {
let s = parse_and_display("sin(x^2 + 1)");
assert!(s.contains("sin"), "got: {s}");
}
#[test]
fn parse_trailing_garbage_error() {
let ctx = Context::new();
assert!(parse(&ctx, "x )").is_err());
}
#[test]
fn parse_unmatched_paren_error() {
let ctx = Context::new();
assert!(parse(&ctx, "(x + 1").is_err());
}
#[test]
fn parse_unexpected_char_error() {
let ctx = Context::new();
assert!(parse(&ctx, "x & y").is_err());
}
#[test]
fn parse_empty_function_call_error() {
let ctx = Context::new();
assert!(parse(&ctx, "sin()").is_err());
}
#[test]
fn parse_abs_function() {
let s = parse_and_display("abs(x)");
assert!(s.contains("x"), "got: {s}");
}
#[test]
fn parse_ln_function() {
let s = parse_and_display("ln(x)");
assert!(s.contains("x"), "got: {s}");
}
#[test]
fn parse_log_alias() {
let s = parse_and_display("log(x)");
assert!(s.contains("x"), "got: {s}");
}
#[test]
fn parse_constant_infinity() {
let s = parse_and_display("inf");
assert!(
s.contains("oo") || s.contains("inf") || s.contains("∞"),
"got: {s}"
);
}
#[test]
fn parse_deeply_nested_parens() {
let s = parse_and_display("((((x))))");
assert_eq!(s, "x");
}
#[test]
fn parse_chained_additions() {
let s = parse_and_display("a + b + c + d");
assert!(
s.contains("a") && s.contains("b") && s.contains("c") && s.contains("d"),
"got: {s}"
);
}
#[test]
fn parse_chained_multiplications() {
let s = parse_and_display("a * b * c");
assert!(
s.contains("a") && s.contains("b") && s.contains("c"),
"got: {s}"
);
}
#[test]
fn parse_mixed_precedence() {
let s = parse_and_display("a + b * c");
assert!(s.contains("b*c") || s.contains("c*b"), "got: {s}");
}
#[test]
fn parse_float() {
let ctx = Context::new();
let result = parse(&ctx, "3.14").unwrap();
let s = format!("{result}");
assert!(
s.contains("157") || s.contains("3.14") || s.contains("314"),
"got: {s}"
);
}
#[test]
fn parse_implicit_mul_number_var() {
let ctx = Context::new();
let result = parse(&ctx, "2x").unwrap();
let s = format!("{result}");
assert!(s.contains("2") && s.contains("x"), "2x should be 2*x: {s}");
}
#[test]
fn parse_implicit_mul_var_paren() {
let ctx = Context::new();
assert!(parse(&ctx, "x(x+1)").is_err());
}
#[test]
fn parse_constant_pi_variants() {
assert_eq!(parse_and_display("pi"), "pi");
assert_eq!(parse_and_display("Pi"), "pi");
assert_eq!(parse_and_display("PI"), "pi");
}
#[test]
fn parse_constant_i_unit() {
let ctx = Context::new();
let result = parse(&ctx, "I").unwrap();
assert_eq!(format!("{result}"), "I");
}
#[test]
fn parse_constant_i_lowercase() {
let ctx = Context::new();
let result = parse(&ctx, "i").unwrap();
assert_eq!(format!("{result}"), "I");
}
#[test]
fn parse_euler_formula() {
let ctx = Context::new();
let result = parse(&ctx, "exp(I*pi)").unwrap();
let s = format!("{result}");
assert!(s.contains("I") && s.contains("pi"), "got: {s}");
}
#[test]
fn parse_float_times_var() {
let ctx = Context::new();
let result = parse(&ctx, "2.5*x").unwrap();
let s = format!("{result}");
assert!(
s.contains("x") && (s.contains("5/2") || s.contains("2.5") || s.contains("5*1/2")),
"got: {s}"
);
}
#[test]
fn parse_log_two_args() {
let ctx = Context::new();
let result = parse(&ctx, "log(x, 2)").unwrap();
let s = format!("{result}");
assert!(s.contains("x"), "log(x,2) should parse: {s}");
}
#[test]
fn parse_implicit_mul_number_paren() {
let ctx = Context::new();
let result = parse(&ctx, "3(x+1)").unwrap();
let s = format!("{result}");
assert!(
s.contains("3") && s.contains("x"),
"3(x+1) should be 3*(x+1): {s}"
);
}
#[test]
fn parse_implicit_mul_paren_paren() {
let ctx = Context::new();
let result = parse(&ctx, "(a)(b)").unwrap();
let s = format!("{result}");
assert!(
s.contains("a") && s.contains("b"),
"(a)(b) should be a*b: {s}"
);
}
#[test]
fn parse_implicit_mul_coeff_pi() {
let ctx = Context::new();
let result = parse(&ctx, "2pi").unwrap();
let s = format!("{result}");
assert!(
s.contains("2") && s.contains("pi"),
"2pi should be 2*pi: {s}"
);
}
#[test]
fn parse_large_integer() {
let ctx = Context::new();
let result = parse(&ctx, "99999999999999999999999999999").unwrap();
let s = format!("{result}");
assert_eq!(s, "99999999999999999999999999999");
}
#[test]
fn parse_integer_beyond_i64_max() {
let ctx = Context::new();
let result = parse(&ctx, "9223372036854775808").unwrap();
let s = format!("{result}");
assert_eq!(
s, "9223372036854775808",
"should handle integers > i64::MAX"
);
}
#[test]
fn parse_integer_beyond_i128_max() {
let ctx = Context::new();
let big = "123456789012345678901234567890123456789012345678901234567890";
let result = parse(&ctx, big).unwrap();
let s = format!("{result}");
assert_eq!(s, big, "should handle integers > i128::MAX");
}
#[test]
fn parse_large_integer_arithmetic() {
let ctx = Context::new();
let result = parse(&ctx, "1000000000000000000000000000000^2").unwrap();
let evaled = result.eval();
let s = format!("{evaled}");
assert_eq!(
s, "1000000000000000000000000000000000000000000000000000000000000",
"large integer exponentiation should be exact"
);
}
#[test]
fn parse_negative_large_integer() {
let ctx = Context::new();
let result = parse(&ctx, "-99999999999999999999999999999").unwrap();
let s = format!("{result}");
assert_eq!(s, "-99999999999999999999999999999");
}
#[test]
fn parse_many_decimal_places() {
let ctx = Context::new();
let result = parse(&ctx, "1.0000000000000000000001");
assert!(
result.is_ok(),
"parsing many decimal places should not panic"
);
}
#[test]
fn parse_decimal_exact_rational_simple() {
let ctx = Context::new();
let result = parse(&ctx, "0.5").unwrap();
let s = format!("{result}");
assert_eq!(s, "1/2", "0.5 should parse as exact rational 1/2, got: {s}");
}
#[test]
fn parse_decimal_exact_rational_quarter() {
let ctx = Context::new();
let result = parse(&ctx, "0.25").unwrap();
let s = format!("{result}");
assert_eq!(
s, "1/4",
"0.25 should parse as exact rational 1/4, got: {s}"
);
}
#[test]
fn parse_decimal_exact_rational_third_approx() {
let ctx = Context::new();
let result = parse(&ctx, "0.333").unwrap();
let s = format!("{result}");
assert_eq!(
s, "333/1000",
"0.333 should parse as exact rational 333/1000, got: {s}"
);
}
#[test]
fn parse_decimal_preserves_all_digits() {
let ctx = Context::new();
let input = "1.00000000000000000000000000000000001";
let result = parse(&ctx, input).unwrap();
let big_denom = parse(&ctx, "100000000000000000000000000000000000").unwrap();
let product = &result * &big_denom;
let s = format!("{}", product.eval());
assert_eq!(
s, "100000000000000000000000000000000001",
"1.00000000000000000000000000000000001 * 10^35 should be exact, got: {s}"
);
}
#[test]
fn parse_decimal_20_places_exact() {
let ctx = Context::new();
let result = parse(&ctx, "3.14159265358979323846").unwrap();
let s = format!("{result}");
assert_eq!(
s, "157079632679489661923/50000000000000000000",
"20-digit decimal should be exact reduced rational"
);
}
#[test]
fn parse_decimal_20_places_multiply_back() {
let ctx = Context::new();
let result = parse(&ctx, "3.14159265358979323846").unwrap();
let denom = parse(&ctx, "50000000000000000000").unwrap();
let product = (&result * &denom).eval();
let s = format!("{product}");
assert_eq!(
s, "157079632679489661923",
"rational * denominator should recover exact numerator"
);
}
#[test]
fn parse_decimal_50_places_exact() {
let ctx = Context::new();
let input = "3.14159265358979323846264338327950288419716939937510";
let result = parse(&ctx, input).unwrap();
let big = parse(&ctx, "100000000000000000000000000000000000000000000000000").unwrap();
let product = (&result * &big).eval();
let s = format!("{product}");
assert_eq!(
s, "314159265358979323846264338327950288419716939937510",
"50-digit decimal * 10^50 must recover exact integer (bc-verified)"
);
}
#[test]
fn parse_decimal_large_integer_part_and_fraction_exact() {
let ctx = Context::new();
let result = parse(
&ctx,
"123456789012345678901234567890.123456789012345678901234567890",
)
.unwrap();
let denom = parse(&ctx, "1000000000000000000000000000000").unwrap(); let product = (&result * &denom).eval();
let s = format!("{product}");
assert_eq!(
s, "123456789012345678901234567890123456789012345678901234567890",
"large decimal * 10^30 must recover exact integer (bc-verified)"
);
}
#[test]
fn parse_decimal_used_in_arithmetic() {
let ctx = Context::new();
let result = parse(&ctx, "0.1 + 0.2").unwrap();
let s = format!("{result}");
assert_eq!(s, "3/10", "0.1 + 0.2 should be exactly 3/10, got: {s}");
}
#[test]
fn parse_decimal_multiplication_exact() {
let ctx = Context::new();
let result = parse(&ctx, "0.1 * 0.1").unwrap();
let s = format!("{result}");
assert_eq!(s, "1/100", "0.1 * 0.1 should be exactly 1/100, got: {s}");
}
#[test]
fn parse_decimal_vs_fraction_equivalence() {
let ctx = Context::new();
let x = ctx.symbol("x");
let from_decimal = parse(&ctx, "2.5 * x").unwrap();
let five_halves = ctx.rational(5, 2);
let from_fraction = &five_halves * &x;
assert_eq!(
from_decimal, from_fraction,
"2.5*x and (5/2)*x should be identical expressions"
);
}
#[test]
fn parse_zero_integer() {
let ctx = Context::new();
let result = parse(&ctx, "0").unwrap();
let s = format!("{result}");
assert_eq!(s, "0");
}
#[test]
fn parse_zero_decimal() {
let ctx = Context::new();
let result = parse(&ctx, "0.0").unwrap();
let s = format!("{result}");
assert_eq!(s, "0", "0.0 should parse as 0, got: {s}");
}
#[test]
fn parse_leading_zeros_integer() {
let ctx = Context::new();
let result = parse(&ctx, "007").unwrap();
let s = format!("{result}");
assert_eq!(s, "7", "007 should parse as 7, got: {s}");
}
#[test]
fn parse_leading_zeros_decimal() {
let ctx = Context::new();
let result = parse(&ctx, "0.00100").unwrap();
let s = format!("{result}");
assert_eq!(s, "1/1000", "0.00100 should parse as 1/1000, got: {s}");
}
#[test]
fn parse_deep_nesting_limit() {
let ctx = Context::new();
let deep = "(".repeat(300) + "1" + &")".repeat(300);
let result = parse(&ctx, &deep);
assert!(
result.is_err(),
"deeply nested input should return error, not stack overflow"
);
}
#[test]
fn parse_moderate_nesting_ok() {
let ctx = Context::new();
let expr = "(".repeat(50) + "x + 1" + &")".repeat(50);
let result = parse(&ctx, &expr);
assert!(result.is_ok(), "50 levels of nesting should be fine");
}
#[test]
fn parse_decimal_with_variable_exact_coeff() {
let ctx = Context::new();
let result = parse(&ctx, "3.14 * x").unwrap();
let s = format!("{result}");
assert_eq!(
s, "157/50*x",
"3.14*x should have exact coefficient 157/50, got: {s}"
);
}
#[test]
fn parse_integer_one() {
let ctx = Context::new();
let result = parse(&ctx, "1").unwrap();
let s = format!("{result}");
assert_eq!(s, "1");
}
#[test]
fn parse_negative_decimal() {
let ctx = Context::new();
let result = parse(&ctx, "-0.5").unwrap();
let s = format!("{result}");
assert_eq!(s, "-1/2", "-0.5 should parse as -1/2, got: {s}");
}
#[test]
fn parse_beyond_i64_arithmetic() {
let ctx = Context::new();
let result = parse(&ctx, "9223372036854775808 + 1").unwrap();
let s = format!("{result}");
assert_eq!(
s, "9223372036854775809",
"i64::MAX+1 + 1 should be exact (bc-verified)"
);
}
#[test]
fn parse_beyond_i128() {
let ctx = Context::new();
let result = parse(&ctx, "170141183460469231731687303715884105728").unwrap();
let s = format!("{result}");
assert_eq!(
s, "170141183460469231731687303715884105728",
"i128::MAX+1 should parse and display exactly"
);
}
#[test]
fn parse_huge_integer_squared() {
let ctx = Context::new();
let base = parse(&ctx, "99999999999999999999999999999").unwrap();
let squared = base.powi(2).eval();
let s = format!("{squared}");
assert_eq!(
s, "9999999999999999999999999999800000000000000000000000000001",
"(10^29 - 1)^2 should be exact (bc-verified)"
);
}
#[test]
fn parse_tiny_number_exact() {
let ctx = Context::new();
let tiny = parse(&ctx, "0.000000000000000000000000000000000001").unwrap();
let big = parse(&ctx, "1000000000000000000000000000000000000").unwrap();
let product = (&tiny * &big).eval();
let s = format!("{product}");
assert_eq!(s, "1", "10^-36 * 10^36 should be exactly 1 (bc-verified)");
}
#[test]
fn parse_tiny_number_display() {
let ctx = Context::new();
let tiny = parse(&ctx, "0.000000000000000000000000000000000001").unwrap();
let s = format!("{tiny}");
assert_eq!(
s, "1/1000000000000000000000000000000000000",
"10^-36 should display as exact fraction (bc-verified)"
);
}
#[test]
fn parse_large_integer_addition() {
let ctx = Context::new();
let result = parse(&ctx, "123456789012345678901234567890 + 1").unwrap();
let s = format!("{result}");
assert_eq!(
s, "123456789012345678901234567891",
"large integer + 1 should be exact (bc-verified)"
);
}
#[test]
fn parse_sympy_power_operator() {
assert_eq!(parse_and_display("x**2"), parse_and_display("x^2"));
}
#[test]
fn parse_sympy_double_star_in_expression() {
assert_eq!(
parse_and_display("3*x**2 + 2"),
parse_and_display("3*x^2 + 2")
);
}
#[test]
fn parse_sympy_abs_capital() {
assert_eq!(parse_and_display("Abs(x)"), parse_and_display("abs(x)"));
}
#[test]
fn parse_sympy_full_expression() {
let ctx = Context::new();
let result = parse(&ctx, "x**2*log(x)/2 - x**2/4").unwrap();
let s = format!("{result}");
assert!(
s.contains("ln") && s.contains("x"),
"should parse SymPy integrate output: {s}"
);
}
#[test]
fn parse_sympy_trig_identity() {
let ctx = Context::new();
let result = parse(&ctx, "sin(x)**2 + cos(x)**2").unwrap();
let simplified = result.simplify();
assert_eq!(format!("{simplified}"), "1");
}
#[test]
fn parse_named_constants() {
let ctx = Context::new();
assert_eq!(parse(&ctx, "EulerGamma").unwrap(), ctx.euler_gamma());
assert_eq!(parse(&ctx, "Catalan").unwrap(), ctx.catalan());
assert_eq!(parse(&ctx, "GoldenRatio").unwrap(), ctx.golden_ratio());
assert_eq!(parse(&ctx, "zoo").unwrap(), ctx.complex_infinity());
assert_eq!(
parse_and_display("EulerGamma + Catalan"),
"EulerGamma + Catalan"
);
}
#[test]
fn parse_complex_functions() {
let ctx = Context::new();
let z = ctx.symbol("z");
assert_eq!(parse(&ctx, "re(z)").unwrap(), z.re());
assert_eq!(parse(&ctx, "im(z)").unwrap(), z.im());
assert_eq!(parse(&ctx, "conjugate(z)").unwrap(), z.conjugate());
assert_eq!(parse(&ctx, "conj(z)").unwrap(), z.conjugate());
assert_eq!(parse(&ctx, "arg(z)").unwrap(), z.arg());
assert_eq!(parse_and_display("Re(3 + 4*I)"), "3");
assert_eq!(parse_and_display("im(3 + 4*I)"), "4");
}
#[test]
fn parse_special_functions() {
let ctx = Context::new();
let x = ctx.symbol("x");
let n = ctx.symbol("n");
assert_eq!(parse(&ctx, "Si(x)").unwrap(), x.si());
assert_eq!(parse(&ctx, "Ci(x)").unwrap(), x.ci());
assert_eq!(parse(&ctx, "Ei(x)").unwrap(), x.ei());
assert_eq!(parse(&ctx, "li(x)").unwrap(), x.li());
assert_eq!(parse(&ctx, "zeta(x)").unwrap(), x.zeta());
assert_eq!(parse(&ctx, "polygamma(n, x)").unwrap(), x.polygamma(&n));
assert_eq!(
parse(&ctx, "KroneckerDelta(n, x)").unwrap(),
x.kronecker_delta(&n)
);
assert_eq!(
parse(&ctx, "kronecker_delta(n, x)").unwrap(),
x.kronecker_delta(&n)
);
assert_eq!(parse_and_display("zeta(2)"), "1/6*pi^2");
assert_eq!(parse_and_display("Si(0)"), "0");
assert_eq!(parse_and_display("atan2(1, 1)"), "atan2(1, 1)");
}
#[test]
fn parse_bool_relations_and_connectives() {
let ctx = Context::new();
let (x, y) = (ctx.symbol("x"), ctx.symbol("y"));
let zero = ctx.int(0);
let one = ctx.int(1);
assert_eq!(parse_bool(&ctx, "x > 0").unwrap(), x.gt(&zero));
assert_eq!(parse_bool(&ctx, "x >= 0").unwrap(), x.ge(&zero));
assert_eq!(parse_bool(&ctx, "x < 1").unwrap(), x.lt(&one));
assert_eq!(parse_bool(&ctx, "x <= 1").unwrap(), x.le(&one));
assert_eq!(parse_bool(&ctx, "x == 1").unwrap(), x.eq_expr(&one));
assert_eq!(parse_bool(&ctx, "x != 1").unwrap(), x.ne_expr(&one));
assert_eq!(
parse_bool(&ctx, "x > 0 & x < 1").unwrap(),
x.gt(&zero).and(&x.lt(&one))
);
assert_eq!(
parse_bool(&ctx, "x > 0 && x < 1").unwrap(),
parse_bool(&ctx, "x > 0 and x < 1").unwrap()
);
assert_eq!(
parse_bool(&ctx, "x > 0 | y > 0").unwrap(),
x.gt(&zero).or(&y.gt(&zero))
);
assert_eq!(
parse_bool(&ctx, "x > 0 || y > 0").unwrap(),
parse_bool(&ctx, "x > 0 or y > 0").unwrap()
);
assert_eq!(parse_bool(&ctx, "~(x > 0)").unwrap(), x.gt(&zero).not());
assert_eq!(parse_bool(&ctx, "!(x > 0)").unwrap(), x.gt(&zero).not());
assert_eq!(parse_bool(&ctx, "not x > 0").unwrap(), x.gt(&zero).not());
assert_eq!(parse_bool(&ctx, "True").unwrap().to_string(), "True");
assert_eq!(parse_bool(&ctx, "false").unwrap().to_string(), "False");
}
#[test]
fn parse_bool_precedence() {
let ctx = Context::new();
let (x, y) = (ctx.symbol("x"), ctx.symbol("y"));
let zero = ctx.int(0);
let one = ctx.int(1);
assert_eq!(
parse_bool(&ctx, "x > 0 & x < 1 | y == 0").unwrap(),
x.gt(&zero).and(&x.lt(&one)).or(&y.eq_expr(&zero))
);
assert_eq!(
parse_bool(&ctx, "x > 0 | x < 1 & y == 0").unwrap(),
x.gt(&zero).or(&x.lt(&one).and(&y.eq_expr(&zero)))
);
assert_eq!(
parse_bool(&ctx, "x + 1 > 2*y").unwrap(),
(&x + 1).gt(&(2 * &y))
);
assert_eq!(
parse_bool(&ctx, "not x > 0 & y > 0").unwrap(),
x.gt(&zero).not().and(&y.gt(&zero))
);
assert_eq!(
parse_bool(&ctx, "x > 0 & y > 0 & x < 1")
.unwrap()
.to_string(),
"x > 0 & y > 0 & 1 > x"
);
assert_eq!(
parse_bool(&ctx, "(x > 0 | y > 0) & x < 1").unwrap(),
x.gt(&zero).or(&y.gt(&zero)).and(&x.lt(&one))
);
}
#[test]
fn parse_bool_function_forms() {
let ctx = Context::new();
let (x, y) = (ctx.symbol("x"), ctx.symbol("y"));
let one = ctx.int(1);
assert_eq!(parse_bool(&ctx, "Eq(x, 1)").unwrap(), x.eq_expr(&one));
assert_eq!(parse_bool(&ctx, "Ne(x, 1)").unwrap(), x.ne_expr(&one));
assert_eq!(parse_bool(&ctx, "Lt(x, 1)").unwrap(), x.lt(&one));
assert_eq!(parse_bool(&ctx, "Le(x, 1)").unwrap(), x.le(&one));
assert_eq!(parse_bool(&ctx, "Gt(x, 1)").unwrap(), x.gt(&one));
assert_eq!(parse_bool(&ctx, "Ge(x, 1)").unwrap(), x.ge(&one));
assert_eq!(
parse_bool(&ctx, "And(x > 1, y > 1)").unwrap(),
x.gt(&one).and(&y.gt(&one))
);
assert_eq!(
parse_bool(&ctx, "Or(x > 1, y > 1)").unwrap(),
x.gt(&one).or(&y.gt(&one))
);
assert_eq!(parse_bool(&ctx, "Not(x > 1)").unwrap(), x.gt(&one).not());
}
#[test]
fn parse_bool_errors() {
let ctx = Context::new();
assert!(parse_bool(&ctx, "x + 1").is_err());
assert!(parse_bool(&ctx, "0 < x < 1").is_err());
assert!(parse_bool(&ctx, "(x > 0) + 1").is_err());
assert!(parse_bool(&ctx, "x & y").is_err());
assert!(parse_bool(&ctx, "not x").is_err());
assert!(parse_bool(&ctx, "(x > 0) > 1").is_err());
assert!(parse_bool(&ctx, "-(x > 0)").is_err());
assert!(parse_bool(&ctx, "x > 0 &").is_err());
}
#[test]
fn parse_strict_rejects_relations_and_keeps_keywords_as_symbols() {
let ctx = Context::new();
assert!(parse(&ctx, "x > 0").is_err());
assert!(parse(&ctx, "x & y").is_err());
assert!(parse(&ctx, "~x").is_err());
assert!(parse(&ctx, "x != 1").is_err());
assert_eq!(parse(&ctx, "and").unwrap(), ctx.symbol("and"));
assert_eq!(parse(&ctx, "True").unwrap(), ctx.symbol("True"));
assert!(parse(&ctx, "Sum(k, k=1..3)").is_ok());
}
#[test]
fn parse_implicit_application() {
let ctx = Context::new();
let (x, y) = (ctx.symbol("x"), ctx.symbol("y"));
assert_eq!(parse_implicit(&ctx, "sin x").unwrap(), x.sin());
assert_eq!(parse_implicit(&ctx, "2 sin x").unwrap(), 2 * &x.sin());
assert_eq!(parse_implicit(&ctx, "sin 2x").unwrap(), (2 * &x).sin());
assert_eq!(parse_implicit(&ctx, "sin x^2").unwrap(), x.powi(2).sin());
assert_eq!(
parse_implicit(&ctx, "sin x cos y").unwrap(),
&x.sin() * &y.cos()
);
assert_eq!(parse_implicit(&ctx, "sin x + 1").unwrap(), &x.sin() + 1);
assert_eq!(parse_implicit(&ctx, "sin x/2").unwrap(), (&x / 2).sin());
assert_eq!(parse_implicit(&ctx, "sin x y").unwrap(), (&x * &y).sin());
assert_eq!(parse_implicit(&ctx, "exp (x) y").unwrap(), &x.exp() * &y);
assert_eq!(parse_implicit(&ctx, "sqrt 2").unwrap(), ctx.int(2).sqrt());
assert_eq!(parse_implicit(&ctx, "ln x^2").unwrap(), x.powi(2).ln());
assert_eq!(
parse_implicit(&ctx, "sin(x cos y)").unwrap(),
(&x * &y.cos()).sin()
);
}
#[test]
fn parse_implicit_products() {
let ctx = Context::new();
let (x, y, z) = (ctx.symbol("x"), ctx.symbol("y"), ctx.symbol("z"));
assert_eq!(
parse_implicit(&ctx, "2x + 3(y-1)").unwrap(),
2 * &x + 3 * (&y - 1)
);
assert_eq!(parse_implicit(&ctx, "x y z").unwrap(), &x * &y * &z);
assert_eq!(parse_implicit(&ctx, "x(x+1)").unwrap(), &x * (&x + 1));
assert_eq!(
parse_implicit(&ctx, "(x+1)(x-1)").unwrap(),
(&x + 1) * (&x - 1)
);
let f = ctx.symbol("f");
assert_eq!(parse_implicit(&ctx, "f(x)").unwrap(), &f * &x);
let c = ctx.symbol("c");
assert_eq!(parse_implicit(&ctx, "c(x+1)").unwrap(), &c * (&x + 1));
assert_eq!(
parse_implicit(&ctx, "binomial(x, 2)").unwrap(),
parse(&ctx, "C(x, 2)").unwrap()
);
assert_eq!(parse_implicit(&ctx, "pi(x)").unwrap(), ctx.pi() * &x);
assert_eq!(parse_implicit(&ctx, "2 pi x").unwrap(), 2 * ctx.pi() * &x);
}
#[test]
fn parse_implicit_errors_and_limits() {
let ctx = Context::new();
assert!(parse_implicit(&ctx, "sin + 1").is_err());
assert!(parse_implicit(&ctx, "sin").is_err());
let (re, x) = (ctx.symbol("re"), ctx.symbol("x"));
assert_eq!(parse_implicit(&ctx, "re x").unwrap(), &re * &x);
assert_eq!(parse_implicit(&ctx, "re(x)").unwrap(), x.re());
assert!(parse_implicit(&ctx, "x > 0").is_err());
}
#[test]
fn known_functions_table_matches_call_tables() {
let ctx = Context::new();
for name in KNOWN_FUNCTIONS {
let ok = [
format!("{name}(x)"),
format!("{name}(x, y)"),
format!("{name}(x, y, z)"),
format!("{name}(x, y, z, w)"),
]
.iter()
.any(|s| parse(&ctx, s).is_ok());
assert!(ok, "`{name}` is listed but no arity parses");
}
for name in KNOWN_FUNCTIONS {
assert!(
!is_implicit_unary_function(name) || parse(&ctx, &format!("{name}(x)")).is_ok(),
"`{name}` is applied implicitly but is not a unary function"
);
}
}
}