diff --git a/libraries/math-parser/src/ast.rs b/libraries/math-parser/src/ast.rs index 73009205991..2020d12fcd2 100644 --- a/libraries/math-parser/src/ast.rs +++ b/libraries/math-parser/src/ast.rs @@ -68,7 +68,7 @@ pub enum UnaryOp { } /// The tree as written, before each subexpression's sort is read from its spelling. -#[derive(Debug, PartialEq)] +#[derive(Debug, Clone, PartialEq)] pub enum Syntax { Lit(Literal), Var(String), @@ -111,15 +111,40 @@ pub enum Syntax { from: Box, to: Box, }, + /// An expression with the names its `where` clause defines, like `a + f(2) where a = 1, f(t) = t^2`. + Where { + body: Box, + bindings: Vec, + }, + /// A call whose parentheses end with a `where` clause, boxed so a tree's every node stays small. + CallWhere(Box), +} + +/// A call whose parentheses end with a `where` clause, like `max(a, b where a = 1)`, whose names every argument may read but +/// the function's name, outside the parentheses, can't. +#[derive(Debug, Clone, PartialEq)] +pub struct CallWhere { + pub name: String, + pub arguments: Vec, + pub bindings: Vec, } /// One case of a piecewise, the value it takes where its condition holds. -#[derive(Debug, PartialEq)] +#[derive(Debug, Clone, PartialEq)] pub struct Case { pub value: Syntax, pub condition: Syntax, } +/// One definition in a `where` clause: a value like `a = 1`, or a function like `f(t) = t^2`. +#[derive(Debug, Clone, PartialEq)] +pub struct Binding { + pub name: String, + /// The function's parameters, of which a value has none. + pub parameters: Vec, + pub value: Syntax, +} + /// A parsed expression, whose every subexpression has the sort its spelling fixes: a value or a matrix. #[derive(Debug)] pub enum Node { @@ -174,6 +199,18 @@ pub enum ValueNode { matrices: Vec, distinct: bool, }, + /// A value a `where` clause defines, or a parameter of the function being called. + Local(Local), + /// A call of a function a `where` clause defines. + Call { + function: Local, + arguments: Vec, + }, + /// An expression within the names its `where` clause defines. + Where { + clause: Box, + body: Box, + }, } /// A subexpression evaluating to a matrix. @@ -214,6 +251,18 @@ pub enum MatrixNode { cases: Vec>, otherwise: Option>, }, + /// A matrix a `where` clause defines, or a parameter of the function being called. + Local(Local), + /// A call of a function a `where` clause defines. + Call { + function: Local, + arguments: Vec, + }, + /// An expression within the names its `where` clause defines. + Where { + clause: Box, + body: Box, + }, } /// One case of a sorted piecewise, whose condition is a value whatever the sort of its cases. @@ -222,3 +271,18 @@ pub struct SortedCase { pub value: T, pub condition: ValueNode, } + +/// Where a name a `where` clause or a function's parameters define lives: its scope, counted outward from the innermost, and +/// its position among that scope's values or functions. +#[derive(Debug, Clone, Copy, PartialEq)] +pub struct Local { + pub depth: usize, + pub index: usize, +} + +/// A sorted `where` clause: each value's definition, and each function's body over its parameters. +#[derive(Debug)] +pub struct Clause { + pub values: Vec, + pub functions: Vec, +} diff --git a/libraries/math-parser/src/executer.rs b/libraries/math-parser/src/executer.rs index 6302b82265b..d90385560b1 100644 --- a/libraries/math-parser/src/executer.rs +++ b/libraries/math-parser/src/executer.rs @@ -1,4 +1,4 @@ -use crate::ast::{BinaryOp, Literal, MatrixNode, Node, SortedCase, UnaryOp, ValueNode}; +use crate::ast::{BinaryOp, Clause, Literal, Local, MatrixNode, Node, SortedCase, UnaryOp, ValueNode}; use crate::constants::{Builtin, MatrixToValue, ValueOfRegions, builtin_function, suffixed_function}; use crate::context::{EvalContext, FunctionProvider, ValueProvider}; use crate::lexer::Constant; @@ -6,6 +6,7 @@ use crate::matrix::{Matrix, Region}; use crate::object::Object; use crate::quaternion::Quaternion; use crate::value::{Number, Value}; +use std::cell::OnceCell; use thiserror::Error; #[derive(Debug, Error)] @@ -105,14 +106,147 @@ fn resolve_matrix(context: &EvalContext { + context: &'a EvalContext, + frame: Option<&'a Frame<'a>>, +} + +impl<'a, V: ValueProvider, F: FunctionProvider> Scope<'a, V, F> { + fn root(context: &'a EvalContext) -> Self { + Self { context, frame: None } + } + + /// The frame `depth` levels out from the innermost, which always exists for a tree from the sort pass. + fn frame(&self, depth: usize) -> Result<&'a Frame<'a>, EvalError> { + let mut frame = self.frame; + for _ in 0..depth { + frame = frame.and_then(|frame| frame.parent); + } + frame.ok_or(EvalError::OperatorTypeError) + } +} + +/// The definitions of one `where` clause, or the arguments of one call, each evaluated on its first read and at most once. +struct Frame<'a> { + /// The frame lexically around this one, which for a call is the frame of the clause defining the function. + parent: Option<&'a Frame<'a>>, + definitions: Definitions<'a>, + /// Each definition's result, once read. + results: &'a [OnceCell], +} + +enum Definitions<'a> { + /// A clause's definitions, evaluated within the clause so each may read the others. + Clause(&'a Clause), + /// A call's arguments, evaluated where the call is. + Arguments { arguments: &'a [Node], caller: Option<&'a Frame<'a>> }, +} + +/// A read through the names `where` clauses define, of the sort `T` evaluates to. +enum Bound<'a, T> { + Local(Local), + Call(Local, &'a [Node]), + Where(&'a Clause, &'a T), +} + +// Each sort reads through `where` definitions by an outlined function, keeping its frames out of the recursive evaluator's, and +// called from several sites since wasm-opt inlines a function with one call site +#[inline(never)] +fn bound_value(scope: &Scope<'_, V, F>, bound: Bound) -> Result { + match bound { + Bound::Local(local) => read(scope, local)?.as_value().copied().ok_or(EvalError::OperatorTypeError), + Bound::Call(function, arguments) => call(scope, function, arguments, |body, scope| match body { + Node::Value(body) => body.eval_in(scope), + Node::Matrix(_) => Err(EvalError::OperatorTypeError), + }), + Bound::Where(clause, body) => enter(scope, scope.frame, Definitions::Clause(clause), clause.values.len(), |scope| body.eval_in(scope)), + } +} + +#[inline(never)] +fn bound_matrix(scope: &Scope<'_, V, F>, bound: Bound) -> Result { + match bound { + Bound::Local(local) => read(scope, local)?.as_matrix().copied().ok_or(EvalError::OperatorTypeError), + Bound::Call(function, arguments) => call(scope, function, arguments, |body, scope| match body { + Node::Matrix(body) => body.eval_in(scope), + Node::Value(_) => Err(EvalError::OperatorTypeError), + }), + Bound::Where(clause, body) => enter(scope, scope.frame, Definitions::Clause(clause), clause.values.len(), |scope| body.eval_in(scope)), + } +} + +/// A definition or argument's result, evaluated on its first read: a clause's definition within the clause, so it can read the +/// others, and an argument where the call is. +fn read<'a, V: ValueProvider, F: FunctionProvider>(scope: &Scope<'a, V, F>, local: Local) -> Result<&'a Object, EvalError> { + let frame = scope.frame(local.depth)?; + let result = frame.results.get(local.index).ok_or(EvalError::OperatorTypeError)?; + if let Some(object) = result.get() { + return Ok(object); + } + + let (definition, within) = match frame.definitions { + Definitions::Clause(clause) => (clause.values.get(local.index), Some(frame)), + Definitions::Arguments { arguments, caller } => (arguments.get(local.index), caller), + }; + let object = definition.ok_or(EvalError::OperatorTypeError)?.eval_in(&Scope { + context: scope.context, + frame: within, + })?; + Ok(result.get_or_init(|| object)) +} + +/// Runs a function a `where` clause defines, whose body reads its parameters within the scope defining the function, lexically. +fn call<'a, V: ValueProvider, F: FunctionProvider, T>( + scope: &Scope<'a, V, F>, + function: Local, + arguments: &'a [Node], + evaluate: impl FnOnce(&Node, &Scope<'_, V, F>) -> Result, +) -> Result { + let definer = scope.frame(function.depth)?; + let Definitions::Clause(clause) = definer.definitions else { + return Err(EvalError::OperatorTypeError); + }; + let body = clause.functions.get(function.index).ok_or(EvalError::OperatorTypeError)?; + + let definitions = Definitions::Arguments { arguments, caller: scope.frame }; + enter(scope, Some(definer), definitions, arguments.len(), |scope| evaluate(body, scope)) +} + +/// Evaluates a body within a new frame of unread definitions, held on the stack when there are few. +fn enter<'a, V: ValueProvider, F: FunctionProvider, T>( + scope: &Scope<'a, V, F>, + parent: Option<&'a Frame<'a>>, + definitions: Definitions<'a>, + count: usize, + body: impl FnOnce(&Scope<'_, V, F>) -> Result, +) -> Result { + const STACK_RESULTS: usize = 4; + let stack_results: [OnceCell; STACK_RESULTS]; + let heap_results: Vec>; + let results: &[OnceCell] = if count <= STACK_RESULTS { + stack_results = Default::default(); + &stack_results[..count] + } else { + heap_results = (0..count).map(|_| OnceCell::new()).collect(); + &heap_results + }; + + let frame = Frame { parent, definitions, results }; + body(&Scope { + context: scope.context, + frame: Some(&frame), + }) +} + /// The case whose condition holds, or `None` where none does: every condition is evaluated and must be a truth value, and /// at most one may hold, since the cases are unordered. #[inline(always)] -fn holding_case<'a, T, V: ValueProvider, F: FunctionProvider>(context: &EvalContext, cases: &'a [SortedCase]) -> Result, EvalError> { +fn holding_case<'a, T, V: ValueProvider, F: FunctionProvider>(scope: &Scope<'_, V, F>, cases: &'a [SortedCase]) -> Result, EvalError> { let mut holding = None; let mut overlapping = false; for case in cases { - let Value::Number(condition) = case.condition.eval(context)?; + let Value::Number(condition) = case.condition.eval_in(scope)?; match condition.as_bool() { Some(false) => {} Some(true) if holding.is_none() => holding = Some(&case.value), @@ -133,10 +267,10 @@ enum Operand { } impl Node { - fn operand(&self, context: &EvalContext) -> Result { + fn operand(&self, scope: &Scope<'_, V, F>) -> Result { match self { - Node::Value(value) => value.eval(context).map(Operand::Value), - Node::Matrix(matrix) => matrix.eval(context).map(Operand::Matrix), + Node::Value(value) => value.eval_in(scope).map(Operand::Value), + Node::Matrix(matrix) => matrix.eval_in(scope).map(Operand::Matrix), } } } @@ -177,17 +311,28 @@ impl Node { Node::Matrix(matrix) => matrix.eval(context).map(Object::from), } } + + fn eval_in(&self, scope: &Scope<'_, V, F>) -> Result { + match self { + Node::Value(value) => value.eval_in(scope).map(Object::Value), + Node::Matrix(matrix) => matrix.eval_in(scope).map(Object::from), + } + } } impl ValueNode { pub fn eval(&self, context: &EvalContext) -> Result { + self.eval_in(&Scope::root(context)) + } + + fn eval_in(&self, scope: &Scope<'_, V, F>) -> Result { match self { ValueNode::Lit(lit) => match lit { Literal::Integer(integer) => Ok(Value::from_i64(*integer)), Literal::Float(float) => Ok(Value::from_f64(*float)), }, - ValueNode::BinOp { lhs, op, rhs } => match (lhs.eval(context)?, rhs.eval(context)?) { + ValueNode::BinOp { lhs, op, rhs } => match (lhs.eval_in(scope)?, rhs.eval_in(scope)?) { (Value::Number(lhs), Value::Number(rhs)) => { // Logic rejects operands that aren't truth values, while the other operators reject operand types they don't support let rejected = if matches!(op, BinaryOp::And | BinaryOp::Or) { @@ -198,17 +343,17 @@ impl ValueNode { settle(Value::Number(lhs.binary_op(*op, rhs).ok_or(rejected)?)) } }, - ValueNode::UnaryOp { expr, op } => match expr.eval(context)? { + ValueNode::UnaryOp { expr, op } => match expr.eval_in(scope)? { Value::Number(num) => { let rejected = if *op == UnaryOp::Not { EvalError::NotATruthValue } else { EvalError::OperatorTypeError }; settle(Value::Number(num.unary_op(*op).ok_or(rejected)?)) } }, ValueNode::Comparison { first, rest } => { - let Value::Number(first) = first.eval(context)?; + let Value::Number(first) = first.eval_in(scope)?; let rest = rest .iter() - .map(|(op, operand)| operand.eval(context).map(|Value::Number(number)| (*op, number))) + .map(|(op, operand)| operand.eval_in(scope).map(|Value::Number(number)| (*op, number))) .collect::, EvalError>>()?; // A `!=` chain asserts every pair distinct, while the ordered chains assert each adjacent pair's relation; every pair is checked so an unsupported comparison errors regardless of the others @@ -227,7 +372,7 @@ impl ValueNode { Ok(Value::from_bool(holds)) } ValueNode::Var(name) => { - let value = resolve_value(context, name).ok_or_else(|| EvalError::MissingValue(name.clone()))?; + let value = resolve_value(scope.context, name).ok_or_else(|| EvalError::MissingValue(name.clone()))?; canonical_host_value(name, value) } ValueNode::FnCall { name, expr } => { @@ -236,11 +381,11 @@ impl ValueNode { let heap_values: Vec; let values: &[Value] = if expr.len() <= stack_values.len() { for (slot, argument) in stack_values.iter_mut().zip(expr) { - *slot = argument.eval(context)?; + *slot = argument.eval_in(scope)?; } &stack_values[..expr.len()] } else { - heap_values = expr.iter().map(|argument| argument.eval(context)).collect::, EvalError>>()?; + heap_values = expr.iter().map(|argument| argument.eval_in(scope)).collect::, EvalError>>()?; &heap_values }; @@ -250,7 +395,7 @@ impl ValueNode { None => (false, name.as_str()), }; - if !prefixed && let Some(value) = context.run_function(bare_name, values) { + if !prefixed && let Some(value) = scope.context.run_function(bare_name, values) { settle(canonical_host_value(bare_name, value)?) } else if let Some(Builtin::Values { function, .. }) = builtin_function(bare_name) { settle(function(values).ok_or(EvalError::TypeError)?) @@ -258,7 +403,7 @@ impl ValueNode { // A base-suffixed call like `log10(x)` runs the two-argument form with the suffix baked in as its second argument let [value] = values else { return Err(EvalError::TypeError) }; settle(function(&[*value, Value::from_f64(base)]).ok_or(EvalError::TypeError)?) - } else if let Some(value) = resolve_value(context, name) + } else if let Some(value) = resolve_value(scope.context, name) && let [Value::Number(argument)] = values { // A known value applied to one argument is implicit multiplication, so `x(2)` matches `2(3)` and `i(16)` @@ -269,15 +414,18 @@ impl ValueNode { } } // Only the chosen value is evaluated, so an error in any other case's value is never raised - ValueNode::Piecewise { cases, otherwise } => match (holding_case(context, cases)?, otherwise) { - (Some(value), _) => value.eval(context), - (None, Some(otherwise)) => otherwise.eval(context), + ValueNode::Piecewise { cases, otherwise } => match (holding_case(scope, cases)?, otherwise) { + (Some(value), _) => value.eval_in(scope), + (None, Some(otherwise)) => otherwise.eval_in(scope), (None, None) => Err(EvalError::NoCaseHolds), }, - ValueNode::Apply { matrix, value } => value_of_matrix(context, MatrixValueCase::Apply(matrix, value)), - ValueNode::OfMatrix { function, matrix } => value_of_matrix(context, MatrixValueCase::OfMatrix(*function, matrix)), - ValueNode::OfValueAndRegions { function, value, regions } => value_of_matrix(context, MatrixValueCase::OfValueAndRegions(*function, value, regions)), - ValueNode::MatrixComparison { matrices, distinct } => value_of_matrix(context, MatrixValueCase::Comparison(matrices, *distinct)), + ValueNode::Apply { matrix, value } => value_of_matrix(scope, MatrixValueCase::Apply(matrix, value)), + ValueNode::OfMatrix { function, matrix } => value_of_matrix(scope, MatrixValueCase::OfMatrix(*function, matrix)), + ValueNode::OfValueAndRegions { function, value, regions } => value_of_matrix(scope, MatrixValueCase::OfValueAndRegions(*function, value, regions)), + ValueNode::MatrixComparison { matrices, distinct } => value_of_matrix(scope, MatrixValueCase::Comparison(matrices, *distinct)), + ValueNode::Local(local) => bound_value(scope, Bound::Local(*local)), + ValueNode::Call { function, arguments } => bound_value(scope, Bound::Call(*function, arguments)), + ValueNode::Where { clause, body } => bound_value(scope, Bound::Where(clause, body)), } } } @@ -293,25 +441,25 @@ enum MatrixValueCase<'a> { // One function for every value taken from a matrix, called from several sites, since wasm-opt inlines a function with one call site back into the evaluator whatever its attributes say #[cold] #[inline(never)] -fn value_of_matrix(context: &EvalContext, case: MatrixValueCase) -> Result { +fn value_of_matrix(scope: &Scope<'_, V, F>, case: MatrixValueCase) -> Result { match case { MatrixValueCase::Apply(matrix, value) => { - let matrix = matrix.eval(context)?; - let Value::Number(value) = value.eval(context)?; + let matrix = matrix.eval_in(scope)?; + let Value::Number(value) = value.eval_in(scope)?; settle(Value::from(matrix.apply(value.to_quaternion()))) } - MatrixValueCase::OfMatrix(function, matrix) => settle(function(matrix.eval(context)?)), + MatrixValueCase::OfMatrix(function, matrix) => settle(function(matrix.eval_in(scope)?)), MatrixValueCase::OfValueAndRegions(function, value, regions) => { - let value = value.eval(context)?; + let value = value.eval_in(scope)?; // A range literal is kept by its corners, which may be infinite where no matrix can hold them let region_of = |region: &MatrixNode| -> Result { match region { MatrixNode::Range { from, to } => { - let (Value::Number(from), Value::Number(to)) = (from.eval(context)?, to.eval(context)?); + let (Value::Number(from), Value::Number(to)) = (from.eval_in(scope)?, to.eval_in(scope)?); Ok(Region::Range(from, to)) } - region => region.eval(context).map(Region::Map), + region => region.eval_in(scope).map(Region::Map), } }; @@ -335,11 +483,11 @@ fn value_of_matrix(context: &EvalContext< let heap_matrices: Vec; let matrices: &[Matrix] = if matrices.len() <= stack_matrices.len() { for (slot, matrix) in stack_matrices.iter_mut().zip(matrices) { - *slot = matrix.eval(context)?; + *slot = matrix.eval_in(scope)?; } &stack_matrices[..matrices.len()] } else { - heap_matrices = matrices.iter().map(|matrix| matrix.eval(context)).collect::, EvalError>>()?; + heap_matrices = matrices.iter().map(|matrix| matrix.eval_in(scope)).collect::, EvalError>>()?; &heap_matrices }; @@ -355,9 +503,13 @@ fn value_of_matrix(context: &EvalContext< impl MatrixNode { pub fn eval(&self, context: &EvalContext) -> Result { + self.eval_in(&Scope::root(context)) + } + + fn eval_in(&self, scope: &Scope<'_, V, F>) -> Result { match self { MatrixNode::Var(name) => { - let matrix = resolve_matrix(context, name).ok_or_else(|| EvalError::MissingValue(name.clone()))?; + let matrix = resolve_matrix(scope.context, name).ok_or_else(|| EvalError::MissingValue(name.clone()))?; if matrix.is_nan() { return Err(EvalError::NotANumber(name.clone())); } @@ -369,7 +521,7 @@ impl MatrixNode { return Err(EvalError::TypeError); } for (slot, entry) in quaternions.iter_mut().zip(entries) { - let Value::Number(number) = entry.eval(context)?; + let Value::Number(number) = entry.eval_in(scope)?; *slot = number.to_quaternion(); } @@ -383,23 +535,23 @@ impl MatrixNode { let heap_values: Vec; let values: &[Value] = if arguments.len() <= stack_values.len() { for (slot, argument) in stack_values.iter_mut().zip(arguments) { - *slot = argument.eval(context)?; + *slot = argument.eval_in(scope)?; } &stack_values[..arguments.len()] } else { - heap_values = arguments.iter().map(|argument| argument.eval(context)).collect::, EvalError>>()?; + heap_values = arguments.iter().map(|argument| argument.eval_in(scope)).collect::, EvalError>>()?; &heap_values }; settle_matrix(function(values).ok_or(EvalError::TypeError)?) } MatrixNode::Range { from, to } => { - let (Value::Number(from), Value::Number(to)) = (from.eval(context)?, to.eval(context)?); + let (Value::Number(from), Value::Number(to)) = (from.eval_in(scope)?, to.eval_in(scope)?); settle_matrix(Matrix::range(from.to_quaternion(), to.to_quaternion())) } - MatrixNode::OfMatrix { function, matrix } => settle_matrix(function(matrix.eval(context)?)), - MatrixNode::BinOp { lhs, op, rhs } => matrix_binary_op(lhs.operand(context)?, *op, rhs.operand(context)?), + MatrixNode::OfMatrix { function, matrix } => settle_matrix(function(matrix.eval_in(scope)?)), + MatrixNode::BinOp { lhs, op, rhs } => matrix_binary_op(lhs.operand(scope)?, *op, rhs.operand(scope)?), MatrixNode::UnaryOp { expr, op } => { - let matrix = expr.eval(context)?; + let matrix = expr.eval_in(scope)?; match op { UnaryOp::Pos => Ok(matrix), UnaryOp::Neg => settle_matrix(-matrix), @@ -408,11 +560,14 @@ impl MatrixNode { UnaryOp::Not | UnaryOp::Fac | UnaryOp::Magnitude => Err(EvalError::OperatorTypeError), } } - MatrixNode::Piecewise { cases, otherwise } => match (holding_case(context, cases)?, otherwise) { - (Some(matrix), _) => matrix.eval(context), - (None, Some(otherwise)) => otherwise.eval(context), + MatrixNode::Piecewise { cases, otherwise } => match (holding_case(scope, cases)?, otherwise) { + (Some(matrix), _) => matrix.eval_in(scope), + (None, Some(otherwise)) => otherwise.eval_in(scope), (None, None) => Err(EvalError::NoCaseHolds), }, + MatrixNode::Local(local) => bound_matrix(scope, Bound::Local(*local)), + MatrixNode::Call { function, arguments } => bound_matrix(scope, Bound::Call(*function, arguments)), + MatrixNode::Where { clause, body } => bound_matrix(scope, Bound::Where(clause, body)), } } } diff --git a/libraries/math-parser/src/lexer.rs b/libraries/math-parser/src/lexer.rs index 0f6f6645db1..26a34151ddb 100644 --- a/libraries/math-parser/src/lexer.rs +++ b/libraries/math-parser/src/lexer.rs @@ -47,10 +47,12 @@ pub enum Token<'src> { Ge, Neq, EqEq, + /// The `=` that defines a name in a `where` clause. + Equals, If, Otherwise, - /// Reserved for `where` bindings, so the parser never matches it yet and no host binding can claim the name first. + /// Begins the clause of bindings for the expression before it. Where, /// Source that is no token, which the parser never matches, forcing a parse error rather than silently truncating the input. @@ -106,6 +108,7 @@ impl<'src> fmt::Display for Token<'src> { Token::Ge => f.write_str(">="), Token::Neq => f.write_str("!="), Token::EqEq => f.write_str("=="), + Token::Equals => f.write_str("="), Token::If => f.write_str("if"), Token::Otherwise => f.write_str("otherwise"), @@ -549,7 +552,7 @@ impl<'a> Lexer<'a> { self.bump(); EqEq } else { - Error(LexError::Unrecognized) + Equals } } diff --git a/libraries/math-parser/src/lib.rs b/libraries/math-parser/src/lib.rs index f7f1a753e81..0244092ed80 100644 --- a/libraries/math-parser/src/lib.rs +++ b/libraries/math-parser/src/lib.rs @@ -57,7 +57,7 @@ mod tests { #[test] fn unrecognized_characters_fail_to_parse() { // Unrecognized trailing input must be rejected rather than silently dropped after a valid prefix - for input in ["2@", "5#", "2 $ 3", "sqrt(4)@", "5 & 3", "5 | 3", "2 = 3", "\\", "2 \\ 3", "\\2", "\\_foo"] { + for input in ["2@", "5#", "2 $ 3", "sqrt(4)@", "5 & 3", "5 | 3", "\\", "2 \\ 3", "\\2", "\\_foo"] { assert!(evaluate(input).is_err(), "expected `{input}` to be a parse error"); } } @@ -225,7 +225,6 @@ mod tests { ("x if 1", "`if` joins a case's value to its condition, like `{a if x > 0, b otherwise}`"), ("{1 if 1 otherwise}", "`otherwise` ends the one case with no condition, like `{a if x > 0, b otherwise}`"), ("{1 otherwise, 2 otherwise}", "A piecewise has at most one `otherwise` case"), - ("where", "`where` is a reserved word, so it can't be a name"), ] { let error = evaluate(input).unwrap_err().to_string(); assert!(error.starts_with(expected), "`{input}` gave the error `{error}`"); @@ -606,6 +605,253 @@ mod tests { assert!(ast::Node::try_parse_from_str("clamp(8)").is_err(), "without the host, `clamp` is the builtin"); } + /// Binds `x` to 3 and `y` to 4, and supplies the function `double`. + struct WhereHost; + + impl context::ValueProvider for WhereHost { + fn get_value(&self, name: &str) -> Option { + match name { + "x" => Some(Value::from_f64(3.)), + "y" => Some(Value::from_f64(4.)), + _ => None, + } + } + } + + impl context::FunctionProvider for WhereHost { + fn run_function(&self, name: &str, args: &[Value]) -> Option { + (name == "double").then(|| Value::from_f64(2. * args[0].as_real().unwrap())) + } + fn provides(&self, name: &str) -> bool { + name == "double" + } + } + + fn evaluate_with_where_host(source: &str) -> Result { + ast::Node::try_parse_with_functions(source, &WhereHost).unwrap().eval(&EvalContext::new(WhereHost, WhereHost)) + } + + #[test] + fn where_defines_names_for_the_expression_before_it() { + let real = |source: &str| evaluate_with_where_host(source).unwrap().as_real(); + + assert_eq!(real("a + b where a = 1, b = 2"), Some(3.)); + assert_eq!(real("{sin(r)/r if r != 0, 1 otherwise} where r = hypot(x, y)"), Some(5_f64.sin() / 5.)); + assert_eq!(real("f(0) + f(1) where f(t) = t^2 + c, c = 3"), Some(7.)); + assert_eq!(real("g(2, 3) where g(a, b) = a b"), Some(6.)); + + // A clause runs to its closing parenthesis, so every comma inside separates definitions + assert_eq!(real("2 (a + b where a = x^2, b = y^2)"), Some(50.)); + assert_eq!(real("(a where a = 1) + (a where a = 2)"), Some(3.)); + assert_eq!(real("a where a = (b where b = 2) + 1"), Some(3.)); + + // Within a call's parentheses, a clause after the last argument is read by every argument, but not by the function's name + assert_eq!(real("sqrt(a where a = 16)"), Some(4.)); + assert_eq!(real("max(a, b where a = 1, b = 2)"), Some(2.)); + assert_eq!(real("f(2 where f(t) = t + 1) where f(t) = 10 t"), Some(20.)); + assert_eq!(real("k(a where a = 3, k = 5) where k = 2"), Some(6.)); + assert_eq!(real("sin(0 where sin = 2)"), Some(0.)); + } + + #[test] + fn where_definitions_are_ordered_by_dependency() { + let real = |source: &str| evaluate(source).unwrap().unwrap().as_real(); + let error = |source: &str| evaluate(source).unwrap_err().to_string(); + + // A definition may read any other in its clause, whatever their order + assert_eq!(real("a where a = b + 1, b = 2"), Some(3.)); + assert_eq!(real("f(1) where f(t) = g(t) + c, g(t) = 2t, c = 1"), Some(3.)); + + // A cycle could never finish, whether through values, functions, or a nested clause + assert_eq!(error("x where x = x + 1"), "`x` is defined in terms of itself"); + assert_eq!(error("f(1) where f(t) = f(t - 1)"), "`f` is defined in terms of itself"); + assert_eq!(error("a where a = b, b = a"), "`a` and `b` are defined in terms of each other"); + assert_eq!(error("f(1) where f(t) = g(t), g(t) = f(t)"), "`f` and `g` are defined in terms of each other"); + assert_eq!(error("a where a = b, b = c, c = f(1), f(t) = a"), "`a`, `b`, `c`, and `f` are defined in terms of each other"); + assert_eq!(error("a where a = (c where c = b), b = a"), "`a` and `b` are defined in terms of each other"); + + // A parameter or an inner definition of the same name is another name, as is a function beside a value + assert_eq!(real("a where f(a) = a + 1, a = f(2)"), Some(3.)); + assert_eq!(real("a where a = (b where b = 1), b = a"), Some(1.)); + assert_eq!(real("f(2) where f = 3, f(t) = f t"), Some(6.)); + } + + #[test] + fn where_scoping_is_lexical() { + let real = |source: &str| evaluate_with_where_host(source).unwrap().as_real(); + + // A clause shadows the host's bindings, the builtins, and enclosing clauses, and a parameter shadows every outer name + assert_eq!(real("x where x = 2"), Some(2.)); + assert_eq!(real("e + \\e where e = 2"), Some(2. + std::f64::consts::E)); + assert_eq!(real("(a where a = 2) + a where a = 5"), Some(7.)); + assert_eq!(real("f(1) where f(a) = a, a = 5"), Some(1.)); + assert_eq!(real("f(1) + x where f(x) = x"), Some(4.)); + + // A function reads the names around its definition, not around its call + assert_eq!(real("(f(1) where a = 5) where f(t) = t + a, a = 2"), Some(3.)); + + // A clause's names are unknown outside its parentheses + assert!(matches!(evaluate_with_where_host("(a where a = 1) + a"), Err(EvalError::MissingValue(name)) if name == "a")); + } + + #[test] + fn where_calls_and_values_are_separate_namespaces() { + let real = |source: &str| evaluate_with_where_host(source).unwrap().as_real(); + + assert_eq!(real("f + f(1) where f = 2, f(t) = t + 1"), Some(4.)); + + // A value leaves any function of its name a call, and multiplies its one argument only where no such function exists + assert_eq!(real("sin(0) + sin where sin = 2"), Some(2.)); + assert_eq!(real("double(1) + double where double = 5"), Some(7.)); + assert_eq!(real("k(x + 1) where k = 2"), Some(8.)); + + // A defined function shadows the host's function and the builtin of its name, which the prefix still reaches + assert_eq!(real("double(1) where double(t) = t + 5"), Some(6.)); + assert_eq!(real("sin(2) where sin(t) = t"), Some(2.)); + assert_eq!(real("\\sin(0) where sin(t) = t + 1"), Some(0.)); + + // A function is only ever called, so its bare name reads a value + assert!(matches!(evaluate("f where f(t) = t").unwrap(), Err(EvalError::MissingValue(name)) if name == "f")); + } + + #[test] + fn where_definitions_are_evaluated_lazily_and_once() { + use std::cell::Cell; + use std::rc::Rc; + + struct Counting(Rc>); + impl context::FunctionProvider for Counting { + fn run_function(&self, name: &str, args: &[Value]) -> Option { + (name == "counted").then(|| { + self.0.set(self.0.get() + 1); + args[0] + }) + } + fn provides(&self, name: &str) -> bool { + name == "counted" + } + } + let calls = Rc::new(Cell::new(0)); + let counted = |source: &str| { + calls.set(0); + let node = ast::Node::try_parse_with_functions(source, &Counting(calls.clone())).unwrap(); + let result = node.eval(&EvalContext::new(context::NothingMap, Counting(calls.clone()))).unwrap().as_real(); + (result, calls.get()) + }; + + // A definition or argument the result never reaches is never evaluated, so its error is never raised + for source in ["1 where a = 0/0", "f(0/0) where f(t) = 1", "{1 if 1 > 0, a otherwise} where a = 0/0"] { + assert_eq!(evaluate(source).unwrap().unwrap().as_real(), Some(1.), "`{source}`"); + } + assert!(matches!(evaluate("a + 1 where a = 0/0").unwrap(), Err(EvalError::Indeterminate))); + + // A definition is evaluated once however often it's read, as is an argument within one call + assert_eq!(counted("a + a + a where a = counted(2)"), (Some(6.), 1)); + assert_eq!(counted("f(counted(2)) where f(t) = t t t"), (Some(8.), 1)); + assert_eq!(counted("0 where a = counted(2)"), (Some(0.), 0)); + + // So reads never multiply the work, however deeply calls nest + let nested = format!("{}1{} where f(t) = t + t", "f(".repeat(40), ")".repeat(40)); + assert_eq!(evaluate(&nested).unwrap().unwrap().as_real(), Some(2_f64.powi(40))); + } + + #[test] + fn where_names_follow_the_case_rule() { + let real = |source: &str| evaluate(source).unwrap().unwrap().as_real(); + let error = |source: &str| evaluate(source).unwrap_err().to_string(); + + // A name beginning with a capital letter defines a matrix and any other a value, parameters included + assert_eq!(real("det(M) where M = 2 I"), Some(16.)); + assert_eq!(real("det(F(3)) where F(t) = t I"), Some(81.)); + assert_eq!(real("f(2 I) where f(T) = det(T)"), Some(16.)); + + // A name of the wrong case is recased in the message where the case exists to flip + assert_eq!(error("m where m = I"), "`m` is defined as a matrix, so rename it to begin with a capital letter, like `M`"); + assert_eq!(error("Mass where Mass = 1"), "`Mass` is defined as a value, so rename it to begin with a lowercase letter, like `mass`"); + assert_eq!(error("f(1) where f(t) = t I"), "`f(t)` is defined as a matrix, so rename `f` to begin with a capital letter, like `F`"); + assert_eq!(error("あ where あ = I"), "`あ` is defined as a matrix, so rename it to begin with a capital letter"); + assert_eq!(error("𝐀 where 𝐀 = 1"), "`𝐀` is defined as a value, so rename it to begin with a lowercase letter"); + + // The definition is checked before any read of it, even one in an earlier definition, so a read never takes the blame + assert_eq!(error("det(m) where m = I"), "`m` is defined as a matrix, so rename it to begin with a capital letter, like `M`"); + assert_eq!(error("a where a = det(m), m = I"), "`m` is defined as a matrix, so rename it to begin with a capital letter, like `M`"); + assert_eq!( + error("det(f(1)) where f(t) = t I"), + "`f(t)` is defined as a matrix, so rename `f` to begin with a capital letter, like `F`" + ); + assert_eq!( + error("Width < 10 where Width = 5"), + "`Width` is defined as a value, so rename it to begin with a lowercase letter, like `width`" + ); + + // A definition agreeing with its name leaves a mismatched read the mistake, like a host's name + assert_eq!(error("det(m) where m = 2"), "A value stands where a matrix is needed"); + assert_eq!(error("det(f(1)) where f(t) = 2 t"), "A value stands where a matrix is needed"); + assert_eq!(error("det(x)"), "A value stands where a matrix is needed"); + + // A parameter has no definition, so a read in its function's body decides its sort + assert_eq!(error("f(1) where f(t) = det(t)"), "`t` is used as a matrix, so rename it to begin with a capital letter, like `T`"); + assert_eq!(error("f(2 I) where f(t) = t^T"), "`t` is used as a matrix, so rename it to begin with a capital letter, like `T`"); + + // A call is blamed on the parameter where the body could take what's passed, and on the argument otherwise + assert_eq!(error("f(2 I) where f(t) = 2 t"), "`t` is passed a matrix, so rename it to begin with a capital letter, like `T`"); + assert_eq!(error("f(1) where f(T) = det(T)"), "`f(T)` is passed a value for `T`, which its body uses as a matrix"); + assert_eq!(error("f(I) where f(t) = sin(t)"), "`f(t)` is passed a matrix for `t`, which its body uses as a value"); + assert!(evaluate("f(1) where f(t) = f(I)").is_err(), "a body testing a call to itself must still finish"); + + // One function serves every rung it's called with + assert_eq!(real("f(2) + f(i) where f(t) = t^2"), Some(3.)); + } + + #[test] + fn where_definitions_are_checked_when_parsed() { + let error = |source: &str| evaluate(source).unwrap_err().to_string(); + + // Each name is defined once per clause, where a function and a value may share one + assert_eq!(error("a where a = 1, a = 2"), "`a` is defined twice in one `where` clause"); + assert_eq!(error("f(1) where f(t) = t, f(t, u) = t"), "`f` is defined twice in one `where` clause"); + assert_eq!(error("f(1, 2) where f(t, t) = t"), "`f` has two parameters named `t`"); + + // A function takes one argument per parameter, and has at least one parameter + assert_eq!(error("f(1, 2) where f(t) = t"), "`f` takes 1 argument"); + assert_eq!(error("f(1) where f(a, b) = a"), "`f` takes 2 arguments"); + assert!(evaluate("f(1) where f() = 1").is_err()); + + // The prefix always reaches the builtin, so no clause can define a name that has it + assert!(error("\\pi where \\pi = 3").starts_with("A `\\` name is always the builtin, so no `where` clause can define one")); + + // Even an unread definition must stand + assert_eq!(error("1 where a = sin(I)"), "A matrix stands where a value is needed"); + } + + #[test] + fn misplaced_equals_and_where_are_told_where_they_belong() { + let error = |source: &str| evaluate(source).unwrap_err().to_string(); + + // A single `=` only defines, so one that compares is pointed to `==` + for source in ["x = 2", "{1 if x = 0, 2 otherwise}", "a where a = 1 = 2"] { + let error = error(source); + assert!( + error.starts_with("`=` names a value in a `where` clause, so equality is written `==`"), + "`{source}` gave the error `{error}`" + ); + } + + // A clause ends the whole expression or stands within parentheses, with commas between its definitions + for source in ["[a where a = 1]", "a where a = 1 where b = 2", "{a where a = 1 if 1, 0 otherwise}", "|a where a = 1|"] { + let error = error(source); + assert!( + error.starts_with("`where` defines names for the whole expression or within parentheses"), + "`{source}` gave the error `{error}`" + ); + } + + // Elsewhere the parser says what it expected + assert_eq!(error("x where"), "Found end of input, expected a name, at 7..7"); + assert_eq!(error("x where a"), "Found end of input, expected `(` or `=`, at 9..9"); + assert_eq!(error("x where a = = 2"), "Found `=`, expected `-`, `+`, `!`, `¬`, or a value, at 12..13"); + } + #[test] fn rename_identifiers_is_token_exact() { let a_to_x = |name: &str| name.eq_ignore_ascii_case("a").then(|| "x".to_string()); diff --git a/libraries/math-parser/src/parser.rs b/libraries/math-parser/src/parser.rs index 59d2f815226..44337a5ef1b 100644 --- a/libraries/math-parser/src/parser.rs +++ b/libraries/math-parser/src/parser.rs @@ -1,9 +1,9 @@ -use crate::ast::{BinaryOp, Case, Literal, Node, Syntax, UnaryOp}; +use crate::ast::{BinaryOp, Binding, CallWhere, Case, Literal, Node, Syntax, UnaryOp}; use crate::context::{FunctionProvider, NothingMap}; use crate::lexer::{LexError, Lexer, Span, Token}; use crate::sort::sorted; use chumsky::cache::{Cache, Cached}; -use chumsky::error::{EmptyErr, LabelError, RichReason}; +use chumsky::error::{EmptyErr, LabelError, RichPattern, RichReason}; use chumsky::input::ValueInput; use chumsky::{Parser, prelude::*}; use std::fmt; @@ -169,53 +169,107 @@ fn parse(src: &str) -> Result { Err(parse_errs) => Err(ParseError( parse_errs .into_iter() - .map(|e| match e.found() { - Some(Token::Percent) => ErrorMessage::from_prose("`%` is reserved for percentages, so the remainder is written `mod(a, b)`").at(e.span()), - Some(Token::If) => ErrorMessage::from_prose("`if` joins a case's value to its condition, like `{a if x > 0, b otherwise}`").at(e.span()), - Some(Token::Otherwise) => ErrorMessage::from_prose("`otherwise` ends the one case with no condition, like `{a if x > 0, b otherwise}`").at(e.span()), - Some(Token::Where) => ErrorMessage::from_prose("`where` is a reserved word, so it can't be a name").at(e.span()), - // The offending source is quoted as its own part, since it may hold anything, backticks included - Some(Token::Error(error)) => { - let text = src.get(e.span().start..e.span().end).unwrap_or_default(); - let reason = match error { - LexError::Unrecognized => "is not recognized", - LexError::MalformedNumber => "is not a valid number", - LexError::NumberAfterNumber => "can't follow another number, so write them as one or put `*` between them", - LexError::LeadingDotAfterOperand => "needs its leading zero after an operand, like `0.5`", - }; - let reason = ErrorMessage::from_prose(&format!(" {reason}")); - let parts = std::iter::once(MessagePart::Code(text.to_string())).chain(reason.parts).collect(); - ErrorMessage { parts, span: None }.at(e.span()) - } - _ => match e.reason() { - RichReason::Custom(message) => ErrorMessage::from_prose(message).at(e.span()), - // Chumsky's own wording is "found ... expected ..." in lowercase, without a comma, with its tokens between single quotes - RichReason::ExpectedFound { .. } => { - let message = e.to_string().replacen(" expected ", ", expected ", 1); - let mut characters = message.chars(); - let sentence_case: String = characters.next().into_iter().flat_map(char::to_uppercase).chain(characters).collect(); - ErrorMessage::with_code_between(&sentence_case, '\'').at(e.span()) + .map(|e| { + // Where `==` could have come next, the unexpected token followed a complete operand + let after_operand = e.expected().any(|pattern| matches!(pattern, RichPattern::Token(token) if **token == Token::EqEq)); + + match e.found() { + Some(Token::Percent) => ErrorMessage::from_prose("`%` is reserved for percentages, so the remainder is written `mod(a, b)`").at(e.span()), + Some(Token::If) => ErrorMessage::from_prose("`if` joins a case's value to its condition, like `{a if x > 0, b otherwise}`").at(e.span()), + Some(Token::Otherwise) => ErrorMessage::from_prose("`otherwise` ends the one case with no condition, like `{a if x > 0, b otherwise}`").at(e.span()), + Some(Token::Equals) if after_operand => ErrorMessage::from_prose("`=` names a value in a `where` clause, so equality is written `==`").at(e.span()), + Some(Token::Where) if after_operand => { + ErrorMessage::from_prose("`where` defines names for the whole expression or within parentheses, like `2 (a + b where a = 1, b = 2)`").at(e.span()) + } + // The offending source is quoted as its own part, since it may hold anything, backticks included + Some(Token::Error(error)) => { + let text = src.get(e.span().start..e.span().end).unwrap_or_default(); + let reason = match error { + LexError::Unrecognized => "is not recognized", + LexError::MalformedNumber => "is not a valid number", + LexError::NumberAfterNumber => "can't follow another number, so write them as one or put `*` between them", + LexError::LeadingDotAfterOperand => "needs its leading zero after an operand, like `0.5`", + }; + let reason = ErrorMessage::from_prose(&format!(" {reason}")); + let parts = std::iter::once(MessagePart::Code(text.to_string())).chain(reason.parts).collect(); + ErrorMessage { parts, span: None }.at(e.span()) } - }, + _ => match e.reason() { + RichReason::Custom(message) => ErrorMessage::from_prose(message).at(e.span()), + // Chumsky's wording: lowercase "found ... expected ...", no comma, tokens in single quotes, and "a, or b" for two alternatives + RichReason::ExpectedFound { .. } => { + let mut message = e.to_string().replacen(" expected ", ", expected ", 1); + if e.expected().len() == 2 { + message = message.replacen(", or ", " or ", 1); + } + let mut characters = message.chars(); + let sentence_case: String = characters.next().into_iter().flat_map(char::to_uppercase).chain(characters).collect(); + ErrorMessage::with_code_between(&sentence_case, '\'').at(e.span()) + } + }, + } }) .collect(), )), } } +/// A `where` clause: bindings separated by commas, each a value like `a = 1` or a function like `f(t) = t^2`, whose defining +/// expression has no clause of its own unless parenthesized. +fn where_clause<'src, I, E>(expression: impl Parser<'src, I, Syntax, E> + Clone) -> impl Parser<'src, I, Vec, E> + Clone +where + I: ValueInput<'src, Token = Token<'src>, Span = Span>, + E: extra::ParserExtra<'src, I>, + E::Error: LabelError<'src, I, &'static str> + CustomError, +{ + // The `\` prefix always reaches the language's own builtin, so no clause can define a name that has it + let name = select! {Token::Ident(name) => name}.labelled("a name").validate(|name: &str, extra, emitter| { + if name.starts_with('\\') { + emitter.emit(CustomError::custom(extra.span(), "A `\\` name is always the builtin, so no `where` clause can define one")); + } + name.to_string() + }); + + let parameters = name.separated_by(just(Token::Comma)).at_least(1).collect::>(); + let binding = name + .then(parameters.delimited_by(just(Token::LParen), just(Token::RParen)).or_not()) + .then_ignore(just(Token::Equals)) + .then(expression) + .map(|((name, parameters), value)| Binding { + name, + parameters: parameters.unwrap_or_default(), + value, + }); + + just(Token::Where).ignore_then(binding.separated_by(just(Token::Comma)).at_least(1).collect()) +} + +/// The expression, within the names its `where` clause defines if it has one. +fn with_bindings((body, bindings): (Syntax, Option>)) -> Syntax { + match bindings { + Some(bindings) => Syntax::Where { body: Box::new(body), bindings }, + None => body, + } +} + pub fn parser<'src, I, E>() -> impl Parser<'src, I, Syntax, E> where I: ValueInput<'src, Token = Token<'src>, Span = Span>, E: extra::ParserExtra<'src, I>, E::Error: LabelError<'src, I, &'static str> + CustomError, { - recursive(|expr| { + let expression = recursive(|expr| { let constant = select! { Token::Integer(integer) => Syntax::Lit(Literal::Integer(integer)), Token::Float(float) => Syntax::Lit(Literal::Float(float)), }; - let args = expr.clone().separated_by(just(Token::Comma)).collect::>().delimited_by(just(Token::LParen), just(Token::RParen)); + let args = expr + .clone() + .separated_by(just(Token::Comma)) + .collect::>() + .then(where_clause(expr.clone()).or_not()) + .delimited_by(just(Token::LParen), just(Token::RParen)); // Each case is a value then its condition, except the one `otherwise` case, which may stand anywhere since case order means nothing let case = expr.clone().then(choice((just(Token::If).ignore_then(expr.clone()).map(Some), just(Token::Otherwise).map(|_| None)))); @@ -259,12 +313,21 @@ where let ident = select! {Token::Ident(s) => s}.labelled("a name"); // An ident followed by parenthesized args is a function call, otherwise a variable - let call_or_var = ident.then(args.or_not()).map(|(name, args): (&str, Option>)| match args { - Some(args) => Syntax::FnCall { name: name.to_string(), expr: args }, + let call_or_var = ident.then(args.or_not()).map(|(name, args)| match args { + Some((args, None)) => Syntax::FnCall { name: name.to_string(), expr: args }, + Some((arguments, Some(bindings))) => Syntax::CallWhere(Box::new(CallWhere { + name: name.to_string(), + arguments, + bindings, + })), None => Syntax::Var(name.to_string()), }); - let parens = expr.clone().delimited_by(just(Token::LParen), just(Token::RParen)); + let parens = expr + .clone() + .then(where_clause(expr.clone()).or_not()) + .map(with_bindings) + .delimited_by(just(Token::LParen), just(Token::RParen)); let magnitude = expr.clone().delimited_by(just(Token::BarOpen), just(Token::BarClose)).map(|expr| Syntax::UnaryOp { op: UnaryOp::Magnitude, expr: Box::new(expr), @@ -388,7 +451,9 @@ where op, rhs: Box::new(rhs), }) - }) + }); + + expression.clone().then(where_clause(expression).or_not()).map(with_bindings) } #[cfg(test)] @@ -482,13 +547,22 @@ mod tests { }, test_parse_ii_call: "ii(16)" => Syntax::FnCall { name: "ii".to_string(), - expr: vec![Syntax::Lit(Literal::Integer(16))] + expr: vec![Syntax::Lit(Literal::Integer(16))], }, // `i` is a name a binding may shadow, so only the evaluator can read this call as `i` times its argument test_parse_i_mul: "i(16)" => Syntax::FnCall { name: "i".to_string(), expr: vec![Syntax::Lit(Literal::Integer(16))], }, + test_call_where_clause: "max(a, b where a = 1)" => Syntax::CallWhere(Box::new(CallWhere { + name: "max".to_string(), + arguments: vec![Syntax::Var("a".to_string()), Syntax::Var("b".to_string())], + bindings: vec![Binding { + name: "a".to_string(), + parameters: vec![], + value: Syntax::Lit(Literal::Integer(1)), + }], + })), test_parse_complex_expr: "(1 + 2) * 3 - 4 ^ 2" => Syntax::BinOp { lhs: Box::new(Syntax::BinOp { lhs: Box::new(Syntax::BinOp { @@ -506,6 +580,21 @@ mod tests { rhs: Box::new(Syntax::Lit(Literal::Integer(2))), }), }, + test_where_clause: "a where a = 1, f(t, u) = t" => Syntax::Where { + body: Box::new(Syntax::Var("a".to_string())), + bindings: vec![ + Binding { + name: "a".to_string(), + parameters: vec![], + value: Syntax::Lit(Literal::Integer(1)), + }, + Binding { + name: "f".to_string(), + parameters: vec!["t".to_string(), "u".to_string()], + value: Syntax::Var("t".to_string()), + }, + ], + }, test_piecewise_expr: "{0 otherwise, x + 3 if x < 0}" => Syntax::Piecewise { cases: vec![Case { value: Syntax::BinOp { diff --git a/libraries/math-parser/src/sort.rs b/libraries/math-parser/src/sort.rs index 2313d31ab0c..f7a7f371894 100644 --- a/libraries/math-parser/src/sort.rs +++ b/libraries/math-parser/src/sort.rs @@ -1,243 +1,678 @@ -use crate::ast::{BinaryOp, Case, Literal, MatrixNode, Node, SortedCase, Syntax, UnaryOp, ValueNode}; -use crate::constants::{Builtin, builtin_function}; +use crate::ast::{BinaryOp, Binding, CallWhere, Case, Clause, Literal, Local, MatrixNode, Node, SortedCase, Syntax, UnaryOp, ValueNode}; +use crate::constants::{Builtin, builtin_function, suffixed_function}; use crate::context::FunctionProvider; use crate::lexer::names_matrix; +use std::borrow::Cow; use std::fmt; -/// A subexpression standing where its sort cannot, which fails the parse, since every sort is fixed by spelling. -#[derive(Debug, Clone, Copy, PartialEq)] -pub struct SortError(&'static str); +/// A subexpression standing where its sort cannot, or a `where` clause with a cycle, a name defined twice, a call of the wrong +/// arity, or a definition of the sort its name's case forbids. Either fails the parse. +#[derive(Debug, Clone, PartialEq)] +pub struct SortError(Cow<'static, str>); impl fmt::Display for SortError { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.write_str(self.0) + f.write_str(&self.0) } } -const MATRIX_AS_VALUE: SortError = SortError("A matrix stands where a value is needed"); -const VALUE_AS_MATRIX: SortError = SortError("A value stands where a matrix is needed"); -const NO_MATRIX_OPERATOR: SortError = SortError("The operator has no meaning for a matrix"); -const MIXED_CASES: SortError = SortError("A piecewise's cases must all be values or all be matrices"); -const INVALID_ARGUMENTS: SortError = SortError("Invalid arguments for function call"); +const MATRIX_AS_VALUE: SortError = SortError(Cow::Borrowed("A matrix stands where a value is needed")); +const VALUE_AS_MATRIX: SortError = SortError(Cow::Borrowed("A value stands where a matrix is needed")); +const NO_MATRIX_OPERATOR: SortError = SortError(Cow::Borrowed("The operator has no meaning for a matrix")); +const MIXED_CASES: SortError = SortError(Cow::Borrowed("A piecewise's cases must all be values or all be matrices")); +const INVALID_ARGUMENTS: SortError = SortError(Cow::Borrowed("Invalid arguments for function call")); -/// Reads the sort of every subexpression from its spelling, so each evaluator takes only trees of its own sort. A call of a -/// function the host provides is a value, whatever builtin shares its name. +/// Reads the sort of every subexpression from its spelling, so each evaluator takes only trees of its own sort, and finds where +/// each name a `where` clause defines lives. A call of a function the host provides is a value, whatever builtin shares its name. pub fn sorted(syntax: Syntax, functions: &dyn FunctionProvider) -> Result { - Ok(match syntax { - Syntax::Lit(literal) => Node::Value(ValueNode::Lit(literal)), - Syntax::Var(name) if names_matrix(&name) => Node::Matrix(MatrixNode::Var(name)), - Syntax::Var(name) => Node::Value(ValueNode::Var(name)), - Syntax::FnCall { name, expr } => call(name, expr, functions)?, - Syntax::BinOp { lhs, op, rhs } => binary(sorted(*lhs, functions)?, op, sorted(*rhs, functions)?)?, - Syntax::Product { first, rest } => product(sorted(*first, functions)?, &mut rest.into_iter(), functions)?, - Syntax::UnaryOp { expr, op } => match (sorted(*expr, functions)?, op) { - (Node::Value(_), UnaryOp::Transpose) => return Err(VALUE_AS_MATRIX), - (Node::Value(expr), op) => Node::Value(ValueNode::UnaryOp { expr: Box::new(expr), op }), - (Node::Matrix(expr), UnaryOp::Pos | UnaryOp::Neg | UnaryOp::Transpose) => Node::Matrix(MatrixNode::UnaryOp { expr: Box::new(expr), op }), - (Node::Matrix(_), _) => return Err(NO_MATRIX_OPERATOR), - }, - Syntax::Comparison { first, rest } => { - let first = sorted(*first, functions)?; - let rest = rest - .into_iter() - .map(|(op, operand)| Ok((op, sorted(operand, functions)?))) - .collect::, SortError>>()?; - match first { - Node::Value(first) => { - let rest = rest.into_iter().map(|(op, operand)| Ok((op, value(operand)?))).collect::, SortError>>()?; - Node::Value(ValueNode::Comparison { first: Box::new(first), rest }) + Sorter { functions, scopes: Vec::new() }.sort(syntax) +} + +/// One definition in a `where` clause, by its position among the clause's values or among its functions. +#[derive(Clone, Copy, PartialEq)] +enum Definition { + Value(usize), + Function(usize), +} + +/// The names one scope defines: a `where` clause's values and functions, or a function's parameters, which are values. +#[derive(Default)] +struct Scope { + values: Vec, + functions: Vec, + /// Whether the values are a function's parameters rather than a clause's definitions. + holds_parameters: bool, + /// The parameter read as the other sort than its case gives, while testing whether the function's body could take that sort. + flipped: Option, + /// The clause's definition being sorted, which depends on every other of the clause's definitions it reads. + defining: Option, + /// Each definition paired with one it reads. + dependencies: Vec<(Definition, Definition)>, +} + +/// A function a `where` clause defines. +struct Function { + name: String, + parameters: Vec, + /// The unsorted body, sorted again only to test a call whose argument is the other sort than its parameter's case gives. + body: Syntax, +} + +impl Scope { + fn parameters(parameters: Vec) -> Self { + Self { + values: parameters, + holds_parameters: true, + ..Self::default() + } + } + + /// Adds a clause's definition, whose name may appear once among the clause's values and once among its functions, with no + /// parameter named twice. + fn define(&mut self, Binding { name, parameters, value }: &Binding) -> Result { + let taken = if parameters.is_empty() { + self.values.contains(name) + } else { + self.functions.iter().any(|function| function.name == *name) + }; + if taken { + return Err(SortError(format!("`{name}` is defined twice in one `where` clause").into())); + } + + if let Some(parameter) = parameters + .iter() + .enumerate() + .find_map(|(index, parameter)| parameters[..index].contains(parameter).then_some(parameter)) + { + return Err(SortError(format!("`{name}` has two parameters named `{parameter}`").into())); + } + + Ok(if parameters.is_empty() { + self.values.push(name.clone()); + Definition::Value(self.values.len() - 1) + } else { + self.functions.push(Function { + name: name.clone(), + parameters: parameters.clone(), + body: value.clone(), + }); + Definition::Function(self.functions.len() - 1) + }) + } + + fn record_read(&mut self, definition: Definition) { + if let Some(defining) = self.defining { + self.dependencies.push((defining, definition)); + } + } + + /// Rejects a definition that depends on itself, directly or through others, since its evaluation could never finish. + fn reject_cycles(&self) -> Result<(), SortError> { + // The values and then the functions, numbered as one list + let position = |definition: Definition| match definition { + Definition::Value(index) => index, + Definition::Function(index) => self.values.len() + index, + }; + let names = self + .values + .iter() + .chain(self.functions.iter().map(|function| &function.name)) + .map(|name| format!("`{name}`")) + .collect::>(); + + let mut reads = vec![Vec::new(); names.len()]; + for &(reader, read) in &self.dependencies { + reads[position(reader)].push(position(read)); + } + + let Some(cycle) = find_cycle(&reads) else { return Ok(()) }; + let message = match cycle.as_slice() { + [name] => format!("{} is defined in terms of itself", names[*name]), + [first, second] => format!("{} and {} are defined in terms of each other", names[*first], names[*second]), + [others @ .., last] => { + let others = others.iter().map(|index| names[*index].as_str()).collect::>().join(", "); + format!("{others}, and {} are defined in terms of each other", names[*last]) + } + [] => return Ok(()), + }; + Err(SortError(message.into())) + } +} + +/// The first cycle found among nodes with the given successors, in the order its members lead to one another. O(nodes + edges). +fn find_cycle(successors: &[Vec]) -> Option> { + #[derive(Clone, Copy, PartialEq)] + enum Visit { + Unvisited, + OnPath, + Finished, + } + + fn visit(node: usize, successors: &[Vec], visits: &mut [Visit], path: &mut Vec) -> Option> { + visits[node] = Visit::OnPath; + path.push(node); + + for &next in &successors[node] { + match visits[next] { + Visit::OnPath => { + let start = path.iter().position(|&member| member == next)?; + return Some(path[start..].to_vec()); } - // Matrices have no order, so a chain over them is `==` or `!=` throughout - Node::Matrix(first) => { - let distinct = rest.iter().all(|(op, _)| *op == BinaryOp::Neq); - if !distinct && !rest.iter().all(|(op, _)| *op == BinaryOp::Eq) { - return Err(NO_MATRIX_OPERATOR); + Visit::Unvisited => { + if let Some(cycle) = visit(next, successors, visits, path) { + return Some(cycle); } - let rest = rest.into_iter().map(|(_, operand)| matrix(operand).map_err(|_| NO_MATRIX_OPERATOR)); - let matrices = std::iter::once(Ok(first)).chain(rest).collect::, SortError>>()?; - Node::Value(ValueNode::MatrixComparison { matrices, distinct }) } + Visit::Finished => {} } } - Syntax::Piecewise { cases, otherwise } => piecewise(cases, otherwise, functions)?, - Syntax::Matrix { entries, by_rows } => Node::Matrix(MatrixNode::Literal { - entries: entries.into_iter().map(|entry| value(sorted(entry, functions)?)).collect::, SortError>>()?, - by_rows, - }), - Syntax::Range { from, to } => Node::Matrix(MatrixNode::Range { - from: Box::new(value(sorted(*from, functions)?)?), - to: Box::new(value(sorted(*to, functions)?)?), - }), + + path.pop(); + visits[node] = Visit::Finished; + None + } + + let mut visits = vec![Visit::Unvisited; successors.len()]; + let mut path = Vec::new(); + (0..successors.len()).find_map(|node| { + if visits[node] == Visit::Unvisited { + visit(node, successors, &mut visits, &mut path) + } else { + None + } }) } -fn value(node: Node) -> Result { - match node { - Node::Value(value) => Ok(value), - Node::Matrix(_) => Err(MATRIX_AS_VALUE), - } +struct Sorter<'a> { + functions: &'a dyn FunctionProvider, + /// The scopes around the subexpression being sorted, innermost last. + scopes: Vec, } -fn matrix(node: Node) -> Result { - match node { - Node::Matrix(matrix) => Ok(matrix), - Node::Value(_) => Err(VALUE_AS_MATRIX), +impl Sorter<'_> { + fn sort(&mut self, syntax: Syntax) -> Result { + Ok(match syntax { + Syntax::Lit(literal) => Node::Value(ValueNode::Lit(literal)), + Syntax::Var(name) => self.name(name, 0), + Syntax::FnCall { name, expr } => self.call(name, expr, 0)?, + Syntax::BinOp { lhs, op, rhs } => { + let lhs = self.sort(*lhs)?; + let rhs = self.sort(*rhs)?; + self.binary(lhs, op, rhs)? + } + Syntax::Product { first, rest } => { + let first = self.sort(*first)?; + self.product(first, &mut rest.into_iter())? + } + Syntax::UnaryOp { expr, op } => match (self.sort(*expr)?, op) { + (expr @ Node::Value(_), UnaryOp::Transpose) => return Err(self.misplaced(&expr, VALUE_AS_MATRIX)), + (Node::Value(expr), op) => Node::Value(ValueNode::UnaryOp { expr: Box::new(expr), op }), + (Node::Matrix(expr), UnaryOp::Pos | UnaryOp::Neg | UnaryOp::Transpose) => Node::Matrix(MatrixNode::UnaryOp { expr: Box::new(expr), op }), + (expr @ Node::Matrix(_), _) => return Err(self.misplaced(&expr, NO_MATRIX_OPERATOR)), + }, + Syntax::Comparison { first, rest } => { + let first = self.sort(*first)?; + let rest = rest + .into_iter() + .map(|(op, operand)| Ok((op, self.sort(operand)?))) + .collect::, SortError>>()?; + match first { + Node::Value(first) => { + let rest = rest + .into_iter() + .map(|(op, operand)| Ok((op, self.value(operand, MATRIX_AS_VALUE)?))) + .collect::, SortError>>()?; + Node::Value(ValueNode::Comparison { first: Box::new(first), rest }) + } + // Matrices have no order, so a chain over them is `==` or `!=` throughout + Node::Matrix(first) => { + let distinct = rest.iter().all(|(op, _)| *op == BinaryOp::Neq); + if !distinct && !rest.iter().all(|(op, _)| *op == BinaryOp::Eq) { + return Err(self.misplaced(&Node::Matrix(first), NO_MATRIX_OPERATOR)); + } + let rest = rest.into_iter().map(|(_, operand)| self.matrix(operand, NO_MATRIX_OPERATOR)); + let matrices = std::iter::once(Ok(first)).chain(rest).collect::, SortError>>()?; + Node::Value(ValueNode::MatrixComparison { matrices, distinct }) + } + } + } + Syntax::Piecewise { cases, otherwise } => self.piecewise(cases, otherwise)?, + Syntax::Matrix { entries, by_rows } => Node::Matrix(MatrixNode::Literal { + entries: entries + .into_iter() + .map(|entry| { + let entry = self.sort(entry)?; + self.value(entry, MATRIX_AS_VALUE) + }) + .collect::, SortError>>()?, + by_rows, + }), + Syntax::Range { from, to } => { + let from = self.sort(*from)?; + let to = self.sort(*to)?; + Node::Matrix(MatrixNode::Range { + from: Box::new(self.value(from, MATRIX_AS_VALUE)?), + to: Box::new(self.value(to, MATRIX_AS_VALUE)?), + }) + } + Syntax::Where { body, bindings } => self.clause(bindings, |sorter| sorter.sort(*body))?, + // The function's name stands outside the parentheses holding the clause, so it looks past the clause's scope + Syntax::CallWhere(call) => { + let CallWhere { name, arguments, bindings } = *call; + self.clause(bindings, |sorter| sorter.call(name, arguments, 1))? + } + }) } -} -fn values(nodes: Vec) -> Result, SortError> { - nodes.into_iter().map(value).collect() -} + /// A name read where the innermost scope defining it puts it, or else the host's binding or builtin of that spelling. The + /// innermost `skip` scopes are passed over. + fn name(&mut self, name: String, skip: usize) -> Node { + match self.local_value(&name, skip) { + Some(local) => local, + None if names_matrix(&name) => Node::Matrix(MatrixNode::Var(name)), + None => Node::Value(ValueNode::Var(name)), + } + } -fn one(arguments: Vec) -> Result { - <[Node; 1]>::try_from(arguments).map(|[argument]| argument).map_err(|_| INVALID_ARGUMENTS) -} + /// A read of the value the innermost scope defining it puts, past the innermost `skip` scopes, which a `\` name never reaches. + fn local_value(&mut self, name: &str, skip: usize) -> Option { + if name.starts_with('\\') { + return None; + } -/// A call's sort follows the builtin's, a matrix's name applying it to its one argument, and any other name taking values alone. -/// A range function like `within(p, R)` takes a value and then matrices. -fn call(name: String, arguments: Vec, functions: &dyn FunctionProvider) -> Result { - let arguments = arguments.into_iter().map(|argument| sorted(argument, functions)).collect::, SortError>>()?; - let (prefixed, bare_name) = match name.strip_prefix('\\') { - Some(bare_name) => (true, bare_name), - None => (false, name.as_str()), - }; + self.scopes.iter_mut().rev().enumerate().skip(skip).find_map(|(depth, scope)| { + let index = scope.values.iter().position(|value| value == name)?; + scope.record_read(Definition::Value(index)); - // A matrix applied to one argument is implicit multiplication, so `M(v)` matches `M v` - if names_matrix(bare_name) { - return binary(Node::Matrix(MatrixNode::Var(name)), BinaryOp::Mul, one(arguments)?); - } - - // A host function shadows the builtin of its spelling unless the `\` prefix asks for the language's own - let builtin = if !prefixed && functions.provides(bare_name) { None } else { builtin_function(bare_name) }; - - Ok(match builtin { - Some(Builtin::OfMatrix(function)) => Node::Value(ValueNode::OfMatrix { - function, - matrix: Box::new(matrix(one(arguments)?)?), - }), - Some(Builtin::MatrixOfMatrix(function)) => Node::Matrix(MatrixNode::OfMatrix { - function, - matrix: Box::new(matrix(one(arguments)?)?), - }), - Some(Builtin::MatrixOfValues { function, arity }) => { - if !arity.contains(&arguments.len()) { - return Err(INVALID_ARGUMENTS); - } - Node::Matrix(MatrixNode::FromValues { - function, - arguments: values(arguments)?, + let local = Local { depth, index }; + Some(if names_matrix(name) != (scope.flipped == Some(index)) { + Node::Matrix(MatrixNode::Local(local)) + } else { + Node::Value(ValueNode::Local(local)) }) + }) + } + + /// Where the innermost scope defining a function puts it, with its parameters, past the innermost `skip` scopes, which a `\` + /// name never reaches. + fn local_function(&mut self, name: &str, skip: usize) -> Option<(Local, Vec)> { + if name.starts_with('\\') { + return None; } - Some(Builtin::OfValueAndRegions { function, regions }) => { - if arguments.len() != regions + 1 { - return Err(INVALID_ARGUMENTS); - } - let mut arguments = arguments.into_iter(); - Node::Value(ValueNode::OfValueAndRegions { + + self.scopes.iter_mut().rev().enumerate().skip(skip).find_map(|(depth, scope)| { + let index = scope.functions.iter().position(|function| function.name == name)?; + scope.record_read(Definition::Function(index)); + Some((Local { depth, index }, scope.functions[index].parameters.clone())) + }) + } + + /// A call's sort follows the function's, where a `where` clause's function comes first and a matrix's name applies to its one + /// argument. The name looks past the innermost `skip` scopes, which only its arguments see. + fn call(&mut self, name: String, arguments: Vec, skip: usize) -> Result { + if let Some((function, parameters)) = self.local_function(&name, skip) { + return self.local_call(&name, function, ¶meters, arguments); + } + + let arguments = arguments.into_iter().map(|argument| self.sort(argument)).collect::, SortError>>()?; + let (prefixed, bare_name) = match name.strip_prefix('\\') { + Some(bare_name) => (true, bare_name), + None => (false, name.as_str()), + }; + + // A matrix applied to one argument is implicit multiplication, so `M(v)` matches `M v` + if names_matrix(bare_name) { + let matrix = self.name(name, skip); + return self.binary(matrix, BinaryOp::Mul, one(arguments)?); + } + + // A host function shadows the builtin of its spelling unless the `\` prefix asks for the language's own + let host_function = !prefixed && self.functions.provides(bare_name); + let builtin = if host_function { None } else { builtin_function(bare_name) }; + + // Where no function has the name, a value a `where` clause defines is implicit multiplication like any value, so `k(x + 1)` is `k (x + 1)` + if !host_function + && builtin.is_none() + && suffixed_function(bare_name).is_none() + && let Some(local) = self.local_value(&name, skip) + { + return self.binary(local, BinaryOp::Mul, one(arguments)?); + } + + Ok(match builtin { + Some(Builtin::OfMatrix(function)) => Node::Value(ValueNode::OfMatrix { + function, + matrix: Box::new(self.matrix(one(arguments)?, VALUE_AS_MATRIX)?), + }), + Some(Builtin::MatrixOfMatrix(function)) => Node::Matrix(MatrixNode::OfMatrix { function, - value: Box::new(value(arguments.next().ok_or(INVALID_ARGUMENTS)?)?), - regions: arguments.map(matrix).collect::, SortError>>()?, + matrix: Box::new(self.matrix(one(arguments)?, VALUE_AS_MATRIX)?), + }), + Some(Builtin::MatrixOfValues { function, arity }) => { + if !arity.contains(&arguments.len()) { + return Err(INVALID_ARGUMENTS); + } + Node::Matrix(MatrixNode::FromValues { + function, + arguments: self.values(arguments)?, + }) + } + Some(Builtin::OfValueAndRegions { function, regions }) => { + if arguments.len() != regions + 1 { + return Err(INVALID_ARGUMENTS); + } + let mut arguments = arguments.into_iter(); + Node::Value(ValueNode::OfValueAndRegions { + function, + value: Box::new(self.value(arguments.next().ok_or(INVALID_ARGUMENTS)?, MATRIX_AS_VALUE)?), + regions: arguments.map(|region| self.matrix(region, VALUE_AS_MATRIX)).collect::, SortError>>()?, + }) + } + _ => Node::Value(ValueNode::FnCall { name, expr: self.values(arguments)? }), + }) + } + + /// A call of a function a `where` clause defines, taking one argument per parameter, each of the sort its parameter's spelling fixes. + fn local_call(&mut self, name: &str, function: Local, parameters: &[String], arguments: Vec) -> Result { + if arguments.len() != parameters.len() { + let count = parameters.len(); + let plural = if count == 1 { "" } else { "s" }; + return Err(SortError(format!("`{name}` takes {count} argument{plural}").into())); + } + + let arguments = arguments + .into_iter() + .zip(parameters) + .enumerate() + .map(|(index, (argument, parameter))| { + let argument = self.sort(argument)?; + let passed_matrix = matches!(argument, Node::Matrix(_)); + if passed_matrix == names_matrix(parameter) { + return Ok(argument); + } + + // The parameter's case is the likelier mistake where the body could take what's passed, and the argument otherwise + if self.takes_other_sort(function, index) { + return Err(misnamed(parameter, &[], "passed", passed_matrix)); + } + let (passed, needed) = if passed_matrix { ("matrix", "value") } else { ("value", "matrix") }; + let signature = format!("{name}({})", parameters.join(", ")); + Err(SortError(format!("`{signature}` is passed a {passed} for `{parameter}`, which its body uses as a {needed}").into())) }) + .collect::, SortError>>()?; + + Ok(if names_matrix(name) { + Node::Matrix(MatrixNode::Call { function, arguments }) + } else { + Node::Value(ValueNode::Call { function, arguments }) + }) + } + + /// Whether a function's body sorts with one parameter read as the other sort than its case gives, as a call passing that sort needs. + /// The body is sorted again where the function is defined, so this serves only a failing call. + fn takes_other_sort(&mut self, function: Local, parameter: usize) -> bool { + // One test at a time, since a tested body may make a failing call that would test again + if self.scopes.iter().any(|scope| scope.flipped.is_some()) { + return false; } - _ => Node::Value(ValueNode::FnCall { name, expr: values(arguments)? }), - }) -} -/// The sort of an operation: value·value and matrix·value are values, matrix·matrix and value·matrix are matrices, and -/// `+` joins like sorts or attaches a value to a matrix as translation. -fn binary(lhs: Node, op: BinaryOp, rhs: Node) -> Result { - use BinaryOp as Op; - Ok(match (lhs, op, rhs) { - (Node::Value(lhs), op, Node::Value(rhs)) => Node::Value(ValueNode::BinOp { - lhs: Box::new(lhs), - op, - rhs: Box::new(rhs), - }), - (Node::Matrix(matrix), Op::Mul, Node::Value(value)) => Node::Value(ValueNode::Apply { - matrix: Box::new(matrix), - value: Box::new(value), - }), - // Division is times-inverse, so a matrix over a value applies to the value's reciprocal - (Node::Matrix(matrix), Op::Div, Node::Value(value)) => Node::Value(ValueNode::Apply { - matrix: Box::new(matrix), - value: Box::new(reciprocal(value)), - }), - (Node::Matrix(lhs), Op::Eq | Op::Neq, Node::Matrix(rhs)) => Node::Value(ValueNode::MatrixComparison { - matrices: vec![lhs, rhs], - distinct: op == Op::Neq, - }), - (lhs @ Node::Value(_), Op::Mul | Op::Div | Op::Add | Op::Sub, rhs @ Node::Matrix(_)) - | (lhs @ Node::Matrix(_), Op::Mul | Op::Div | Op::Add | Op::Sub, rhs @ Node::Matrix(_)) - | (lhs @ Node::Matrix(_), Op::Add | Op::Sub | Op::Pow, rhs @ Node::Value(_)) => Node::Matrix(MatrixNode::BinOp { - lhs: Box::new(lhs), - op, - rhs: Box::new(rhs), - }), - _ => return Err(NO_MATRIX_OPERATOR), - }) -} + let Some(definer) = self.scopes.len().checked_sub(function.depth + 1) else { return false }; + let Some(Function { parameters, body, .. }) = self.scopes[definer].functions.get(function.index) else { + return false; + }; + let (parameters, body) = (parameters.clone(), body.clone()); -fn reciprocal(value: ValueNode) -> ValueNode { - ValueNode::BinOp { - lhs: Box::new(ValueNode::Lit(Literal::Integer(1))), - op: BinaryOp::Div, - rhs: Box::new(value), + let inner_scopes = self.scopes.split_off(definer + 1); + self.scopes.push(Scope { + flipped: Some(parameter), + ..Scope::parameters(parameters) + }); + let sorts = self.sort(body).is_ok(); + self.scopes.truncate(definer + 1); + self.scopes.extend(inner_scopes); + + sorts } -} -/// A product folds left, except that a matrix meeting a value applies to the whole rest of the product, so `M 2 v` is `M (2 v)` -/// as in linear algebra rather than `(M 2) v`. -fn product(first: Node, rest: &mut std::vec::IntoIter<(BinaryOp, Syntax)>, functions: &dyn FunctionProvider) -> Result { - let mut accumulated = first; + /// A product folds left, except that a matrix meeting a value applies to the whole rest of the product, so `M 2 v` is `M (2 v)` + /// as in linear algebra rather than `(M 2) v`. + fn product(&mut self, first: Node, rest: &mut std::vec::IntoIter<(BinaryOp, Syntax)>) -> Result { + let mut accumulated = first; - while let Some((op, factor)) = rest.next() { - accumulated = match (accumulated, sorted(factor, functions)?) { - (Node::Matrix(matrix), Node::Value(factor)) => { - // Dividing by a value applies the matrix to its reciprocal times the rest, so `M / v w` is `M (v⁻¹ w)` - let factor = if op == BinaryOp::Div { reciprocal(factor) } else { factor }; + while let Some((op, factor)) = rest.next() { + accumulated = match (accumulated, self.sort(factor)?) { + (Node::Matrix(matrix), Node::Value(factor)) => { + // Dividing by a value applies the matrix to its reciprocal times the rest, so `M / v w` is `M (v⁻¹ w)` + let factor = if op == BinaryOp::Div { reciprocal(factor) } else { factor }; - let argument = product(Node::Value(factor), rest, functions)?; - return binary(Node::Matrix(matrix), BinaryOp::Mul, argument); - } - (accumulated, factor) => binary(accumulated, op, factor)?, - }; + let argument = self.product(Node::Value(factor), rest)?; + return self.binary(Node::Matrix(matrix), BinaryOp::Mul, argument); + } + (accumulated, factor) => self.binary(accumulated, op, factor)?, + }; + } + + Ok(accumulated) } - Ok(accumulated) -} + /// A piecewise takes the sort of its cases, which must agree, under conditions that are values. + fn piecewise(&mut self, cases: Vec, otherwise: Option>) -> Result { + let cases = cases + .into_iter() + .map(|Case { value: case, condition }| { + let case = self.sort(case)?; + let condition = self.sort(condition)?; + Ok((case, self.value(condition, MATRIX_AS_VALUE)?)) + }) + .collect::, SortError>>()?; + let otherwise = otherwise.map(|otherwise| self.sort(*otherwise)).transpose()?; -/// A piecewise takes the sort of its cases, which must agree, under conditions that are values. -fn piecewise(cases: Vec, otherwise: Option>, functions: &dyn FunctionProvider) -> Result { - let cases = cases - .into_iter() - .map(|Case { value: case, condition }| Ok((sorted(case, functions)?, value(sorted(condition, functions)?)?))) - .collect::, SortError>>()?; - let otherwise = otherwise.map(|otherwise| sorted(*otherwise, functions)).transpose()?; + let first = cases.first().map(|(case, _)| case).or(otherwise.as_ref()); + if first.is_some_and(|first| matches!(first, Node::Matrix(_))) { + let cases = cases + .into_iter() + .map(|(case, condition)| { + Ok(SortedCase { + value: self.matrix(case, MIXED_CASES)?, + condition, + }) + }) + .collect::, SortError>>()?; + let otherwise = otherwise.map(|otherwise| self.matrix(otherwise, MIXED_CASES)).transpose()?.map(Box::new); + return Ok(Node::Matrix(MatrixNode::Piecewise { cases, otherwise })); + } - let first = cases.first().map(|(case, _)| case).or(otherwise.as_ref()); - if first.is_some_and(|first| matches!(first, Node::Matrix(_))) { let cases = cases .into_iter() .map(|(case, condition)| { Ok(SortedCase { - value: matrix(case).map_err(|_| MIXED_CASES)?, + value: self.value(case, MIXED_CASES)?, condition, }) }) .collect::, SortError>>()?; - let otherwise = otherwise.map(|otherwise| matrix(otherwise).map_err(|_| MIXED_CASES)).transpose()?.map(Box::new); - return Ok(Node::Matrix(MatrixNode::Piecewise { cases, otherwise })); + let otherwise = otherwise.map(|otherwise| self.value(otherwise, MIXED_CASES)).transpose()?.map(Box::new); + Ok(Node::Value(ValueNode::Piecewise { cases, otherwise })) } - let cases = cases - .into_iter() - .map(|(case, condition)| { - Ok(SortedCase { - value: value(case).map_err(|_| MIXED_CASES)?, - condition, - }) + /// A `where` clause and the body it ends, sorted within a scope holding every definition of the clause at once, so each may + /// read any other whatever their order, as long as none depends on itself. The whole takes the sort of its body. + fn clause(&mut self, bindings: Vec, body: impl FnOnce(&mut Self) -> Result) -> Result { + let mut scope = Scope::default(); + let definitions = bindings.iter().map(|binding| scope.define(binding)).collect::, SortError>>()?; + self.scopes.push(scope); + let depth = self.scopes.len(); + + // Every definition is checked against its name's case before any read of it can fail, so a mismatch is blamed on the name, + // not the read. A definition that fails to sort waits until the others are checked, since it may read a misnamed one. + let mut clause = Clause { + values: Vec::new(), + functions: Vec::new(), + }; + let mut failure = None; + for (Binding { name, parameters, value }, definition) in bindings.into_iter().zip(definitions) { + if let Some(scope) = self.scopes.last_mut() { + scope.defining = Some(definition); + } + + // A function's parameters are the innermost scope of its body, shadowing every other name + let sorted = if parameters.is_empty() { + self.sort(value).map(|sorted| (sorted, parameters)) + } else { + self.scopes.push(Scope::parameters(parameters)); + self.sort(value).map(|sorted| (sorted, self.scopes.pop().unwrap_or_default().values)) + }; + + match sorted { + Ok((sorted, parameters)) => match (definition, defined_as(&name, ¶meters, sorted)?) { + (Definition::Value(_), sorted) => clause.values.push(sorted), + (Definition::Function(_), sorted) => clause.functions.push(sorted), + }, + Err(error) => { + self.scopes.truncate(depth); + failure.get_or_insert(error); + } + } + } + if let Some(error) = failure { + return Err(error); + } + + // The body reads the definitions without being one of them + if let Some(scope) = self.scopes.last_mut() { + scope.defining = None; + scope.reject_cycles()?; + } + let body = body(self)?; + self.scopes.pop(); + + let clause = Box::new(clause); + Ok(match body { + Node::Value(body) => Node::Value(ValueNode::Where { clause, body: Box::new(body) }), + Node::Matrix(body) => Node::Matrix(MatrixNode::Where { clause, body: Box::new(body) }), + }) + } + + /// The error for a node of the wrong sort, which names the parameter the node reads, if it's one, instead of `rejected`. + /// A `where` definition's name is left to its clause, which blames it only where its definition disagrees with it. + fn misplaced(&self, node: &Node, rejected: SortError) -> SortError { + let (Node::Value(ValueNode::Local(local)) | Node::Matrix(MatrixNode::Local(local))) = node else { + return rejected; + }; + let parameter = self + .scopes + .iter() + .rev() + .nth(local.depth) + .filter(|scope| scope.holds_parameters) + .and_then(|scope| scope.values.get(local.index)); + + match parameter { + Some(parameter) => misnamed(parameter, &[], "used as", matches!(node, Node::Value(_))), + None => rejected, + } + } + + fn value(&self, node: Node, rejected: SortError) -> Result { + match node { + Node::Value(value) => Ok(value), + node => Err(self.misplaced(&node, rejected)), + } + } + + fn matrix(&self, node: Node, rejected: SortError) -> Result { + match node { + Node::Matrix(matrix) => Ok(matrix), + node => Err(self.misplaced(&node, rejected)), + } + } + + fn values(&self, nodes: Vec) -> Result, SortError> { + nodes.into_iter().map(|node| self.value(node, MATRIX_AS_VALUE)).collect() + } + + /// The sort of an operation: value·value and matrix·value are values, matrix·matrix and value·matrix are matrices, and + /// `+` joins like sorts or attaches a value to a matrix as translation. + fn binary(&self, lhs: Node, op: BinaryOp, rhs: Node) -> Result { + use BinaryOp as Op; + Ok(match (lhs, op, rhs) { + (Node::Value(lhs), op, Node::Value(rhs)) => Node::Value(ValueNode::BinOp { + lhs: Box::new(lhs), + op, + rhs: Box::new(rhs), + }), + (Node::Matrix(matrix), Op::Mul, Node::Value(value)) => Node::Value(ValueNode::Apply { + matrix: Box::new(matrix), + value: Box::new(value), + }), + // Division is times-inverse, so a matrix over a value applies to the value's reciprocal + (Node::Matrix(matrix), Op::Div, Node::Value(value)) => Node::Value(ValueNode::Apply { + matrix: Box::new(matrix), + value: Box::new(reciprocal(value)), + }), + (Node::Matrix(lhs), Op::Eq | Op::Neq, Node::Matrix(rhs)) => Node::Value(ValueNode::MatrixComparison { + matrices: vec![lhs, rhs], + distinct: op == Op::Neq, + }), + (lhs @ Node::Value(_), Op::Mul | Op::Div | Op::Add | Op::Sub, rhs @ Node::Matrix(_)) + | (lhs @ Node::Matrix(_), Op::Mul | Op::Div | Op::Add | Op::Sub, rhs @ Node::Matrix(_)) + | (lhs @ Node::Matrix(_), Op::Add | Op::Sub | Op::Pow, rhs @ Node::Value(_)) => Node::Matrix(MatrixNode::BinOp { + lhs: Box::new(lhs), + op, + rhs: Box::new(rhs), + }), + // Where only one operand is a matrix, it's the one the operator can't take + (lhs, _, rhs) => { + let matrix = match (lhs, rhs) { + (matrix @ Node::Matrix(_), Node::Value(_)) | (Node::Value(_), matrix @ Node::Matrix(_)) => matrix, + _ => return Err(NO_MATRIX_OPERATOR), + }; + return Err(self.misplaced(&matrix, NO_MATRIX_OPERATOR)); + } }) - .collect::, SortError>>()?; - let otherwise = otherwise.map(|otherwise| value(otherwise).map_err(|_| MIXED_CASES)).transpose()?.map(Box::new); - Ok(Node::Value(ValueNode::Piecewise { cases, otherwise })) + } +} + +/// A definition of the sort its name's spelling fixes, a matrix for a name beginning with a capital letter and a value otherwise. +fn defined_as(name: &str, parameters: &[String], definition: Node) -> Result { + let is_matrix = matches!(definition, Node::Matrix(_)); + if names_matrix(name) == is_matrix { + return Ok(definition); + } + Err(misnamed(name, parameters, "defined as", is_matrix)) +} + +/// Advice to recase a `where` name or parameter defined, used, or passed as the other sort, since its case is the likelier +/// mistake, with a function quoted with its parameters. The recased name is offered where recasing gives that sort. +fn misnamed(name: &str, parameters: &[String], treated: &str, as_matrix: bool) -> SortError { + let (sort, letter) = if as_matrix { ("matrix", "capital") } else { ("value", "lowercase") }; + let mut message = if parameters.is_empty() { + format!("`{name}` is {treated} a {sort}, so rename it to begin with a {letter} letter") + } else { + format!("`{name}({})` is {treated} a {sort}, so rename `{name}` to begin with a {letter} letter", parameters.join(", ")) + }; + + let mut characters = name.chars(); + let recased: String = match characters.next() { + Some(first) if as_matrix => first.to_uppercase().chain(characters).collect(), + Some(first) => first.to_lowercase().chain(characters).collect(), + None => String::new(), + }; + if names_matrix(&recased) == as_matrix { + message.push_str(&format!(", like `{recased}`")); + } + + SortError(message.into()) +} + +fn one(arguments: Vec) -> Result { + <[Node; 1]>::try_from(arguments).map(|[argument]| argument).map_err(|_| INVALID_ARGUMENTS) +} + +fn reciprocal(value: ValueNode) -> ValueNode { + ValueNode::BinOp { + lhs: Box::new(ValueNode::Lit(Literal::Integer(1))), + op: BinaryOp::Div, + rhs: Box::new(value), + } }