use std::any::TypeId;
use std::sync::Arc;
use anyhow::bail;
use crate::func::def::{
FunctionDef, ParameterBinding, ParameterBindingProvider, Parameters, SpecialPosition,
};
use crate::tree::ast::expression::{Expression, ExpressionKind, IntLiteral};
use crate::tree::ast::identifier::SimpleIdentifier;
use crate::tree::ast::ops::UnaryPrefixOp;
use crate::tree::typed_ast::expression::TypedExpression;
use crate::types::array::Array;
use crate::types::map::Map;
use crate::types::matcher::{
numeric_or_interval_matcher, AnyMatcher, BaseMatcher, ExactMatcher, MapKeyMatcher,
NumericMatcher, OrMatcher,
};
use crate::types::struct_type::Struct;
use crate::types::{Type, BOOLEAN, DOUBLE, INT, INTERVAL, STRING, TIMESTAMP, VARIANT};
#[derive(Default, Clone)]
pub struct CountStar;
impl FunctionDef for CountStar {
fn name(&self) -> &'static str {
"count"
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(INT)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct CountAny;
impl FunctionDef for CountAny {
fn name(&self) -> &'static str {
"count"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", AnyMatcher::default())
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(INT)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct CountDistinct;
impl FunctionDef for CountDistinct {
fn name(&self) -> &'static str {
"count_distinct"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", AnyMatcher::default())
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(INT)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct ApproxDistinct;
impl FunctionDef for ApproxDistinct {
fn name(&self) -> &'static str {
"approx_distinct"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", AnyMatcher::default())
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(INT)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct ApproxTopK;
impl ApproxTopK {
pub const DEFAULT_K: i64 = 5;
pub const DEFAULT_MAX_ITEMS_TRACKED: i64 = 10_000;
pub const MAX_K: i64 = 100_000;
pub const MAX_ITEMS_TRACKED: i64 = 100_000;
pub fn validate_config(k: i64, max_items_tracked: i64) -> anyhow::Result<(usize, usize)> {
if k <= 0 {
bail!("approx_top_k: k must be greater than 0, got {k}");
}
if k > Self::MAX_K {
bail!(
"approx_top_k: k must be less than or equal to {}, got {k}",
Self::MAX_K
);
}
if max_items_tracked < k {
bail!(
"approx_top_k: max_items_tracked must be greater than or equal to k ({k}), got {max_items_tracked}"
);
}
if max_items_tracked > Self::MAX_ITEMS_TRACKED {
bail!(
"approx_top_k: max_items_tracked must be less than or equal to {}, got {max_items_tracked}",
Self::MAX_ITEMS_TRACKED
);
}
let k = usize::try_from(k)
.map_err(|_| anyhow::anyhow!("approx_top_k: k does not fit this platform"))?;
let max_items_tracked = usize::try_from(max_items_tracked).map_err(|_| {
anyhow::anyhow!("approx_top_k: max_items_tracked does not fit this platform")
})?;
Ok((k, max_items_tracked))
}
fn integer_literal(name: &str, expression: &Expression) -> anyhow::Result<i64> {
match &expression.kind {
ExpressionKind::IntLiteral(IntLiteral { int }) => Ok(*int),
ExpressionKind::UnaryPrefixOperator(operator) => {
let ExpressionKind::IntLiteral(IntLiteral { int }) = &operator.operand.kind else {
bail!("approx_top_k: {name} must be an integer literal");
};
match operator.operator {
UnaryPrefixOp::Plus => Ok(*int),
UnaryPrefixOp::Minus => int.checked_neg().ok_or_else(|| {
anyhow::anyhow!("approx_top_k: {name} integer literal is out of range")
}),
_ => bail!("approx_top_k: {name} must be an integer literal"),
}
}
_ => bail!("approx_top_k: {name} must be an integer literal"),
}
}
}
impl FunctionDef for ApproxTopK {
fn name(&self) -> &'static str {
"approx_top_k"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("value", BaseMatcher)
.with_default(
"k",
ExactMatcher::of(INT),
Expression::from_kind(IntLiteral {
int: Self::DEFAULT_K,
}),
)
.with_default(
"max_items_tracked",
ExactMatcher::of(INT),
Expression::from_kind(IntLiteral {
int: Self::DEFAULT_MAX_ITEMS_TRACKED,
}),
)
}
fn refine_binding(
&self,
binding: ParameterBinding<Arc<TypedExpression>>,
) -> anyhow::Result<ParameterBinding<Arc<TypedExpression>>> {
let k = Self::integer_literal("k", binding.get_by_name("k")?.ast.as_ref())?;
let max_items_tracked = Self::integer_literal(
"max_items_tracked",
binding.get_by_name("max_items_tracked")?.ast.as_ref(),
)?;
Self::validate_config(k, max_items_tracked)?;
Ok(binding)
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
let item_type = bindings.get_by_name("value")?.typ().clone();
let result = Struct::new([
(SimpleIdentifier::new("item"), item_type),
(SimpleIdentifier::new("count"), INT),
]);
Ok(Array::new(result.into()).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
fn manages_window_clause(&self) -> bool {
true
}
}
#[derive(Default, Clone)]
pub struct CountIf;
impl FunctionDef for CountIf {
fn name(&self) -> &'static str {
"count_if"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("condition", ExactMatcher::of(BOOLEAN))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(INT)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct Sum;
impl FunctionDef for Sum {
fn name(&self) -> &'static str {
"sum"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", numeric_or_interval_matcher())
}
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>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct Avg;
impl FunctionDef for Avg {
fn name(&self) -> &'static str {
"avg"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", numeric_or_interval_matcher())
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
match bindings.get_by_index(0)?.typ() {
Type::Int => Ok(DOUBLE),
other => Ok(other.clone()),
}
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct Stddev;
impl FunctionDef for Stddev {
fn name(&self) -> &'static str {
"stddev"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", NumericMatcher::default())
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(DOUBLE)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct ApproxPercentile;
impl FunctionDef for ApproxPercentile {
fn name(&self) -> &'static str {
"approx_percentile"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("x", NumericMatcher::default())
.with("percentile", NumericMatcher::default())
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
match bindings.get_by_name("x")?.typ() {
Type::Int => Ok(INT.into()),
_ => Ok(DOUBLE.clone()),
}
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct AggMin;
impl FunctionDef for AggMin {
fn name(&self) -> &'static str {
"min"
}
fn parameters(&self) -> Parameters {
Parameters::new().with(
"x",
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> {
Ok(bindings.get_by_index(0)?.typ().clone())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct AggMax;
impl FunctionDef for AggMax {
fn name(&self) -> &'static str {
"max"
}
fn parameters(&self) -> Parameters {
Parameters::new().with(
"x",
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> {
Ok(bindings.get_by_index(0)?.typ().clone())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct AnyValue;
impl FunctionDef for AnyValue {
fn name(&self) -> &'static str {
"any_value"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", AnyMatcher::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>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct SchemaOfVariantAgg;
impl FunctionDef for SchemaOfVariantAgg {
fn name(&self) -> &'static str {
"schema_of_variant_agg"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("value", ExactMatcher::of(VARIANT))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(STRING)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct ArrayAgg;
impl FunctionDef for ArrayAgg {
fn name(&self) -> &'static str {
"array_agg"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", AnyMatcher::default())
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(Array::new(bindings.get_by_index(0)?.typ().clone()).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
fn sortable_input(&self) -> bool {
true
}
}
#[derive(Default, Clone)]
pub struct SetAgg;
impl FunctionDef for SetAgg {
fn name(&self) -> &'static str {
"set_agg"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", AnyMatcher::default())
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(Array::new(bindings.get_by_index(0)?.typ().clone()).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct MapAgg;
impl FunctionDef for MapAgg {
fn name(&self) -> &'static str {
"map_agg"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("key", MapKeyMatcher::default())
.with("value", AnyMatcher::default())
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
let key = bindings.get_by_index(0)?.typ().clone();
let value = bindings.get_by_index(1)?.typ().clone();
Ok(Map::new(key, value).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
fn sortable_input(&self) -> bool {
true
}
}
#[derive(Default, Clone)]
pub struct MultimapAgg;
impl FunctionDef for MultimapAgg {
fn name(&self) -> &'static str {
"multimap_agg"
}
fn parameters(&self) -> Parameters {
Parameters::new()
.with("key", MapKeyMatcher::default())
.with("value", AnyMatcher::default())
}
fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
let key = bindings.get_by_index(0)?.typ().clone();
let value = bindings.get_by_index(1)?.typ().clone();
Ok(Map::new(key, Array::new(value).into()).into())
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
fn sortable_input(&self) -> bool {
true
}
}
#[derive(Default, Clone)]
pub struct AggAny;
impl FunctionDef for AggAny {
fn name(&self) -> &'static str {
"any"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", ExactMatcher::of(BOOLEAN))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(BOOLEAN)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[derive(Default, Clone)]
pub struct AggAll;
impl FunctionDef for AggAll {
fn name(&self) -> &'static str {
"all"
}
fn parameters(&self) -> Parameters {
Parameters::new().with("x", ExactMatcher::of(BOOLEAN))
}
fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
Ok(BOOLEAN)
}
fn type_id(&self) -> TypeId {
TypeId::of::<Self>()
}
fn special_position(&self) -> Option<SpecialPosition> {
Some(SpecialPosition::Agg)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tree::ast::ParseWithErrors;
fn typed_expression(source: &str) -> Arc<TypedExpression> {
let expression = Expression::parse_result(source).expect("expression must parse");
expression.into()
}
#[test]
fn approx_top_k_validation_preserves_unary_plus_arguments() {
let value = typed_expression("'value'");
let k = typed_expression("+3");
let max_items_tracked = typed_expression("+1000");
let binding = ParameterBinding::from_named([
("value".to_string(), value),
("k".to_string(), k.clone()),
("max_items_tracked".to_string(), max_items_tracked.clone()),
]);
let validated = ApproxTopK
.refine_binding(binding)
.expect("valid configuration must pass validation");
assert!(Arc::ptr_eq(validated.get_by_name("k").unwrap(), &k));
assert!(Arc::ptr_eq(
validated.get_by_name("max_items_tracked").unwrap(),
&max_items_tracked
));
}
}