1use std::convert::TryFrom;
2use std::sync::Arc;
3
4use crate::errors::{ConnectorXOutError, OutResult};
5use crate::source_router::{SourceConn, SourceType};
6#[cfg(feature = "src_bigquery")]
7use crate::sources::bigquery::BigQueryDialect;
8#[cfg(feature = "src_clickhouse")]
9use crate::sources::clickhouse::{ClickHouseSource, ClickHouseSourceError};
10#[cfg(feature = "src_mssql_tiberius")]
11use crate::sources::mssql::{mssql_config, FloatN, IntN, MsSQLTypeSystem};
12#[cfg(feature = "src_mysql")]
13use crate::sources::mysql::{build_opts, MySQLTypeSystem};
14#[cfg(feature = "src_oracle")]
15use crate::sources::oracle::{OracleDialect, OracleSource};
16#[cfg(feature = "src_postgres")]
17use crate::sources::postgres::{rewrite_tls_args, PostgresTypeSystem};
18#[cfg(feature = "src_trino")]
19use crate::sources::trino::TrinoDialect;
20#[cfg(feature = "src_sqlite")]
21use crate::sql::get_partition_range_query_sep;
22use crate::sql::{get_partition_range_query, single_col_partition_query, CXQuery};
23use anyhow::anyhow;
24use fehler::{throw, throws};
25#[cfg(feature = "src_bigquery")]
26use gcp_bigquery_client;
27#[cfg(feature = "src_mysql")]
28use r2d2_mysql::mysql::{prelude::Queryable, Pool, Row};
29#[cfg(feature = "src_sqlite")]
30use rusqlite::{types::Type, Connection};
31#[cfg(feature = "src_postgres")]
32use rust_decimal::{prelude::ToPrimitive, Decimal};
33#[cfg(feature = "src_postgres")]
34use rust_decimal_macros::dec;
35#[cfg(feature = "src_clickhouse")]
36use serde::Deserialize;
37#[cfg(feature = "src_clickhouse")]
38use serde_json::Value as JsonValue;
39#[cfg(feature = "src_clickhouse")]
40use sqlparser::dialect::ClickHouseDialect;
41#[cfg(feature = "src_mssql_common")]
42use sqlparser::dialect::MsSqlDialect;
43#[cfg(feature = "src_mysql")]
44use sqlparser::dialect::MySqlDialect;
45#[cfg(feature = "src_postgres")]
46use sqlparser::dialect::PostgreSqlDialect;
47#[cfg(feature = "src_sqlite")]
48use sqlparser::dialect::SQLiteDialect;
49#[cfg(feature = "src_mssql_tiberius")]
50use tiberius::Client;
51#[cfg(any(
52 feature = "src_bigquery",
53 feature = "src_mssql_tiberius",
54 feature = "src_trino"
55))]
56use tokio::{net::TcpStream, runtime::Runtime};
57#[cfg(feature = "src_mssql_tiberius")]
58use tokio_util::compat::TokioAsyncWriteCompatExt;
59use url::Url;
60
61pub struct PartitionQuery {
62 query: String,
63 column: String,
64 min: Option<i64>,
65 max: Option<i64>,
66 num: usize,
67}
68
69impl PartitionQuery {
70 pub fn new(query: &str, column: &str, min: Option<i64>, max: Option<i64>, num: usize) -> Self {
71 Self {
72 query: query.into(),
73 column: column.into(),
74 min,
75 max,
76 num,
77 }
78 }
79}
80
81pub fn partition(part: &PartitionQuery, source_conn: &SourceConn) -> OutResult<Vec<CXQuery>> {
82 let mut queries = vec![];
83
84 if part.num == 0 {
85 throw!(anyhow!("partition count (num) must be greater than zero"));
86 }
87
88 let num = i64::try_from(part.num).map_err(|_| {
89 anyhow!(
90 "partition count (num) is too large to represent safely: {}",
91 part.num
92 )
93 })?;
94
95 let (min, max) = match (part.min, part.max) {
96 (None, None) => get_col_range(source_conn, &part.query, &part.column)?,
97 (Some(min), Some(max)) => (min, max),
98 _ => throw!(anyhow!(
99 "partition_query range can not be partially specified",
100 )),
101 };
102
103 if max < min {
104 throw!(anyhow!(
105 "partition range is invalid: max ({}) must be greater than or equal to min ({})",
106 max,
107 min
108 ));
109 }
110
111 let range_len = max
112 .checked_sub(min)
113 .and_then(|value| value.checked_add(1))
114 .ok_or_else(|| {
115 anyhow!(
116 "partition range overflow: min={}, max={} is too large",
117 min,
118 max
119 )
120 })?;
121
122 let partition_size = range_len / num;
123
124 let final_upper = max.checked_add(1).ok_or_else(|| {
125 anyhow!(
126 "partition upper bound overflow: max={} cannot be incremented safely",
127 max
128 )
129 })?;
130
131 for i in 0i64..num {
132 let lower = i
133 .checked_mul(partition_size)
134 .and_then(|offset| min.checked_add(offset))
135 .ok_or_else(|| {
136 anyhow!(
137 "partition lower bound overflow: min={}, step={}, partition_size={}",
138 min,
139 i,
140 partition_size
141 )
142 })?;
143 let upper = if i == num - 1 {
144 final_upper
145 } else {
146 (i + 1)
147 .checked_mul(partition_size)
148 .and_then(|offset| min.checked_add(offset))
149 .ok_or_else(|| {
150 anyhow!(
151 "partition upper bound overflow: min={}, step={}, partition_size={}",
152 min,
153 i + 1,
154 partition_size
155 )
156 })?
157 };
158 let partition_query = get_part_query(source_conn, &part.query, &part.column, lower, upper)?;
159 queries.push(partition_query);
160 }
161 Ok(queries)
162}
163
164pub fn get_col_range(source_conn: &SourceConn, query: &str, col: &str) -> OutResult<(i64, i64)> {
165 match source_conn.ty {
166 #[cfg(feature = "src_postgres")]
167 SourceType::Postgres => pg_get_partition_range(&source_conn.conn, query, col),
168 #[cfg(feature = "src_sqlite")]
169 SourceType::SQLite => sqlite_get_partition_range(&source_conn.conn, query, col),
170 #[cfg(feature = "src_mysql")]
171 SourceType::MySQL => mysql_get_partition_range(&source_conn.conn, query, col),
172 #[cfg(all(feature = "src_mssql_tiberius", not(feature = "src_mssql_tds")))]
173 SourceType::MsSQL => mssql_get_partition_range(&source_conn.conn, query, col),
174 #[cfg(all(feature = "src_mssql_tiberius", feature = "src_mssql_tds"))]
175 SourceType::MsSQL => match crate::sources::mssql::active_driver() {
176 crate::sources::mssql::MsSQLDriverKind::Tiberius => {
177 mssql_get_partition_range(&source_conn.conn, query, col)
178 }
179 crate::sources::mssql::MsSQLDriverKind::MssqlTds => Ok(
180 crate::sources::mssql::tds_get_partition_range(&source_conn.conn, query, col)?,
181 ),
182 },
183 #[cfg(all(feature = "src_mssql_tds", not(feature = "src_mssql_tiberius")))]
184 SourceType::MsSQL => Ok(crate::sources::mssql::tds_get_partition_range(
185 &source_conn.conn,
186 query,
187 col,
188 )?),
189 #[cfg(feature = "src_oracle")]
190 SourceType::Oracle => oracle_get_partition_range(&source_conn.conn, query, col),
191 #[cfg(feature = "src_bigquery")]
192 SourceType::BigQuery => bigquery_get_partition_range(&source_conn.conn, query, col),
193 #[cfg(feature = "src_trino")]
194 SourceType::Trino => trino_get_partition_range(&source_conn.conn, query, col),
195 #[cfg(feature = "src_clickhouse")]
196 SourceType::ClickHouse => clickhouse_get_partition_range(&source_conn.conn, query, col),
197 _ => unimplemented!("{:?} not implemented!", source_conn.ty),
198 }
199}
200
201#[throws(ConnectorXOutError)]
202pub fn get_part_query(
203 source_conn: &SourceConn,
204 query: &str,
205 col: &str,
206 lower: i64,
207 upper: i64,
208) -> CXQuery<String> {
209 let query = match source_conn.ty {
210 #[cfg(feature = "src_postgres")]
211 SourceType::Postgres => {
212 single_col_partition_query(query, col, lower, upper, &PostgreSqlDialect {})?
213 }
214 #[cfg(feature = "src_sqlite")]
215 SourceType::SQLite => {
216 single_col_partition_query(query, col, lower, upper, &SQLiteDialect {})?
217 }
218 #[cfg(feature = "src_mysql")]
219 SourceType::MySQL => {
220 single_col_partition_query(query, col, lower, upper, &MySqlDialect {})?
221 }
222 #[cfg(feature = "src_mssql_common")]
223 SourceType::MsSQL => {
224 single_col_partition_query(query, col, lower, upper, &MsSqlDialect {})?
225 }
226 #[cfg(feature = "src_oracle")]
227 SourceType::Oracle => {
228 single_col_partition_query(query, col, lower, upper, &OracleDialect {})?
229 }
230 #[cfg(feature = "src_bigquery")]
231 SourceType::BigQuery => {
232 single_col_partition_query(query, col, lower, upper, &BigQueryDialect {})?
233 }
234 #[cfg(feature = "src_trino")]
235 SourceType::Trino => {
236 single_col_partition_query(query, col, lower, upper, &TrinoDialect {})?
237 }
238 #[cfg(feature = "src_clickhouse")]
239 SourceType::ClickHouse => {
240 single_col_partition_query(query, col, lower, upper, &ClickHouseDialect {})?
241 }
242 _ => unimplemented!("{:?} not implemented!", source_conn.ty),
243 };
244 CXQuery::Wrapped(query)
245}
246
247#[cfg(feature = "src_postgres")]
248#[throws(ConnectorXOutError)]
249fn pg_get_partition_range(conn: &Url, query: &str, col: &str) -> (i64, i64) {
250 let (config, tls) = rewrite_tls_args(conn)?;
251 let mut client = match tls {
252 None => config.connect(postgres::NoTls)?,
253 Some(tls_conn) => config.connect(tls_conn)?,
254 };
255 let range_query = get_partition_range_query(query, col, &PostgreSqlDialect {})?;
256 let row = client.query_one(range_query.as_str(), &[])?;
257
258 let col_type = PostgresTypeSystem::from(row.columns()[0].type_());
259 let (min_v, max_v) = match col_type {
260 PostgresTypeSystem::Int2(_) => {
261 let min_v: Option<i16> = row.get(0);
262 let max_v: Option<i16> = row.get(1);
263 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
264 }
265 PostgresTypeSystem::Int4(_) => {
266 let min_v: Option<i32> = row.get(0);
267 let max_v: Option<i32> = row.get(1);
268 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
269 }
270 PostgresTypeSystem::Int8(_) => {
271 let min_v: Option<i64> = row.get(0);
272 let max_v: Option<i64> = row.get(1);
273 (min_v.unwrap_or(0), max_v.unwrap_or(0))
274 }
275 PostgresTypeSystem::Float4(_) => {
276 let min_v: Option<f32> = row.get(0);
277 let max_v: Option<f32> = row.get(1);
278 (min_v.unwrap_or(0.0) as i64, max_v.unwrap_or(0.0) as i64)
279 }
280 PostgresTypeSystem::Float8(_) => {
281 let min_v: Option<f64> = row.get(0);
282 let max_v: Option<f64> = row.get(1);
283 (min_v.unwrap_or(0.0) as i64, max_v.unwrap_or(0.0) as i64)
284 }
285 PostgresTypeSystem::Numeric(_) => {
286 let min_v: Option<Decimal> = row.get(0);
287 let max_v: Option<Decimal> = row.get(1);
288 (
289 min_v.unwrap_or(dec!(0.0)).to_i64().unwrap_or(0),
290 max_v.unwrap_or(dec!(0.0)).to_i64().unwrap_or(0),
291 )
292 }
293 _ => throw!(anyhow!(
294 "Partition can only be done on int or float columns"
295 )),
296 };
297
298 (min_v, max_v)
299}
300
301#[cfg(feature = "src_sqlite")]
302#[throws(ConnectorXOutError)]
303fn sqlite_get_partition_range(conn: &Url, query: &str, col: &str) -> (i64, i64) {
304 let conn = Connection::open(&conn.as_str()[9..])?;
306 let (min_query, max_query) = get_partition_range_query_sep(query, col, &SQLiteDialect {})?;
309 let mut error = None;
310 let min_v = conn.query_row(min_query.as_str(), [], |row| {
311 let col_type = row.get_ref(0)?.data_type();
313 match col_type {
314 Type::Integer => row.get(0),
315 Type::Real => {
316 let v: f64 = row.get(0)?;
317 Ok(v as i64)
318 }
319 Type::Null => Ok(0),
320 _ => {
321 error = Some(anyhow!("Partition can only be done on integer columns"));
322 Ok(0)
323 }
324 }
325 })?;
326 match error {
327 None => {}
328 Some(e) => throw!(e),
329 }
330 let max_v = conn.query_row(max_query.as_str(), [], |row| {
331 let col_type = row.get_ref(0)?.data_type();
332 match col_type {
333 Type::Integer => row.get(0),
334 Type::Real => {
335 let v: f64 = row.get(0)?;
336 Ok(v as i64)
337 }
338 Type::Null => Ok(0),
339 _ => {
340 error = Some(anyhow!("Partition can only be done on integer columns"));
341 Ok(0)
342 }
343 }
344 })?;
345 match error {
346 None => {}
347 Some(e) => throw!(e),
348 }
349
350 (min_v, max_v)
351}
352
353#[cfg(feature = "src_mysql")]
354#[throws(ConnectorXOutError)]
355fn mysql_get_partition_range(conn: &Url, query: &str, col: &str) -> (i64, i64) {
356 let pool = Pool::new(build_opts(conn.as_str())?)?;
357 let mut conn = pool.get_conn()?;
358 let range_query = get_partition_range_query(query, col, &MySqlDialect {})?;
359 let row: Row = conn
360 .query_first(range_query)?
361 .ok_or_else(|| anyhow!("mysql range: no row returns"))?;
362
363 let col_type = MySQLTypeSystem::from((
364 &row.columns()[0].column_type(),
365 &row.columns()[0].flags(),
366 row.columns()[0].character_set(),
367 ));
368
369 let (min_v, max_v) = match col_type {
370 MySQLTypeSystem::Tiny(_) => {
371 let min_v: Option<i8> = row
372 .get(0)
373 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
374 let max_v: Option<i8> = row
375 .get(1)
376 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
377 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
378 }
379 MySQLTypeSystem::Short(_) => {
380 let min_v: Option<i16> = row
381 .get(0)
382 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
383 let max_v: Option<i16> = row
384 .get(1)
385 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
386 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
387 }
388 MySQLTypeSystem::Int24(_) => {
389 let min_v: Option<i32> = row
390 .get(0)
391 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
392 let max_v: Option<i32> = row
393 .get(1)
394 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
395 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
396 }
397 MySQLTypeSystem::Long(_) => {
398 let min_v: Option<i64> = row
399 .get(0)
400 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
401 let max_v: Option<i64> = row
402 .get(1)
403 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
404 (min_v.unwrap_or(0), max_v.unwrap_or(0))
405 }
406 MySQLTypeSystem::LongLong(_) => {
407 let min_v: Option<i64> = row
408 .get(0)
409 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
410 let max_v: Option<i64> = row
411 .get(1)
412 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
413 (min_v.unwrap_or(0), max_v.unwrap_or(0))
414 }
415 MySQLTypeSystem::UTiny(_) => {
416 let min_v: Option<u8> = row
417 .get(0)
418 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
419 let max_v: Option<u8> = row
420 .get(1)
421 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
422 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
423 }
424 MySQLTypeSystem::UShort(_) => {
425 let min_v: Option<u16> = row
426 .get(0)
427 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
428 let max_v: Option<u16> = row
429 .get(1)
430 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
431 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
432 }
433 MySQLTypeSystem::UInt24(_) => {
434 let min_v: Option<u32> = row
435 .get(0)
436 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
437 let max_v: Option<u32> = row
438 .get(1)
439 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
440 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
441 }
442 MySQLTypeSystem::ULong(_) => {
443 let min_v: Option<u32> = row
444 .get(0)
445 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
446 let max_v: Option<u32> = row
447 .get(1)
448 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
449 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
450 }
451 MySQLTypeSystem::ULongLong(_) => {
452 let min_v: Option<u64> = row
453 .get(0)
454 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
455 let max_v: Option<u64> = row
456 .get(1)
457 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
458 (min_v.unwrap_or(0) as i64, max_v.unwrap_or(0) as i64)
459 }
460 MySQLTypeSystem::Float(_) => {
461 let min_v: Option<f32> = row
462 .get(0)
463 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
464 let max_v: Option<f32> = row
465 .get(1)
466 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
467 (min_v.unwrap_or(0.0) as i64, max_v.unwrap_or(0.0) as i64)
468 }
469 MySQLTypeSystem::Double(_) => {
470 let min_v: Option<f64> = row
471 .get(0)
472 .ok_or_else(|| anyhow!("mysql range: cannot get min value"))?;
473 let max_v: Option<f64> = row
474 .get(1)
475 .ok_or_else(|| anyhow!("mysql range: cannot get max value"))?;
476 (min_v.unwrap_or(0.0) as i64, max_v.unwrap_or(0.0) as i64)
477 }
478 _ => throw!(anyhow!("Partition can only be done on int columns")),
479 };
480
481 (min_v, max_v)
482}
483
484#[cfg(feature = "src_mssql_tiberius")]
485#[throws(ConnectorXOutError)]
486fn mssql_get_partition_range(conn: &Url, query: &str, col: &str) -> (i64, i64) {
487 let rt = Runtime::new().expect("Failed to create runtime");
488 let config = mssql_config(conn)?;
489 let tcp = rt.block_on(TcpStream::connect(config.get_addr()))?;
490 tcp.set_nodelay(true)?;
491
492 let mut client = rt.block_on(Client::connect(config, tcp.compat_write()))?;
493
494 let range_query = get_partition_range_query(query, col, &MsSqlDialect {})?;
495 let query_result = rt.block_on(client.query(range_query.as_str(), &[]))?;
496 let row = rt.block_on(query_result.into_row())?.unwrap();
497
498 let col_type = MsSQLTypeSystem::from(&row.columns()[0].column_type());
499 let (min_v, max_v) = match col_type {
500 MsSQLTypeSystem::Tinyint(_) => {
501 let min_v: u8 = row.get(0).unwrap_or(0);
502 let max_v: u8 = row.get(1).unwrap_or(0);
503 (min_v as i64, max_v as i64)
504 }
505 MsSQLTypeSystem::Smallint(_) => {
506 let min_v: i16 = row.get(0).unwrap_or(0);
507 let max_v: i16 = row.get(1).unwrap_or(0);
508 (min_v as i64, max_v as i64)
509 }
510 MsSQLTypeSystem::Int(_) => {
511 let min_v: i32 = row.get(0).unwrap_or(0);
512 let max_v: i32 = row.get(1).unwrap_or(0);
513 (min_v as i64, max_v as i64)
514 }
515 MsSQLTypeSystem::Bigint(_) => {
516 let min_v: i64 = row.get(0).unwrap_or(0);
517 let max_v: i64 = row.get(1).unwrap_or(0);
518 (min_v, max_v)
519 }
520 MsSQLTypeSystem::Intn(_) => {
521 let min_v: IntN = row.get(0).unwrap_or(IntN(0));
522 let max_v: IntN = row.get(1).unwrap_or(IntN(0));
523 (min_v.0, max_v.0)
524 }
525 MsSQLTypeSystem::Float24(_) => {
526 let min_v: f32 = row.get(0).unwrap_or(0.0);
527 let max_v: f32 = row.get(1).unwrap_or(0.0);
528 (min_v as i64, max_v as i64)
529 }
530 MsSQLTypeSystem::Float53(_) => {
531 let min_v: f64 = row.get(0).unwrap_or(0.0);
532 let max_v: f64 = row.get(1).unwrap_or(0.0);
533 (min_v as i64, max_v as i64)
534 }
535 MsSQLTypeSystem::Floatn(_) => {
536 let min_v: FloatN = row.get(0).unwrap_or(FloatN(0.0));
537 let max_v: FloatN = row.get(1).unwrap_or(FloatN(0.0));
538 (min_v.0 as i64, max_v.0 as i64)
539 }
540 _ => throw!(anyhow!(
541 "Partition can only be done on int or float columns"
542 )),
543 };
544
545 (min_v, max_v)
546}
547
548#[cfg(feature = "src_oracle")]
549#[throws(ConnectorXOutError)]
550fn oracle_get_partition_range(conn: &Url, query: &str, col: &str) -> (i64, i64) {
551 let source = OracleSource::new(conn.as_str(), 1)?;
552 let conn = source.get_conn()?;
553 let range_query = get_partition_range_query(query, col, &OracleDialect {})?;
554 let row = conn.query_row(range_query.as_str(), &[])?;
555 let min_v: i64 = row.get(0).unwrap_or(0);
556 let max_v: i64 = row.get(1).unwrap_or(0);
557 (min_v, max_v)
558}
559
560#[cfg(feature = "src_bigquery")]
561#[throws(ConnectorXOutError)] fn bigquery_get_partition_range(conn: &Url, query: &str, col: &str) -> (i64, i64) {
563 let rt = Runtime::new().expect("Failed to create runtime");
564 let url = Url::parse(conn.as_str())?;
565 let sa_key_path = url.path();
566 let client = rt.block_on(gcp_bigquery_client::Client::from_service_account_key_file(
567 sa_key_path,
568 ))?;
569
570 let auth_data = std::fs::read_to_string(sa_key_path)?;
571 let auth_json: serde_json::Value = serde_json::from_str(&auth_data)?;
572 let project_id = auth_json
573 .get("project_id")
574 .ok_or_else(|| anyhow!("Cannot get project_id from auth file"))?
575 .as_str()
576 .ok_or_else(|| anyhow!("Cannot get project_id as string from auth file"))?;
577 let range_query = get_partition_range_query(query, col, &BigQueryDialect {})?;
578
579 let query_result = rt.block_on(client.job().query(
580 project_id,
581 gcp_bigquery_client::model::query_request::QueryRequest::new(range_query.as_str()),
582 ))?;
583 let mut rs = gcp_bigquery_client::model::query_response::ResultSet::new_from_query_response(
584 query_result,
585 );
586 rs.next_row();
587 let min_v = rs.get_i64(0)?.unwrap_or(0);
588 let max_v = rs.get_i64(1)?.unwrap_or(0);
589
590 (min_v, max_v)
591}
592
593#[cfg(feature = "src_trino")]
594#[throws(ConnectorXOutError)]
595fn trino_get_partition_range(conn: &Url, query: &str, col: &str) -> (i64, i64) {
596 use crate::sources::trino::{build_client_from_url, TrinoDialect, TrinoPartitionQueryResult};
597
598 let rt = Runtime::new().expect("Failed to create runtime");
599
600 let client =
601 build_client_from_url(conn).map_err(|e| anyhow!("Failed to build Trino client: {}", e))?;
602
603 let range_query = get_partition_range_query(query, col, &TrinoDialect {})?;
604 let query_result = rt.block_on(client.get_all::<TrinoPartitionQueryResult>(range_query));
605
606 let query_result = match query_result {
607 Ok(query_result) => Ok(query_result.into_vec()),
608 Err(e) => match e {
609 prusto::error::Error::EmptyData => {
610 Ok(vec![TrinoPartitionQueryResult { _col0: 0, _col1: 0 }])
611 }
612 _ => Err(anyhow!("Failed to get query result: {}", e)),
613 },
614 }?;
615
616 let result = query_result
617 .first()
618 .unwrap_or(&TrinoPartitionQueryResult { _col0: 0, _col1: 0 });
619
620 (result._col0, result._col1)
621}
622
623#[cfg(feature = "src_clickhouse")]
624#[throws(ConnectorXOutError)]
625fn clickhouse_get_partition_range(conn: &Url, query: &str, col: &str) -> (i64, i64) {
626 use sqlparser::dialect::ClickHouseDialect;
627
628 let rt = Arc::new(tokio::runtime::Runtime::new().expect("Failed to create runtime"));
629 let clickhouse_source = ClickHouseSource::new(rt.clone(), conn.as_str())
630 .expect("Failed to create ClickHouse client");
631
632 let range_query = get_partition_range_query(query, col, &ClickHouseDialect {})?;
633
634 let response = rt.block_on(async {
635 let mut cursor = clickhouse_source
636 .client
637 .query(range_query.as_str())
638 .fetch_bytes("JSONCompact")
639 .map_err(|e| anyhow!("ClickHouse error: {}", e))?;
640 let bytes = cursor
641 .collect()
642 .await
643 .map_err(|e| anyhow!("ClickHouse error: {}", e))?;
644 Ok::<_, ClickHouseSourceError>(bytes)
645 })?;
646
647 #[derive(Debug, Deserialize)]
648 struct MinMaxResponse {
649 data: Vec<Vec<JsonValue>>,
650 }
651
652 let parsed: MinMaxResponse = serde_json::from_slice(&response)
653 .map_err(|e| anyhow!("Failed to parse min max response: {}", e))?;
654
655 let (min_v, max_v) = if let Some(row) = parsed.data.first() {
656 let min_v = row.first().and_then(|v| v.as_i64()).unwrap_or(0);
657 let max_v = row.get(1).and_then(|v| v.as_i64()).unwrap_or(0);
658
659 (min_v, max_v)
660 } else {
661 (0, 0)
662 };
663
664 (min_v, max_v)
665}
666
667#[cfg(test)]
668mod tests {
669 use super::*;
670
671 fn sqlite_source_conn() -> SourceConn {
675 SourceConn::new(
676 SourceType::SQLite,
677 Url::parse("sqlite://test.db").unwrap(),
678 "binary".to_string(),
679 )
680 }
681
682 fn queries_as_strings(queries: &[CXQuery]) -> Vec<String> {
683 queries.iter().map(|q| q.to_string()).collect()
684 }
685
686 #[test]
687 fn partitions_evenly_divisible_range() {
688 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(0), Some(9), 2);
689 let source_conn = sqlite_source_conn();
690 let queries = partition(&part, &source_conn).unwrap();
691
692 assert_eq!(queries.len(), 2);
693 let strings = queries_as_strings(&queries);
694 assert!(strings[0].contains("0 <=") && strings[0].contains("< 5"));
695 assert!(strings[1].contains("5 <=") && strings[1].contains("< 10"));
696 }
697
698 #[test]
699 fn partitions_range_not_evenly_divisible() {
700 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(0), Some(10), 3);
703 let source_conn = sqlite_source_conn();
704 let queries = partition(&part, &source_conn).unwrap();
705
706 assert_eq!(queries.len(), 3);
707 let strings = queries_as_strings(&queries);
708 assert!(strings[0].contains("0 <=") && strings[0].contains("< 3"));
709 assert!(strings[1].contains("3 <=") && strings[1].contains("< 6"));
710 assert!(strings[2].contains("6 <=") && strings[2].contains("< 11"));
712 }
713
714 #[test]
715 fn single_partition_covers_whole_range() {
716 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(5), Some(15), 1);
717 let source_conn = sqlite_source_conn();
718 let queries = partition(&part, &source_conn).unwrap();
719
720 assert_eq!(queries.len(), 1);
721 let strings = queries_as_strings(&queries);
722 assert!(strings[0].contains("5 <=") && strings[0].contains("< 16"));
723 }
724
725 #[test]
726 fn min_equals_max_single_row_range() {
727 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(7), Some(7), 1);
728 let source_conn = sqlite_source_conn();
729 let queries = partition(&part, &source_conn).unwrap();
730
731 assert_eq!(queries.len(), 1);
732 let strings = queries_as_strings(&queries);
733 assert!(strings[0].contains("7 <=") && strings[0].contains("< 8"));
734 }
735
736 #[test]
737 fn negative_range_partitions_correctly() {
738 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(-10), Some(-1), 2);
739 let source_conn = sqlite_source_conn();
740 let queries = partition(&part, &source_conn).unwrap();
741
742 assert_eq!(queries.len(), 2);
743 let strings = queries_as_strings(&queries);
744 assert!(strings[0].contains("-10 <=") && strings[0].contains("< -5"));
745 assert!(strings[1].contains("-5 <=") && strings[1].contains("< 0"));
746 }
747
748 #[test]
749 fn partially_specified_range_returns_error() {
750 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(0), None, 2);
751 let source_conn = sqlite_source_conn();
752 let result = partition(&part, &source_conn);
753
754 assert!(result.is_err());
755 }
756
757 #[test]
758 fn zero_partitions_returns_error_on_division_by_zero() {
759 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(0), Some(9), 0);
760 let source_conn = sqlite_source_conn();
761 let result = partition(&part, &source_conn);
762 assert!(result.is_err());
763 let err = result.unwrap_err().to_string();
764 assert!(err.contains("partition count (num) must be greater than zero"));
765 }
766
767 #[test]
768 fn inverted_range_returns_error() {
769 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(10), Some(0), 2);
770 let source_conn = sqlite_source_conn();
771 let result = partition(&part, &source_conn);
772 assert!(result.is_err());
773 let err = result.unwrap_err().to_string();
774 assert!(
775 err.contains("partition range is invalid"),
776 "unexpected error message: {}",
777 err
778 );
779 }
780
781 #[test]
782 fn more_partitions_than_range_values_preserves_partition_count() {
783 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(0), Some(1), 4);
784 let source_conn = sqlite_source_conn();
785 let queries = partition(&part, &source_conn).unwrap();
786
787 assert_eq!(queries.len(), 4);
788 let strings = queries_as_strings(&queries);
789 assert!(strings[0].contains("0 <=") && strings[0].contains("< 0"));
790 assert!(strings[1].contains("0 <=") && strings[1].contains("< 0"));
791 assert!(strings[2].contains("0 <=") && strings[2].contains("< 0"));
792 assert!(strings[3].contains("0 <=") && strings[3].contains("< 2"));
793 }
794
795 #[test]
796 fn large_range_does_not_overflow_with_small_num() {
797 let part = PartitionQuery::new(
800 "SELECT * FROM test",
801 "id",
802 Some(0),
803 Some(1_000_000_000_000),
804 4,
805 );
806 let source_conn = sqlite_source_conn();
807 let queries = partition(&part, &source_conn).unwrap();
808
809 assert_eq!(queries.len(), 4);
810 let strings = queries_as_strings(&queries);
811 assert!(strings[0].contains("0 <="));
812 assert!(strings[3].contains("< 1000000000001"));
813 }
814
815 #[test]
816 fn many_small_partitions_produce_correct_count_and_bounds() {
817 let part = PartitionQuery::new("SELECT * FROM test", "id", Some(1), Some(100), 10);
818 let source_conn = sqlite_source_conn();
819 let queries = partition(&part, &source_conn).unwrap();
820
821 assert_eq!(queries.len(), 10);
822 let strings = queries_as_strings(&queries);
823 assert!(strings[0].contains("1 <=") && strings[0].contains("< 11"));
824 assert!(strings[9].contains("91 <=") && strings[9].contains("< 101"));
825 }
826
827 #[test]
828 fn partition_range_overflow() {
829 let part = PartitionQuery::new(
831 "SELECT * FROM test",
832 "id",
833 Some(i64::MIN),
834 Some(i64::MAX),
835 2,
836 );
837 let source_conn = sqlite_source_conn();
838 let res = partition(&part, &source_conn);
839 assert!(res.is_err());
840 let err = res.unwrap_err().to_string();
841 assert!(
842 err.contains("partition range overflow"),
843 "unexpected error message: {}",
844 err
845 );
846 }
847
848 #[test]
849 fn partition_upper_overflow_at_i64_max() {
850 let part = PartitionQuery::new(
853 "SELECT * FROM test",
854 "id",
855 Some(i64::MAX - 5),
856 Some(i64::MAX),
857 1,
858 );
859 let source_conn = sqlite_source_conn();
860 let res = partition(&part, &source_conn);
861 assert!(res.is_err());
862 let err = res.unwrap_err().to_string();
863 assert!(
864 err.contains("partition upper bound overflow"),
865 "unexpected error message: {}",
866 err
867 );
868 }
869
870 #[test]
871 fn partition_upper_overflow_multi_partition() {
872 let part = PartitionQuery::new(
874 "SELECT * FROM test",
875 "id",
876 Some(i64::MAX - 5),
877 Some(i64::MAX),
878 2,
879 );
880 let source_conn = sqlite_source_conn();
881 let res = partition(&part, &source_conn);
882 assert!(res.is_err());
883 let err = res.unwrap_err().to_string();
884 assert!(
885 err.contains("partition upper bound overflow"),
886 "unexpected error message: {}",
887 err
888 );
889 }
890}