laminar_sql/parser/window_rewriter/
mod.rs1use sqlparser::ast::{
9 Expr, FunctionArg, FunctionArgExpr, FunctionArguments, Ident, Query, Select, SelectItem,
10 SetExpr, Statement,
11};
12
13use super::{ParseError, WindowFunction};
14
15pub struct WindowRewriter;
17
18impl WindowRewriter {
19 pub fn rewrite_statement(stmt: &mut Statement) -> Result<(), ParseError> {
41 if let Statement::Query(query) = stmt {
42 Self::rewrite_query(query)?;
43 }
44 Ok(())
45 }
46
47 fn rewrite_query(query: &mut Query) -> Result<(), ParseError> {
49 if let SetExpr::Select(select) = &mut *query.body {
50 Self::rewrite_select(select)?;
51 }
52 Ok(())
53 }
54
55 fn rewrite_select(select: &mut Select) -> Result<(), ParseError> {
60 let window_func = Self::find_window_in_group_by(select)?;
62
63 if let Some(_window) = window_func {
64 Self::ensure_window_columns_in_projection(select);
66 }
67
68 Ok(())
69 }
70
71 fn find_window_in_group_by(select: &Select) -> Result<Option<WindowFunction>, ParseError> {
73 match &select.group_by {
75 sqlparser::ast::GroupByExpr::Expressions(exprs, _modifiers) => {
76 for expr in exprs {
77 if let Some(window) = Self::extract_window_function(expr)? {
78 return Ok(Some(window));
79 }
80 }
81 }
82 sqlparser::ast::GroupByExpr::All(_) => {}
83 }
84 Ok(None)
85 }
86
87 fn ensure_window_columns_in_projection(select: &mut Select) {
89 let has_window_start = Self::has_projection_column(select, "window_start");
90 let has_window_end = Self::has_projection_column(select, "window_end");
91
92 if !has_window_start {
94 select.projection.insert(
95 0,
96 SelectItem::UnnamedExpr(Expr::Identifier(Ident::new("window_start"))),
97 );
98 }
99
100 if !has_window_end {
102 select.projection.insert(
103 1,
104 SelectItem::UnnamedExpr(Expr::Identifier(Ident::new("window_end"))),
105 );
106 }
107 }
108
109 fn has_projection_column(select: &Select, name: &str) -> bool {
111 select.projection.iter().any(|item| {
112 if let SelectItem::UnnamedExpr(Expr::Identifier(ident)) = item {
113 ident.value.eq_ignore_ascii_case(name)
114 } else if let SelectItem::ExprWithAlias { alias, .. } = item {
115 alias.value.eq_ignore_ascii_case(name)
116 } else {
117 false
118 }
119 })
120 }
121
122 #[must_use]
124 pub fn contains_window_function(expr: &Expr) -> bool {
125 match expr {
126 Expr::Function(func) => {
127 if let Some(name) = func.name.0.last() {
128 let func_name = name.to_string().to_uppercase();
129 matches!(
130 func_name.as_str(),
131 "TUMBLE" | "HOP" | "SLIDE" | "SESSION" | "CUMULATE"
132 )
133 } else {
134 false
135 }
136 }
137 _ => false,
138 }
139 }
140
141 pub fn extract_window_function(expr: &Expr) -> Result<Option<WindowFunction>, ParseError> {
159 match expr {
160 Expr::Function(func) => {
161 let name =
162 func.name.0.last().ok_or_else(|| {
163 ParseError::WindowError("Empty function name".to_string())
164 })?;
165
166 let func_name = name.to_string().to_uppercase();
167
168 let args = Self::extract_function_args(&func.args)?;
170
171 match func_name.as_str() {
172 "TUMBLE" => Self::parse_tumble_args(&args),
173 "HOP" | "SLIDE" => Self::parse_hop_args(&args),
174 "SESSION" => Self::parse_session_args(&args),
175 "CUMULATE" => Self::parse_cumulate_args(&args),
176 _ => Ok(None),
177 }
178 }
179 _ => Ok(None),
180 }
181 }
182
183 fn extract_function_args(args: &FunctionArguments) -> Result<Vec<Expr>, ParseError> {
185 match args {
186 FunctionArguments::List(arg_list) => {
187 let mut result = Vec::new();
188 for arg in &arg_list.args {
189 if let Some(expr) = Self::extract_arg_expr(arg) {
190 result.push(expr);
191 }
192 }
193 Ok(result)
194 }
195 FunctionArguments::None => Ok(vec![]),
196 FunctionArguments::Subquery(_) => Err(ParseError::WindowError(
197 "Subquery arguments not supported for window functions".to_string(),
198 )),
199 }
200 }
201
202 fn extract_arg_expr(arg: &FunctionArg) -> Option<Expr> {
204 match arg {
205 FunctionArg::Unnamed(arg_expr) => match arg_expr {
206 FunctionArgExpr::Expr(expr) => Some(expr.clone()),
207 FunctionArgExpr::Wildcard | FunctionArgExpr::QualifiedWildcard(_) => None,
208 },
209 FunctionArg::Named { arg, .. } | FunctionArg::ExprNamed { arg, .. } => match arg {
210 FunctionArgExpr::Expr(expr) => Some(expr.clone()),
211 FunctionArgExpr::Wildcard | FunctionArgExpr::QualifiedWildcard(_) => None,
212 },
213 }
214 }
215
216 fn parse_tumble_args(args: &[Expr]) -> Result<Option<WindowFunction>, ParseError> {
218 if args.len() < 2 || args.len() > 3 {
219 return Err(ParseError::WindowError(format!(
220 "TUMBLE requires 2-3 arguments (time_column, interval [, offset]), got {}",
221 args.len()
222 )));
223 }
224
225 Ok(Some(WindowFunction::Tumble {
226 time_column: Box::new(args[0].clone()),
227 interval: Box::new(args[1].clone()),
228 offset: args.get(2).map(|e| Box::new(e.clone())),
229 }))
230 }
231
232 fn parse_hop_args(args: &[Expr]) -> Result<Option<WindowFunction>, ParseError> {
234 if args.len() < 3 || args.len() > 4 {
235 return Err(ParseError::WindowError(format!(
236 "HOP/SLIDE requires 3-4 arguments (time_column, slide_interval, window_size [, offset]), got {}",
237 args.len()
238 )));
239 }
240
241 Ok(Some(WindowFunction::Hop {
242 time_column: Box::new(args[0].clone()),
243 slide_interval: Box::new(args[1].clone()),
244 window_interval: Box::new(args[2].clone()),
245 offset: args.get(3).map(|e| Box::new(e.clone())),
246 }))
247 }
248
249 fn parse_session_args(args: &[Expr]) -> Result<Option<WindowFunction>, ParseError> {
251 if args.len() != 2 {
252 return Err(ParseError::WindowError(format!(
253 "SESSION requires 2 arguments (time_column, gap_interval), got {}",
254 args.len()
255 )));
256 }
257
258 Ok(Some(WindowFunction::Session {
259 time_column: Box::new(args[0].clone()),
260 gap_interval: Box::new(args[1].clone()),
261 }))
262 }
263
264 fn parse_cumulate_args(args: &[Expr]) -> Result<Option<WindowFunction>, ParseError> {
266 if args.len() != 3 {
267 return Err(ParseError::WindowError(format!(
268 "CUMULATE requires 3 arguments (time_column, step_interval, max_size_interval), got {}",
269 args.len()
270 )));
271 }
272
273 Ok(Some(WindowFunction::Cumulate {
274 time_column: Box::new(args[0].clone()),
275 step_interval: Box::new(args[1].clone()),
276 max_size_interval: Box::new(args[2].clone()),
277 }))
278 }
279
280 #[must_use]
284 pub fn get_time_column_name(window: &WindowFunction) -> Option<String> {
285 let expr = match window {
286 WindowFunction::Tumble { time_column, .. }
287 | WindowFunction::Hop { time_column, .. }
288 | WindowFunction::Session { time_column, .. }
289 | WindowFunction::Cumulate { time_column, .. } => time_column.as_ref(),
290 };
291
292 match expr {
293 Expr::Identifier(ident) => Some(ident.value.clone()),
294 Expr::CompoundIdentifier(parts) => parts.last().map(|p| p.value.clone()),
295 _ => None,
296 }
297 }
298
299 pub fn parse_interval_to_duration(expr: &Expr) -> Result<std::time::Duration, ParseError> {
307 match expr {
308 Expr::Interval(interval) => {
309 let value = Self::extract_interval_value(&interval.value)?;
311
312 let unit = interval
314 .leading_field
315 .clone()
316 .unwrap_or(sqlparser::ast::DateTimeField::Second);
317
318 match unit {
319 sqlparser::ast::DateTimeField::Millisecond
320 | sqlparser::ast::DateTimeField::Milliseconds => {
321 return Ok(std::time::Duration::from_millis(value));
322 }
323 _ => {}
324 }
325
326 let seconds =
327 match unit {
328 sqlparser::ast::DateTimeField::Second
329 | sqlparser::ast::DateTimeField::Seconds => value,
330 sqlparser::ast::DateTimeField::Minute
331 | sqlparser::ast::DateTimeField::Minutes => value * 60,
332 sqlparser::ast::DateTimeField::Hour
333 | sqlparser::ast::DateTimeField::Hours => value * 3600,
334 sqlparser::ast::DateTimeField::Day
335 | sqlparser::ast::DateTimeField::Days => value * 86400,
336 _ => {
337 return Err(ParseError::WindowError(format!(
338 "Unsupported interval unit: {unit:?}"
339 )))
340 }
341 };
342
343 Ok(std::time::Duration::from_secs(seconds))
344 }
345 Expr::Value(value_with_span) => {
347 use sqlparser::ast::Value;
348 if let Value::SingleQuotedString(s) = &value_with_span.value {
349 Self::parse_interval_string(s)
350 } else {
351 Err(ParseError::WindowError(format!(
352 "Expected string value, got: {value_with_span:?}"
353 )))
354 }
355 }
356 Expr::Identifier(ident) => Self::parse_interval_string(&ident.value),
358 _ => Err(ParseError::WindowError(format!(
359 "Expected INTERVAL expression, got: {expr:?}"
360 ))),
361 }
362 }
363
364 fn extract_interval_value(expr: &Expr) -> Result<u64, ParseError> {
366 match expr {
367 Expr::Value(value_with_span) => {
368 use sqlparser::ast::Value;
369 match &value_with_span.value {
370 Value::Number(n, _) => n.parse::<u64>().map_err(|_| {
371 ParseError::WindowError(format!("Invalid interval value: {n}"))
372 }),
373 Value::SingleQuotedString(s) => {
374 let num_str = s.split_whitespace().next().unwrap_or(s);
376 num_str.parse::<u64>().map_err(|_| {
377 ParseError::WindowError(format!("Invalid interval value: {s}"))
378 })
379 }
380 _ => Err(ParseError::WindowError(format!(
381 "Unsupported value type in interval: {value_with_span:?}"
382 ))),
383 }
384 }
385 _ => Err(ParseError::WindowError(format!(
386 "Cannot extract interval value from: {expr:?}"
387 ))),
388 }
389 }
390
391 fn parse_interval_string(s: &str) -> Result<std::time::Duration, ParseError> {
393 let parts: Vec<&str> = s.split_whitespace().collect();
394 if parts.is_empty() {
395 return Err(ParseError::WindowError("Empty interval string".to_string()));
396 }
397
398 let value: u64 = parts[0].parse().map_err(|_| {
399 ParseError::WindowError(format!("Invalid interval value: {}", parts[0]))
400 })?;
401
402 let unit = if parts.len() > 1 {
403 parts[1].to_uppercase()
404 } else {
405 "SECOND".to_string()
406 };
407
408 if matches!(unit.as_str(), "MILLISECOND" | "MILLISECONDS" | "MS") {
409 return Ok(std::time::Duration::from_millis(value));
410 }
411
412 let seconds = match unit.as_str() {
413 "SECOND" | "SECONDS" | "S" => value,
414 "MINUTE" | "MINUTES" | "M" => value * 60,
415 "HOUR" | "HOURS" | "H" => value * 3600,
416 "DAY" | "DAYS" | "D" => value * 86400,
417 _ => {
418 return Err(ParseError::WindowError(format!(
419 "Unsupported interval unit: {unit}"
420 )))
421 }
422 };
423
424 Ok(std::time::Duration::from_secs(seconds))
425 }
426}
427
428#[cfg(test)]
429mod tests;