1mod 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
64fn 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 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 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);