use std::any::TypeId;
use std::sync::Arc;
use anyhow::{anyhow, bail};
use crate::func::def::{
FunctionDef, MatchTestFailure, ParameterBinding, ParameterBindingProvider, Parameters,
};
use crate::operator::Operator;
use crate::tree::ast::expression::{Cast, Expression, ExpressionKind};
use crate::tree::typed_ast::context::ExpressionTranslationContext;
use crate::tree::typed_ast::expression::{CastKind, TypedCast, TypedExpression};
use crate::types::array::Array;
use crate::types::matcher::{
numeric_or_interval_matcher, AnyMatcher, ArrayMatcher, ExactMatcher, MapMatcher,
NumericMatcher, OrMatcher,
};
use crate::types::tuple::Tuple;
use crate::types::{Type, BOOLEAN, DOUBLE, INT, INTERVAL, STRING, TIMESTAMP};
#[derive(Default, Clone)]
pub struct ArrayConcat;
impl FunctionDef for ArrayConcat {
fn name(&self) -> &'static str {
Operator::Plus.str()
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("left", ArrayMatcher::default())
.with("right", ArrayMatcher::default())
}
fn refine_binding(
&self,
binding: ParameterBinding<Arc<TypedExpression>>,
) -> anyhow::Result<ParameterBinding<Arc<TypedExpression>>> {
let Ok(typed_left) = binding.get_by_name("left").cloned() else {
return Ok(binding);
};
let Ok(typed_right) = binding.get_by_name("right").cloned() else {
return Ok(binding);
};
let Type::Array(_) = typed_left.resolved_type.as_ref() else {
return Ok(binding);
};
let Type::Array(_) = typed_right.resolved_type.as_ref() else {
return Ok(binding);
};
let Ok(merged_type) = typed_left
.resolved_type
.as_ref()
.clone()
.merge(typed_right.resolved_type.as_ref().clone())
else {
return Ok(binding);
};
let final_left = wrap_with_cast_if_needed(typed_left, &merged_type);
let final_right = wrap_with_cast_if_needed(typed_right, &merged_type);
Ok(ParameterBinding::from_named([
("left".to_string(), final_left),
("right".to_string(), final_right),
]))
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
let left = bindings.get_by_index(0)?.typ().clone();
let right = bindings.get_by_index(1)?.typ().clone();
left.merge(right)
.map_err(|e| anyhow!(MatchTestFailure::wrap(e)))
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
fn wrap_with_cast_if_needed(
expr: Arc<TypedExpression>,
target_type: &Type,
) -> Arc<TypedExpression> {
if expr.resolved_type.as_ref() == target_type {
return expr;
}
let Some(cast_kind) = CastKind::from_ast_with_context(&expr, target_type) else {
return expr;
};
if matches!(cast_kind, CastKind::Identity) {
return expr;
}
let cast_ast = Arc::new(Expression {
span: expr.ast.span.clone(),
kind: ExpressionKind::Cast(Cast {
expression: expr.ast.clone(),
target_type: Arc::new(target_type.clone()),
}),
});
Arc::new(TypedExpression {
ast: cast_ast,
resolved_type: Arc::new(target_type.clone()),
kind: TypedCast {
value: expr,
target_type: target_type.clone(),
cast_kind,
}
.into(),
})
}
#[derive(Default, Clone)]
pub struct FilterNull;
impl FunctionDef for FilterNull {
fn name(&self) -> &'static str {
"filter_null"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", ArrayMatcher::default())
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(bindings.get_by_index(0)?.typ().clone())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArrayOrMapLen;
impl FunctionDef for ArrayOrMapLen {
fn name(&self) -> &'static str {
"len"
}
fn parameters(&self) -> Parameters {
Parameters::new().with(
"x",
OrMatcher::default()
.with(ArrayMatcher::default())
.with(MapMatcher::default()),
)
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(INT)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArrayDistinct;
impl FunctionDef for ArrayDistinct {
fn name(&self) -> &'static str {
"array_distinct"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", ArrayMatcher::default())
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(bindings.get_by_index(0)?.typ().clone())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct Slice;
impl FunctionDef for Slice {
fn name(&self) -> &'static str {
"slice"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("array", ArrayMatcher::default())
.with("start", ExactMatcher::of(INT))
.with("end", ExactMatcher::of(INT))
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(bindings.get_by_index(0)?.typ().clone())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct Split;
impl FunctionDef for Split {
fn name(&self) -> &'static str {
"split"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("string", ExactMatcher::of(STRING))
.with("delimiter", ExactMatcher::of(STRING))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(Array::new(STRING).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArrayJoin2;
impl FunctionDef for ArrayJoin2 {
fn name(&self) -> &'static str {
"array_join"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("array", ArrayMatcher::of(ExactMatcher::of(STRING)))
.with("delimiter", ExactMatcher::of(STRING))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(STRING)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArrayJoin3;
impl FunctionDef for ArrayJoin3 {
fn name(&self) -> &'static str {
"array_join"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("array", ArrayMatcher::of(ExactMatcher::of(STRING)))
.with("delimiter", ExactMatcher::of(STRING))
.with("null_replacement", ExactMatcher::of(STRING))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(STRING)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct Flatten;
impl FunctionDef for Flatten {
fn name(&self) -> &'static str {
"flatten"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", ArrayMatcher::of(ArrayMatcher::default()))
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
let array_of_arrays = match bindings.get_by_index(0)?.typ() {
Type::Array(outer_array) => match outer_array.element_type.as_ref() {
Type::Array(inner_array) => inner_array.element_type.as_ref().clone(),
_ => bail!("flatten requires an array of arrays"),
},
_ => bail!("flatten requires an array of arrays"),
};
Ok(Array::new(array_of_arrays).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArrayAny;
impl FunctionDef for ArrayAny {
fn name(&self) -> &'static str {
"any"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", ArrayMatcher::of(ExactMatcher::of(BOOLEAN)))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(BOOLEAN)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArrayAll;
impl FunctionDef for ArrayAll {
fn name(&self) -> &'static str {
"all"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", ArrayMatcher::of(ExactMatcher::of(BOOLEAN)))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(BOOLEAN)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArrayMax;
impl FunctionDef for ArrayMax {
fn name(&self) -> &'static str {
"max"
}
fn parameters(&self) -> Parameters {
Parameters::new().with(
"x",
ArrayMatcher::of(
OrMatcher::default()
.with(NumericMatcher::default())
.with(ExactMatcher::of(STRING))
.with(ExactMatcher::of(TIMESTAMP))
.with(ExactMatcher::of(INTERVAL)),
),
)
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
match bindings.get_by_index(0)?.typ() {
Type::Array(array) => Ok((*array.element_type).clone()),
_ => bail!("parameter should be an array, checked by matcher"),
}
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArrayMin;
impl FunctionDef for ArrayMin {
fn name(&self) -> &'static str {
"min"
}
fn parameters(&self) -> Parameters {
Parameters::new().with(
"x",
ArrayMatcher::of(
OrMatcher::default()
.with(NumericMatcher::default())
.with(ExactMatcher::of(STRING))
.with(ExactMatcher::of(TIMESTAMP))
.with(ExactMatcher::of(INTERVAL)),
),
)
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
match bindings.get_by_index(0)?.typ() {
Type::Array(array) => Ok((*array.element_type).clone()),
_ => bail!("parameter should be an array, checked by matcher"),
}
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArraySum;
impl FunctionDef for ArraySum {
fn name(&self) -> &'static str {
"sum"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", ArrayMatcher::of(numeric_or_interval_matcher()))
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
match bindings.get_by_index(0)?.typ() {
Type::Array(array) => Ok((*array.element_type).clone()),
_ => bail!("parameter should be an array, checked by matcher"),
}
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct ArrayAvg;
impl FunctionDef for ArrayAvg {
fn name(&self) -> &'static str {
"avg"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", ArrayMatcher::of(numeric_or_interval_matcher()))
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
match bindings.get_by_index(0)?.typ() {
Type::Array(array) => match *array.element_type {
Type::Int => Ok(DOUBLE),
ref other => Ok(other.clone()),
},
_ => bail!("parameter should be an array, checked by matcher"),
}
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct Sequence2;
impl FunctionDef for Sequence2 {
fn name(&self) -> &'static str {
"sequence"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("start", ExactMatcher::of(INT))
.with("stop", ExactMatcher::of(INT))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(Array::new(INT).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct Sequence3;
impl FunctionDef for Sequence3 {
fn name(&self) -> &'static str {
"sequence"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("start", ExactMatcher::of(INT))
.with("stop", ExactMatcher::of(INT))
.with("step", ExactMatcher::of(INT))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(Array::new(INT).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct Zip;
impl FunctionDef for Zip {
fn name(&self) -> &'static str {
"zip"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("left", ArrayMatcher::default())
.with("right", ArrayMatcher::default())
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
let left_elem = match bindings.get_by_index(0)?.typ() {
Type::Array(a) => (*a.element_type).clone(),
_ => bail!("zip requires array arguments"),
};
let right_elem = match bindings.get_by_index(1)?.typ() {
Type::Array(a) => (*a.element_type).clone(),
_ => bail!("zip requires array arguments"),
};
Ok(Array::new(Tuple::new(vec![left_elem, right_elem]).into()).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}
#[derive(Default, Clone)]
pub struct Filter;
impl FunctionDef for Filter {
fn name(&self) -> &'static str {
"filter"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("array", AnyMatcher)
.with("lambda", AnyMatcher)
}
fn custom_bind(
&self,
ast_binding: &ParameterBinding<Arc<Expression>>,
ctx: &mut ExpressionTranslationContext,
) -> anyhow::Result<Option<ParameterBinding<Arc<TypedExpression>>>> {
let array_ast = ast_binding.get_by_name("array")?.clone();
let typed_array = Arc::new(TypedExpression::from_ast_with_context(array_ast, ctx));
let element_type = match typed_array.resolved_type.as_ref() {
Type::Array(arr) => arr.element_type.clone(),
other => bail!("Expected array type, got {}", other),
};
let lambda_ast = ast_binding.get_by_name("lambda")?.clone();
let ExpressionKind::Lambda(lambda) = &lambda_ast.kind else {
bail!("Expected lambda expression");
};
let lambda_with_hints = lambda.with_param_types(&[element_type])?;
let lambda_expr = Expression {
kind: lambda_with_hints.into(),
span: lambda_ast.span.clone(),
};
let typed_lambda = Arc::new(TypedExpression::from_ast_with_context(
Arc::new(lambda_expr),
ctx,
));
match typed_lambda.resolved_type.as_ref() {
Type::Function(f) => {
if *f.return_type != BOOLEAN {
bail!("filter lambda must return boolean, got {}", f.return_type);
}
}
Type::Unknown => {}
other => bail!("Expected function type for lambda, got {}", other),
}
Ok(Some(ParameterBinding::from_named([
("array".to_string(), typed_array),
("lambda".to_string(), typed_lambda),
])))
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(bindings.get_by_name("array")?.typ().clone())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
}