Skip to main content

connectorx/
partition.rs

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    // remove the first "sqlite://" manually since url.path is not correct for windows and for relative path
305    let conn = Connection::open(&conn.as_str()[9..])?;
306    // SQLite only optimize min max queries when there is only one aggregation
307    // https://www.sqlite.org/optoverview.html#minmax
308    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        // declare type for count query will be None, only need to check the returned value type
312        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)] // TODO
562fn 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    /// Builds a `SourceConn` backed by SQLite, which does not require a live
672    /// database connection to generate partition queries (the SQL is
673    /// generated purely from the parsed dialect).
674    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        // Range [0, 10] has 11 values split across 3 partitions: sizes 3, 3, 3 with
701        // remainder handled by the last partition absorbing everything up to max + 1.
702        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        // Last partition always goes up to max + 1, regardless of even division.
711        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        // Sanity check that reasonably large but safe ranges partition correctly
798        // without overflow for a small number of partitions.
799        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        // (max - min + 1) overflows i64 when min is i64::MIN and max is i64::MAX
830        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        // max is i64::MAX; range_len is 6 (does not overflow),
851        // but max + 1 for the partition upper bound overflows i64.
852        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        // Multi-partition where first partition succeeds but last partition overflows max + 1
873        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}