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), Wrapped(Q), }
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
100fn 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 #[cfg(feature = "src_oracle")]
190 if dialect.type_id() == (OracleDialect {}.type_id()) {
191 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![]; }
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#[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 }
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 #[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 }
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 lb = Expr::BinaryOp {
406 left: Box::new(Expr::Value(Value::Number(lower.to_string(), false))),
407 op: BinaryOperator::LtEq,
408 right: cid.clone(),
409 };
410
411 let ub = Expr::BinaryOp {
412 left: cid,
413 op: BinaryOperator::Lt,
414 right: Box::new(Expr::Value(Value::Number(upper.to_string(), false))),
415 };
416
417 let selection = Expr::BinaryOp {
418 left: Box::new(lb),
419 op: BinaryOperator::And,
420 right: Box::new(ub),
421 };
422
423 if query.limit.is_none() && select.top.is_none() && !query.order_by.is_empty() {
424 query.order_by.clear();
428 }
429
430 let ast_part: Statement = wrap_query(
431 &mut query,
432 vec![SelectItem::Wildcard(WildcardAdditionalOptions::default())],
433 Some(selection),
434 table_alias,
435 );
436 format!("{}", ast_part)
437 }
438 Err(e) => {
439 warn!("parser error: {:?}, manually compose query string", e);
440 format!("SELECT * FROM ({}) AS CXTMPTAB_PART WHERE CXTMPTAB_PART.{} >= {} AND CXTMPTAB_PART.{} < {}", sql, col, lower, col, upper)
441 }
442 };
443
444 debug!("Transformed single column partition query: {}", tsql);
445 tsql
446}
447
448#[throws(ConnectorXError)]
449pub fn get_partition_range_query<T: Dialect>(sql: &str, col: &str, dialect: &T) -> String {
450 trace!("Incoming query: {}", sql);
451 const RANGE_TMP_TAB_NAME: &str = "CXTMPTAB_RANGE";
452
453 #[allow(unused_mut)]
454 let mut table_alias = RANGE_TMP_TAB_NAME;
455 #[allow(unused_mut)]
456 let mut args = vec![FunctionArg::Unnamed(FunctionArgExpr::Expr(
457 Expr::CompoundIdentifier(vec![
458 Ident {
459 value: RANGE_TMP_TAB_NAME.to_string(),
460 quote_style: None,
461 },
462 Ident {
463 value: col.to_string(),
464 quote_style: None,
465 },
466 ]),
467 ))];
468
469 #[cfg(feature = "src_oracle")]
471 if dialect.type_id() == (OracleDialect {}.type_id()) {
472 return format!(
473 "SELECT MIN({}.{}) as min, MAX({}.{}) as max FROM ({}) {}",
474 RANGE_TMP_TAB_NAME, col, RANGE_TMP_TAB_NAME, col, sql, RANGE_TMP_TAB_NAME
475 );
476 }
482
483 let tsql = match Parser::parse_sql(dialect, sql) {
484 Ok(ast) => {
485 if ast.len() != 1 {
486 throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
487 }
488
489 let mut query = ast[0]
490 .as_query()
491 .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?
492 .clone();
493
494 if query.limit.is_none() && query.offset.is_none() {
495 query.order_by = vec![]; }
497 let projection = vec![
498 SelectItem::UnnamedExpr(Expr::Function(Function {
499 name: ObjectName(vec![Ident {
500 value: "min".to_string(),
501 quote_style: None,
502 }]),
503 args: args.clone(),
504 over: None,
505 distinct: false,
506 order_by: vec![],
507 special: false,
508 })),
509 SelectItem::UnnamedExpr(Expr::Function(Function {
510 name: ObjectName(vec![Ident {
511 value: "max".to_string(),
512 quote_style: None,
513 }]),
514 args,
515 over: None,
516 distinct: false,
517 order_by: vec![],
518 special: false,
519 })),
520 ];
521 let ast_range: Statement = wrap_query(&mut query, projection, None, table_alias);
522 format!("{}", ast_range)
523 }
524 Err(e) => {
525 warn!("parser error: {:?}, manually compose query string", e);
526 format!(
527 "SELECT MIN({}.{}) as min, MAX({}.{}) as max FROM ({}) AS {}",
528 RANGE_TMP_TAB_NAME, col, RANGE_TMP_TAB_NAME, col, sql, RANGE_TMP_TAB_NAME
529 )
530 }
531 };
532
533 debug!("Transformed partition range query: {}", tsql);
534 tsql
535}
536
537#[throws(ConnectorXError)]
538pub fn get_partition_range_query_sep<T: Dialect>(
539 sql: &str,
540 col: &str,
541 dialect: &T,
542) -> (String, String) {
543 trace!("Incoming query: {}", sql);
544 const RANGE_TMP_TAB_NAME: &str = "CXTMPTAB_RANGE";
545
546 let (sql_min, sql_max) = match Parser::parse_sql(dialect, sql) {
547 Ok(ast) => {
548 if ast.len() != 1 {
549 throw!(ConnectorXError::SqlQueryNotSupported(sql.to_string()));
550 }
551
552 let mut query = ast[0]
553 .as_query()
554 .ok_or_else(|| ConnectorXError::SqlQueryNotSupported(sql.to_string()))?
555 .clone();
556
557 query.order_by = vec![];
558 let min_proj = vec![SelectItem::UnnamedExpr(Expr::Function(Function {
559 name: ObjectName(vec![Ident {
560 value: "min".to_string(),
561 quote_style: None,
562 }]),
563 args: vec![FunctionArg::Unnamed(FunctionArgExpr::Expr(
564 Expr::CompoundIdentifier(vec![
565 Ident {
566 value: RANGE_TMP_TAB_NAME.to_string(),
567 quote_style: None,
568 },
569 Ident {
570 value: col.to_string(),
571 quote_style: None,
572 },
573 ]),
574 ))],
575 over: None,
576 distinct: false,
577 order_by: vec![],
578 special: false,
579 }))];
580 let max_proj = vec![SelectItem::UnnamedExpr(Expr::Function(Function {
581 name: ObjectName(vec![Ident {
582 value: "max".to_string(),
583 quote_style: None,
584 }]),
585 args: vec![FunctionArg::Unnamed(FunctionArgExpr::Expr(
586 Expr::CompoundIdentifier(vec![
587 Ident {
588 value: RANGE_TMP_TAB_NAME.into(),
589 quote_style: None,
590 },
591 Ident {
592 value: col.into(),
593 quote_style: None,
594 },
595 ]),
596 ))],
597 over: None,
598 distinct: false,
599 order_by: vec![],
600 special: false,
601 }))];
602 let ast_range_min: Statement =
603 wrap_query(&mut query.clone(), min_proj, None, RANGE_TMP_TAB_NAME);
604 let ast_range_max: Statement =
605 wrap_query(&mut query, max_proj, None, RANGE_TMP_TAB_NAME);
606 (format!("{}", ast_range_min), format!("{}", ast_range_max))
607 }
608 Err(e) => {
609 warn!("parser error: {:?}, manually compose query string", e);
610 (
611 format!(
612 "SELECT MIN({}.{}) as min FROM ({}) AS {}",
613 RANGE_TMP_TAB_NAME, col, sql, RANGE_TMP_TAB_NAME
614 ),
615 format!(
616 "SELECT MAX({}.{}) as max FROM ({}) AS {}",
617 RANGE_TMP_TAB_NAME, col, sql, RANGE_TMP_TAB_NAME
618 ),
619 )
620 }
621 };
622 debug!(
623 "Transformed separated partition range query: {}, {}",
624 sql_min, sql_max
625 );
626 (sql_min, sql_max)
627}
628
629#[cfg(test)]
630mod tests {
631 use super::{
632 count_query, get_partition_range_query, get_partition_range_query_sep, limit0_query,
633 limit1_query, single_col_partition_query, CXQuery,
634 };
635 use crate::errors::ConnectorXError;
636 use sqlparser::dialect::{GenericDialect, MsSqlDialect};
637
638 #[test]
639 fn preserves_query_variant_and_maps_values() {
640 let naked = CXQuery::naked("select 1");
641 assert_eq!(naked.as_str(), "select 1");
642 assert!(matches!(naked.map(|query| query.len()), CXQuery::Naked(8)));
643 assert_eq!(
644 CXQuery::Wrapped("select 1".to_string()).to_string(),
645 "select 1"
646 );
647 }
648
649 #[test]
650 fn unwraps_successful_and_failed_query_results() {
651 let naked: CXQuery<Result<&str, &str>> = CXQuery::Naked(Ok("select 1"));
652 assert_eq!(naked.result().unwrap().as_str(), "select 1");
653
654 let wrapped: CXQuery<Result<&str, &str>> = CXQuery::Wrapped(Err("bad"));
655 assert!(matches!(wrapped.result(), Err("bad")));
656 }
657
658 #[test]
659 fn transforms_count_and_limit_queries() {
660 let query = CXQuery::naked("SELECT id FROM items ORDER BY id");
661 let dialect = GenericDialect {};
662 assert!(count_query(&query, &dialect)
663 .unwrap()
664 .as_str()
665 .to_ascii_lowercase()
666 .contains("count"));
667 assert!(limit0_query(&query, &dialect)
668 .unwrap()
669 .as_str()
670 .contains("LIMIT 0"));
671 assert!(limit1_query(&query, &dialect)
672 .unwrap()
673 .as_str()
674 .contains("LIMIT 1"));
675 }
676
677 #[test]
678 fn uses_fallback_for_unparseable_limit_query() {
679 let query = CXQuery::naked("SELECT FROM");
680 assert_eq!(
681 limit0_query(&query, &GenericDialect {}).unwrap().as_str(),
682 "SELECT FROM LIMIT 0"
683 );
684 assert_eq!(
685 limit1_query(&query, &GenericDialect {}).unwrap().as_str(),
686 "SELECT FROM LIMIT 1"
687 );
688 }
689
690 #[test]
691 fn rejects_multiple_and_non_select_queries() {
692 let dialect = GenericDialect {};
693 assert!(matches!(
694 count_query(&CXQuery::naked("SELECT 1; SELECT 2"), &dialect),
695 Err(ConnectorXError::SqlQueryNotSupported(_))
696 ));
697 assert!(matches!(
698 limit0_query(&CXQuery::naked("INSERT INTO items VALUES (1)"), &dialect),
699 Err(ConnectorXError::SqlQueryNotSupported(_))
700 ));
701 }
702
703 #[test]
704 fn creates_partition_and_range_queries() {
705 let dialect = GenericDialect {};
706 let partition = single_col_partition_query(
707 "SELECT id, value FROM items ORDER BY id",
708 "id",
709 10,
710 20,
711 &dialect,
712 )
713 .unwrap();
714 assert!(partition.contains("CXTMPTAB_PART"));
715 assert!(partition.contains("10"));
716 assert!(partition.contains("20"));
717 assert!(!partition.contains("ORDER BY id"));
718
719 let range =
720 get_partition_range_query("SELECT id FROM items ORDER BY id", "id", &dialect).unwrap();
721 let range_lower = range.to_ascii_lowercase();
722 assert!(range_lower.contains("min"));
723 assert!(range_lower.contains("max"));
724 assert!(range.contains("CXTMPTAB_RANGE"));
725
726 let (min, max) =
727 get_partition_range_query_sep("SELECT id FROM items ORDER BY id", "id", &dialect)
728 .unwrap();
729 assert!(min.to_ascii_lowercase().contains("min"));
730 assert!(min.contains("CXTMPTAB_RANGE"));
731 assert!(max.to_ascii_lowercase().contains("max"));
732 assert!(max.contains("CXTMPTAB_RANGE"));
733 }
734
735 #[test]
736 fn clears_mssql_ordering_only_when_safe() {
737 let query = single_col_partition_query(
738 "SELECT id FROM items ORDER BY id",
739 "id",
740 0,
741 1,
742 &MsSqlDialect {},
743 )
744 .unwrap();
745 assert!(!query.contains("ORDER BY id"));
746 }
747}