@@ -16,7 +16,6 @@ use std::collections::HashSet;
1616
1717use super :: { Binder , QueryBindStep } ;
1818use crate :: errors:: DatabaseError ;
19- use crate :: expression:: function:: scala:: ScalarFunction ;
2019use crate :: expression:: visitor:: { walk_expr, ExprVisitor } ;
2120use crate :: expression:: visitor_mut:: { walk_mut_expr, ExprVisitorMut } ;
2221use crate :: planner:: LogicalPlan ;
@@ -27,6 +26,22 @@ use crate::{
2726 planner:: operator:: { aggregate:: AggregateOperator , sort:: SortField } ,
2827} ;
2928
29+ struct AggregateCallCollector < ' a > {
30+ agg_calls : & ' a mut Vec < ScalarExpression > ,
31+ }
32+
33+ impl < ' expr > ExprVisitor < ' expr > for AggregateCallCollector < ' _ > {
34+ fn visit ( & mut self , expr : & ' expr ScalarExpression ) -> Result < ( ) , DatabaseError > {
35+ match expr {
36+ ScalarExpression :: AggCall { .. } => self . agg_calls . push ( expr. clone ( ) ) ,
37+ ScalarExpression :: Alias { expr, .. } => self . visit ( expr) ?,
38+ ScalarExpression :: Empty | ScalarExpression :: TableFunction ( _) => unreachable ! ( ) ,
39+ _ => walk_expr ( self , expr) ?,
40+ }
41+ Ok ( ( ) )
42+ }
43+ }
44+
3045impl < T : Transaction , A : AsRef < [ ( & ' static str , DataValue ) ] > > Binder < ' _ , ' _ , T , A > {
3146 pub fn bind_aggregate (
3247 & mut self ,
@@ -48,7 +63,7 @@ impl<T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'_, '_, T, A>
4863 select_items : & mut [ ScalarExpression ] ,
4964 ) -> Result < ( ) , DatabaseError > {
5065 for column in select_items {
51- self . visit_column_agg_expr ( column) ?;
66+ self . collect_aggregate_calls ( column) ?;
5267 }
5368 Ok ( ( ) )
5469 }
@@ -77,14 +92,14 @@ impl<T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'_, '_, T, A>
7792 F : FnMut ( & mut Self , I :: Item ) -> Result < SortField , DatabaseError > ,
7893 {
7994 if let Some ( having) = having. as_mut ( ) {
80- self . visit_column_agg_expr ( having) ?;
95+ self . collect_aggregate_calls ( having) ?;
8196 }
8297 let mut return_orderby = None ;
8398 if let Some ( orderby) = orderby {
8499 let mut fields = Vec :: new ( ) ;
85100 for orderby in orderby {
86- let mut field = bind_sort_field ( self , orderby) ?;
87- self . visit_column_agg_expr ( & mut field. expr ) ?;
101+ let field = bind_sort_field ( self , orderby) ?;
102+ self . collect_aggregate_calls ( & field. expr ) ?;
88103 fields. push ( field) ;
89104 }
90105 return_orderby = Some ( fields) ;
@@ -119,119 +134,11 @@ impl<T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'_, '_, T, A>
119134 Ok ( ( ) )
120135 }
121136
122- fn visit_column_agg_expr ( & mut self , expr : & mut ScalarExpression ) -> Result < ( ) , DatabaseError > {
123- match expr {
124- ScalarExpression :: AggCall { .. } => {
125- self . context . agg_calls . push ( expr. clone ( ) ) ;
126- }
127- ScalarExpression :: TypeCast { expr, .. } => self . visit_column_agg_expr ( expr) ?,
128- ScalarExpression :: IsNull { expr, .. } => self . visit_column_agg_expr ( expr) ?,
129- ScalarExpression :: Unary { expr, .. } => self . visit_column_agg_expr ( expr) ?,
130- ScalarExpression :: Alias { expr, .. } => self . visit_column_agg_expr ( expr) ?,
131- ScalarExpression :: Binary {
132- left_expr,
133- right_expr,
134- ..
135- } => {
136- self . visit_column_agg_expr ( left_expr) ?;
137- self . visit_column_agg_expr ( right_expr) ?;
138- }
139- ScalarExpression :: In { expr, args, .. } => {
140- self . visit_column_agg_expr ( expr) ?;
141- for arg in args {
142- self . visit_column_agg_expr ( arg) ?;
143- }
144- }
145- ScalarExpression :: Between {
146- expr,
147- left_expr,
148- right_expr,
149- ..
150- } => {
151- self . visit_column_agg_expr ( expr) ?;
152- self . visit_column_agg_expr ( left_expr) ?;
153- self . visit_column_agg_expr ( right_expr) ?;
154- }
155- ScalarExpression :: SubString {
156- expr,
157- for_expr,
158- from_expr,
159- } => {
160- self . visit_column_agg_expr ( expr) ?;
161- if let Some ( expr) = for_expr {
162- self . visit_column_agg_expr ( expr) ?;
163- }
164- if let Some ( expr) = from_expr {
165- self . visit_column_agg_expr ( expr) ?;
166- }
167- }
168- ScalarExpression :: Position { expr, in_expr } => {
169- self . visit_column_agg_expr ( expr) ?;
170- self . visit_column_agg_expr ( in_expr) ?;
171- }
172- ScalarExpression :: Trim {
173- expr,
174- trim_what_expr,
175- ..
176- } => {
177- self . visit_column_agg_expr ( expr) ?;
178- if let Some ( trim_what_expr) = trim_what_expr {
179- self . visit_column_agg_expr ( trim_what_expr) ?;
180- }
181- }
182- ScalarExpression :: Constant ( _) | ScalarExpression :: ColumnRef { .. } => ( ) ,
183- ScalarExpression :: Empty => unreachable ! ( ) ,
184- ScalarExpression :: Tuple ( args)
185- | ScalarExpression :: ScalaFunction ( ScalarFunction { args, .. } )
186- | ScalarExpression :: Coalesce { exprs : args, .. } => {
187- for expr in args {
188- self . visit_column_agg_expr ( expr) ?;
189- }
190- }
191- ScalarExpression :: If {
192- condition,
193- left_expr,
194- right_expr,
195- ..
196- } => {
197- self . visit_column_agg_expr ( condition) ?;
198- self . visit_column_agg_expr ( left_expr) ?;
199- self . visit_column_agg_expr ( right_expr) ?;
200- }
201- ScalarExpression :: IfNull {
202- left_expr,
203- right_expr,
204- ..
205- }
206- | ScalarExpression :: NullIf {
207- left_expr,
208- right_expr,
209- ..
210- } => {
211- self . visit_column_agg_expr ( left_expr) ?;
212- self . visit_column_agg_expr ( right_expr) ?;
213- }
214- ScalarExpression :: CaseWhen {
215- operand_expr,
216- expr_pairs,
217- else_expr,
218- ..
219- } => {
220- if let Some ( expr) = operand_expr {
221- self . visit_column_agg_expr ( expr) ?;
222- }
223- for ( expr_1, expr_2) in expr_pairs {
224- self . visit_column_agg_expr ( expr_1) ?;
225- self . visit_column_agg_expr ( expr_2) ?;
226- }
227- if let Some ( expr) = else_expr {
228- self . visit_column_agg_expr ( expr) ?;
229- }
230- }
231- ScalarExpression :: TableFunction ( _) => unreachable ! ( ) ,
137+ fn collect_aggregate_calls ( & mut self , expr : & ScalarExpression ) -> Result < ( ) , DatabaseError > {
138+ AggregateCallCollector {
139+ agg_calls : & mut self . context . agg_calls ,
232140 }
233-
234- Ok ( ( ) )
141+ . visit ( expr)
235142 }
236143
237144 /// Validate select exprs must appear in the GROUP BY clause or be used in
@@ -269,6 +176,10 @@ impl<T: Transaction, A: AsRef<[(&'static str, DataValue)]>> Binder<'_, '_, T, A>
269176 HashSet :: from_iter ( group_raw_exprs. iter ( ) . copied ( ) ) ;
270177
271178 for expr in select_items {
179+ if expr. has_window_call ( ) ? {
180+ HavingOrderByValidator :: new ( groupby, & self . context . agg_calls ) . visit ( expr) ?;
181+ continue ;
182+ }
272183 if expr. has_agg_call ( ) ? {
273184 continue ;
274185 }
0 commit comments