Skip to main content

connectorx/
sql.rs

1use crate::errors::ConnectorXError;
2#[cfg(feature = "src_oracle")]
3use crate::sources::oracle::OracleDialect;
4use fehler::{throw, throws};
5use log::{debug, trace, warn};
6use sqlparser::ast::{
7    BinaryOperator, Expr, Function, FunctionArg, FunctionArgExpr, Ident, ObjectName, Query, Select,
8    SelectItem, SetExpr, Statement, TableAlias, TableFactor, TableWithJoins, Value,
9    WildcardAdditionalOptions,
10};
11use sqlparser::dialect::Dialect;
12use sqlparser::parser::Parser;
13#[cfg(feature = "src_oracle")]
14use std::any::Any;
15
16#[derive(Debug, Clone)]
17pub enum CXQuery<Q = String> {
18    Naked(Q),   // The query directly comes from the user
19    Wrapped(Q), // The user query is already wrapped in a subquery
20}
21
22impl<Q: std::fmt::Display> std::fmt::Display for CXQuery<Q> {
23    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
24        match self {
25            CXQuery::Naked(q) => write!(f, "{}", q),
26            CXQuery::Wrapped(q) => write!(f, "{}", q),
27        }
28    }
29}
30
31impl<Q: AsRef<str>> CXQuery<Q> {
32    pub fn as_str(&self) -> &str {
33        match self {
34            CXQuery::Naked(q) => q.as_ref(),
35            CXQuery::Wrapped(q) => q.as_ref(),
36        }
37    }
38}
39
40impl From<&str> for CXQuery {
41    fn from(s: &str) -> CXQuery<String> {
42        CXQuery::Naked(s.to_string())
43    }
44}
45
46impl From<&&str> for CXQuery {
47    fn from(s: &&str) -> CXQuery<String> {
48        CXQuery::Naked(s.to_string())
49    }
50}
51
52impl From<&String> for CXQuery {
53    fn from(s: &String) -> CXQuery {
54        CXQuery::Naked(s.clone())
55    }
56}
57
58impl From<&CXQuery> for CXQuery {
59    fn from(q: &CXQuery) -> CXQuery {
60        q.clone()
61    }
62}
63
64impl CXQuery<String> {
65    pub fn naked<Q: AsRef<str>>(q: Q) -> Self {
66        CXQuery::Naked(q.as_ref().to_string())
67    }
68}
69
70impl<Q: AsRef<str>> AsRef<str> for CXQuery<Q> {
71    fn as_ref(&self) -> &str {
72        match self {
73            CXQuery::Naked(q) => q.as_ref(),
74            CXQuery::Wrapped(q) => q.as_ref(),
75        }
76    }
77}
78
79impl<Q> CXQuery<Q> {
80    pub fn map<F, U>(&self, f: F) -> CXQuery<U>
81    where
82        F: Fn(&Q) -> U,
83    {
84        match self {
85            CXQuery::Naked(q) => CXQuery::Naked(f(q)),
86            CXQuery::Wrapped(q) => CXQuery::Wrapped(f(q)),
87        }
88    }
89}
90
91impl<Q, E> CXQuery<Result<Q, E>> {
92    pub fn result(self) -> Result<CXQuery<Q>, E> {
93        match self {
94            CXQuery::Naked(q) => q.map(CXQuery::Naked),
95            CXQuery::Wrapped(q) => q.map(CXQuery::Wrapped),
96        }
97    }
98}
99
100// wrap a query into a derived table
101fn wrap_query(
102    query: &mut Query,
103    projection: Vec<SelectItem>,
104    selection: Option<Expr>,
105    tmp_tab_name: &str,
106) -> Statement {
107    let with = query.with.clone();
108    query.with = None;
109    let alias = if tmp_tab_name.is_empty() {
110        None
111    } else {
112        Some(TableAlias {
113            name: Ident {
114                value: tmp_tab_name.into(),
115                quote_style: None,
116            },
117            columns: vec![],
118        })
119    };
120    Statement::Query(Box::new(Query {
121        with,
122        locks: vec![],
123        body: Box::new(SetExpr::Select(Box::new(Select {
124            distinct: None,
125            top: None,
126            projection,
127            from: vec![TableWithJoins {
128                relation: TableFactor::Derived {
129                    lateral: false,
130                    subquery: Box::new(query.clone()),
131                    alias,
132                },
133                joins: vec![],
134            }],
135            lateral_views: vec![],
136            selection,
137            group_by: vec![],
138            cluster_by: vec![],
139            distribute_by: vec![],
140            sort_by: vec![],
141            having: None,
142            into: None,
143            named_window: vec![],
144            qualify: None,
145        }))),
146        order_by: vec![],
147        limit: None,
148        offset: None,
149        fetch: None,
150    }))
151}
152
153trait StatementExt {
154    fn as_query(&self) -> Option<&Query>;
155}
156
157impl StatementExt for Statement {
158    fn as_query(&self) -> Option<&Query> {
159        match self {
160            Statement::Query(q) => Some(q),
161            _ => None,
162        }
163    }
164}
165
166trait QueryExt {
167    fn as_select_mut(&mut self) -> Option<&mut Select>;
168}
169
170impl QueryExt for Query {
171    fn as_select_mut(&mut self) -> Option<&mut Select> {
172        match *self.body {
173            SetExpr::Select(ref mut select) => Some(select),
174            _ => None,
175        }
176    }
177}
178
179#[throws(ConnectorXError)]
180pub fn count_query<T: Dialect>(sql: &CXQuery<String>, dialect: &T) -> CXQuery<String> {
181    trace!("Incoming query: {}", sql);
182
183    const COUNT_TMP_TAB_NAME: &str = "CXTMPTAB_COUNT";
184
185    #[allow(unused_mut)]
186    let mut table_alias = COUNT_TMP_TAB_NAME;
187
188    // HACK: Some dialect (e.g. Oracle) does not support "AS" for alias
189    #[cfg(feature = "src_oracle")]
190    if dialect.type_id() == (OracleDialect {}.type_id()) {
191        // table_alias = "";
192        return CXQuery::Wrapped(format!(
193            "SELECT COUNT(*) FROM ({}) {}",
194            sql.as_str(),
195            COUNT_TMP_TAB_NAME
196        ));
197    }
198
199    let tsql = match sql.map(|sql| Parser::parse_sql(dialect, sql)).result() {
200        Ok(ast) => {
201            let projection = vec![SelectItem::UnnamedExpr(Expr::Function(Function {
202                name: ObjectName(vec![Ident {
203                    value: "count".to_string(),
204                    quote_style: None,
205                }]),
206                args: vec![FunctionArg::Unnamed(FunctionArgExpr::Wildcard)],
207                over: None,
208                distinct: false,
209                order_by: vec![],
210                special: false,
211            }))];
212            let ast_count: Statement = match ast {
213                CXQuery::Naked(ast) => {
214                    if ast.len() != 1 {
215                        throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
216                    }
217                    let mut query = ast[0]
218                        .as_query()
219                        .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?
220                        .clone();
221                    if query.offset.is_none() {
222                        query.order_by = vec![]; // mssql offset must appear with order by
223                    }
224                    let select = query
225                        .as_select_mut()
226                        .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?;
227                    select.sort_by = vec![];
228                    wrap_query(&mut query, projection, None, table_alias)
229                }
230                CXQuery::Wrapped(ast) => {
231                    if ast.len() != 1 {
232                        throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
233                    }
234                    let mut query = ast[0]
235                        .as_query()
236                        .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?
237                        .clone();
238                    let select = query
239                        .as_select_mut()
240                        .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?;
241                    select.projection = projection;
242                    Statement::Query(Box::new(query))
243                }
244            };
245            format!("{}", ast_count)
246        }
247        Err(e) => {
248            warn!("parser error: {:?}, manually compose query string", e);
249            format!(
250                "SELECT COUNT(*) FROM ({}) as {}",
251                sql.as_str(),
252                COUNT_TMP_TAB_NAME
253            )
254        }
255    };
256
257    debug!("Transformed count query: {}", tsql);
258    CXQuery::Wrapped(tsql)
259}
260
261#[throws(ConnectorXError)]
262pub fn limit0_query<T: Dialect>(sql: &CXQuery<String>, dialect: &T) -> CXQuery<String> {
263    trace!("Incoming query: {}", sql);
264
265    let sql = match Parser::parse_sql(dialect, sql.as_str()) {
266        Ok(mut ast) => {
267            if ast.len() != 1 {
268                throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
269            }
270
271            match &mut ast[0] {
272                Statement::Query(q) => {
273                    q.limit = Some(Expr::Value(Value::Number("0".to_string(), false)));
274                }
275                _ => throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string())),
276            };
277
278            format!("{}", ast[0])
279        }
280        Err(e) => {
281            warn!("parser error: {:?}, manually compose query string", e);
282            format!("{} LIMIT 0", sql.as_str())
283        }
284    };
285
286    debug!("Transformed limit 0 query: {}", sql);
287    CXQuery::Wrapped(sql)
288}
289
290/// Like limit0_query but with LIMIT 1. Used by SQLite where schema inference
291/// requires at least one row when decl_type is not available.
292#[throws(ConnectorXError)]
293pub fn limit1_query<T: Dialect>(sql: &CXQuery<String>, dialect: &T) -> CXQuery<String> {
294    trace!("Incoming query: {}", sql);
295
296    let sql = match Parser::parse_sql(dialect, sql.as_str()) {
297        Ok(mut ast) => {
298            if ast.len() != 1 {
299                throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
300            }
301
302            match &mut ast[0] {
303                Statement::Query(q) => {
304                    q.limit = Some(Expr::Value(Value::Number("1".to_string(), false)));
305                }
306                _ => throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string())),
307            };
308
309            format!("{}", ast[0])
310        }
311        Err(e) => {
312            warn!("parser error: {:?}, manually compose query string", e);
313            format!("{} LIMIT 1", sql.as_str())
314        }
315    };
316
317    debug!("Transformed limit 1 query: {}", sql);
318    CXQuery::Wrapped(sql)
319}
320
321#[throws(ConnectorXError)]
322#[cfg(feature = "src_oracle")]
323pub fn limit0_query_oracle(sql: &CXQuery<String>) -> CXQuery<String> {
324    trace!("Incoming oracle query: {}", sql);
325
326    CXQuery::Wrapped(format!("SELECT * FROM ({}) WHERE 1=0", sql))
327
328    // let ast = Parser::parse_sql(&OracleDialect {}, sql.as_str())?;
329    // if ast.len() != 1 {
330    //     throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
331    // }
332    // let ast_part: Statement;
333    // let mut query = ast[0]
334    //     .as_query()
335    //     .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?
336    //     .clone();
337
338    // let selection = Expr::BinaryOp {
339    //     left: Box::new(Expr::CompoundIdentifier(vec![Ident {
340    //         value: "rownum".to_string(),
341    //         quote_style: None,
342    //     }])),
343    //     op: BinaryOperator::Eq,
344    //     right: Box::new(Expr::Value(Value::Number("1".to_string(), false))),
345    // };
346    // ast_part = wrap_query(&mut query, vec![SelectItem::Wildcard], Some(selection), "");
347
348    // let tsql = format!("{}", ast_part);
349    // debug!("Transformed limit 1 query: {}", tsql);
350    // CXQuery::Wrapped(tsql)
351}
352
353#[throws(ConnectorXError)]
354pub fn single_col_partition_query<T: Dialect>(
355    sql: &str,
356    col: &str,
357    lower: i64,
358    upper: i64,
359    dialect: &T,
360) -> String {
361    trace!("Incoming query: {}", sql);
362    const PART_TMP_TAB_NAME: &str = "CXTMPTAB_PART";
363
364    #[allow(unused_mut)]
365    let mut table_alias = PART_TMP_TAB_NAME;
366    #[allow(unused_mut)]
367    let mut cid = Box::new(Expr::CompoundIdentifier(vec![
368        Ident {
369            value: PART_TMP_TAB_NAME.to_string(),
370            quote_style: None,
371        },
372        Ident {
373            value: col.to_string(),
374            quote_style: None,
375        },
376    ]));
377
378    // HACK: Some dialect (e.g. Oracle) does not support "AS" for alias
379    #[cfg(feature = "src_oracle")]
380    if dialect.type_id() == (OracleDialect {}.type_id()) {
381        return format!("SELECT * FROM ({}) CXTMPTAB_PART WHERE CXTMPTAB_PART.{} >= {} AND CXTMPTAB_PART.{} < {}", sql, col, lower, col, upper);
382        // table_alias = "";
383        // cid = Box::new(Expr::Identifier(Ident {
384        //     value: col.to_string(),
385        //     quote_style: None,
386        // }));
387    }
388
389    let tsql = match Parser::parse_sql(dialect, sql) {
390        Ok(ast) => {
391            if ast.len() != 1 {
392                throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
393            }
394
395            let mut query = ast[0]
396                .as_query()
397                .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?
398                .clone();
399
400            let select = query
401                .as_select_mut()
402                .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?
403                .clone();
404
405            let ast_part: Statement;
406
407            let lb = Expr::BinaryOp {
408                left: Box::new(Expr::Value(Value::Number(lower.to_string(), false))),
409                op: BinaryOperator::LtEq,
410                right: cid.clone(),
411            };
412
413            let ub = Expr::BinaryOp {
414                left: cid,
415                op: BinaryOperator::Lt,
416                right: Box::new(Expr::Value(Value::Number(upper.to_string(), false))),
417            };
418
419            let selection = Expr::BinaryOp {
420                left: Box::new(lb),
421                op: BinaryOperator::And,
422                right: Box::new(ub),
423            };
424
425            if query.limit.is_none() && select.top.is_none() && !query.order_by.is_empty() {
426                // order by in a partition query does not make sense because partition is unordered.
427                // clear the order by beceause mssql does not support order by in a derived table.
428                // also order by in the derived table does not make any difference.
429                query.order_by.clear();
430            }
431
432            ast_part = wrap_query(
433                &mut query,
434                vec![SelectItem::Wildcard(WildcardAdditionalOptions::default())],
435                Some(selection),
436                table_alias,
437            );
438            format!("{}", ast_part)
439        }
440        Err(e) => {
441            warn!("parser error: {:?}, manually compose query string", e);
442            format!("SELECT * FROM ({}) AS CXTMPTAB_PART WHERE CXTMPTAB_PART.{} >= {} AND CXTMPTAB_PART.{} < {}", sql, col, lower, col, upper)
443        }
444    };
445
446    debug!("Transformed single column partition query: {}", tsql);
447    tsql
448}
449
450#[throws(ConnectorXError)]
451pub fn get_partition_range_query<T: Dialect>(sql: &str, col: &str, dialect: &T) -> String {
452    trace!("Incoming query: {}", sql);
453    const RANGE_TMP_TAB_NAME: &str = "CXTMPTAB_RANGE";
454
455    #[allow(unused_mut)]
456    let mut table_alias = RANGE_TMP_TAB_NAME;
457    #[allow(unused_mut)]
458    let mut args = vec![FunctionArg::Unnamed(FunctionArgExpr::Expr(
459        Expr::CompoundIdentifier(vec![
460            Ident {
461                value: RANGE_TMP_TAB_NAME.to_string(),
462                quote_style: None,
463            },
464            Ident {
465                value: col.to_string(),
466                quote_style: None,
467            },
468        ]),
469    ))];
470
471    // HACK: Some dialect (e.g. Oracle) does not support "AS" for alias
472    #[cfg(feature = "src_oracle")]
473    if dialect.type_id() == (OracleDialect {}.type_id()) {
474        return format!(
475            "SELECT MIN({}.{}) as min, MAX({}.{}) as max FROM ({}) {}",
476            RANGE_TMP_TAB_NAME, col, RANGE_TMP_TAB_NAME, col, sql, RANGE_TMP_TAB_NAME
477        );
478        // table_alias = "";
479        // args = vec![FunctionArg::Unnamed(Expr::Identifier(Ident {
480        //     value: col.to_string(),
481        //     quote_style: None,
482        // }))];
483    }
484
485    let tsql = match Parser::parse_sql(dialect, sql) {
486        Ok(ast) => {
487            if ast.len() != 1 {
488                throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
489            }
490
491            let mut query = ast[0]
492                .as_query()
493                .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?
494                .clone();
495            let ast_range: Statement;
496
497            if query.limit.is_none() && query.offset.is_none() {
498                query.order_by = vec![]; // only omit orderby when there is no limit and offset in the query
499            }
500            let projection = vec![
501                SelectItem::UnnamedExpr(Expr::Function(Function {
502                    name: ObjectName(vec![Ident {
503                        value: "min".to_string(),
504                        quote_style: None,
505                    }]),
506                    args: args.clone(),
507                    over: None,
508                    distinct: false,
509                    order_by: vec![],
510                    special: false,
511                })),
512                SelectItem::UnnamedExpr(Expr::Function(Function {
513                    name: ObjectName(vec![Ident {
514                        value: "max".to_string(),
515                        quote_style: None,
516                    }]),
517                    args,
518                    over: None,
519                    distinct: false,
520                    order_by: vec![],
521                    special: false,
522                })),
523            ];
524            ast_range = wrap_query(&mut query, projection, None, table_alias);
525            format!("{}", ast_range)
526        }
527        Err(e) => {
528            warn!("parser error: {:?}, manually compose query string", e);
529            format!(
530                "SELECT MIN({}.{}) as min, MAX({}.{}) as max FROM ({}) AS {}",
531                RANGE_TMP_TAB_NAME, col, RANGE_TMP_TAB_NAME, col, sql, RANGE_TMP_TAB_NAME
532            )
533        }
534    };
535
536    debug!("Transformed partition range query: {}", tsql);
537    tsql
538}
539
540#[throws(ConnectorXError)]
541pub fn get_partition_range_query_sep<T: Dialect>(
542    sql: &str,
543    col: &str,
544    dialect: &T,
545) -> (String, String) {
546    trace!("Incoming query: {}", sql);
547    const RANGE_TMP_TAB_NAME: &str = "CXTMPTAB_RANGE";
548
549    let (sql_min, sql_max) = match Parser::parse_sql(dialect, sql) {
550        Ok(ast) => {
551            if ast.len() != 1 {
552                throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
553            }
554
555            let mut query = ast[0]
556                .as_query()
557                .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?
558                .clone();
559
560            let ast_range_min: Statement;
561            let ast_range_max: Statement;
562
563            query.order_by = vec![];
564            let min_proj = vec![SelectItem::UnnamedExpr(Expr::Function(Function {
565                name: ObjectName(vec![Ident {
566                    value: "min".to_string(),
567                    quote_style: None,
568                }]),
569                args: vec![FunctionArg::Unnamed(FunctionArgExpr::Expr(
570                    Expr::CompoundIdentifier(vec![
571                        Ident {
572                            value: RANGE_TMP_TAB_NAME.to_string(),
573                            quote_style: None,
574                        },
575                        Ident {
576                            value: col.to_string(),
577                            quote_style: None,
578                        },
579                    ]),
580                ))],
581                over: None,
582                distinct: false,
583                order_by: vec![],
584                special: false,
585            }))];
586            let max_proj = vec![SelectItem::UnnamedExpr(Expr::Function(Function {
587                name: ObjectName(vec![Ident {
588                    value: "max".to_string(),
589                    quote_style: None,
590                }]),
591                args: vec![FunctionArg::Unnamed(FunctionArgExpr::Expr(
592                    Expr::CompoundIdentifier(vec![
593                        Ident {
594                            value: RANGE_TMP_TAB_NAME.into(),
595                            quote_style: None,
596                        },
597                        Ident {
598                            value: col.into(),
599                            quote_style: None,
600                        },
601                    ]),
602                ))],
603                over: None,
604                distinct: false,
605                order_by: vec![],
606                special: false,
607            }))];
608            ast_range_min = wrap_query(&mut query.clone(), min_proj, None, RANGE_TMP_TAB_NAME);
609            ast_range_max = wrap_query(&mut query, max_proj, None, RANGE_TMP_TAB_NAME);
610            (format!("{}", ast_range_min), format!("{}", ast_range_max))
611        }
612        Err(e) => {
613            warn!("parser error: {:?}, manually compose query string", e);
614            (
615                format!(
616                    "SELECT MIN({}.{}) as min FROM ({}) AS {}",
617                    RANGE_TMP_TAB_NAME, col, sql, RANGE_TMP_TAB_NAME
618                ),
619                format!(
620                    "SELECT MAX({}.{}) as max FROM ({}) AS {}",
621                    RANGE_TMP_TAB_NAME, col, sql, RANGE_TMP_TAB_NAME
622                ),
623            )
624        }
625    };
626    debug!(
627        "Transformed separated partition range query: {}, {}",
628        sql_min, sql_max
629    );
630    (sql_min, sql_max)
631}