use std::collections::HashMap;
use std::error::Error;
use std::sync::Arc;
use ordermap::OrderMap;
use crate::completion::{CompletionItem, CompletionItemKind};
use crate::func::def::{
BroadcastResolution, DirectResolution, FunctionDef, FunctionParameterBindingFailure,
FunctionParameterBindingFailures, FunctionResolution, MatchTestFailure, ParameterBinding,
Parameters,
};
use crate::func::defs::*;
use crate::func::err::{ApplyAttempt, ApplyFailure};
use crate::tree::ast::expression::{Expression, ExpressionKind};
use crate::tree::typed_ast::context::ExpressionTranslationContext;
use crate::tree::typed_ast::expression::TypedExpression;
use crate::types::Type;
#[derive(Clone)]
pub struct FunctionRegistry {
pub unary_prefix_operation_defs: HashMap<String, Vec<Arc<dyn FunctionDef>>>,
pub unary_postfix_operation_defs: HashMap<String, Vec<Arc<dyn FunctionDef>>>,
pub binary_operation_defs: HashMap<String, Vec<Arc<dyn FunctionDef>>>,
pub function_defs: HashMap<String, Vec<Arc<dyn FunctionDef>>>,
}
impl std::fmt::Debug for FunctionRegistry {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FunctionRegistry")
.field(
"unary_prefix_count",
&self.unary_prefix_operation_defs.len(),
)
.field(
"unary_postfix_count",
&self.unary_postfix_operation_defs.len(),
)
.field("binary_count", &self.binary_operation_defs.len())
.field("function_count", &self.function_defs.len())
.finish()
}
}
impl FunctionRegistry {
pub fn empty() -> Self {
Self {
unary_prefix_operation_defs: HashMap::new(),
unary_postfix_operation_defs: HashMap::new(),
binary_operation_defs: HashMap::new(),
function_defs: HashMap::new(),
}
}
pub fn register_unary_prefix(&mut self, op: impl FunctionDef) {
self.unary_prefix_operation_defs
.entry(op.name().to_lowercase())
.or_default()
.push(Arc::new(op));
}
pub fn register_unary_postfix(&mut self, op: impl FunctionDef) {
self.unary_postfix_operation_defs
.entry(op.name().to_lowercase())
.or_default()
.push(Arc::new(op));
}
pub fn register_binary(&mut self, op: impl FunctionDef) {
self.binary_operation_defs
.entry(op.name().to_lowercase())
.or_default()
.push(Arc::new(op));
}
pub fn register_function(&mut self, func: impl FunctionDef) {
self.function_defs
.entry(func.name().to_lowercase())
.or_default()
.push(Arc::new(func));
}
pub fn autocomplete_suggestions(&self) -> Vec<CompletionItem> {
self.function_defs
.iter()
.flat_map(|(name, defs)| {
defs.iter().map(|def| {
let params = def.parameters();
CompletionItem::new(name.clone())
.with_detail(format!("({})", params))
.with_snippet(format!(
"{}({}) ${{}}",
name.clone(),
params.autocomplete_snippet()
))
.with_kind(CompletionItemKind::Function)
.with_section("Functions".to_string())
})
})
.collect()
}
pub fn bind(
&self,
function_name: &str,
positional: Vec<Arc<Expression>>,
named: OrderMap<String, Arc<Expression>>,
function_type: InvocationType,
ctx: &mut ExpressionTranslationContext,
) -> Result<FunctionResolution<Arc<TypedExpression>>, FunctionBindingError> {
let function_map = match &function_type {
InvocationType::UnaryPrefix => &self.unary_prefix_operation_defs,
InvocationType::UnaryPostfix => &self.unary_postfix_operation_defs,
InvocationType::Binary => &self.binary_operation_defs,
InvocationType::Function => &self.function_defs,
};
if matches!(function_type, InvocationType::Function)
&& function_name.chars().any(|c| c.is_uppercase())
{
let lowered = function_name.to_lowercase();
let suggestion = function_map.contains_key(&lowered).then_some(lowered);
return Err(FunctionBindingError::FunctionNameNotLowercase {
function_name: function_name.into(),
suggestion,
});
}
let lookup_key = match &function_type {
InvocationType::Function => function_name.to_string(),
_ => function_name.to_lowercase(),
};
let candidates = function_map.get(&lookup_key).ok_or_else(|| {
FunctionBindingError::FunctionNotFound {
function_name: function_name.into(),
}
})?;
let has_lambdas = positional
.iter()
.any(|e| matches!(e.kind, ExpressionKind::Lambda(_)))
|| named
.values()
.any(|e| matches!(e.kind, ExpressionKind::Lambda(_)));
let mut binding_attempts = Vec::new();
let mut broadcast_match: Option<FunctionResolution<Arc<TypedExpression>>> = None;
let mut standard_candidates: Vec<(
Arc<dyn FunctionDef>,
ParameterBinding<Arc<Expression>>,
)> = Vec::new();
for function_def in candidates {
if let Some(special_position) = function_def.special_position() {
if !ctx.fctx.specials_allowed.contains(&special_position) {
binding_attempts.push(ApplyAttempt {
function_def: format!("{}", function_def),
binding_failures: FunctionParameterBindingFailures(vec![
FunctionParameterBindingFailure::DoesNotMatch(
format!(
"{} function not allowed in this context",
special_position
)
.into(),
),
]),
});
continue;
}
}
let params = function_def.parameters();
let ast_binding: ParameterBinding<Arc<Expression>> =
match params.bind(positional.clone(), named.clone()) {
Ok(b) => b,
Err(failures) => {
binding_attempts.push(ApplyAttempt {
function_def: format!("{}", function_def),
binding_failures: FunctionParameterBindingFailures(failures),
});
continue;
}
};
match function_def.custom_bind(&ast_binding, ctx) {
Ok(None) => {
standard_candidates.push((function_def.clone(), ast_binding));
}
Ok(Some(binding)) => {
match function_def.return_type(&binding) {
Ok(typ) => {
return Ok(DirectResolution {
function_def: function_def.clone(),
binding,
typ,
}
.into());
}
Err(e) => {
binding_attempts.push(ApplyAttempt {
function_def: format!("{}", function_def),
binding_failures: FunctionParameterBindingFailures(vec![
FunctionParameterBindingFailure::DoesNotMatch(
format!("{}", e).into(),
),
]),
});
}
}
}
Err(e) => {
binding_attempts.push(ApplyAttempt {
function_def: format!("{}", function_def),
binding_failures: FunctionParameterBindingFailures(vec![
FunctionParameterBindingFailure::DoesNotMatch(format!("{}", e).into()),
]),
});
}
}
}
if has_lambdas {
return Err(FunctionBindingError::NoMatchingSignature {
function_name: function_name.to_string(),
positional,
named,
apply_failure: ApplyFailure(binding_attempts),
});
}
let typed_positional: Vec<Arc<TypedExpression>> = positional
.iter()
.map(|ast| Arc::new(TypedExpression::from_ast_with_context(ast.clone(), ctx)))
.collect();
let typed_named: OrderMap<String, Arc<TypedExpression>> = named
.iter()
.map(|(k, ast)| {
(
k.clone(),
Arc::new(TypedExpression::from_ast_with_context(ast.clone(), ctx)),
)
})
.collect();
for (function_def, _ast_binding) in standard_candidates {
let params = function_def.parameters();
let result = self.try_standard_bind_typed(
&function_def,
¶ms,
typed_positional.clone(),
typed_named.clone(),
);
match result {
Ok(resolution) => {
match &resolution {
FunctionResolution::Direct(_) => {
return Ok(resolution);
}
FunctionResolution::Broadcast(_) => {
if broadcast_match.is_some() {
return Err(FunctionBindingError::AmbiguousBroadcast {
function_name: function_name.to_string(),
});
}
broadcast_match = Some(resolution);
}
}
}
Err(attempt) => {
binding_attempts.push(attempt);
}
}
}
if let Some(resolution) = broadcast_match {
return Ok(resolution);
}
Err(FunctionBindingError::NoMatchingSignature {
function_name: function_name.to_string(),
positional,
named,
apply_failure: ApplyFailure(binding_attempts),
})
}
fn try_standard_bind_typed(
&self,
function_def: &Arc<dyn FunctionDef>,
params: &Parameters,
typed_positional: Vec<Arc<TypedExpression>>,
typed_named: OrderMap<String, Arc<TypedExpression>>,
) -> Result<FunctionResolution<Arc<TypedExpression>>, ApplyAttempt> {
let typed_binding: ParameterBinding<Arc<TypedExpression>> =
match params.bind(typed_positional, typed_named) {
Ok(b) => b,
Err(failures) => {
return Err(ApplyAttempt {
function_def: format!("{}", function_def),
binding_failures: FunctionParameterBindingFailures(failures),
});
}
};
if let Err(failures) = params.check(&typed_binding) {
return self
.try_broadcast(function_def, params, &typed_binding)
.ok_or_else(|| ApplyAttempt {
function_def: format!("{}", function_def),
binding_failures: FunctionParameterBindingFailures(failures),
});
}
let typed_binding = function_def.refine_binding(typed_binding);
let typ = match function_def.return_type(&typed_binding) {
Ok(t) => t,
Err(e) => {
if let Some(broadcast) = self.try_broadcast(function_def, params, &typed_binding) {
return Ok(broadcast);
}
return Err(match e.downcast::<MatchTestFailure>() {
Ok(mtf) => ApplyAttempt {
function_def: format!("{}", function_def),
binding_failures: FunctionParameterBindingFailures(vec![
FunctionParameterBindingFailure::DoesNotMatch(mtf.0),
]),
},
Err(other) => ApplyAttempt {
function_def: format!("{}", function_def),
binding_failures: FunctionParameterBindingFailures(vec![
FunctionParameterBindingFailure::DoesNotMatch(
format!("{}", other).into(),
),
]),
},
});
}
};
Ok(DirectResolution {
function_def: function_def.clone(),
binding: typed_binding,
typ,
}
.into())
}
fn try_broadcast(
&self,
function_def: &Arc<dyn FunctionDef>,
params: &Parameters,
typed_binding: &ParameterBinding<Arc<TypedExpression>>,
) -> Option<FunctionResolution<Arc<TypedExpression>>> {
use crate::types::array::Array;
let binding_len = typed_binding.len();
for i in 0..binding_len {
let arg = typed_binding.get_by_index(i).ok()?;
let element_type = match arg.resolved_type.as_ref() {
Type::Array(arr) => arr.element_type.as_ref().clone(),
_ => continue,
};
let element_expr = Arc::new(TypedExpression {
ast: arg.ast.clone(),
resolved_type: Arc::new(element_type),
kind: crate::tree::typed_ast::expression::TypedExpressionKind::Leaf,
});
let test_binding = match typed_binding.clone().replace_by_index(i, element_expr) {
Ok(b) => b,
Err(_) => continue,
};
if params.check(&test_binding).is_err() {
continue;
}
let element_result_type = match function_def.return_type(&test_binding) {
Ok(t) => t,
Err(_) => continue,
};
let result_type: Type = Array::new(element_result_type).into();
return Some(
BroadcastResolution {
function_def: function_def.clone(),
binding: typed_binding.clone(),
typ: result_type,
broadcast_position: i,
}
.into(),
);
}
None
}
}
#[derive(Debug, Clone)]
pub enum InvocationType {
UnaryPrefix,
UnaryPostfix,
Binary,
Function,
}
#[derive(thiserror::Error, Debug)]
pub enum FunctionBindingError {
#[error("Function not found: {function_name}")]
FunctionNotFound { function_name: String },
#[error("Function name '{function_name}' must be lowercase{}", suggestion.as_ref().map(|s| format!(": use '{}'", s)).unwrap_or_default())]
FunctionNameNotLowercase {
function_name: String,
suggestion: Option<String>,
},
#[error(
"No matching signature found for function '{function_name}'. Attempted bindings:\n{apply_failure}"
)]
NoMatchingSignature {
function_name: String,
positional: Vec<Arc<Expression>>,
named: OrderMap<String, Arc<Expression>>,
apply_failure: ApplyFailure,
},
#[error(
"Ambiguous broadcast for function '{function_name}': multiple array arguments could be broadcast"
)]
AmbiguousBroadcast { function_name: String },
#[error("Fatal failure during type matching: {0}")]
Fatal(Box<dyn Error + Send + Sync + 'static>),
}
impl Default for FunctionRegistry {
fn default() -> Self {
let mut registry = Self::empty();
registry.register_unary_prefix(Not);
registry.register_unary_prefix(UnaryMinus);
registry.register_unary_prefix(UnaryPlus);
registry.register_unary_prefix(UnaryRangePrefix);
registry.register_unary_prefix(UnaryRangePrefixInclusive);
registry.register_unary_postfix(UnaryRangePostfix);
registry.register_binary(And);
registry.register_binary(Or);
registry.register_binary(NumericPlus);
registry.register_binary(NumericMinus);
registry.register_binary(NumericMultiply);
registry.register_binary(NumericDivide);
registry.register_binary(NumericModulo);
registry.register_binary(StringConcat);
registry.register_binary(ArrayConcat);
registry.register_binary(NumericEq);
registry.register_binary(NumericNeq);
registry.register_binary(NumericLt);
registry.register_binary(NumericLte);
registry.register_binary(NumericGt);
registry.register_binary(NumericGte);
registry.register_binary(StringEq);
registry.register_binary(StringNeq);
registry.register_binary(BooleanEq);
registry.register_binary(BooleanNeq);
registry.register_binary(TimestampEq);
registry.register_binary(TimestampNeq);
registry.register_binary(TimestampLt);
registry.register_binary(TimestampLte);
registry.register_binary(TimestampGt);
registry.register_binary(TimestampGte);
registry.register_binary(IntervalEq);
registry.register_binary(IntervalNeq);
registry.register_binary(IntervalLt);
registry.register_binary(IntervalLte);
registry.register_binary(IntervalGt);
registry.register_binary(IntervalGte);
registry.register_binary(Is);
registry.register_binary(IsNot);
registry.register_binary(InArray);
registry.register_binary(NotInArray);
registry.register_binary(InMap);
registry.register_binary(NotInMap);
registry.register_binary(NumericRange);
registry.register_binary(NumericRangeInclusive);
registry.register_binary(TimestampRange);
registry.register_binary(TimestampRangeInclusive);
registry.register_binary(InRange);
registry.register_binary(NotInRange);
registry.register_binary(InTimestampTimestamp);
registry.register_binary(NotInTimestampTimestamp);
registry.register_binary(InTimestampInterval);
registry.register_binary(NotInTimestampInterval);
registry.register_binary(InTuple);
registry.register_binary(NotInTuple);
registry.register_binary(IntervalMultiplyNumeric);
registry.register_binary(NumericMultiplyInterval);
registry.register_binary(IntervalDivideNumeric);
registry.register_binary(TimestampPlusInterval);
registry.register_binary(TimestampMinusInterval);
registry.register_binary(IntervalPlusTimestamp);
registry.register_binary(TimestampMinusTimestamp);
registry.register_binary(IntervalMinusInterval);
registry.register_binary(IntervalPlusInterval);
registry.register_binary(CalendarIntervalPlusCalendarInterval);
registry.register_binary(CalendarIntervalMinusCalendarInterval);
registry.register_binary(CalendarIntervalMultiplyInt);
registry.register_binary(IntMultiplyCalendarInterval);
registry.register_binary(CalendarIntervalDivideInt);
registry.register_binary(TimestampPlusCalendarInterval);
registry.register_binary(CalendarIntervalPlusTimestamp);
registry.register_binary(TimestampMinusCalendarInterval);
registry.register_function(Abs);
registry.register_function(Cbrt);
registry.register_function(Ceil);
registry.register_function(Degrees);
registry.register_function(Euler);
registry.register_function(Exp);
registry.register_function(Floor);
registry.register_function(Ln);
registry.register_function(Log);
registry.register_function(Log10);
registry.register_function(Log2);
registry.register_function(Pi);
registry.register_function(Pow);
registry.register_function(Radians);
registry.register_function(Round1);
registry.register_function(Round2);
registry.register_function(Sign);
registry.register_function(Sqrt);
registry.register_function(Truncate);
registry.register_function(WidthBucket4);
registry.register_function(WidthBucket2);
registry.register_function(If2);
registry.register_function(If3);
registry.register_function(Case);
registry.register_function(ArrayCoalesce);
registry.register_function(Coalesce);
registry.register_function(RegexpCount);
registry.register_function(RegexpExtractAll2);
registry.register_function(RegexpExtractAll3);
registry.register_function(RegexpExtract2);
registry.register_function(RegexpExtract3);
registry.register_function(RegexpLike);
registry.register_function(RegexpPosition2);
registry.register_function(RegexpPosition3);
registry.register_function(RegexpPosition4);
registry.register_function(RegexpReplace2);
registry.register_function(RegexpReplace3);
registry.register_function(RegexpSplit);
registry.register_function(Replace2);
registry.register_function(Replace3);
registry.register_function(Substr2);
registry.register_function(Substr3);
registry.register_function(StartsWith);
registry.register_function(EndsWith);
registry.register_function(Contains);
registry.register_function(CidrContains);
registry.register_function(IsIpv4);
registry.register_function(IsIpv6);
registry.register_function(Lower);
registry.register_function(Upper);
registry.register_function(UrlEncode);
registry.register_function(UrlDecode);
registry.register_function(StringLen);
registry.register_function(Uuid);
registry.register_function(Uuid5);
registry.register_function(ParseJsonString);
registry.register_function(ParseJsonVariant);
registry.register_function(ToJsonString);
registry.register_function(ArrayVariantToJson);
registry.register_function(Typeof);
registry.register_function(MapFromPairs);
registry.register_function(MapFromArrays);
registry.register_function(MapEmpty);
registry.register_function(MapFromKeyValue);
registry.register_function(MapKeys);
registry.register_function(MapValues);
registry.register_function(TransformValues);
registry.register_function(MapVariantGetValues);
registry.register_function(MapVariantToJsonValues);
registry.register_function(Now);
registry.register_function(Today);
registry.register_function(Yesterday);
registry.register_function(Tomorrow);
registry.register_function(Ts);
registry.register_function(Year);
registry.register_function(Month);
registry.register_function(Day);
registry.register_function(DayOfWeek);
registry.register_function(Hour);
registry.register_function(Minute);
registry.register_function(Second);
registry.register_function(AtTimezone);
registry.register_function(ToMillis);
registry.register_function(ToNanos);
registry.register_function(FromMillis);
registry.register_function(FromNanos);
registry.register_function(FromUnixtimeSeconds);
registry.register_function(FromUnixtimeMillis);
registry.register_function(FromUnixtimeMicros);
registry.register_function(FromUnixtimeNanos);
registry.register_function(ToUnixtime);
registry.register_function(CountStar);
registry.register_function(CountAny);
registry.register_function(CountDistinct);
registry.register_function(ApproxDistinct);
registry.register_function(CountIf);
registry.register_function(Sum);
registry.register_function(Avg);
registry.register_function(Stddev);
registry.register_function(ApproxPercentile);
registry.register_function(AggMin);
registry.register_function(AggMax);
registry.register_function(AnyValue);
registry.register_function(ArrayAgg);
registry.register_function(SetAgg);
registry.register_function(MapAgg);
registry.register_function(MultimapAgg);
registry.register_function(AggAny);
registry.register_function(AggAll);
registry.register_function(RowNumber);
registry.register_function(Rank);
registry.register_function(DenseRank);
registry.register_function(Lag);
registry.register_function(Lead);
registry.register_function(FirstValue);
registry.register_function(LastValue);
registry.register_function(NthValue);
registry.register_function(CumeDist);
registry.register_function(PercentRank);
registry.register_function(MatchFirst);
registry.register_function(MatchLast);
registry.register_function(MatchCountStar);
registry.register_function(MatchCountAny);
registry.register_function(MatchSum);
registry.register_function(MatchAvg);
registry.register_function(MatchMin);
registry.register_function(MatchMax);
registry.register_function(MatchArrayAgg);
registry.register_function(MatchCountDistinct);
registry.register_function(InternalMatchLength);
registry.register_function(InternalMatchCountNonNull);
registry.register_function(InternalMatchSum);
registry.register_function(InternalMatchAvg);
registry.register_function(InternalMatchLast);
registry.register_function(InternalMatchMin);
registry.register_function(InternalMatchMax);
registry.register_function(FilterNull);
registry.register_function(ArrayOrMapLen);
registry.register_function(ArrayDistinct);
registry.register_function(Slice);
registry.register_function(Split);
registry.register_function(ArrayJoin2);
registry.register_function(ArrayJoin3);
registry.register_function(Flatten);
registry.register_function(ArrayAny);
registry.register_function(ArrayAll);
registry.register_function(ArrayMax);
registry.register_function(ArrayMin);
registry.register_function(ArraySum);
registry.register_function(ArrayAvg);
registry.register_function(Sequence2);
registry.register_function(Sequence3);
registry.register_function(Zip);
registry.register_function(Transform);
registry.register_function(Filter);
registry.register_function(GetArray);
registry.register_function(GetMap);
registry.register_function(ArrayVariantGet);
registry
}
}