mirror of
https://github.com/quickwit-oss/tantivy.git
synced 2026-10-06 20:02:45 +00:00
Added a way to build a function call untyped expression in a validated way
This commit is contained in:
@@ -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,
|
||||
|
||||
@@ -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};
|
||||
|
||||
@@ -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<UntypedExpr>,
|
||||
) -> Result<UntypedExpr, InvalidFunctionCall> {
|
||||
function.call(args)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<Literal> for UntypedExpr {
|
||||
|
||||
@@ -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<cranelift_module::ModuleError> for CompileError {
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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());
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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(
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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 {
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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<Option<JoinArguments>, 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..]
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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 {
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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,
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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,
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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())?;
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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()
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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)?;
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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)?;
|
||||
|
||||
@@ -22,22 +22,38 @@ pub(crate) struct LeftFnCall {
|
||||
length: usize,
|
||||
}
|
||||
|
||||
fn constant_length(expression: &UntypedExpr) -> Option<usize> {
|
||||
fn constant_length(expression: &UntypedExpr) -> Result<Option<usize>, 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<TypedExpr, CompileError> {
|
||||
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 {
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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());
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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,
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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,
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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(
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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(
|
||||
|
||||
@@ -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<UntypedExpr>) -> Result<UntypedExpr, InvalidFunctionCall> {
|
||||
match self {
|
||||
Function::Abs => <AbsFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::And => <AndFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Ceil => <CeilFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Concat => <ConcatFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Add => <AddFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Divide => <DivideFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Eq => <EqFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Floor => <FloorFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Gt => <GtFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::GtEq => <GtEqFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::If => <IfFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::IntMod => <IntModFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Left => <LeftFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Lt => <LtFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::LtEq => <LtEqFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::IsNotNull => <IsNotNullFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::IsNull => <IsNullFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Lower => <LowerFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Max => <MaxFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Min => <MinFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Multiply => <MultiplyFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Neq => <NeqFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Not => <NotFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Or => <OrFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Pow => <PowFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Sqrt => <SqrtFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::RegexpExtract => <RegexpExtractFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::RegexpLike => <RegexpLikeFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Right => <RightFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Round => <RoundFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::SplitAfter => <SplitAfterFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::SplitBefore => <SplitBeforeFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Subtract => <SubtractFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Substring => <SubstringFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::SubstringCount => <SubstringCountFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::TextJoin => <TextJoinFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Trim => <TrimFnCall as FnCall>::validate_args(&args)?,
|
||||
Function::Upper => <UpperFnCall as FnCall>::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<FnCallEnum> {
|
||||
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<FnCallEnum> {
|
||||
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<FnCallEnum> {
|
||||
builder: &mut FunctionBuilder<'_>,
|
||||
) -> Result<LoweredValue, CompileError>;
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::compile::compile;
|
||||
|
||||
fn call_error(function: Function, args: Vec<UntypedExpr>) -> 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,
|
||||
})
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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(
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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()
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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)?;
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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()
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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 {
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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 =
|
||||
|
||||
@@ -23,22 +23,38 @@ pub(crate) struct RightFnCall {
|
||||
length: usize,
|
||||
}
|
||||
|
||||
fn constant_length(expression: &UntypedExpr) -> Option<usize> {
|
||||
fn constant_length(expression: &UntypedExpr) -> Result<Option<usize>, 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<TypedExpr, CompileError> {
|
||||
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 {
|
||||
|
||||
@@ -35,14 +35,25 @@ pub(crate) struct RoundFnCall {
|
||||
precision: i64,
|
||||
}
|
||||
|
||||
fn constant_precision(expression: Option<&UntypedExpr>) -> Option<i64> {
|
||||
fn constant_precision(
|
||||
expression: Option<&UntypedExpr>,
|
||||
) -> Result<Option<i64>, 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<i64> {
|
||||
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<TypedExpr, CompileError> {
|
||||
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);
|
||||
|
||||
@@ -31,25 +31,49 @@ pub(crate) struct SplitAfterFnCall {
|
||||
occurrence: usize,
|
||||
}
|
||||
|
||||
fn constant_occurrence(expression: Option<&UntypedExpr>) -> Option<usize> {
|
||||
fn constant_occurrence(
|
||||
expression: Option<&UntypedExpr>,
|
||||
) -> Result<Option<usize>, 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<TypedExpr, CompileError> {
|
||||
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 {
|
||||
|
||||
@@ -30,25 +30,49 @@ pub(crate) struct SplitBeforeFnCall {
|
||||
occurrence: usize,
|
||||
}
|
||||
|
||||
fn constant_occurrence(expression: Option<&UntypedExpr>) -> Option<usize> {
|
||||
fn constant_occurrence(
|
||||
expression: Option<&UntypedExpr>,
|
||||
) -> Result<Option<usize>, 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<TypedExpr, CompileError> {
|
||||
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 {
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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 {
|
||||
|
||||
@@ -31,11 +31,23 @@ pub(crate) struct SubstringFnCall {
|
||||
length: usize,
|
||||
}
|
||||
|
||||
fn constant_usize(expression: &UntypedExpr) -> Option<usize> {
|
||||
fn constant_usize(
|
||||
expression: &UntypedExpr,
|
||||
argument: usize,
|
||||
) -> Result<Option<usize>, 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> {
|
||||
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<TypedExpr, CompileError> {
|
||||
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());
|
||||
};
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
Self::ARG_COUNT.validate(args)?;
|
||||
let args = args
|
||||
.iter()
|
||||
.map(|arg| context.apply_types(arg, InferredTypeSet::STRING))
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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(
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
Self::ARG_COUNT.validate(args)?;
|
||||
let Some(arguments) = concat::apply_join_types("TEXT_JOIN", args, context)? else {
|
||||
return Ok(TypedExpr::none());
|
||||
};
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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,
|
||||
|
||||
@@ -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<TypedExpr, CompileError> {
|
||||
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());
|
||||
|
||||
Reference in New Issue
Block a user