From 6bdca278f35e81ef82bb80035e2f4e107ab804a2 Mon Sep 17 00:00:00 2001 From: Paul Masurel Date: Mon, 31 Aug 2026 17:57:51 +0200 Subject: [PATCH] Added a way to build a function call untyped expression in a validated way --- jitexpr/src/ast/infer_types.rs | 3 + jitexpr/src/ast/mod.rs | 2 +- jitexpr/src/ast/untyped_expr.rs | 11 + jitexpr/src/compile/error.rs | 4 +- jitexpr/src/functions/abs.rs | 4 +- jitexpr/src/functions/add.rs | 3 + jitexpr/src/functions/and.rs | 4 +- jitexpr/src/functions/ceil.rs | 4 +- jitexpr/src/functions/concat.rs | 35 ++- jitexpr/src/functions/divide.rs | 4 +- jitexpr/src/functions/eq.rs | 4 +- jitexpr/src/functions/floor.rs | 4 +- jitexpr/src/functions/gt.rs | 4 +- jitexpr/src/functions/gt_eq.rs | 4 +- jitexpr/src/functions/if_fn.rs | 4 +- jitexpr/src/functions/int_mod.rs | 4 +- jitexpr/src/functions/is_not_null.rs | 4 +- jitexpr/src/functions/is_null.rs | 4 +- jitexpr/src/functions/left.rs | 34 ++- jitexpr/src/functions/lower.rs | 4 +- jitexpr/src/functions/lt.rs | 4 +- jitexpr/src/functions/lt_eq.rs | 4 +- jitexpr/src/functions/max.rs | 4 +- jitexpr/src/functions/min.rs | 4 +- jitexpr/src/functions/mod.rs | 262 ++++++++++++++++++++++- jitexpr/src/functions/multiply.rs | 5 +- jitexpr/src/functions/neq.rs | 4 +- jitexpr/src/functions/not.rs | 4 +- jitexpr/src/functions/or.rs | 4 +- jitexpr/src/functions/pow.rs | 4 +- jitexpr/src/functions/regexp_extract.rs | 37 +++- jitexpr/src/functions/regexp_like.rs | 17 +- jitexpr/src/functions/right.rs | 34 ++- jitexpr/src/functions/round.rs | 46 ++-- jitexpr/src/functions/split_after.rs | 55 +++-- jitexpr/src/functions/split_before.rs | 55 +++-- jitexpr/src/functions/sqrt.rs | 4 +- jitexpr/src/functions/substring.rs | 42 +++- jitexpr/src/functions/substring_count.rs | 3 + jitexpr/src/functions/subtract.rs | 5 +- jitexpr/src/functions/text_join.rs | 8 + jitexpr/src/functions/trim.rs | 46 +++- jitexpr/src/functions/upper.rs | 4 +- 43 files changed, 678 insertions(+), 121 deletions(-) diff --git a/jitexpr/src/ast/infer_types.rs b/jitexpr/src/ast/infer_types.rs index 42a30acc7..9d5ea5647 100644 --- a/jitexpr/src/ast/infer_types.rs +++ b/jitexpr/src/ast/infer_types.rs @@ -2,6 +2,7 @@ use std::collections::HashMap; use std::collections::hash_map::Entry; use crate::ast::{Function, Literal, UntypedExpr}; +use crate::functions::InvalidFunctionCall; use crate::types::VarType; #[derive(Default, Copy, Clone, Debug, Eq, PartialEq)] @@ -137,6 +138,8 @@ impl std::fmt::Display for InferredTypeSet { #[derive(Debug, thiserror::Error)] pub enum TypeError { + #[error(transparent)] + InvalidFunctionCall(#[from] InvalidFunctionCall), #[error("function `{function:?}` returns `{got}`, expected `{expected}`")] WrongFunctionReturnType { function: Function, diff --git a/jitexpr/src/ast/mod.rs b/jitexpr/src/ast/mod.rs index b872a349d..e7c3a396e 100644 --- a/jitexpr/src/ast/mod.rs +++ b/jitexpr/src/ast/mod.rs @@ -9,4 +9,4 @@ pub use literal::Literal; pub use serialize::{DeserializeError, deserialize, serialize}; pub use untyped_expr::UntypedExpr; -pub use crate::functions::Function; +pub use crate::functions::{Function, InvalidFunctionCall}; diff --git a/jitexpr/src/ast/untyped_expr.rs b/jitexpr/src/ast/untyped_expr.rs index 0127639a9..c19933ff6 100644 --- a/jitexpr/src/ast/untyped_expr.rs +++ b/jitexpr/src/ast/untyped_expr.rs @@ -1,6 +1,7 @@ use std::sync::Arc; use crate::ast::{Function, Literal}; +use crate::functions::InvalidFunctionCall; /// An expression AST. /// @@ -23,6 +24,16 @@ impl UntypedExpr { pub fn variable(variable_name: impl ToString) -> UntypedExpr { UntypedExpr::Variable(Arc::from(variable_name.to_string())) } + + /// Creates an untyped expression that is a function over different arguments. + /// + /// This call will validate the arguments and + pub fn call( + function: Function, + args: Vec, + ) -> Result { + function.call(args) + } } impl From for UntypedExpr { diff --git a/jitexpr/src/compile/error.rs b/jitexpr/src/compile/error.rs index 13336c354..c63d2da59 100644 --- a/jitexpr/src/compile/error.rs +++ b/jitexpr/src/compile/error.rs @@ -1,4 +1,4 @@ -use crate::ast::{Function, TypeError}; +use crate::ast::{Function, InvalidFunctionCall, TypeError}; use crate::types::VarType; #[derive(Debug, thiserror::Error)] @@ -20,6 +20,8 @@ pub enum CompileError { #[source] source: regex::Error, }, + #[error("arguments do not match the function {0}")] + InvalidArguments(#[from] InvalidFunctionCall), } impl From for CompileError { diff --git a/jitexpr/src/functions/abs.rs b/jitexpr/src/functions/abs.rs index 2823df49e..a8917fbc4 100644 --- a/jitexpr/src/functions/abs.rs +++ b/jitexpr/src/functions/abs.rs @@ -26,6 +26,8 @@ pub(crate) struct AbsFnCall { } impl FnCall for AbsFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -63,7 +65,7 @@ impl FnCall for AbsFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 1, "expected 1 arg for ABS"); + Self::ARG_COUNT.validate(args)?; let target_types = match &args[0] { UntypedExpr::Literal(literal) => { let declared = InferredTypeSet::singleton(literal.r#type()); diff --git a/jitexpr/src/functions/add.rs b/jitexpr/src/functions/add.rs index e9e78b509..5a97825ae 100644 --- a/jitexpr/src/functions/add.rs +++ b/jitexpr/src/functions/add.rs @@ -30,6 +30,8 @@ pub(crate) struct AddFnCall { } impl FnCall for AddFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Any; + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -65,6 +67,7 @@ impl FnCall for AddFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { + Self::ARG_COUNT.validate(args)?; let mut return_types = InferredTypeSet::NUMERICAL.intersect(target_type_set); for arg in args { let arg_types = crate::ast::infer_type_with_variable_types( diff --git a/jitexpr/src/functions/and.rs b/jitexpr/src/functions/and.rs index 1e2fd8da5..801f97562 100644 --- a/jitexpr/src/functions/and.rs +++ b/jitexpr/src/functions/and.rs @@ -22,6 +22,8 @@ pub(crate) struct AndFnCall { } impl FnCall for AndFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::AtLeast(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -53,7 +55,7 @@ impl FnCall for AndFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert!(!args.is_empty(), "expected at least 1 arg for AND"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); let args = args diff --git a/jitexpr/src/functions/ceil.rs b/jitexpr/src/functions/ceil.rs index 6e9e6d739..39cdd7e4d 100644 --- a/jitexpr/src/functions/ceil.rs +++ b/jitexpr/src/functions/ceil.rs @@ -23,6 +23,8 @@ pub(crate) struct CeilFnCall { } impl FnCall for CeilFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -51,7 +53,7 @@ impl FnCall for CeilFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 1, "expected 1 arg for CEIL"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::I64)); let arg = context.apply_types(&args[0], InferredTypeSet::NUMERICAL)?; if arg.return_type == VarType::None { diff --git a/jitexpr/src/functions/concat.rs b/jitexpr/src/functions/concat.rs index 2745e67ec..a5e87fed0 100644 --- a/jitexpr/src/functions/concat.rs +++ b/jitexpr/src/functions/concat.rs @@ -48,6 +48,13 @@ pub(super) struct JoinArguments { } impl FnCall for ConcatFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::AtLeast(4); + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + validate_join_args(args) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -61,6 +68,7 @@ impl FnCall for ConcatFnCall { _target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { + Self::ARG_COUNT.validate(args)?; let Some(arguments) = apply_join_types("CONCAT", args, context)? else { return Ok(TypedExpr::none()); }; @@ -89,6 +97,15 @@ impl FnCall for ConcatFnCall { } } +pub(super) fn validate_join_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + super::validate_literal(args, 0, VarType::Str, |literal| { + matches!(literal, Literal::String(_)) + })?; + super::validate_literal(args, 1, VarType::Str, |literal| { + matches!(literal, Literal::String(_)) + }) +} + pub(super) fn infer_join_types<'a>( function: Function, args: &'a [UntypedExpr], @@ -116,19 +133,23 @@ pub(super) fn infer_join_types<'a>( } pub(super) fn apply_join_types( - function_name: &str, + _function_name: &str, args: &[UntypedExpr], context: &mut CompileFnBuilder<'_, '_>, ) -> Result, CompileError> { - assert!( - args.len() >= 4, - "expected at least 4 args for {function_name}" - ); let UntypedExpr::Literal(Literal::String(delimiter)) = &args[0] else { - panic!("{function_name} delimiter must be a string literal"); + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 1, + expected: VarType::Str, + } + .into()); }; let UntypedExpr::Literal(Literal::String(ignore_empty)) = &args[1] else { - panic!("{function_name} ignore-empty flag must be a string literal"); + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::Str, + } + .into()); }; let values = args[2..] diff --git a/jitexpr/src/functions/divide.rs b/jitexpr/src/functions/divide.rs index bca9f3008..74408a95c 100644 --- a/jitexpr/src/functions/divide.rs +++ b/jitexpr/src/functions/divide.rs @@ -24,6 +24,8 @@ pub(crate) struct DivideFnCall { } impl FnCall for DivideFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -55,7 +57,7 @@ impl FnCall for DivideFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for DIVIDE"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::F64)); let typed_args = args diff --git a/jitexpr/src/functions/eq.rs b/jitexpr/src/functions/eq.rs index 87c70f10e..2624a1510 100644 --- a/jitexpr/src/functions/eq.rs +++ b/jitexpr/src/functions/eq.rs @@ -36,6 +36,8 @@ pub(crate) struct EqFnCall { } impl FnCall for EqFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -68,7 +70,7 @@ impl FnCall for EqFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "Expected 2 args for EQ"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); // EQ must retain each literal's declared type. In contrast with ADD, it diff --git a/jitexpr/src/functions/floor.rs b/jitexpr/src/functions/floor.rs index bd7634339..3de105646 100644 --- a/jitexpr/src/functions/floor.rs +++ b/jitexpr/src/functions/floor.rs @@ -23,6 +23,8 @@ pub(crate) struct FloorFnCall { } impl FnCall for FloorFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -51,7 +53,7 @@ impl FnCall for FloorFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 1, "expected 1 arg for FLOOR"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::I64)); let arg = context.apply_types(&args[0], InferredTypeSet::NUMERICAL)?; if arg.return_type == VarType::None { diff --git a/jitexpr/src/functions/gt.rs b/jitexpr/src/functions/gt.rs index f541ad391..d62ac25ab 100644 --- a/jitexpr/src/functions/gt.rs +++ b/jitexpr/src/functions/gt.rs @@ -25,6 +25,8 @@ pub(crate) struct GtFnCall { } impl FnCall for GtFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -38,7 +40,7 @@ impl FnCall for GtFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for GT"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); Ok(TypedExpr { return_type: VarType::Bool, diff --git a/jitexpr/src/functions/gt_eq.rs b/jitexpr/src/functions/gt_eq.rs index b04c7a8d5..6cd46656f 100644 --- a/jitexpr/src/functions/gt_eq.rs +++ b/jitexpr/src/functions/gt_eq.rs @@ -25,6 +25,8 @@ pub(crate) struct GtEqFnCall { } impl FnCall for GtEqFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -38,7 +40,7 @@ impl FnCall for GtEqFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for GT_EQ"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); Ok(TypedExpr { return_type: VarType::Bool, diff --git a/jitexpr/src/functions/if_fn.rs b/jitexpr/src/functions/if_fn.rs index abdbd49c9..9b3ee47a5 100644 --- a/jitexpr/src/functions/if_fn.rs +++ b/jitexpr/src/functions/if_fn.rs @@ -47,6 +47,8 @@ fn select_type(types: InferredTypeSet) -> VarType { } impl FnCall for IfFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(3); + fn infer_types<'a>( args: &'a [UntypedExpr], target: InferredTypeSet, @@ -78,7 +80,7 @@ impl FnCall for IfFnCall { target: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 3, "expected 3 args for IF"); + Self::ARG_COUNT.validate(args)?; let condition = context.apply_types(&args[0], InferredTypeSet::BOOLEAN)?; let left = crate::ast::infer_type_with_variable_types(&args[1], target, context.variable_types())?; diff --git a/jitexpr/src/functions/int_mod.rs b/jitexpr/src/functions/int_mod.rs index bca5f672c..ec9f025f4 100644 --- a/jitexpr/src/functions/int_mod.rs +++ b/jitexpr/src/functions/int_mod.rs @@ -31,6 +31,8 @@ pub(crate) struct IntModFnCall { } impl FnCall for IntModFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -61,7 +63,7 @@ impl FnCall for IntModFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for INT_MOD"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::F64)); let args = args .iter() diff --git a/jitexpr/src/functions/is_not_null.rs b/jitexpr/src/functions/is_not_null.rs index c263fdebf..e58f83a2a 100644 --- a/jitexpr/src/functions/is_not_null.rs +++ b/jitexpr/src/functions/is_not_null.rs @@ -22,6 +22,8 @@ pub(crate) struct IsNotNullFnCall { } impl FnCall for IsNotNullFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -51,7 +53,7 @@ impl FnCall for IsNotNullFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 1, "Expected 1 arg for IS_NOT_NULL"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); let arg = context.apply_types(&args[0], InferredTypeSet::ALL)?; diff --git a/jitexpr/src/functions/is_null.rs b/jitexpr/src/functions/is_null.rs index 71d5e570d..4f3ef6657 100644 --- a/jitexpr/src/functions/is_null.rs +++ b/jitexpr/src/functions/is_null.rs @@ -21,6 +21,8 @@ pub(crate) struct IsNullFnCall { } impl FnCall for IsNullFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -50,7 +52,7 @@ impl FnCall for IsNullFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 1, "expected 1 arg for IS_NULL"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); let arg = context.apply_types(&args[0], InferredTypeSet::ALL)?; diff --git a/jitexpr/src/functions/left.rs b/jitexpr/src/functions/left.rs index 02750b1d9..c3be387a6 100644 --- a/jitexpr/src/functions/left.rs +++ b/jitexpr/src/functions/left.rs @@ -22,22 +22,38 @@ pub(crate) struct LeftFnCall { length: usize, } -fn constant_length(expression: &UntypedExpr) -> Option { +fn constant_length(expression: &UntypedExpr) -> Result, super::InvalidFunctionCall> { let UntypedExpr::Literal(literal) = expression else { - panic!("LEFT length must be constant"); + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::I64, + }); }; - match literal { + if !literal.is_none() && !literal.types().contains(VarType::I64) { + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::I64, + }); + } + Ok(match literal { Literal::I64(value) => usize::try_from(*value).ok(), Literal::U64(value) => usize::try_from(*value).ok(), Literal::F64(value) => usize::try_from(*value as i64).ok(), Literal::None => None, - Literal::Bool(_) | Literal::String(_) => { - unreachable!("type inference constrains LEFT length to an integer") - } - } + Literal::Bool(_) | Literal::String(_) => None, + }) } impl FnCall for LeftFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + super::validate_literal(args, 1, VarType::I64, |literal| { + literal.is_none() || literal.types().contains(VarType::I64) + }) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -67,9 +83,9 @@ impl FnCall for LeftFnCall { _target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for LEFT"); + Self::ARG_COUNT.validate(args)?; let input = context.apply_types(&args[0], InferredTypeSet::STRING)?; - let Some(length) = constant_length(&args[1]) else { + let Some(length) = constant_length(&args[1])? else { return Ok(TypedExpr::none()); }; Ok(TypedExpr { diff --git a/jitexpr/src/functions/lower.rs b/jitexpr/src/functions/lower.rs index 8c8b27eaf..be3a5894c 100644 --- a/jitexpr/src/functions/lower.rs +++ b/jitexpr/src/functions/lower.rs @@ -32,6 +32,8 @@ pub(crate) struct LowerFnCall { } impl FnCall for LowerFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -60,7 +62,7 @@ impl FnCall for LowerFnCall { _target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 1, "Expected 1 arg for LOWER"); + Self::ARG_COUNT.validate(args)?; let arg = context.apply_types(&args[0], InferredTypeSet::STRING)?; if arg.return_type == VarType::None { return Ok(TypedExpr::none()); diff --git a/jitexpr/src/functions/lt.rs b/jitexpr/src/functions/lt.rs index f78f6dd49..8539303a8 100644 --- a/jitexpr/src/functions/lt.rs +++ b/jitexpr/src/functions/lt.rs @@ -25,6 +25,8 @@ pub(crate) struct LtFnCall { } impl FnCall for LtFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -38,7 +40,7 @@ impl FnCall for LtFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for LT"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); Ok(TypedExpr { return_type: VarType::Bool, diff --git a/jitexpr/src/functions/lt_eq.rs b/jitexpr/src/functions/lt_eq.rs index 45dc52f64..723f69e2c 100644 --- a/jitexpr/src/functions/lt_eq.rs +++ b/jitexpr/src/functions/lt_eq.rs @@ -25,6 +25,8 @@ pub(crate) struct LtEqFnCall { } impl FnCall for LtEqFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -38,7 +40,7 @@ impl FnCall for LtEqFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for LT_EQ"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); Ok(TypedExpr { return_type: VarType::Bool, diff --git a/jitexpr/src/functions/max.rs b/jitexpr/src/functions/max.rs index 7e56fb42e..71ffefcc8 100644 --- a/jitexpr/src/functions/max.rs +++ b/jitexpr/src/functions/max.rs @@ -24,6 +24,8 @@ pub(crate) struct MaxFnCall { } impl FnCall for MaxFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::AtLeast(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -67,7 +69,7 @@ impl FnCall for MaxFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert!(!args.is_empty(), "expected at least 1 arg for MAX"); + Self::ARG_COUNT.validate(args)?; let mut return_types = InferredTypeSet::NUMERICAL.intersect(target_type_set); for arg in args { return_types = return_types.intersect(crate::ast::infer_type_with_variable_types( diff --git a/jitexpr/src/functions/min.rs b/jitexpr/src/functions/min.rs index 80fc0fb54..a6bf8da95 100644 --- a/jitexpr/src/functions/min.rs +++ b/jitexpr/src/functions/min.rs @@ -24,6 +24,8 @@ pub(crate) struct MinFnCall { } impl FnCall for MinFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::AtLeast(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -67,7 +69,7 @@ impl FnCall for MinFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert!(!args.is_empty(), "expected at least 1 arg for MIN"); + Self::ARG_COUNT.validate(args)?; let mut return_types = InferredTypeSet::NUMERICAL.intersect(target_type_set); for arg in args { return_types = return_types.intersect(crate::ast::infer_type_with_variable_types( diff --git a/jitexpr/src/functions/mod.rs b/jitexpr/src/functions/mod.rs index 14446334e..8a56a4754 100644 --- a/jitexpr/src/functions/mod.rs +++ b/jitexpr/src/functions/mod.rs @@ -84,7 +84,7 @@ pub(crate) use self::subtract::SubtractFnCall; pub(crate) use self::text_join::TextJoinFnCall; pub(crate) use self::trim::TrimFnCall; pub(crate) use self::upper::UpperFnCall; -use crate::ast::{InferredTypeSet, TypeError, UntypedExpr}; +use crate::ast::{InferredTypeSet, Literal, TypeError, UntypedExpr}; use crate::compile::{CompileError, CompileFnBuilder, LoweredValue, LoweringContext, TypedExpr}; use crate::types::VarType; @@ -170,6 +170,54 @@ pub enum Function { } impl Function { + pub(crate) fn call(self, args: Vec) -> Result { + match self { + Function::Abs => ::validate_args(&args)?, + Function::And => ::validate_args(&args)?, + Function::Ceil => ::validate_args(&args)?, + Function::Concat => ::validate_args(&args)?, + Function::Add => ::validate_args(&args)?, + Function::Divide => ::validate_args(&args)?, + Function::Eq => ::validate_args(&args)?, + Function::Floor => ::validate_args(&args)?, + Function::Gt => ::validate_args(&args)?, + Function::GtEq => ::validate_args(&args)?, + Function::If => ::validate_args(&args)?, + Function::IntMod => ::validate_args(&args)?, + Function::Left => ::validate_args(&args)?, + Function::Lt => ::validate_args(&args)?, + Function::LtEq => ::validate_args(&args)?, + Function::IsNotNull => ::validate_args(&args)?, + Function::IsNull => ::validate_args(&args)?, + Function::Lower => ::validate_args(&args)?, + Function::Max => ::validate_args(&args)?, + Function::Min => ::validate_args(&args)?, + Function::Multiply => ::validate_args(&args)?, + Function::Neq => ::validate_args(&args)?, + Function::Not => ::validate_args(&args)?, + Function::Or => ::validate_args(&args)?, + Function::Pow => ::validate_args(&args)?, + Function::Sqrt => ::validate_args(&args)?, + Function::RegexpExtract => ::validate_args(&args)?, + Function::RegexpLike => ::validate_args(&args)?, + Function::Right => ::validate_args(&args)?, + Function::Round => ::validate_args(&args)?, + Function::SplitAfter => ::validate_args(&args)?, + Function::SplitBefore => ::validate_args(&args)?, + Function::Subtract => ::validate_args(&args)?, + Function::Substring => ::validate_args(&args)?, + Function::SubstringCount => ::validate_args(&args)?, + Function::TextJoin => ::validate_args(&args)?, + Function::Trim => ::validate_args(&args)?, + Function::Upper => ::validate_args(&args)?, + } + + Ok(UntypedExpr::Call { + function: self, + args, + }) + } + pub(crate) fn call_with_types( self, args: &[UntypedExpr], @@ -557,12 +605,104 @@ impl FnCallEnum { } } +/// Error representing an invalid function call. +#[derive(Debug, Eq, PartialEq, thiserror::Error)] +pub enum InvalidFunctionCall { + #[error("invalid number of arguments: expected {expected}, got {provided}")] + InvalidNumberOfArguments { + expected: ArgumentCount, + provided: usize, + }, + #[error("argument {argument} must be a {expected:?} literal")] + ExpectedLiteral { argument: usize, expected: VarType }, + #[error("invalid value for argument {argument}: expected {expected}")] + InvalidLiteralValue { + argument: usize, + expected: &'static str, + }, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum ArgumentCount { + Any, + Exactly(usize), + AtLeast(usize), + Between { min: usize, max: usize }, +} + +impl std::fmt::Display for ArgumentCount { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + match *self { + ArgumentCount::Any => formatter.write_str("any number of arguments"), + ArgumentCount::Exactly(1) => formatter.write_str("exactly 1 argument"), + ArgumentCount::Exactly(count) => { + write!(formatter, "exactly {count} arguments") + } + ArgumentCount::AtLeast(1) => formatter.write_str("at least 1 argument"), + ArgumentCount::AtLeast(count) => { + write!(formatter, "at least {count} arguments") + } + ArgumentCount::Between { min: 1, max: 1 } => formatter.write_str("exactly 1 argument"), + ArgumentCount::Between { min, max } if min == max => { + write!(formatter, "exactly {min} arguments") + } + ArgumentCount::Between { min, max } => { + write!(formatter, "between {min} and {max} arguments") + } + } + } +} + +impl ArgumentCount { + fn validate(self, args: &[UntypedExpr]) -> Result<(), InvalidFunctionCall> { + let provided = args.len(); + let is_valid = match self { + ArgumentCount::Any => true, + ArgumentCount::Exactly(expected) => provided == expected, + ArgumentCount::AtLeast(expected) => provided >= expected, + ArgumentCount::Between { min, max } => (min..=max).contains(&provided), + }; + if is_valid { + Ok(()) + } else { + Err(InvalidFunctionCall::InvalidNumberOfArguments { + expected: self, + provided, + }) + } + } +} + +pub(crate) fn validate_literal( + args: &[UntypedExpr], + index: usize, + expected: VarType, + is_valid: impl FnOnce(&Literal) -> bool, +) -> Result<(), InvalidFunctionCall> { + let Some(UntypedExpr::Literal(literal)) = args.get(index) else { + return Err(InvalidFunctionCall::ExpectedLiteral { + argument: index + 1, + expected, + }); + }; + if is_valid(literal) { + Ok(()) + } else { + Err(InvalidFunctionCall::ExpectedLiteral { + argument: index + 1, + expected, + }) + } +} + /// Implements the type-inference, typed-AST, and lowering phases of a function call. /// /// The static methods operate on an [`UntypedExpr`] call before a concrete call node exists. /// Once [`FnCall::call_with_types`] has produced that node, [`FnCall::args_mut`] and /// [`FnCall::lower`] operate on its typed representation. pub(crate) trait FnCall: std::fmt::Debug + Into { + const ARG_COUNT: ArgumentCount; + /// Constrains the call and its arguments to the types accepted by its parent expression. /// /// Implementations validate their signature, recursively infer every argument, update @@ -576,6 +716,11 @@ pub(crate) trait FnCall: std::fmt::Debug + Into { where Self: Sized; + fn validate_args(args: &[UntypedExpr]) -> Result<(), InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + Ok(()) + } + /// Builds the typed call after concrete variable types have been supplied. /// /// `target_type_set` communicates the result types set accepted by the parent call. The @@ -613,3 +758,118 @@ pub(crate) trait FnCall: std::fmt::Debug + Into { builder: &mut FunctionBuilder<'_>, ) -> Result; } + +#[cfg(test)] +mod tests { + use super::*; + use crate::compile::compile; + + fn call_error(function: Function, args: Vec) -> InvalidFunctionCall { + match UntypedExpr::call(function, args) { + Ok(_) => panic!("expected the function call to be rejected"), + Err(error) => error, + } + } + + #[test] + fn test_argument_count_display() { + assert_eq!(ArgumentCount::Any.to_string(), "any number of arguments"); + assert_eq!(ArgumentCount::Exactly(1).to_string(), "exactly 1 argument"); + assert_eq!(ArgumentCount::Exactly(2).to_string(), "exactly 2 arguments"); + assert_eq!(ArgumentCount::AtLeast(1).to_string(), "at least 1 argument"); + assert_eq!( + ArgumentCount::Between { min: 2, max: 3 }.to_string(), + "between 2 and 3 arguments" + ); + } + + #[test] + fn test_argument_count_validation() { + assert_eq!( + call_error(Function::Abs, Vec::new()), + InvalidFunctionCall::InvalidNumberOfArguments { + expected: ArgumentCount::Exactly(1), + provided: 0, + } + ); + assert_eq!( + call_error(Function::And, Vec::new()), + InvalidFunctionCall::InvalidNumberOfArguments { + expected: ArgumentCount::AtLeast(1), + provided: 0, + } + ); + assert_eq!( + call_error(Function::Round, Vec::new()), + InvalidFunctionCall::InvalidNumberOfArguments { + expected: ArgumentCount::Between { min: 1, max: 2 }, + provided: 0, + } + ); + assert!(UntypedExpr::call(Function::Add, Vec::new()).is_ok()); + } + + #[test] + fn test_literal_argument_validation() { + assert_eq!( + call_error( + Function::RegexpLike, + vec![ + UntypedExpr::variable("input"), + UntypedExpr::variable("pattern") + ], + ), + InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::Str, + } + ); + assert_eq!( + call_error( + Function::RegexpExtract, + vec![ + UntypedExpr::variable("input"), + UntypedExpr::literal("pattern"), + UntypedExpr::literal(1i64), + ], + ), + InvalidFunctionCall::ExpectedLiteral { + argument: 3, + expected: VarType::U64, + } + ); + } + + #[test] + fn test_typed_construction_validates_unchecked_ast() { + let expression = Function::Abs.call_untyped_expr(Vec::new()); + let error = match compile(&expression, &HashMap::new()) { + Ok(_) => panic!("expected compilation to reject the unchecked AST"), + Err(error) => error, + }; + assert!(matches!( + error, + CompileError::InvalidArguments(InvalidFunctionCall::InvalidNumberOfArguments { + expected: ArgumentCount::Exactly(1), + provided: 0, + }) + )); + + let expression = Function::RegexpLike.call_untyped_expr(vec![ + UntypedExpr::variable("input"), + UntypedExpr::variable("pattern"), + ]); + let variable_types = HashMap::from([("input", VarType::Str), ("pattern", VarType::Str)]); + let error = match compile(&expression, &variable_types) { + Ok(_) => panic!("expected compilation to reject the non-literal pattern"), + Err(error) => error, + }; + assert!(matches!( + error, + CompileError::InvalidArguments(InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::Str, + }) + )); + } +} diff --git a/jitexpr/src/functions/multiply.rs b/jitexpr/src/functions/multiply.rs index b98dcd15c..24c791241 100644 --- a/jitexpr/src/functions/multiply.rs +++ b/jitexpr/src/functions/multiply.rs @@ -25,6 +25,8 @@ pub(crate) struct MultiplyFnCall { } impl FnCall for MultiplyFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -68,8 +70,7 @@ impl FnCall for MultiplyFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for MULTIPLY"); - + Self::ARG_COUNT.validate(args)?; let mut return_types = InferredTypeSet::NUMERICAL.intersect(target_type_set); for arg in args { let arg_types = crate::ast::infer_type_with_variable_types( diff --git a/jitexpr/src/functions/neq.rs b/jitexpr/src/functions/neq.rs index b735f56b1..dc9e44ec0 100644 --- a/jitexpr/src/functions/neq.rs +++ b/jitexpr/src/functions/neq.rs @@ -26,6 +26,8 @@ pub(crate) struct NeqFnCall { } impl FnCall for NeqFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -56,7 +58,7 @@ impl FnCall for NeqFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for NEQ"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); let typed_args = args .iter() diff --git a/jitexpr/src/functions/not.rs b/jitexpr/src/functions/not.rs index 7253c4d23..9cd44ff4b 100644 --- a/jitexpr/src/functions/not.rs +++ b/jitexpr/src/functions/not.rs @@ -25,6 +25,8 @@ pub(crate) struct NotFnCall { } impl FnCall for NotFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -54,7 +56,7 @@ impl FnCall for NotFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 1, "expected 1 arg for NOT"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); let arg = context.apply_types(&args[0], InferredTypeSet::BOOLEAN)?; diff --git a/jitexpr/src/functions/or.rs b/jitexpr/src/functions/or.rs index 652004d65..293edeba4 100644 --- a/jitexpr/src/functions/or.rs +++ b/jitexpr/src/functions/or.rs @@ -24,6 +24,8 @@ pub(crate) struct OrFnCall { } impl FnCall for OrFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::AtLeast(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -55,7 +57,7 @@ impl FnCall for OrFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert!(!args.is_empty(), "expected at least 1 arg for OR"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::Bool)); let args = args diff --git a/jitexpr/src/functions/pow.rs b/jitexpr/src/functions/pow.rs index e414e001d..6c6d079ce 100644 --- a/jitexpr/src/functions/pow.rs +++ b/jitexpr/src/functions/pow.rs @@ -30,6 +30,8 @@ pub(crate) struct PowFnCall { } impl FnCall for PowFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -60,7 +62,7 @@ impl FnCall for PowFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for POW"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::F64)); let args = args .iter() diff --git a/jitexpr/src/functions/regexp_extract.rs b/jitexpr/src/functions/regexp_extract.rs index ca8e9fe3d..fe091e4da 100644 --- a/jitexpr/src/functions/regexp_extract.rs +++ b/jitexpr/src/functions/regexp_extract.rs @@ -25,7 +25,7 @@ use crate::ast::{Function, InferredTypeSet, Literal, TypeError, UntypedExpr}; use crate::compile::{ CompileError, CompileFnBuilder, LoweredValue, LoweringContext, TypedExpr, TypedExprAst, }; -use crate::functions::{FnCall, FnCallEnum}; +use crate::functions::{FnCall, FnCallEnum, InvalidFunctionCall}; use crate::types::VarType; const SYMBOL: &str = "jitexpr_regexp_extract"; @@ -46,6 +46,21 @@ impl PartialEq for RegexpExtractFnCall { } impl FnCall for RegexpExtractFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Between { min: 2, max: 3 }; + + fn validate_args(args: &[UntypedExpr]) -> Result<(), InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + super::validate_literal(args, 1, VarType::Str, |literal| { + matches!(literal, Literal::String(_)) + })?; + if args.len() == 3 { + super::validate_literal(args, 2, VarType::U64, |literal| { + matches!(literal, Literal::U64(_)) + })?; + } + Ok(()) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -78,11 +93,7 @@ impl FnCall for RegexpExtractFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert!( - (2..=3).contains(&args.len()), - "Expected 2 or 3 args for regexp_extract" - ); - + Self::ARG_COUNT.validate(args)?; let haystack = context.apply_types(&args[0], target_type_set)?; if haystack.return_type == VarType::None { return Ok(TypedExpr::none()); @@ -90,7 +101,11 @@ impl FnCall for RegexpExtractFnCall { assert_eq!(haystack.return_type, VarType::Str); let UntypedExpr::Literal(Literal::String(pattern)) = &args[1] else { - panic!("regexp_extract pattern must be a string literal"); + return Err(InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::Str, + } + .into()); }; let regex = Arc::new( Regex::new(pattern).map_err(|source| CompileError::InvalidRegex { @@ -102,7 +117,13 @@ impl FnCall for RegexpExtractFnCall { let capture_index = match args.get(2) { None => 0, Some(UntypedExpr::Literal(Literal::U64(capture_index))) => *capture_index, - Some(_) => panic!("regexp_extract capture index must be a u64 literal"), + Some(_) => { + return Err(InvalidFunctionCall::ExpectedLiteral { + argument: 3, + expected: VarType::U64, + } + .into()); + } }; Ok(TypedExpr { diff --git a/jitexpr/src/functions/regexp_like.rs b/jitexpr/src/functions/regexp_like.rs index 38c61d567..32503acfc 100644 --- a/jitexpr/src/functions/regexp_like.rs +++ b/jitexpr/src/functions/regexp_like.rs @@ -41,6 +41,15 @@ fn convert_pattern(pattern: &str) -> String { } impl FnCall for RegexpLikeFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + super::validate_literal(args, 1, VarType::Str, |literal| { + matches!(literal, Literal::String(_)) + }) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target: InferredTypeSet, @@ -69,14 +78,18 @@ impl FnCall for RegexpLikeFnCall { _target: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for REGEXP_LIKE"); + Self::ARG_COUNT.validate(args)?; let input_target = match &args[0] { UntypedExpr::Literal(literal) => InferredTypeSet::singleton(literal.r#type()), _ => InferredTypeSet::ALL, }; let input = context.apply_types(&args[0], input_target)?; let UntypedExpr::Literal(Literal::String(pattern)) = &args[1] else { - panic!("REGEXP_LIKE pattern must be a string literal") + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::Str, + } + .into()); }; let converted = convert_pattern(pattern); let regex = diff --git a/jitexpr/src/functions/right.rs b/jitexpr/src/functions/right.rs index 53ef6743d..ab25e460f 100644 --- a/jitexpr/src/functions/right.rs +++ b/jitexpr/src/functions/right.rs @@ -23,22 +23,38 @@ pub(crate) struct RightFnCall { length: usize, } -fn constant_length(expression: &UntypedExpr) -> Option { +fn constant_length(expression: &UntypedExpr) -> Result, super::InvalidFunctionCall> { let UntypedExpr::Literal(literal) = expression else { - panic!("RIGHT length must be constant"); + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::I64, + }); }; - match literal { + if !literal.is_none() && !literal.types().contains(VarType::I64) { + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::I64, + }); + } + Ok(match literal { Literal::I64(value) => usize::try_from(*value).ok(), Literal::U64(value) => usize::try_from(*value).ok(), Literal::F64(value) => usize::try_from(*value as i64).ok(), Literal::None => None, - Literal::Bool(_) | Literal::String(_) => { - unreachable!("type inference constrains RIGHT length to an integer") - } - } + Literal::Bool(_) | Literal::String(_) => None, + }) } impl FnCall for RightFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + super::validate_literal(args, 1, VarType::I64, |literal| { + literal.is_none() || literal.types().contains(VarType::I64) + }) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -68,9 +84,9 @@ impl FnCall for RightFnCall { _target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for RIGHT"); + Self::ARG_COUNT.validate(args)?; let input = context.apply_types(&args[0], InferredTypeSet::STRING)?; - let Some(length) = constant_length(&args[1]) else { + let Some(length) = constant_length(&args[1])? else { return Ok(TypedExpr::none()); }; Ok(TypedExpr { diff --git a/jitexpr/src/functions/round.rs b/jitexpr/src/functions/round.rs index 00bc36cf8..ecf1b2928 100644 --- a/jitexpr/src/functions/round.rs +++ b/jitexpr/src/functions/round.rs @@ -35,14 +35,25 @@ pub(crate) struct RoundFnCall { precision: i64, } -fn constant_precision(expression: Option<&UntypedExpr>) -> Option { +fn constant_precision( + expression: Option<&UntypedExpr>, +) -> Result, super::InvalidFunctionCall> { let Some(expression) = expression else { - return Some(0); + return Ok(Some(0)); }; let UntypedExpr::Literal(literal) = expression else { - panic!("ROUND precision must be constant"); + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::I64, + }); }; - match literal { + if !literal.is_none() && !literal.types().contains(VarType::I64) { + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::I64, + }); + } + Ok(match literal { Literal::I64(value) => Some(*value), Literal::U64(value) => i64::try_from(*value).ok(), Literal::F64(value) @@ -54,10 +65,8 @@ fn constant_precision(expression: Option<&UntypedExpr>) -> Option { Some(*value as i64) } Literal::None => None, - Literal::F64(_) | Literal::Bool(_) | Literal::String(_) => { - unreachable!("type inference constrains ROUND precision to an integer") - } - } + Literal::F64(_) | Literal::Bool(_) | Literal::String(_) => None, + }) } fn return_type_for_precision(precision: i64) -> VarType { @@ -69,6 +78,18 @@ fn return_type_for_precision(precision: i64) -> VarType { } impl FnCall for RoundFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Between { min: 1, max: 2 }; + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + if args.len() == 2 { + super::validate_literal(args, 1, VarType::I64, |literal| { + literal.is_none() || literal.types().contains(VarType::I64) + })?; + } + Ok(()) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -85,7 +106,7 @@ impl FnCall for RoundFnCall { if let Some(precision) = args.get(1) { crate::ast::infer_types_aux(precision, InferredTypeSet::I64, inferred_types)?; } - let precision = constant_precision(args.get(1)).unwrap_or(0); + let precision = constant_precision(args.get(1))?.unwrap_or(0); let return_types = InferredTypeSet::singleton(return_type_for_precision(precision)); if target_type.intersect(return_types).is_none() { return Err(TypeError::WrongFunctionReturnType { @@ -102,11 +123,8 @@ impl FnCall for RoundFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert!( - (1..=2).contains(&args.len()), - "expected 1 or 2 args for ROUND" - ); - let Some(precision) = constant_precision(args.get(1)) else { + Self::ARG_COUNT.validate(args)?; + let Some(precision) = constant_precision(args.get(1))? else { return Ok(TypedExpr::none()); }; let return_type = return_type_for_precision(precision); diff --git a/jitexpr/src/functions/split_after.rs b/jitexpr/src/functions/split_after.rs index 4a393fe09..74adffc1a 100644 --- a/jitexpr/src/functions/split_after.rs +++ b/jitexpr/src/functions/split_after.rs @@ -31,25 +31,49 @@ pub(crate) struct SplitAfterFnCall { occurrence: usize, } -fn constant_occurrence(expression: Option<&UntypedExpr>) -> Option { +fn constant_occurrence( + expression: Option<&UntypedExpr>, +) -> Result, super::InvalidFunctionCall> { let Some(expression) = expression else { - return Some(0); + return Ok(Some(0)); }; let UntypedExpr::Literal(literal) = expression else { - panic!("SPLIT_AFTER occurrence must be constant"); + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 3, + expected: VarType::I64, + }); }; - match literal { + if !literal.is_none() && !literal.types().contains(VarType::I64) { + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 3, + expected: VarType::I64, + }); + } + Ok(match literal { Literal::I64(value) => usize::try_from(*value).ok(), Literal::U64(value) => usize::try_from(*value).ok(), Literal::F64(value) => usize::try_from(*value as i64).ok(), Literal::None => None, - Literal::Bool(_) | Literal::String(_) => { - unreachable!("type inference constrains SPLIT_AFTER occurrence to an integer") - } - } + Literal::Bool(_) | Literal::String(_) => None, + }) } impl FnCall for SplitAfterFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Between { min: 2, max: 3 }; + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + super::validate_literal(args, 1, VarType::Str, |literal| { + matches!(literal, Literal::String(_) | Literal::None) + })?; + if args.len() == 3 { + super::validate_literal(args, 2, VarType::I64, |literal| { + literal.is_none() || literal.types().contains(VarType::I64) + })?; + } + Ok(()) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -82,17 +106,20 @@ impl FnCall for SplitAfterFnCall { _target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert!( - (2..=3).contains(&args.len()), - "expected 2 or 3 args for SPLIT_AFTER" - ); + Self::ARG_COUNT.validate(args)?; let input = context.apply_types(&args[0], InferredTypeSet::STRING)?; let separator = match &args[1] { UntypedExpr::Literal(Literal::String(separator)) => Arc::clone(separator), UntypedExpr::Literal(Literal::None) => return Ok(TypedExpr::none()), - _ => panic!("SPLIT_AFTER separator must be a string constant"), + _ => { + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::Str, + } + .into()); + } }; - let Some(occurrence) = constant_occurrence(args.get(2)) else { + let Some(occurrence) = constant_occurrence(args.get(2))? else { return Ok(TypedExpr::none()); }; Ok(TypedExpr { diff --git a/jitexpr/src/functions/split_before.rs b/jitexpr/src/functions/split_before.rs index 6b913bba7..42947b82d 100644 --- a/jitexpr/src/functions/split_before.rs +++ b/jitexpr/src/functions/split_before.rs @@ -30,25 +30,49 @@ pub(crate) struct SplitBeforeFnCall { occurrence: usize, } -fn constant_occurrence(expression: Option<&UntypedExpr>) -> Option { +fn constant_occurrence( + expression: Option<&UntypedExpr>, +) -> Result, super::InvalidFunctionCall> { let Some(expression) = expression else { - return Some(0); + return Ok(Some(0)); }; let UntypedExpr::Literal(literal) = expression else { - panic!("SPLIT_BEFORE occurrence must be constant"); + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 3, + expected: VarType::I64, + }); }; - match literal { + if !literal.is_none() && !literal.types().contains(VarType::I64) { + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 3, + expected: VarType::I64, + }); + } + Ok(match literal { Literal::I64(value) => usize::try_from(*value).ok(), Literal::U64(value) => usize::try_from(*value).ok(), Literal::F64(value) => usize::try_from(*value as i64).ok(), Literal::None => None, - Literal::Bool(_) | Literal::String(_) => { - unreachable!("type inference constrains SPLIT_BEFORE occurrence to an integer") - } - } + Literal::Bool(_) | Literal::String(_) => None, + }) } impl FnCall for SplitBeforeFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Between { min: 2, max: 3 }; + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + super::validate_literal(args, 1, VarType::Str, |literal| { + matches!(literal, Literal::String(_) | Literal::None) + })?; + if args.len() == 3 { + super::validate_literal(args, 2, VarType::I64, |literal| { + literal.is_none() || literal.types().contains(VarType::I64) + })?; + } + Ok(()) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -81,17 +105,20 @@ impl FnCall for SplitBeforeFnCall { _target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert!( - (2..=3).contains(&args.len()), - "expected 2 or 3 args for SPLIT_BEFORE" - ); + Self::ARG_COUNT.validate(args)?; let input = context.apply_types(&args[0], InferredTypeSet::STRING)?; let separator = match &args[1] { UntypedExpr::Literal(Literal::String(separator)) => Arc::clone(separator), UntypedExpr::Literal(Literal::None) => return Ok(TypedExpr::none()), - _ => panic!("SPLIT_BEFORE separator must be a string constant"), + _ => { + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::Str, + } + .into()); + } }; - let Some(occurrence) = constant_occurrence(args.get(2)) else { + let Some(occurrence) = constant_occurrence(args.get(2))? else { return Ok(TypedExpr::none()); }; Ok(TypedExpr { diff --git a/jitexpr/src/functions/sqrt.rs b/jitexpr/src/functions/sqrt.rs index 8663ff0e3..a10d3dda9 100644 --- a/jitexpr/src/functions/sqrt.rs +++ b/jitexpr/src/functions/sqrt.rs @@ -21,6 +21,8 @@ pub(crate) struct SqrtFnCall { } impl FnCall for SqrtFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -49,7 +51,7 @@ impl FnCall for SqrtFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 1, "expected 1 arg for SQRT"); + Self::ARG_COUNT.validate(args)?; debug_assert!(target_type_set.contains(VarType::F64)); let arg = context.apply_types(&args[0], InferredTypeSet::F64)?; if arg.return_type == VarType::None { diff --git a/jitexpr/src/functions/substring.rs b/jitexpr/src/functions/substring.rs index 0f8c8b328..203d010ce 100644 --- a/jitexpr/src/functions/substring.rs +++ b/jitexpr/src/functions/substring.rs @@ -31,11 +31,23 @@ pub(crate) struct SubstringFnCall { length: usize, } -fn constant_usize(expression: &UntypedExpr) -> Option { +fn constant_usize( + expression: &UntypedExpr, + argument: usize, +) -> Result, super::InvalidFunctionCall> { let UntypedExpr::Literal(literal) = expression else { - panic!("SUBSTRING bounds must be constants"); + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument, + expected: VarType::I64, + }); }; - match literal { + if !literal.is_none() && !literal.types().contains(VarType::I64) { + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument, + expected: VarType::I64, + }); + } + Ok(match literal { Literal::I64(value) => usize::try_from(*value).ok(), Literal::U64(value) => i64::try_from(*value) .ok() @@ -44,13 +56,23 @@ fn constant_usize(expression: &UntypedExpr) -> Option { usize::try_from(*value as i64).ok() } Literal::None => None, - Literal::Bool(_) | Literal::F64(_) | Literal::String(_) => { - panic!("SUBSTRING bounds must be integer constants") - } - } + Literal::Bool(_) | Literal::F64(_) | Literal::String(_) => None, + }) } impl FnCall for SubstringFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(3); + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + for index in 1..=2 { + super::validate_literal(args, index, VarType::I64, |literal| { + literal.is_none() || literal.types().contains(VarType::I64) + })?; + } + Ok(()) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -81,10 +103,10 @@ impl FnCall for SubstringFnCall { _target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 3, "expected 3 args for SUBSTRING"); + Self::ARG_COUNT.validate(args)?; let input = context.apply_types(&args[0], InferredTypeSet::STRING)?; - let start = constant_usize(&args[1]); - let length = constant_usize(&args[2]); + let start = constant_usize(&args[1], 2)?; + let length = constant_usize(&args[2], 3)?; let (Some(start), Some(length)) = (start, length) else { return Ok(TypedExpr::none()); }; diff --git a/jitexpr/src/functions/substring_count.rs b/jitexpr/src/functions/substring_count.rs index fddf617e1..d74777687 100644 --- a/jitexpr/src/functions/substring_count.rs +++ b/jitexpr/src/functions/substring_count.rs @@ -26,6 +26,8 @@ pub(crate) struct SubstringCountFnCall { } impl FnCall for SubstringCountFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target: InferredTypeSet, @@ -55,6 +57,7 @@ impl FnCall for SubstringCountFnCall { _target: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { + Self::ARG_COUNT.validate(args)?; let args = args .iter() .map(|arg| context.apply_types(arg, InferredTypeSet::STRING)) diff --git a/jitexpr/src/functions/subtract.rs b/jitexpr/src/functions/subtract.rs index e0062cfd0..36d284a7a 100644 --- a/jitexpr/src/functions/subtract.rs +++ b/jitexpr/src/functions/subtract.rs @@ -26,6 +26,8 @@ pub(crate) struct SubtractFnCall { } impl FnCall for SubtractFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(2); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -69,8 +71,7 @@ impl FnCall for SubtractFnCall { target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 2, "expected 2 args for SUBTRACT"); - + Self::ARG_COUNT.validate(args)?; let mut return_types = InferredTypeSet::NUMERICAL.intersect(target_type_set); for arg in args { let arg_types = crate::ast::infer_type_with_variable_types( diff --git a/jitexpr/src/functions/text_join.rs b/jitexpr/src/functions/text_join.rs index c1e88af23..fbf99f473 100644 --- a/jitexpr/src/functions/text_join.rs +++ b/jitexpr/src/functions/text_join.rs @@ -26,6 +26,13 @@ pub(crate) struct TextJoinFnCall { } impl FnCall for TextJoinFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::AtLeast(4); + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + concat::validate_join_args(args) + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -39,6 +46,7 @@ impl FnCall for TextJoinFnCall { _target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { + Self::ARG_COUNT.validate(args)?; let Some(arguments) = concat::apply_join_types("TEXT_JOIN", args, context)? else { return Ok(TypedExpr::none()); }; diff --git a/jitexpr/src/functions/trim.rs b/jitexpr/src/functions/trim.rs index ecbbf4f20..9558b3ef6 100644 --- a/jitexpr/src/functions/trim.rs +++ b/jitexpr/src/functions/trim.rs @@ -48,6 +48,32 @@ pub(crate) struct TrimFnCall { } impl FnCall for TrimFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(3); + + fn validate_args(args: &[UntypedExpr]) -> Result<(), super::InvalidFunctionCall> { + Self::ARG_COUNT.validate(args)?; + super::validate_literal(args, 1, VarType::Str, |literal| { + matches!(literal, Literal::String(_)) + })?; + let UntypedExpr::Literal(Literal::String(mode)) = &args[2] else { + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 3, + expected: VarType::Str, + }); + }; + if ["leading", "trailing", "both"] + .iter() + .any(|valid_mode| mode.eq_ignore_ascii_case(valid_mode)) + { + Ok(()) + } else { + Err(super::InvalidFunctionCall::InvalidLiteralValue { + argument: 3, + expected: "leading, trailing, or both", + }) + } + } + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -78,16 +104,24 @@ impl FnCall for TrimFnCall { _target: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 3, "expected 3 args for TRIM"); + Self::ARG_COUNT.validate(args)?; let input = context.apply_types(&args[0], InferredTypeSet::STRING)?; if input.return_type == VarType::None { return Ok(TypedExpr::none()); } let UntypedExpr::Literal(Literal::String(delimiter)) = &args[1] else { - panic!("TRIM delimiter must be a string literal") + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 2, + expected: VarType::Str, + } + .into()); }; let UntypedExpr::Literal(Literal::String(mode)) = &args[2] else { - panic!("TRIM mode must be a string literal") + return Err(super::InvalidFunctionCall::ExpectedLiteral { + argument: 3, + expected: VarType::Str, + } + .into()); }; let mode = if mode.eq_ignore_ascii_case("leading") { TrimMode::Leading @@ -96,7 +130,11 @@ impl FnCall for TrimFnCall { } else if mode.eq_ignore_ascii_case("both") { TrimMode::Both } else { - panic!("TRIM mode must be leading, trailing, or both") + return Err(super::InvalidFunctionCall::InvalidLiteralValue { + argument: 3, + expected: "leading, trailing, or both", + } + .into()); }; Ok(TypedExpr { return_type: VarType::Str, diff --git a/jitexpr/src/functions/upper.rs b/jitexpr/src/functions/upper.rs index 944cc1824..0270273df 100644 --- a/jitexpr/src/functions/upper.rs +++ b/jitexpr/src/functions/upper.rs @@ -32,6 +32,8 @@ pub(crate) struct UpperFnCall { } impl FnCall for UpperFnCall { + const ARG_COUNT: super::ArgumentCount = super::ArgumentCount::Exactly(1); + fn infer_types<'a>( args: &'a [UntypedExpr], target_type: InferredTypeSet, @@ -60,7 +62,7 @@ impl FnCall for UpperFnCall { _target_type_set: InferredTypeSet, context: &mut CompileFnBuilder<'_, '_>, ) -> Result { - assert_eq!(args.len(), 1, "expected 1 arg for UPPER"); + Self::ARG_COUNT.validate(args)?; let arg = context.apply_types(&args[0], InferredTypeSet::STRING)?; if arg.return_type == VarType::None { return Ok(TypedExpr::none());