Skip to main content

connectorx/sources/mysql/
mod.rs

1//! Source implementation for MySQL database.
2
3mod connection;
4mod errors;
5mod typesystem;
6
7pub use self::connection::build_opts;
8pub use self::errors::MySQLSourceError;
9use crate::constants::DB_BUFFER_SIZE;
10use crate::{
11    data_order::DataOrder,
12    errors::ConnectorXError,
13    sources::{PartitionParser, Produce, Source, SourcePartition},
14    sql::{count_query, limit0_query, CXQuery},
15};
16use anyhow::anyhow;
17use chrono::{NaiveDate, NaiveDateTime, NaiveTime};
18use fehler::{throw, throws};
19use log::{debug, warn};
20use r2d2::{Pool, PooledConnection};
21use r2d2_mysql::{
22    mysql::{prelude::Queryable, Binary, Column, QueryResult, Row, Text},
23    MySqlConnectionManager,
24};
25use rust_decimal::Decimal;
26use serde_json::Value;
27use sqlparser::dialect::MySqlDialect;
28use std::marker::PhantomData;
29pub use typesystem::MySQLTypeSystem;
30
31type MysqlConn = PooledConnection<MySqlConnectionManager>;
32
33pub enum BinaryProtocol {}
34pub enum TextProtocol {}
35
36#[throws(MySQLSourceError)]
37fn get_total_rows(conn: &mut MysqlConn, query: &CXQuery<String>) -> usize {
38    conn.query_first(&count_query(query, &MySqlDialect {})?)?
39        .ok_or_else(|| anyhow!("mysql failed to get the count of query: {}", query))?
40}
41
42pub struct MySQLSource<P> {
43    pool: Pool<MySqlConnectionManager>,
44    origin_query: Option<String>,
45    queries: Vec<CXQuery<String>>,
46    names: Vec<String>,
47    schema: Vec<MySQLTypeSystem>,
48    pre_execution_queries: Option<Vec<String>>,
49    _protocol: PhantomData<P>,
50}
51
52fn column_metadata(columns: &[Column]) -> (Vec<String>, Vec<MySQLTypeSystem>) {
53    columns
54        .iter()
55        .map(|col| {
56            (
57                col.name_str().to_string(),
58                MySQLTypeSystem::from((&col.column_type(), &col.flags(), col.character_set())),
59            )
60        })
61        .unzip()
62}
63
64/// Derive the schema by running each query with `LIMIT 0`. MySQL sends the
65/// column definitions even when the result set is empty, so no row is read.
66fn limit0_metadata(
67    conn: &mut MysqlConn,
68    queries: &[CXQuery<String>],
69) -> Result<(Vec<String>, Vec<MySQLTypeSystem>), MySQLSourceError> {
70    for (i, query) in queries.iter().enumerate() {
71        // assuming all the partition queries yield same schema
72        match conn.query_iter(limit0_query(query, &MySqlDialect {})?.as_str()) {
73            Ok(iter) => {
74                let (names, types) = column_metadata(iter.columns().as_ref());
75                if !names.is_empty() {
76                    return Ok((names, types));
77                }
78                debug!("no result columns for '{}', try next query", query);
79            }
80            Err(e) if i == queries.len() - 1 => {
81                // tried the last query but still get an error
82                debug!("cannot get metadata for '{}': {}", query, e);
83                return Err(e.into());
84            }
85            Err(e) => {
86                debug!("cannot get metadata for '{}', try next query: {}", query, e);
87            }
88        }
89    }
90    Err(anyhow!("cannot get metadata: no query returned a result set with columns").into())
91}
92
93impl<P> MySQLSource<P> {
94    #[throws(MySQLSourceError)]
95    pub fn new(conn: &str, nconn: usize) -> Self {
96        let manager = MySqlConnectionManager::new(build_opts(conn)?);
97        let pool = r2d2::Pool::builder()
98            .max_size(nconn as u32)
99            .build(manager)?;
100
101        Self {
102            pool,
103            origin_query: None,
104            queries: vec![],
105            names: vec![],
106            schema: vec![],
107            pre_execution_queries: None,
108            _protocol: PhantomData,
109        }
110    }
111}
112
113impl<P> Source for MySQLSource<P>
114where
115    MySQLSourcePartition<P>:
116        SourcePartition<TypeSystem = MySQLTypeSystem, Error = MySQLSourceError>,
117    P: Send,
118{
119    const DATA_ORDERS: &'static [DataOrder] = &[DataOrder::RowMajor];
120    type Partition = MySQLSourcePartition<P>;
121    type TypeSystem = MySQLTypeSystem;
122    type Error = MySQLSourceError;
123
124    #[throws(MySQLSourceError)]
125    fn set_data_order(&mut self, data_order: DataOrder) {
126        if !matches!(data_order, DataOrder::RowMajor) {
127            throw!(ConnectorXError::UnsupportedDataOrder(data_order));
128        }
129    }
130
131    fn set_queries<Q: ToString>(&mut self, queries: &[CXQuery<Q>]) {
132        self.queries = queries.iter().map(|q| q.map(Q::to_string)).collect();
133    }
134
135    fn set_origin_query(&mut self, query: Option<String>) {
136        self.origin_query = query;
137    }
138
139    fn set_pre_execution_queries(&mut self, pre_execution_queries: Option<&[String]>) {
140        self.pre_execution_queries = pre_execution_queries.map(|s| s.to_vec());
141    }
142
143    #[throws(MySQLSourceError)]
144    fn fetch_metadata(&mut self) {
145        assert!(!self.queries.is_empty());
146
147        let mut conn = self.pool.get()?;
148        let first_query = &self.queries[0];
149
150        let (names, types) = match conn.prep(first_query) {
151            Ok(stmt) => column_metadata(stmt.columns()),
152            Err(e) => {
153                warn!(
154                    "mysql prepared statement error: {:?}, switch to limit 0 method",
155                    e
156                );
157                limit0_metadata(&mut conn, &self.queries)?
158            }
159        };
160        self.names = names;
161        self.schema = types;
162    }
163
164    #[throws(MySQLSourceError)]
165    fn result_rows(&mut self) -> Option<usize> {
166        match &self.origin_query {
167            Some(q) => {
168                let cxq = CXQuery::Naked(q.clone());
169                let mut conn = self.pool.get()?;
170                let nrows = get_total_rows(&mut conn, &cxq)?;
171                Some(nrows)
172            }
173            None => None,
174        }
175    }
176
177    fn names(&self) -> Vec<String> {
178        self.names.clone()
179    }
180
181    fn schema(&self) -> Vec<Self::TypeSystem> {
182        self.schema.clone()
183    }
184
185    #[throws(MySQLSourceError)]
186    fn partition(self) -> Vec<Self::Partition> {
187        let mut ret = vec![];
188        for query in self.queries {
189            let mut conn = self.pool.get()?;
190
191            if let Some(pre_queries) = &self.pre_execution_queries {
192                for pre_query in pre_queries {
193                    conn.query_drop(pre_query)?;
194                }
195            }
196
197            ret.push(MySQLSourcePartition::new(conn, &query, &self.schema));
198        }
199        ret
200    }
201}
202
203pub struct MySQLSourcePartition<P> {
204    conn: MysqlConn,
205    query: CXQuery<String>,
206    schema: Vec<MySQLTypeSystem>,
207    nrows: usize,
208    ncols: usize,
209    _protocol: PhantomData<P>,
210}
211
212impl<P> MySQLSourcePartition<P> {
213    pub fn new(conn: MysqlConn, query: &CXQuery<String>, schema: &[MySQLTypeSystem]) -> Self {
214        Self {
215            conn,
216            query: query.clone(),
217            schema: schema.to_vec(),
218            nrows: 0,
219            ncols: schema.len(),
220            _protocol: PhantomData,
221        }
222    }
223}
224
225impl SourcePartition for MySQLSourcePartition<BinaryProtocol> {
226    type TypeSystem = MySQLTypeSystem;
227    type Parser<'a> = MySQLBinarySourceParser<'a>;
228    type Error = MySQLSourceError;
229
230    #[throws(MySQLSourceError)]
231    fn result_rows(&mut self) {
232        self.nrows = get_total_rows(&mut self.conn, &self.query)?;
233    }
234
235    #[throws(MySQLSourceError)]
236    fn parser(&mut self) -> Self::Parser<'_> {
237        let stmt = self.conn.prep(self.query.as_str())?;
238        let iter = self.conn.exec_iter(stmt, ())?;
239        MySQLBinarySourceParser::new(iter, &self.schema)
240    }
241
242    fn nrows(&self) -> usize {
243        self.nrows
244    }
245
246    fn ncols(&self) -> usize {
247        self.ncols
248    }
249}
250
251impl SourcePartition for MySQLSourcePartition<TextProtocol> {
252    type TypeSystem = MySQLTypeSystem;
253    type Parser<'a> = MySQLTextSourceParser<'a>;
254    type Error = MySQLSourceError;
255
256    #[throws(MySQLSourceError)]
257    fn result_rows(&mut self) {
258        self.nrows = get_total_rows(&mut self.conn, &self.query)?;
259    }
260
261    #[throws(MySQLSourceError)]
262    fn parser(&mut self) -> Self::Parser<'_> {
263        let query = self.query.clone();
264        let iter = self.conn.query_iter(query)?;
265        MySQLTextSourceParser::new(iter, &self.schema)
266    }
267
268    fn nrows(&self) -> usize {
269        self.nrows
270    }
271
272    fn ncols(&self) -> usize {
273        self.ncols
274    }
275}
276
277pub struct MySQLBinarySourceParser<'a> {
278    iter: QueryResult<'a, 'a, 'a, Binary>,
279    rowbuf: Vec<Row>,
280    ncols: usize,
281    current_col: usize,
282    current_row: usize,
283    is_finished: bool,
284}
285
286impl<'a> MySQLBinarySourceParser<'a> {
287    pub fn new(iter: QueryResult<'a, 'a, 'a, Binary>, schema: &[MySQLTypeSystem]) -> Self {
288        Self {
289            iter,
290            rowbuf: Vec::with_capacity(DB_BUFFER_SIZE),
291            ncols: schema.len(),
292            current_row: 0,
293            current_col: 0,
294            is_finished: false,
295        }
296    }
297
298    #[throws(MySQLSourceError)]
299    fn next_loc(&mut self) -> (usize, usize) {
300        let ret = (self.current_row, self.current_col);
301        self.current_row += (self.current_col + 1) / self.ncols;
302        self.current_col = (self.current_col + 1) % self.ncols;
303        ret
304    }
305}
306
307impl<'a> PartitionParser<'a> for MySQLBinarySourceParser<'a> {
308    type TypeSystem = MySQLTypeSystem;
309    type Error = MySQLSourceError;
310
311    #[throws(MySQLSourceError)]
312    fn fetch_next(&mut self) -> (usize, bool) {
313        assert!(self.current_col == 0);
314        let remaining_rows = self.rowbuf.len() - self.current_row;
315        if remaining_rows > 0 {
316            return (remaining_rows, self.is_finished);
317        } else if self.is_finished {
318            return (0, self.is_finished);
319        }
320
321        if !self.rowbuf.is_empty() {
322            self.rowbuf.drain(..);
323        }
324
325        for _ in 0..DB_BUFFER_SIZE {
326            if let Some(item) = self.iter.next() {
327                self.rowbuf.push(item?);
328            } else {
329                self.is_finished = true;
330                break;
331            }
332        }
333        self.current_row = 0;
334        self.current_col = 0;
335
336        (self.rowbuf.len(), self.is_finished)
337    }
338}
339
340macro_rules! impl_produce_binary {
341    ($($t: ty,)+) => {
342        $(
343            impl<'r, 'a> Produce<'r, $t> for MySQLBinarySourceParser<'a> {
344                type Error = MySQLSourceError;
345
346                #[throws(MySQLSourceError)]
347                fn produce(&'r mut self) -> $t {
348                    let (ridx, cidx) = self.next_loc()?;
349                    let res = self.rowbuf[ridx].take(cidx).ok_or_else(|| anyhow!("mysql cannot parse at position: ({}, {})", ridx, cidx))?;
350                    res
351                }
352            }
353
354            impl<'r, 'a> Produce<'r, Option<$t>> for MySQLBinarySourceParser<'a> {
355                type Error = MySQLSourceError;
356
357                #[throws(MySQLSourceError)]
358                fn produce(&'r mut self) -> Option<$t> {
359                    let (ridx, cidx) = self.next_loc()?;
360                    let res = self.rowbuf[ridx].take(cidx).ok_or_else(|| anyhow!("mysql cannot parse at position: ({}, {})", ridx, cidx))?;
361                    res
362                }
363            }
364        )+
365    };
366}
367
368impl_produce_binary!(
369    i8,
370    i16,
371    i32,
372    i64,
373    u8,
374    u16,
375    u32,
376    u64,
377    f32,
378    f64,
379    NaiveDate,
380    NaiveTime,
381    NaiveDateTime,
382    Decimal,
383    String,
384    Vec<u8>,
385    Value,
386);
387
388pub struct MySQLTextSourceParser<'a> {
389    iter: QueryResult<'a, 'a, 'a, Text>,
390    rowbuf: Vec<Row>,
391    ncols: usize,
392    current_col: usize,
393    current_row: usize,
394    is_finished: bool,
395}
396
397impl<'a> MySQLTextSourceParser<'a> {
398    pub fn new(iter: QueryResult<'a, 'a, 'a, Text>, schema: &[MySQLTypeSystem]) -> Self {
399        Self {
400            iter,
401            rowbuf: Vec::with_capacity(DB_BUFFER_SIZE),
402            ncols: schema.len(),
403            current_row: 0,
404            current_col: 0,
405            is_finished: false,
406        }
407    }
408
409    #[throws(MySQLSourceError)]
410    fn next_loc(&mut self) -> (usize, usize) {
411        let ret = (self.current_row, self.current_col);
412        self.current_row += (self.current_col + 1) / self.ncols;
413        self.current_col = (self.current_col + 1) % self.ncols;
414        ret
415    }
416}
417
418impl<'a> PartitionParser<'a> for MySQLTextSourceParser<'a> {
419    type TypeSystem = MySQLTypeSystem;
420    type Error = MySQLSourceError;
421
422    #[throws(MySQLSourceError)]
423    fn fetch_next(&mut self) -> (usize, bool) {
424        assert!(self.current_col == 0);
425        let remaining_rows = self.rowbuf.len() - self.current_row;
426        if remaining_rows > 0 {
427            return (remaining_rows, self.is_finished);
428        } else if self.is_finished {
429            return (0, self.is_finished);
430        }
431
432        if !self.rowbuf.is_empty() {
433            self.rowbuf.drain(..);
434        }
435        for _ in 0..DB_BUFFER_SIZE {
436            if let Some(item) = self.iter.next() {
437                self.rowbuf.push(item?);
438            } else {
439                self.is_finished = true;
440                break;
441            }
442        }
443        self.current_row = 0;
444        self.current_col = 0;
445        (self.rowbuf.len(), self.is_finished)
446    }
447}
448
449macro_rules! impl_produce_text {
450    ($($t: ty,)+) => {
451        $(
452            impl<'r, 'a> Produce<'r, $t> for MySQLTextSourceParser<'a> {
453                type Error = MySQLSourceError;
454
455                #[throws(MySQLSourceError)]
456                fn produce(&'r mut self) -> $t {
457                    let (ridx, cidx) = self.next_loc()?;
458                    let res = self.rowbuf[ridx].take(cidx).ok_or_else(|| anyhow!("mysql cannot parse at position: ({}, {})", ridx, cidx))?;
459                    res
460                }
461            }
462
463            impl<'r, 'a> Produce<'r, Option<$t>> for MySQLTextSourceParser<'a> {
464                type Error = MySQLSourceError;
465
466                #[throws(MySQLSourceError)]
467                fn produce(&'r mut self) -> Option<$t> {
468                    let (ridx, cidx) = self.next_loc()?;
469                    let res = self.rowbuf[ridx].take(cidx).ok_or_else(|| anyhow!("mysql cannot parse at position: ({}, {})", ridx, cidx))?;
470                    res
471                }
472            }
473        )+
474    };
475}
476
477impl_produce_text!(
478    i8,
479    i16,
480    i32,
481    i64,
482    u8,
483    u16,
484    u32,
485    u64,
486    f32,
487    f64,
488    NaiveDate,
489    NaiveTime,
490    NaiveDateTime,
491    Decimal,
492    String,
493    Vec<u8>,
494    Value,
495);