krishiv_sql/
pipe_syntax.rs1use std::fmt;
30
31#[derive(Debug)]
33pub enum PipeSyntaxError {
34 InvalidSyntax(String),
36 UnsupportedOperator(String),
38}
39
40impl fmt::Display for PipeSyntaxError {
41 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
42 match self {
43 Self::InvalidSyntax(msg) => write!(f, "invalid pipe syntax: {msg}"),
44 Self::UnsupportedOperator(msg) => write!(f, "unsupported pipe operator: {msg}"),
45 }
46 }
47}
48
49impl std::error::Error for PipeSyntaxError {}
50
51pub fn process_pipe_syntax(sql: &str) -> Result<String, PipeSyntaxError> {
55 let trimmed = sql.trim();
56
57 if !trimmed.to_uppercase().starts_with("FROM ") || !trimmed.contains("|>") {
59 return Ok(trimmed.to_string());
60 }
61
62 let parts: Vec<&str> = trimmed.split("|>").collect();
66 let Some((from_clause, stages)) = parts.split_first() else {
67 return Ok(trimmed.to_string());
68 };
69 if stages.is_empty() {
70 return Ok(trimmed.to_string());
71 }
72
73 let from_clause = from_clause.trim();
75 if !from_clause.to_uppercase().starts_with("FROM ") {
76 return Err(PipeSyntaxError::InvalidSyntax(
77 "pipe syntax must start with FROM".into(),
78 ));
79 }
80
81 let table_name = from_clause[5..].trim();
82
83 let mut predicates: Vec<String> = Vec::new();
94 let mut having: Vec<String> = Vec::new();
95 let mut select_clause: Option<String> = None;
96 let mut group_by_clause: Option<String> = None;
97 let mut order_by_clause: Option<String> = None;
98 let mut limit_clause: Option<String> = None;
99 let mut join_clauses = Vec::new();
100
101 fn set_once(
103 slot: &mut Option<String>,
104 value: String,
105 operator: &str,
106 ) -> Result<(), PipeSyntaxError> {
107 if slot.is_some() {
108 return Err(PipeSyntaxError::UnsupportedOperator(format!(
109 "repeated |> {operator}: a pipeline with more than one {operator} stage cannot \
110 be expressed as a single SELECT; write it as a subquery"
111 )));
112 }
113 *slot = Some(value);
114 Ok(())
115 }
116
117 for part in stages {
118 let part = part.trim();
119 let upper = part.to_uppercase();
120
121 if let Some(rest) = strip_kw(part, &upper, "WHERE ") {
122 if group_by_clause.is_some() {
124 having.push(rest.to_string());
125 } else {
126 predicates.push(rest.to_string());
127 }
128 } else if let Some(rest) = strip_kw(part, &upper, "SELECT ") {
129 set_once(&mut select_clause, rest.to_string(), "SELECT")?;
130 } else if let Some(rest) = strip_kw(part, &upper, "GROUP BY ") {
131 set_once(&mut group_by_clause, rest.to_string(), "GROUP BY")?;
132 } else if let Some(rest) = strip_kw(part, &upper, "ORDER BY ") {
133 set_once(&mut order_by_clause, rest.to_string(), "ORDER BY")?;
134 } else if let Some(rest) = strip_kw(part, &upper, "LIMIT ") {
135 set_once(&mut limit_clause, rest.to_string(), "LIMIT")?;
136 } else if upper.starts_with("JOIN ")
137 || upper.starts_with("INNER JOIN ")
138 || upper.starts_with("LEFT JOIN ")
139 || upper.starts_with("RIGHT JOIN ")
140 || upper.starts_with("CROSS JOIN ")
141 {
142 join_clauses.push(part.to_string());
143 } else {
144 return Err(PipeSyntaxError::UnsupportedOperator(part.to_string()));
145 }
146 }
147
148 let mut sql = String::new();
150 match &select_clause {
151 Some(projection) => {
152 sql.push_str("SELECT ");
153 sql.push_str(projection);
154 }
155 None => sql.push_str("SELECT *"),
156 }
157
158 sql.push_str(" FROM ");
159 sql.push_str(table_name);
160
161 for join in &join_clauses {
162 sql.push(' ');
163 sql.push_str(join);
164 }
165
166 if !predicates.is_empty() {
169 sql.push_str(" WHERE ");
170 sql.push_str(&join_predicates(&predicates));
171 }
172 if let Some(group_by) = &group_by_clause {
173 sql.push_str(" GROUP BY ");
174 sql.push_str(group_by);
175 }
176 if !having.is_empty() {
177 sql.push_str(" HAVING ");
178 sql.push_str(&join_predicates(&having));
179 }
180 if let Some(order_by) = &order_by_clause {
181 sql.push_str(" ORDER BY ");
182 sql.push_str(order_by);
183 }
184 if let Some(limit) = &limit_clause {
185 sql.push_str(" LIMIT ");
186 sql.push_str(limit);
187 }
188
189 Ok(sql)
190}
191
192fn strip_kw<'a>(part: &'a str, upper: &str, keyword: &str) -> Option<&'a str> {
196 upper
197 .starts_with(keyword)
198 .then(|| part.get(keyword.len()..))
199 .flatten()
200 .map(str::trim)
201 .filter(|rest| !rest.is_empty())
202}
203
204fn join_predicates(parts: &[String]) -> String {
207 if let [single] = parts {
208 return single.clone();
209 }
210 parts
211 .iter()
212 .map(|p| format!("({p})"))
213 .collect::<Vec<_>>()
214 .join(" AND ")
215}
216
217pub fn has_pipe_syntax(sql: &str) -> bool {
219 let trimmed = sql.trim();
220 trimmed.to_uppercase().starts_with("FROM ") && trimmed.contains("|>")
221}
222
223#[cfg(test)]
224mod tests {
225 use super::*;
226
227 #[test]
228 fn simple_pipe_syntax() {
229 let sql = "FROM orders |> WHERE amount > 100 |> SELECT customer_id, amount";
230 let result = process_pipe_syntax(sql).unwrap();
231 assert_eq!(
232 result,
233 "SELECT customer_id, amount FROM orders WHERE amount > 100"
234 );
235 }
236
237 #[test]
238 fn pipe_syntax_with_group_by() {
239 let sql = "FROM orders |> GROUP BY region |> SELECT region, SUM(amount) as total";
240 let result = process_pipe_syntax(sql).unwrap();
241 assert_eq!(
242 result,
243 "SELECT region, SUM(amount) as total FROM orders GROUP BY region"
244 );
245 }
246
247 #[test]
248 fn pipe_syntax_with_order_by_and_limit() {
249 let sql = "FROM orders |> ORDER BY amount DESC |> LIMIT 10";
250 let result = process_pipe_syntax(sql).unwrap();
251 assert_eq!(result, "SELECT * FROM orders ORDER BY amount DESC LIMIT 10");
252 }
253
254 #[test]
255 fn pipe_syntax_with_join() {
256 let sql = "FROM orders |> JOIN customers ON orders.customer_id = customers.id |> SELECT *";
257 let result = process_pipe_syntax(sql).unwrap();
258 assert_eq!(
259 result,
260 "SELECT * FROM orders JOIN customers ON orders.customer_id = customers.id"
261 );
262 }
263
264 #[test]
268 fn successive_where_stages_are_conjoined_not_overwritten() {
269 let sql = "FROM orders |> WHERE amount > 100 |> WHERE region = 'US' |> SELECT id";
270 let result = process_pipe_syntax(sql).unwrap();
271 assert_eq!(
272 result,
273 "SELECT id FROM orders WHERE (amount > 100) AND (region = 'US')"
274 );
275 }
276
277 #[test]
279 fn conjoined_filters_are_parenthesised() {
280 let sql = "FROM t |> WHERE a OR b |> WHERE c";
281 let result = process_pipe_syntax(sql).unwrap();
282 assert_eq!(result, "SELECT * FROM t WHERE (a OR b) AND (c)");
283 }
284
285 #[test]
288 fn a_filter_after_group_by_becomes_having() {
289 let sql = "FROM orders |> GROUP BY region |> WHERE SUM(amount) > 500 \
290 |> SELECT region, SUM(amount)";
291 let result = process_pipe_syntax(sql).unwrap();
292 assert_eq!(
293 result,
294 "SELECT region, SUM(amount) FROM orders GROUP BY region HAVING SUM(amount) > 500"
295 );
296 }
297
298 #[test]
300 fn a_filter_before_group_by_stays_a_where() {
301 let sql = "FROM orders |> WHERE amount > 0 |> GROUP BY region";
302 let result = process_pipe_syntax(sql).unwrap();
303 assert_eq!(result, "SELECT * FROM orders WHERE amount > 0 GROUP BY region");
304 }
305
306 #[test]
309 fn repeated_unflattenable_stages_are_rejected() {
310 for sql in [
311 "FROM t |> SELECT a |> SELECT b",
312 "FROM t |> GROUP BY a |> GROUP BY b",
313 "FROM t |> ORDER BY a |> ORDER BY b",
314 "FROM t |> LIMIT 1 |> LIMIT 2",
315 ] {
316 let error = process_pipe_syntax(sql)
317 .expect_err("a repeated stage must be refused, not silently dropped");
318 assert!(
319 error.to_string().contains("repeated |>"),
320 "unexpected error for {sql:?}: {error}"
321 );
322 }
323 }
324
325 #[test]
326 fn standard_sql_unchanged() {
327 let sql = "SELECT * FROM orders WHERE amount > 100";
328 let result = process_pipe_syntax(sql).unwrap();
329 assert_eq!(result, sql);
330 }
331
332 #[tokio::test(flavor = "multi_thread", worker_threads = 2)]
336 async fn pipe_syntax_runs_through_the_engine() {
337 use arrow::array::Int64Array;
338 use arrow::datatypes::{DataType, Field, Schema};
339 use std::sync::Arc;
340
341 let engine = crate::SqlEngine::new();
342 let schema = Arc::new(Schema::new(vec![Field::new("amount", DataType::Int64, false)]));
343 let batch = arrow::record_batch::RecordBatch::try_new(
344 schema,
345 vec![Arc::new(Int64Array::from(vec![50_i64, 150, 250]))],
346 )
347 .unwrap();
348 engine
349 .register_record_batches("orders", vec![batch])
350 .await
351 .unwrap();
352
353 let batches = engine
354 .sql("FROM orders |> WHERE amount > 100 |> WHERE amount < 200 |> SELECT amount")
355 .await
356 .expect("piped query should plan")
357 .collect()
358 .await
359 .expect("piped query should execute");
360
361 let rows: usize = batches.iter().map(|b| b.num_rows()).sum();
362 assert_eq!(
363 rows, 1,
364 "both filters must apply: only amount=150 is >100 and <200"
365 );
366 }
367
368 #[test]
369 fn has_pipe_syntax_detection() {
370 assert!(has_pipe_syntax("FROM t |> SELECT *"));
371 assert!(!has_pipe_syntax("SELECT * FROM t"));
372 assert!(!has_pipe_syntax("FROM t"));
373 }
374}