Skip to main content

hamelin_lib/func/defs/
aggregate.rs

1//! Aggregate functions
2
3use std::any::TypeId;
4use std::sync::Arc;
5
6use anyhow::bail;
7
8use crate::func::def::{
9    FunctionDef, ParameterBinding, ParameterBindingProvider, Parameters, SpecialPosition,
10};
11use crate::tree::ast::expression::{Expression, ExpressionKind, IntLiteral};
12use crate::tree::ast::identifier::SimpleIdentifier;
13use crate::tree::ast::ops::UnaryPrefixOp;
14use crate::tree::typed_ast::expression::TypedExpression;
15use crate::types::array::Array;
16use crate::types::map::Map;
17use crate::types::matcher::{
18    numeric_or_interval_matcher, AnyMatcher, BaseMatcher, ExactMatcher, MapKeyMatcher,
19    NumericMatcher, OrMatcher,
20};
21use crate::types::struct_type::Struct;
22use crate::types::{Type, BOOLEAN, DOUBLE, INT, INTERVAL, STRING, TIMESTAMP, VARIANT};
23
24/// count() - Count all rows (COUNT(*)).
25#[derive(Default, Clone)]
26pub struct CountStar;
27
28impl FunctionDef for CountStar {
29    fn name(&self) -> &'static str {
30        "count"
31    }
32
33    fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
34        Ok(INT)
35    }
36
37    fn type_id(&self) -> TypeId {
38        TypeId::of::<Self>()
39    }
40
41    fn special_position(&self) -> Option<SpecialPosition> {
42        Some(SpecialPosition::Agg)
43    }
44}
45
46/// count(x) - Count non-null values.
47#[derive(Default, Clone)]
48pub struct CountAny;
49
50impl FunctionDef for CountAny {
51    fn name(&self) -> &'static str {
52        "count"
53    }
54
55    fn parameters(&self) -> Parameters {
56        Parameters::new().with("x", AnyMatcher::default())
57    }
58
59    fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
60        Ok(INT)
61    }
62
63    fn type_id(&self) -> TypeId {
64        TypeId::of::<Self>()
65    }
66
67    fn special_position(&self) -> Option<SpecialPosition> {
68        Some(SpecialPosition::Agg)
69    }
70}
71
72/// count_distinct(x) - Count distinct non-null values.
73#[derive(Default, Clone)]
74pub struct CountDistinct;
75
76impl FunctionDef for CountDistinct {
77    fn name(&self) -> &'static str {
78        "count_distinct"
79    }
80
81    fn parameters(&self) -> Parameters {
82        Parameters::new().with("x", AnyMatcher::default())
83    }
84
85    fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
86        Ok(INT)
87    }
88
89    fn type_id(&self) -> TypeId {
90        TypeId::of::<Self>()
91    }
92
93    fn special_position(&self) -> Option<SpecialPosition> {
94        Some(SpecialPosition::Agg)
95    }
96}
97
98/// approx_distinct(x) - Approximate count distinct.
99#[derive(Default, Clone)]
100pub struct ApproxDistinct;
101
102impl FunctionDef for ApproxDistinct {
103    fn name(&self) -> &'static str {
104        "approx_distinct"
105    }
106
107    fn parameters(&self) -> Parameters {
108        Parameters::new().with("x", AnyMatcher::default())
109    }
110
111    fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
112        Ok(INT)
113    }
114
115    fn type_id(&self) -> TypeId {
116        TypeId::of::<Self>()
117    }
118
119    fn special_position(&self) -> Option<SpecialPosition> {
120        Some(SpecialPosition::Agg)
121    }
122}
123
124/// approx_top_k(value, k, max_items_tracked) - Approximate most frequent values.
125#[derive(Default, Clone)]
126pub struct ApproxTopK;
127
128impl ApproxTopK {
129    /// Default number of items returned.
130    pub const DEFAULT_K: i64 = 5;
131    /// Default candidate capacity.
132    pub const DEFAULT_MAX_ITEMS_TRACKED: i64 = 10_000;
133    /// Largest supported `k` value.
134    pub const MAX_K: i64 = 100_000;
135    /// Largest supported candidate capacity.
136    pub const MAX_ITEMS_TRACKED: i64 = 100_000;
137
138    /// Validate semantic arguments and convert them to runtime sizes.
139    pub fn validate_config(k: i64, max_items_tracked: i64) -> anyhow::Result<(usize, usize)> {
140        if k <= 0 {
141            bail!("approx_top_k: k must be greater than 0, got {k}");
142        }
143        if k > Self::MAX_K {
144            bail!(
145                "approx_top_k: k must be less than or equal to {}, got {k}",
146                Self::MAX_K
147            );
148        }
149        if max_items_tracked < k {
150            bail!(
151                "approx_top_k: max_items_tracked must be greater than or equal to k ({k}), got {max_items_tracked}"
152            );
153        }
154        if max_items_tracked > Self::MAX_ITEMS_TRACKED {
155            bail!(
156                "approx_top_k: max_items_tracked must be less than or equal to {}, got {max_items_tracked}",
157                Self::MAX_ITEMS_TRACKED
158            );
159        }
160
161        let k = usize::try_from(k)
162            .map_err(|_| anyhow::anyhow!("approx_top_k: k does not fit this platform"))?;
163        let max_items_tracked = usize::try_from(max_items_tracked).map_err(|_| {
164            anyhow::anyhow!("approx_top_k: max_items_tracked does not fit this platform")
165        })?;
166
167        Ok((k, max_items_tracked))
168    }
169
170    fn integer_literal(name: &str, expression: &Expression) -> anyhow::Result<i64> {
171        match &expression.kind {
172            ExpressionKind::IntLiteral(IntLiteral { int }) => Ok(*int),
173            ExpressionKind::UnaryPrefixOperator(operator) => {
174                let ExpressionKind::IntLiteral(IntLiteral { int }) = &operator.operand.kind else {
175                    bail!("approx_top_k: {name} must be an integer literal");
176                };
177                match operator.operator {
178                    UnaryPrefixOp::Plus => Ok(*int),
179                    UnaryPrefixOp::Minus => int.checked_neg().ok_or_else(|| {
180                        anyhow::anyhow!("approx_top_k: {name} integer literal is out of range")
181                    }),
182                    _ => bail!("approx_top_k: {name} must be an integer literal"),
183                }
184            }
185            _ => bail!("approx_top_k: {name} must be an integer literal"),
186        }
187    }
188}
189
190impl FunctionDef for ApproxTopK {
191    fn name(&self) -> &'static str {
192        "approx_top_k"
193    }
194
195    fn parameters(&self) -> Parameters {
196        Parameters::new()
197            .with("value", BaseMatcher)
198            .with_default(
199                "k",
200                ExactMatcher::of(INT),
201                Expression::from_kind(IntLiteral {
202                    int: Self::DEFAULT_K,
203                }),
204            )
205            .with_default(
206                "max_items_tracked",
207                ExactMatcher::of(INT),
208                Expression::from_kind(IntLiteral {
209                    int: Self::DEFAULT_MAX_ITEMS_TRACKED,
210                }),
211            )
212    }
213
214    fn refine_binding(
215        &self,
216        binding: ParameterBinding<Arc<TypedExpression>>,
217    ) -> anyhow::Result<ParameterBinding<Arc<TypedExpression>>> {
218        let k = Self::integer_literal("k", binding.get_by_name("k")?.ast.as_ref())?;
219        let max_items_tracked = Self::integer_literal(
220            "max_items_tracked",
221            binding.get_by_name("max_items_tracked")?.ast.as_ref(),
222        )?;
223        Self::validate_config(k, max_items_tracked)?;
224        Ok(binding)
225    }
226
227    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
228        let item_type = bindings.get_by_name("value")?.typ().clone();
229        let result = Struct::new([
230            (SimpleIdentifier::new("item"), item_type),
231            (SimpleIdentifier::new("count"), INT),
232        ]);
233        Ok(Array::new(result.into()).into())
234    }
235
236    fn type_id(&self) -> TypeId {
237        TypeId::of::<Self>()
238    }
239
240    fn special_position(&self) -> Option<SpecialPosition> {
241        Some(SpecialPosition::Agg)
242    }
243
244    fn manages_window_clause(&self) -> bool {
245        true
246    }
247}
248
249/// count_if(condition) - Count rows where condition is true.
250#[derive(Default, Clone)]
251pub struct CountIf;
252
253impl FunctionDef for CountIf {
254    fn name(&self) -> &'static str {
255        "count_if"
256    }
257
258    fn parameters(&self) -> Parameters {
259        Parameters::new().with("condition", ExactMatcher::of(BOOLEAN))
260    }
261
262    fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
263        Ok(INT)
264    }
265
266    fn type_id(&self) -> TypeId {
267        TypeId::of::<Self>()
268    }
269
270    fn special_position(&self) -> Option<SpecialPosition> {
271        Some(SpecialPosition::Agg)
272    }
273}
274
275/// sum(x) - Sum of numeric values.
276#[derive(Default, Clone)]
277pub struct Sum;
278
279impl FunctionDef for Sum {
280    fn name(&self) -> &'static str {
281        "sum"
282    }
283
284    fn parameters(&self) -> Parameters {
285        Parameters::new().with("x", numeric_or_interval_matcher())
286    }
287
288    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
289        // Return type matches input type
290        Ok(bindings.get_by_index(0)?.typ().clone())
291    }
292
293    fn type_id(&self) -> TypeId {
294        TypeId::of::<Self>()
295    }
296
297    fn special_position(&self) -> Option<SpecialPosition> {
298        Some(SpecialPosition::Agg)
299    }
300}
301
302/// avg(x) - Average of numeric values.
303#[derive(Default, Clone)]
304pub struct Avg;
305
306impl FunctionDef for Avg {
307    fn name(&self) -> &'static str {
308        "avg"
309    }
310
311    fn parameters(&self) -> Parameters {
312        Parameters::new().with("x", numeric_or_interval_matcher())
313    }
314
315    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
316        match bindings.get_by_index(0)?.typ() {
317            Type::Int => Ok(DOUBLE),
318            other => Ok(other.clone()),
319        }
320    }
321
322    fn type_id(&self) -> TypeId {
323        TypeId::of::<Self>()
324    }
325
326    fn special_position(&self) -> Option<SpecialPosition> {
327        Some(SpecialPosition::Agg)
328    }
329}
330
331/// stddev(x) - Standard deviation of numeric values.
332#[derive(Default, Clone)]
333pub struct Stddev;
334
335impl FunctionDef for Stddev {
336    fn name(&self) -> &'static str {
337        "stddev"
338    }
339
340    fn parameters(&self) -> Parameters {
341        Parameters::new().with("x", NumericMatcher::default())
342    }
343
344    fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
345        Ok(DOUBLE)
346    }
347
348    fn type_id(&self) -> TypeId {
349        TypeId::of::<Self>()
350    }
351
352    fn special_position(&self) -> Option<SpecialPosition> {
353        Some(SpecialPosition::Agg)
354    }
355}
356
357/// approx_percentile(x, percentile) - Approximate percentile.
358#[derive(Default, Clone)]
359pub struct ApproxPercentile;
360
361impl FunctionDef for ApproxPercentile {
362    fn name(&self) -> &'static str {
363        "approx_percentile"
364    }
365
366    fn parameters(&self) -> Parameters {
367        Parameters::new()
368            .with("x", NumericMatcher::default())
369            .with("percentile", NumericMatcher::default())
370    }
371
372    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
373        match bindings.get_by_name("x")?.typ() {
374            Type::Int => Ok(INT.into()),
375            _ => Ok(DOUBLE.clone()),
376        }
377    }
378
379    fn type_id(&self) -> TypeId {
380        TypeId::of::<Self>()
381    }
382
383    fn special_position(&self) -> Option<SpecialPosition> {
384        Some(SpecialPosition::Agg)
385    }
386}
387
388/// min(x) - Minimum value (aggregate).
389#[derive(Default, Clone)]
390pub struct AggMin;
391
392impl FunctionDef for AggMin {
393    fn name(&self) -> &'static str {
394        "min"
395    }
396
397    fn parameters(&self) -> Parameters {
398        Parameters::new().with(
399            "x",
400            OrMatcher::default()
401                .with(NumericMatcher::default())
402                .with(ExactMatcher::of(STRING))
403                .with(ExactMatcher::of(TIMESTAMP))
404                .with(ExactMatcher::of(INTERVAL)),
405        )
406    }
407
408    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
409        Ok(bindings.get_by_index(0)?.typ().clone())
410    }
411
412    fn type_id(&self) -> TypeId {
413        TypeId::of::<Self>()
414    }
415
416    fn special_position(&self) -> Option<SpecialPosition> {
417        Some(SpecialPosition::Agg)
418    }
419}
420
421/// max(x) - Maximum value (aggregate).
422#[derive(Default, Clone)]
423pub struct AggMax;
424
425impl FunctionDef for AggMax {
426    fn name(&self) -> &'static str {
427        "max"
428    }
429
430    fn parameters(&self) -> Parameters {
431        Parameters::new().with(
432            "x",
433            OrMatcher::default()
434                .with(NumericMatcher::default())
435                .with(ExactMatcher::of(STRING))
436                .with(ExactMatcher::of(TIMESTAMP))
437                .with(ExactMatcher::of(INTERVAL)),
438        )
439    }
440
441    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
442        Ok(bindings.get_by_index(0)?.typ().clone())
443    }
444
445    fn type_id(&self) -> TypeId {
446        TypeId::of::<Self>()
447    }
448
449    fn special_position(&self) -> Option<SpecialPosition> {
450        Some(SpecialPosition::Agg)
451    }
452}
453
454/// any_value(x) - Return any value from the group.
455#[derive(Default, Clone)]
456pub struct AnyValue;
457
458impl FunctionDef for AnyValue {
459    fn name(&self) -> &'static str {
460        "any_value"
461    }
462
463    fn parameters(&self) -> Parameters {
464        Parameters::new().with("x", AnyMatcher::default())
465    }
466
467    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
468        Ok(bindings.get_by_index(0)?.typ().clone())
469    }
470
471    fn type_id(&self) -> TypeId {
472        TypeId::of::<Self>()
473    }
474
475    fn special_position(&self) -> Option<SpecialPosition> {
476        Some(SpecialPosition::Agg)
477    }
478}
479
480/// schema_of_variant_agg(value) - Infer and merge the Hamelin types of Variant values.
481#[derive(Default, Clone)]
482pub struct SchemaOfVariantAgg;
483
484impl FunctionDef for SchemaOfVariantAgg {
485    fn name(&self) -> &'static str {
486        "schema_of_variant_agg"
487    }
488
489    fn parameters(&self) -> Parameters {
490        Parameters::new().with("value", ExactMatcher::of(VARIANT))
491    }
492
493    fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
494        Ok(STRING)
495    }
496
497    fn type_id(&self) -> TypeId {
498        TypeId::of::<Self>()
499    }
500
501    fn special_position(&self) -> Option<SpecialPosition> {
502        Some(SpecialPosition::Agg)
503    }
504}
505
506/// array_agg(x) - Aggregate values into an array.
507#[derive(Default, Clone)]
508pub struct ArrayAgg;
509
510impl FunctionDef for ArrayAgg {
511    fn name(&self) -> &'static str {
512        "array_agg"
513    }
514
515    fn parameters(&self) -> Parameters {
516        Parameters::new().with("x", AnyMatcher::default())
517    }
518
519    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
520        Ok(Array::new(bindings.get_by_index(0)?.typ().clone()).into())
521    }
522
523    fn type_id(&self) -> TypeId {
524        TypeId::of::<Self>()
525    }
526
527    fn special_position(&self) -> Option<SpecialPosition> {
528        Some(SpecialPosition::Agg)
529    }
530
531    fn sortable_input(&self) -> bool {
532        true
533    }
534}
535
536/// set_agg(x) - Aggregate distinct values into an array.
537#[derive(Default, Clone)]
538pub struct SetAgg;
539
540impl FunctionDef for SetAgg {
541    fn name(&self) -> &'static str {
542        "set_agg"
543    }
544
545    fn parameters(&self) -> Parameters {
546        Parameters::new().with("x", AnyMatcher::default())
547    }
548
549    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
550        Ok(Array::new(bindings.get_by_index(0)?.typ().clone()).into())
551    }
552
553    fn type_id(&self) -> TypeId {
554        TypeId::of::<Self>()
555    }
556
557    fn special_position(&self) -> Option<SpecialPosition> {
558        Some(SpecialPosition::Agg)
559    }
560}
561
562/// map_agg(key, value) - Aggregate key-value pairs into a map.
563#[derive(Default, Clone)]
564pub struct MapAgg;
565
566impl FunctionDef for MapAgg {
567    fn name(&self) -> &'static str {
568        "map_agg"
569    }
570
571    fn parameters(&self) -> Parameters {
572        Parameters::new()
573            .with("key", MapKeyMatcher::default())
574            .with("value", AnyMatcher::default())
575    }
576
577    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
578        let key = bindings.get_by_index(0)?.typ().clone();
579        let value = bindings.get_by_index(1)?.typ().clone();
580        Ok(Map::new(key, value).into())
581    }
582
583    fn type_id(&self) -> TypeId {
584        TypeId::of::<Self>()
585    }
586
587    fn special_position(&self) -> Option<SpecialPosition> {
588        Some(SpecialPosition::Agg)
589    }
590
591    fn sortable_input(&self) -> bool {
592        true
593    }
594}
595
596/// multimap_agg(key, value) - Aggregate key-value pairs into a multimap.
597#[derive(Default, Clone)]
598pub struct MultimapAgg;
599
600impl FunctionDef for MultimapAgg {
601    fn name(&self) -> &'static str {
602        "multimap_agg"
603    }
604
605    fn parameters(&self) -> Parameters {
606        Parameters::new()
607            .with("key", MapKeyMatcher::default())
608            .with("value", AnyMatcher::default())
609    }
610
611    fn return_type(&self, bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
612        let key = bindings.get_by_index(0)?.typ().clone();
613        let value = bindings.get_by_index(1)?.typ().clone();
614        Ok(Map::new(key, Array::new(value).into()).into())
615    }
616
617    fn type_id(&self) -> TypeId {
618        TypeId::of::<Self>()
619    }
620
621    fn special_position(&self) -> Option<SpecialPosition> {
622        Some(SpecialPosition::Agg)
623    }
624
625    fn sortable_input(&self) -> bool {
626        true
627    }
628}
629
630/// any(x) - Aggregate boolean OR (any true).
631#[derive(Default, Clone)]
632pub struct AggAny;
633
634impl FunctionDef for AggAny {
635    fn name(&self) -> &'static str {
636        "any"
637    }
638
639    fn parameters(&self) -> Parameters {
640        Parameters::new().with("x", ExactMatcher::of(BOOLEAN))
641    }
642
643    fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
644        Ok(BOOLEAN)
645    }
646
647    fn type_id(&self) -> TypeId {
648        TypeId::of::<Self>()
649    }
650
651    fn special_position(&self) -> Option<SpecialPosition> {
652        Some(SpecialPosition::Agg)
653    }
654}
655
656/// all(x) - Aggregate boolean AND (all true).
657#[derive(Default, Clone)]
658pub struct AggAll;
659
660impl FunctionDef for AggAll {
661    fn name(&self) -> &'static str {
662        "all"
663    }
664
665    fn parameters(&self) -> Parameters {
666        Parameters::new().with("x", ExactMatcher::of(BOOLEAN))
667    }
668
669    fn return_type(&self, _bindings: &dyn ParameterBindingProvider) -> anyhow::Result<Type> {
670        Ok(BOOLEAN)
671    }
672
673    fn type_id(&self) -> TypeId {
674        TypeId::of::<Self>()
675    }
676
677    fn special_position(&self) -> Option<SpecialPosition> {
678        Some(SpecialPosition::Agg)
679    }
680}
681
682#[cfg(test)]
683mod tests {
684    use super::*;
685    use crate::tree::ast::ParseWithErrors;
686
687    fn typed_expression(source: &str) -> Arc<TypedExpression> {
688        let expression = Expression::parse_result(source).expect("expression must parse");
689        expression.into()
690    }
691
692    #[test]
693    fn approx_top_k_validation_preserves_unary_plus_arguments() {
694        let value = typed_expression("'value'");
695        let k = typed_expression("+3");
696        let max_items_tracked = typed_expression("+1000");
697        let binding = ParameterBinding::from_named([
698            ("value".to_string(), value),
699            ("k".to_string(), k.clone()),
700            ("max_items_tracked".to_string(), max_items_tracked.clone()),
701        ]);
702
703        let validated = ApproxTopK
704            .refine_binding(binding)
705            .expect("valid configuration must pass validation");
706
707        assert!(Arc::ptr_eq(validated.get_by_name("k").unwrap(), &k));
708        assert!(Arc::ptr_eq(
709            validated.get_by_name("max_items_tracked").unwrap(),
710            &max_items_tracked
711        ));
712    }
713}