use std::sync::Arc;
use num_bigint::BigInt;
use num_rational::Ratio;
use smallvec::SmallVec;
use crate::api::context::Context;
use crate::api::expr::Ex;
use crate::base::arena::{
Arena, FN_AIRYAI, FN_AIRYAIPRIME, FN_AIRYBI, FN_AIRYBIPRIME, FN_ASSOC_LAGUERRE,
FN_ASSOC_LEGENDRE, 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 mut parser = Parser::new(input);
let id = 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,
});
}
Ok(result)
})?;
Ok(Ex::from_raw_parts(ctx.id, Arc::clone(&ctx.inner), id))
}
#[derive(Debug, Clone, PartialEq)]
enum Token {
Int(BigInt),
Rational(Ratio<BigInt>),
Ident(String),
Plus,
Minus,
Star,
Slash,
Caret,
LParen,
RParen,
Comma,
Bang,
Eq,
DotDot,
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 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;
Ok(Token::Bang)
}
b'=' => {
self.pos += 1;
Ok(Token::Eq)
}
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,
}),
}
}
}
struct Parser<'a> {
lexer: Lexer<'a>,
current: Token,
depth: usize,
}
impl<'a> Parser<'a> {
fn new(input: &'a str) -> Self {
let mut lexer = Lexer::new(input);
let current = lexer.next_token().unwrap_or(Token::Eof);
Parser {
lexer,
current,
depth: 0,
}
}
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 (op, l_bp, r_bp, implicit) = match &self.current {
Token::Plus => ('+', 1, 2, false),
Token::Minus => ('-', 1, 2, false),
Token::Star => ('*', 3, 4, false),
Token::Slash => ('/', 3, 4, false),
Token::Caret => ('^', 8, 7, false), Token::Int(_) | Token::Rational(_) | Token::Ident(_) | Token::LParen => {
('*', 3, 4, true)
}
_ => break,
};
if l_bp < min_bp {
break;
}
if !implicit {
self.advance()?;
}
let rhs = self.parse_expr(arena, r_bp)?;
lhs = match op {
'+' => arena.add(&[lhs, rhs]),
'-' => arena.sub(lhs, rhs),
'*' => arena.mul(&[lhs, rhs]),
'/' => arena.div(lhs, rhs),
'^' => arena.pow(lhs, rhs),
_ => unreachable!(),
};
}
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()?;
if self.current == Token::LParen {
return self.parse_function_call(arena, &name);
}
match name.as_str() {
"pi" | "Pi" | "PI" => Ok(arena.pi),
"e" | "E" => Ok(arena.e_const),
"I" | "i" => Ok(arena.i_unit),
"inf" | "oo" | "Inf" => Ok(arena.infinity),
"zoo" => Ok(arena.complex_infinity),
"nan" => Ok(arena.nan),
"EulerGamma" | "euler_gamma" => Ok(arena.euler_gamma),
"Catalan" => Ok(arena.catalan),
"GoldenRatio" | "golden_ratio" => Ok(arena.golden_ratio),
_ => Ok(arena.symbol(&name)),
}
}
Token::Minus => {
self.advance()?;
let operand = self.parse_expr(arena, 5)?;
Ok(arena.neg(operand))
}
Token::LParen => {
self.advance()?;
let inner = self.parse_expr(arena, 0)?;
self.expect(&Token::RParen)?;
Ok(inner)
}
other => Err(ParseError {
message: format!("expected expression, got {:?}", other),
position: self.lexer.pos,
}),
}
}
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 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,
}),
}
}
#[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])),
_ => Err(ParseError {
message: format!(
"unknown 4-argument function '{}'. Supported: Series, Sum, Product, Integral, \
jacobi",
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, 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)");
}
}