Skip to main content

connectorx/
get_arrow.rs

1#[cfg(feature = "src_mysql")]
2use crate::sources::mysql::{BinaryProtocol as MySQLBinaryProtocol, TextProtocol};
3#[cfg(feature = "src_postgres")]
4use crate::sources::postgres::{
5    rewrite_tls_args, BinaryProtocol as PgBinaryProtocol, CSVProtocol, CursorProtocol,
6    SimpleProtocol,
7};
8use crate::{
9    arrow_batch_iter::{ArrowBatchIter, RecordBatchIterator},
10    prelude::*,
11    sql::CXQuery,
12};
13use fehler::{throw, throws};
14use log::debug;
15#[cfg(feature = "src_postgres")]
16use postgres::NoTls;
17#[cfg(feature = "src_postgres")]
18use postgres_openssl::MakeTlsConnector;
19#[allow(unused_imports)]
20use std::sync::Arc;
21
22#[allow(unreachable_code, unreachable_patterns, unused_variables, unused_mut)]
23#[throws(ConnectorXOutError)]
24pub fn get_arrow(
25    source_conn: &SourceConn,
26    origin_query: Option<String>,
27    queries: &[CXQuery<String>],
28    pre_execution_queries: Option<&[String]>,
29) -> ArrowDestination {
30    let mut destination = ArrowDestination::new();
31    let protocol = source_conn.proto.as_str();
32    debug!("Protocol: {}", protocol);
33
34    match source_conn.ty {
35        #[cfg(feature = "src_postgres")]
36        SourceType::Postgres => {
37            let (config, tls) = rewrite_tls_args(&source_conn.conn)?;
38            match (protocol, tls) {
39                ("csv", Some(tls_conn)) => {
40                    let source = PostgresSource::<CSVProtocol, MakeTlsConnector>::new(
41                        config,
42                        tls_conn,
43                        queries.len(),
44                    )?;
45                    let mut dispatcher = Dispatcher::<
46                        _,
47                        _,
48                        PostgresArrowTransport<CSVProtocol, MakeTlsConnector>,
49                    >::new(
50                        source, &mut destination, queries, origin_query
51                    );
52                    dispatcher.set_pre_execution_queries(pre_execution_queries);
53                    dispatcher.run()?;
54                }
55                ("csv", None) => {
56                    let source =
57                        PostgresSource::<CSVProtocol, NoTls>::new(config, NoTls, queries.len())?;
58                    let mut dispatcher = Dispatcher::<
59                        _,
60                        _,
61                        PostgresArrowTransport<CSVProtocol, NoTls>,
62                    >::new(
63                        source, &mut destination, queries, origin_query
64                    );
65                    dispatcher.set_pre_execution_queries(pre_execution_queries);
66                    dispatcher.run()?;
67                }
68                ("binary", Some(tls_conn)) => {
69                    let source = PostgresSource::<PgBinaryProtocol, MakeTlsConnector>::new(
70                        config,
71                        tls_conn,
72                        queries.len(),
73                    )?;
74                    let mut dispatcher = Dispatcher::<
75                        _,
76                        _,
77                        PostgresArrowTransport<PgBinaryProtocol, MakeTlsConnector>,
78                    >::new(
79                        source, &mut destination, queries, origin_query
80                    );
81                    dispatcher.set_pre_execution_queries(pre_execution_queries);
82                    dispatcher.run()?;
83                }
84                ("binary", None) => {
85                    let source = PostgresSource::<PgBinaryProtocol, NoTls>::new(
86                        config,
87                        NoTls,
88                        queries.len(),
89                    )?;
90                    let mut dispatcher = Dispatcher::<
91                        _,
92                        _,
93                        PostgresArrowTransport<PgBinaryProtocol, NoTls>,
94                    >::new(
95                        source, &mut destination, queries, origin_query
96                    );
97                    dispatcher.set_pre_execution_queries(pre_execution_queries);
98                    dispatcher.run()?;
99                }
100                ("cursor", Some(tls_conn)) => {
101                    let source = PostgresSource::<CursorProtocol, MakeTlsConnector>::new(
102                        config,
103                        tls_conn,
104                        queries.len(),
105                    )?;
106                    let mut dispatcher = Dispatcher::<
107                        _,
108                        _,
109                        PostgresArrowTransport<CursorProtocol, MakeTlsConnector>,
110                    >::new(
111                        source, &mut destination, queries, origin_query
112                    );
113                    dispatcher.set_pre_execution_queries(pre_execution_queries);
114                    dispatcher.run()?;
115                }
116                ("cursor", None) => {
117                    let source =
118                        PostgresSource::<CursorProtocol, NoTls>::new(config, NoTls, queries.len())?;
119                    let mut dispatcher = Dispatcher::<
120                        _,
121                        _,
122                        PostgresArrowTransport<CursorProtocol, NoTls>,
123                    >::new(
124                        source, &mut destination, queries, origin_query
125                    );
126                    dispatcher.set_pre_execution_queries(pre_execution_queries);
127                    dispatcher.run()?;
128                }
129                ("simple", Some(tls_conn)) => {
130                    let sb = PostgresSource::<SimpleProtocol, MakeTlsConnector>::new(
131                        config,
132                        tls_conn,
133                        queries.len(),
134                    )?;
135                    let mut dispatcher = Dispatcher::<
136                        _,
137                        _,
138                        PostgresArrowTransport<SimpleProtocol, MakeTlsConnector>,
139                    >::new(
140                        sb, &mut destination, queries, origin_query
141                    );
142                    debug!("Running dispatcher");
143                    dispatcher.set_pre_execution_queries(pre_execution_queries);
144                    dispatcher.run()?;
145                }
146                ("simple", None) => {
147                    let sb =
148                        PostgresSource::<SimpleProtocol, NoTls>::new(config, NoTls, queries.len())?;
149                    let mut dispatcher = Dispatcher::<
150                        _,
151                        _,
152                        PostgresArrowTransport<SimpleProtocol, NoTls>,
153                    >::new(
154                        sb, &mut destination, queries, origin_query
155                    );
156                    debug!("Running dispatcher");
157                    dispatcher.set_pre_execution_queries(pre_execution_queries);
158                    dispatcher.run()?;
159                }
160                _ => unimplemented!("{} protocol not supported", protocol),
161            }
162        }
163        #[cfg(feature = "src_mysql")]
164        SourceType::MySQL => match protocol {
165            "binary" => {
166                let source =
167                    MySQLSource::<MySQLBinaryProtocol>::new(&source_conn.conn[..], queries.len())?;
168                let mut dispatcher =
169                    Dispatcher::<_, _, MySQLArrowTransport<MySQLBinaryProtocol>>::new(
170                        source,
171                        &mut destination,
172                        queries,
173                        origin_query,
174                    );
175                dispatcher.set_pre_execution_queries(pre_execution_queries);
176                dispatcher.run()?;
177            }
178            "text" => {
179                let source =
180                    MySQLSource::<TextProtocol>::new(&source_conn.conn[..], queries.len())?;
181                let mut dispatcher = Dispatcher::<_, _, MySQLArrowTransport<TextProtocol>>::new(
182                    source,
183                    &mut destination,
184                    queries,
185                    origin_query,
186                );
187                dispatcher.set_pre_execution_queries(pre_execution_queries);
188                dispatcher.run()?;
189            }
190            _ => unimplemented!("{} protocol not supported", protocol),
191        },
192        #[cfg(feature = "src_sqlite")]
193        SourceType::SQLite => {
194            // remove the first "sqlite://" manually since url.path is not correct for windows
195            let path = &source_conn.conn.as_str()[9..];
196            let source = SQLiteSource::new(path, queries.len())?;
197            let dispatcher = Dispatcher::<_, _, SQLiteArrowTransport>::new(
198                source,
199                &mut destination,
200                queries,
201                origin_query,
202            );
203            dispatcher.run()?;
204        }
205        #[cfg(feature = "src_mssql_common")]
206        SourceType::MsSQL => {
207            let rt = Arc::new(tokio::runtime::Runtime::new().expect("Failed to create runtime"));
208            let source = MsSQLSource::new(rt, &source_conn.conn[..], queries.len())?;
209            let dispatcher = Dispatcher::<_, _, MsSQLArrowTransport>::new(
210                source,
211                &mut destination,
212                queries,
213                origin_query,
214            );
215            dispatcher.run()?;
216        }
217        #[cfg(feature = "src_oracle")]
218        SourceType::Oracle => {
219            let source = OracleSource::new(&source_conn.conn[..], queries.len())?;
220            let dispatcher = Dispatcher::<_, _, OracleArrowTransport>::new(
221                source,
222                &mut destination,
223                queries,
224                origin_query,
225            );
226            dispatcher.run()?;
227        }
228        #[cfg(feature = "src_bigquery")]
229        SourceType::BigQuery => {
230            let rt = Arc::new(tokio::runtime::Runtime::new().expect("Failed to create runtime"));
231            let source = BigQuerySource::new(rt, &source_conn.conn[..])?;
232            let dispatcher = Dispatcher::<_, _, BigQueryArrowTransport>::new(
233                source,
234                &mut destination,
235                queries,
236                origin_query,
237            );
238            dispatcher.run()?;
239        }
240        #[cfg(feature = "src_trino")]
241        SourceType::Trino => {
242            let rt = Arc::new(tokio::runtime::Runtime::new().expect("Failed to create runtime"));
243            let source = TrinoSource::new(rt, &source_conn.conn[..])?;
244            let dispatcher = Dispatcher::<_, _, TrinoArrowTransport>::new(
245                source,
246                &mut destination,
247                queries,
248                origin_query,
249            );
250            dispatcher.run()?;
251        }
252        #[cfg(feature = "src_clickhouse")]
253        SourceType::ClickHouse => {
254            let rt = Arc::new(tokio::runtime::Runtime::new().expect("Failed to create runtime"));
255            let source = ClickHouseSource::new(rt, &source_conn.conn[..])?;
256            let dispatcher = Dispatcher::<_, _, ClickHouseArrowTransport>::new(
257                source,
258                &mut destination,
259                queries,
260                origin_query,
261            );
262            dispatcher.run()?;
263        }
264        _ => throw!(ConnectorXOutError::SourceNotSupport(format!(
265            "{:?}",
266            source_conn.ty
267        ))),
268    }
269
270    destination
271}
272
273#[allow(unreachable_code, unreachable_patterns, unused_variables, unused_mut)]
274#[throws(ConnectorXOutError)]
275pub fn new_record_batch_iter(
276    source_conn: &SourceConn,
277    origin_query: Option<String>,
278    queries: &[CXQuery<String>],
279    batch_size: usize,
280    pre_execution_queries: Option<&[String]>,
281) -> Box<dyn RecordBatchIterator> {
282    let destination = ArrowStreamDestination::new_with_batch_size(batch_size);
283    let protocol = source_conn.proto.as_str();
284    debug!("Protocol: {}", protocol);
285
286    match source_conn.ty {
287        #[cfg(feature = "src_postgres")]
288        SourceType::Postgres => {
289            let (config, tls) = rewrite_tls_args(&source_conn.conn)?;
290            match (protocol, tls) {
291                ("csv", Some(tls_conn)) => {
292                    let mut source = PostgresSource::<CSVProtocol, MakeTlsConnector>::new(
293                        config,
294                        tls_conn,
295                        queries.len(),
296                    )?;
297
298                    source.set_pre_execution_queries(pre_execution_queries);
299
300                    let batch_iter =
301                        ArrowBatchIter::<
302                            _,
303                            PostgresArrowStreamTransport<CSVProtocol, MakeTlsConnector>,
304                        >::new(source, destination, origin_query, queries)?;
305                    let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
306                    return iter;
307                }
308                ("csv", None) => {
309                    let mut source =
310                        PostgresSource::<CSVProtocol, NoTls>::new(config, NoTls, queries.len())?;
311
312                    source.set_pre_execution_queries(pre_execution_queries);
313
314                    let batch_iter = ArrowBatchIter::<
315                        _,
316                        PostgresArrowStreamTransport<CSVProtocol, NoTls>,
317                    >::new(
318                        source, destination, origin_query, queries
319                    )?;
320                    let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
321                    return iter;
322                }
323                ("binary", Some(tls_conn)) => {
324                    let mut source = PostgresSource::<PgBinaryProtocol, MakeTlsConnector>::new(
325                        config,
326                        tls_conn,
327                        queries.len(),
328                    )?;
329
330                    source.set_pre_execution_queries(pre_execution_queries);
331
332                    let batch_iter =
333                        ArrowBatchIter::<
334                            _,
335                            PostgresArrowStreamTransport<PgBinaryProtocol, MakeTlsConnector>,
336                        >::new(source, destination, origin_query, queries)?;
337                    let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
338                    return iter;
339                }
340                ("binary", None) => {
341                    let mut source = PostgresSource::<PgBinaryProtocol, NoTls>::new(
342                        config,
343                        NoTls,
344                        queries.len(),
345                    )?;
346
347                    source.set_pre_execution_queries(pre_execution_queries);
348
349                    let batch_iter = ArrowBatchIter::<
350                        _,
351                        PostgresArrowStreamTransport<PgBinaryProtocol, NoTls>,
352                    >::new(
353                        source, destination, origin_query, queries
354                    )?;
355                    let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
356                    return iter;
357                }
358                ("cursor", Some(tls_conn)) => {
359                    let mut source = PostgresSource::<CursorProtocol, MakeTlsConnector>::new(
360                        config,
361                        tls_conn,
362                        queries.len(),
363                    )?;
364
365                    source.set_pre_execution_queries(pre_execution_queries);
366
367                    let batch_iter =
368                        ArrowBatchIter::<
369                            _,
370                            PostgresArrowStreamTransport<CursorProtocol, MakeTlsConnector>,
371                        >::new(source, destination, origin_query, queries)?;
372                    let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
373                    return iter;
374                }
375                ("cursor", None) => {
376                    let mut source =
377                        PostgresSource::<CursorProtocol, NoTls>::new(config, NoTls, queries.len())?;
378
379                    source.set_pre_execution_queries(pre_execution_queries);
380
381                    let batch_iter = ArrowBatchIter::<
382                        _,
383                        PostgresArrowStreamTransport<CursorProtocol, NoTls>,
384                    >::new(
385                        source, destination, origin_query, queries
386                    )?;
387                    let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
388                    return iter;
389                }
390                _ => unimplemented!("{} protocol not supported", protocol),
391            }
392        }
393        #[cfg(feature = "src_mysql")]
394        SourceType::MySQL => match protocol {
395            "binary" => {
396                let mut source =
397                    MySQLSource::<MySQLBinaryProtocol>::new(&source_conn.conn[..], queries.len())?;
398
399                source.set_pre_execution_queries(pre_execution_queries);
400
401                let batch_iter =
402                    ArrowBatchIter::<_, MySQLArrowStreamTransport<MySQLBinaryProtocol>>::new(
403                        source,
404                        destination,
405                        origin_query,
406                        queries,
407                    )?;
408                let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
409                return iter;
410            }
411            "text" => {
412                let mut source =
413                    MySQLSource::<TextProtocol>::new(&source_conn.conn[..], queries.len())?;
414
415                source.set_pre_execution_queries(pre_execution_queries);
416
417                let batch_iter = ArrowBatchIter::<_, MySQLArrowStreamTransport<TextProtocol>>::new(
418                    source,
419                    destination,
420                    origin_query,
421                    queries,
422                )?;
423                let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
424                return iter;
425            }
426            _ => unimplemented!("{} protocol not supported", protocol),
427        },
428        #[cfg(feature = "src_sqlite")]
429        SourceType::SQLite => {
430            // remove the first "sqlite://" manually since url.path is not correct for windows
431            let path = &source_conn.conn.as_str()[9..];
432            let source = SQLiteSource::new(path, queries.len())?;
433            let batch_iter = ArrowBatchIter::<_, SQLiteArrowStreamTransport>::new(
434                source,
435                destination,
436                origin_query,
437                queries,
438            )?;
439            let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
440            return iter;
441        }
442        #[cfg(feature = "src_mssql_common")]
443        SourceType::MsSQL => {
444            let rt = Arc::new(tokio::runtime::Runtime::new().expect("Failed to create runtime"));
445            let source = MsSQLSource::new(rt, &source_conn.conn[..], queries.len())?;
446            let batch_iter = ArrowBatchIter::<_, MsSQLArrowStreamTransport>::new(
447                source,
448                destination,
449                origin_query,
450                queries,
451            )?;
452            let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
453            return iter;
454        }
455        #[cfg(feature = "src_oracle")]
456        SourceType::Oracle => {
457            let source = OracleSource::new(&source_conn.conn[..], queries.len())?;
458            let batch_iter = ArrowBatchIter::<_, OracleArrowStreamTransport>::new(
459                source,
460                destination,
461                origin_query,
462                queries,
463            )?;
464            let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
465            return iter;
466        }
467        #[cfg(feature = "src_bigquery")]
468        SourceType::BigQuery => {
469            let rt = Arc::new(tokio::runtime::Runtime::new().expect("Failed to create runtime"));
470            let source = BigQuerySource::new(rt, &source_conn.conn[..])?;
471            let batch_iter = ArrowBatchIter::<_, BigQueryArrowStreamTransport>::new(
472                source,
473                destination,
474                origin_query,
475                queries,
476            )?;
477            let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
478            return iter;
479        }
480        #[cfg(feature = "src_clickhouse")]
481        SourceType::ClickHouse => {
482            let rt = Arc::new(tokio::runtime::Runtime::new().expect("Failed to create runtime"));
483            let source = ClickHouseSource::new(rt, &source_conn.conn[..])?;
484            let batch_iter = ArrowBatchIter::<_, ClickHouseArrowStreamTransport>::new(
485                source,
486                destination,
487                origin_query,
488                queries,
489            )?;
490            let iter: Box<dyn RecordBatchIterator> = Box::new(batch_iter);
491            return iter;
492        }
493        _ => throw!(ConnectorXOutError::SourceNotSupport(format!(
494            "{:?}",
495            source_conn.ty
496        ))),
497    }
498}