Skip to main content

connectorx/sources/mssql/
tiberius_impl.rs

1//! Tiberius-backed implementation of the MsSQL source.
2//!
3//! This is the original, default implementation (Phase 0/1 of the migration
4//! plan tracked against sfu-db/connector-x#942). Compiled only when the
5//! `src_mssql_tiberius` feature is active.
6
7use super::driver;
8use super::errors::MsSQLSourceError;
9use super::typesystem::{FloatN, IntN, MsSQLTypeSystem};
10use crate::constants::DB_BUFFER_SIZE;
11use crate::{
12    data_order::DataOrder,
13    errors::ConnectorXError,
14    sources::{PartitionParser, Produce, Source, SourcePartition},
15    sql::{count_query, CXQuery},
16    utils::DummyBox,
17};
18use anyhow::anyhow;
19use bb8::{Pool, PooledConnection};
20use bb8_tiberius::ConnectionManager;
21use chrono::{DateTime, Utc};
22use chrono::{NaiveDate, NaiveDateTime, NaiveTime};
23use fehler::{throw, throws};
24use futures::StreamExt;
25use log::debug;
26use owning_ref::OwningHandle;
27use rust_decimal::Decimal;
28use sqlparser::dialect::MsSqlDialect;
29use std::collections::HashMap;
30use std::sync::Arc;
31use tiberius::{AuthMethod, Config, EncryptionLevel, QueryItem, QueryStream, Row};
32use tokio::runtime::{Handle, Runtime};
33use url::Url;
34use urlencoding::decode;
35use uuid_old::Uuid;
36
37type Conn<'a> = PooledConnection<'a, ConnectionManager>;
38pub struct MsSQLSource {
39    rt: Arc<Runtime>,
40    pool: Pool<ConnectionManager>,
41    origin_query: Option<String>,
42    queries: Vec<CXQuery<String>>,
43    names: Vec<String>,
44    schema: Vec<MsSQLTypeSystem>,
45}
46
47#[throws(MsSQLSourceError)]
48pub fn mssql_config(url: &Url) -> Config {
49    let mut config = Config::new();
50
51    let host = decode(url.host_str().unwrap_or("localhost"))?.into_owned();
52    let hosts: Vec<&str> = host.split('\\').collect();
53    match hosts.len() {
54        1 => config.host(host),
55        2 => {
56            // SQL Server support instance name: `server\instance:port`
57            config.host(hosts[0]);
58            config.instance_name(hosts[1]);
59        }
60        _ => throw!(anyhow!("MsSQL hostname parse error: {}", host)),
61    }
62    config.port(url.port().unwrap_or(1433));
63    // remove the leading "/"
64    config.database(decode(&url.path()[1..])?.to_owned());
65    // Using SQL Server authentication.
66    #[allow(unused)]
67    let params: HashMap<String, String> = url.query_pairs().into_owned().collect();
68    #[cfg(any(windows, feature = "integrated-auth-gssapi"))]
69    match params.get("trusted_connection") {
70        // pefer trusted_connection if set to true
71        Some(v) if v == "true" => {
72            debug!("mssql auth through trusted connection!");
73            config.authentication(AuthMethod::Integrated);
74        }
75        _ => {
76            debug!("mssql auth through sqlserver authentication");
77            config.authentication(AuthMethod::sql_server(
78                decode(url.username())?.to_owned(),
79                decode(url.password().unwrap_or(""))?.to_owned(),
80            ));
81        }
82    };
83    #[cfg(all(not(windows), not(feature = "integrated-auth-gssapi")))]
84    config.authentication(AuthMethod::sql_server(
85        decode(url.username())?.to_owned(),
86        decode(url.password().unwrap_or(""))?.to_owned(),
87    ));
88
89    match params.get("trust_server_certificate") {
90        Some(v) if v.to_lowercase() == "true" => config.trust_cert(),
91        _ => {}
92    };
93
94    match params.get("trust_server_certificate_ca") {
95        Some(v) => config.trust_cert_ca(v),
96        _ => {}
97    };
98
99    match params.get("encrypt") {
100        Some(v) if v.to_lowercase() == "true" => config.encryption(EncryptionLevel::Required),
101        Some(v) if v.to_lowercase() == "false" => config.encryption(EncryptionLevel::Off),
102        _ => config.encryption(EncryptionLevel::NotSupported),
103    };
104
105    match params.get("appname") {
106        Some(appname) => config.application_name(decode(appname)?.to_owned()),
107        _ => {}
108    };
109
110    config
111}
112
113impl MsSQLSource {
114    #[throws(MsSQLSourceError)]
115    pub fn new(rt: Arc<Runtime>, conn: &str, nconn: usize) -> Self {
116        debug!("mssql source using driver: {:?}", driver::active_driver());
117        let url = Url::parse(conn)?;
118        let config = mssql_config(&url)?;
119        let manager = bb8_tiberius::ConnectionManager::new(config);
120        let pool = rt.block_on(Pool::builder().max_size(nconn as u32).build(manager))?;
121
122        Self {
123            rt,
124            pool,
125            origin_query: None,
126            queries: vec![],
127            names: vec![],
128            schema: vec![],
129        }
130    }
131}
132
133impl Source for MsSQLSource
134where
135    MsSQLSourcePartition: SourcePartition<TypeSystem = MsSQLTypeSystem, Error = MsSQLSourceError>,
136{
137    const DATA_ORDERS: &'static [DataOrder] = &[DataOrder::RowMajor];
138    type Partition = MsSQLSourcePartition;
139    type TypeSystem = MsSQLTypeSystem;
140    type Error = MsSQLSourceError;
141
142    #[throws(MsSQLSourceError)]
143    fn set_data_order(&mut self, data_order: DataOrder) {
144        if !matches!(data_order, DataOrder::RowMajor) {
145            throw!(ConnectorXError::UnsupportedDataOrder(data_order));
146        }
147    }
148
149    fn set_queries<Q: ToString>(&mut self, queries: &[CXQuery<Q>]) {
150        self.queries = queries.iter().map(|q| q.map(Q::to_string)).collect();
151    }
152
153    fn set_origin_query(&mut self, query: Option<String>) {
154        self.origin_query = query;
155    }
156
157    #[throws(MsSQLSourceError)]
158    fn fetch_metadata(&mut self) {
159        assert!(!self.queries.is_empty());
160
161        let mut conn = self.rt.block_on(self.pool.get())?;
162        let first_query = &self.queries[0];
163        let (names, types) = match self.rt.block_on(conn.query(first_query.as_str(), &[])) {
164            Ok(mut stream) => match self.rt.block_on(async { stream.columns().await }) {
165                Ok(Some(columns)) => columns
166                    .iter()
167                    .map(|col| {
168                        (
169                            col.name().to_string(),
170                            MsSQLTypeSystem::from(&col.column_type()),
171                        )
172                    })
173                    .unzip(),
174                Ok(None) => {
175                    throw!(anyhow!(
176                        "MsSQL returned no columns for query: {}",
177                        first_query
178                    ));
179                }
180                Err(e) => {
181                    throw!(anyhow!("Error fetching columns: {}", e));
182                }
183            },
184            Err(e) => {
185                debug!(
186                    "cannot get metadata for '{}', try next query: {}",
187                    first_query, e
188                );
189                throw!(e);
190            }
191        };
192
193        self.names = names;
194        self.schema = types;
195    }
196
197    #[throws(MsSQLSourceError)]
198    fn result_rows(&mut self) -> Option<usize> {
199        match &self.origin_query {
200            Some(q) => {
201                let cxq = CXQuery::Naked(q.clone());
202                let cquery = count_query(&cxq, &MsSqlDialect {})?;
203                let mut conn = self.rt.block_on(self.pool.get())?;
204
205                let stream = self.rt.block_on(conn.query(cquery.as_str(), &[]))?;
206                let row = self
207                    .rt
208                    .block_on(stream.into_row())?
209                    .ok_or_else(|| anyhow!("MsSQL failed to get the count of query: {}", q))?;
210
211                let row: i32 = row.get(0).ok_or(MsSQLSourceError::GetNRowsFailed)?; // the count in mssql is i32
212                Some(row as usize)
213            }
214            None => None,
215        }
216    }
217
218    fn names(&self) -> Vec<String> {
219        self.names.clone()
220    }
221
222    fn schema(&self) -> Vec<Self::TypeSystem> {
223        self.schema.clone()
224    }
225
226    #[throws(MsSQLSourceError)]
227    fn partition(self) -> Vec<Self::Partition> {
228        let mut ret = vec![];
229        for query in self.queries {
230            ret.push(MsSQLSourcePartition::new(
231                self.pool.clone(),
232                self.rt.clone(),
233                &query,
234                &self.schema,
235            ));
236        }
237        ret
238    }
239}
240
241pub struct MsSQLSourcePartition {
242    pool: Pool<ConnectionManager>,
243    rt: Arc<Runtime>,
244    query: CXQuery<String>,
245    schema: Vec<MsSQLTypeSystem>,
246    nrows: usize,
247    ncols: usize,
248}
249
250impl MsSQLSourcePartition {
251    pub fn new(
252        pool: Pool<ConnectionManager>,
253        handle: Arc<Runtime>,
254        query: &CXQuery<String>,
255        schema: &[MsSQLTypeSystem],
256    ) -> Self {
257        Self {
258            rt: handle,
259            pool,
260            query: query.clone(),
261            schema: schema.to_vec(),
262            nrows: 0,
263            ncols: schema.len(),
264        }
265    }
266}
267
268impl SourcePartition for MsSQLSourcePartition {
269    type TypeSystem = MsSQLTypeSystem;
270    type Parser<'a> = MsSQLSourceParser<'a>;
271    type Error = MsSQLSourceError;
272
273    #[throws(MsSQLSourceError)]
274    fn result_rows(&mut self) {
275        let cquery = count_query(&self.query, &MsSqlDialect {})?;
276        let mut conn = self.rt.block_on(self.pool.get())?;
277
278        let stream = self.rt.block_on(conn.query(cquery.as_str(), &[]))?;
279        let row = self
280            .rt
281            .block_on(stream.into_row())?
282            .ok_or_else(|| anyhow!("MsSQL failed to get the count of query: {}", self.query))?;
283
284        let row: i32 = row.get(0).ok_or(MsSQLSourceError::GetNRowsFailed)?; // the count in mssql is i32
285        self.nrows = row as usize;
286    }
287
288    #[throws(MsSQLSourceError)]
289    fn parser<'a>(&'a mut self) -> Self::Parser<'a> {
290        let conn = self.rt.block_on(self.pool.get())?;
291        let rows: OwningHandle<Box<Conn<'a>>, DummyBox<QueryStream<'a>>> =
292            OwningHandle::new_with_fn(Box::new(conn), |conn: *const Conn<'a>| unsafe {
293                let conn = &mut *(conn as *mut Conn<'a>);
294
295                DummyBox(
296                    self.rt
297                        .block_on(conn.query(self.query.as_str(), &[]))
298                        .unwrap(),
299                )
300            });
301
302        MsSQLSourceParser::new(self.rt.handle(), rows, &self.schema)
303    }
304
305    fn nrows(&self) -> usize {
306        self.nrows
307    }
308
309    fn ncols(&self) -> usize {
310        self.ncols
311    }
312}
313
314pub struct MsSQLSourceParser<'a> {
315    rt: &'a Handle,
316    iter: OwningHandle<Box<Conn<'a>>, DummyBox<QueryStream<'a>>>,
317    rowbuf: Vec<Row>,
318    ncols: usize,
319    current_col: usize,
320    current_row: usize,
321    is_finished: bool,
322}
323
324impl<'a> MsSQLSourceParser<'a> {
325    fn new(
326        rt: &'a Handle,
327        iter: OwningHandle<Box<Conn<'a>>, DummyBox<QueryStream<'a>>>,
328        schema: &[MsSQLTypeSystem],
329    ) -> Self {
330        Self {
331            rt,
332            iter,
333            rowbuf: Vec::with_capacity(DB_BUFFER_SIZE),
334            ncols: schema.len(),
335            current_row: 0,
336            current_col: 0,
337            is_finished: false,
338        }
339    }
340
341    #[throws(MsSQLSourceError)]
342    fn next_loc(&mut self) -> (usize, usize) {
343        let ret = (self.current_row, self.current_col);
344        self.current_row += (self.current_col + 1) / self.ncols;
345        self.current_col = (self.current_col + 1) % self.ncols;
346        ret
347    }
348}
349
350impl<'a> PartitionParser<'a> for MsSQLSourceParser<'a> {
351    type TypeSystem = MsSQLTypeSystem;
352    type Error = MsSQLSourceError;
353
354    #[throws(MsSQLSourceError)]
355    fn fetch_next(&mut self) -> (usize, bool) {
356        assert!(self.current_col == 0);
357        let remaining_rows = self.rowbuf.len() - self.current_row;
358        if remaining_rows > 0 {
359            return (remaining_rows, self.is_finished);
360        } else if self.is_finished {
361            return (0, self.is_finished);
362        }
363
364        if !self.rowbuf.is_empty() {
365            self.rowbuf.drain(..);
366        }
367
368        for _ in 0..DB_BUFFER_SIZE {
369            if let Some(item) = self.rt.block_on(self.iter.next()) {
370                match item.map_err(MsSQLSourceError::MsSQLError)? {
371                    QueryItem::Row(row) => self.rowbuf.push(row),
372                    _ => continue,
373                }
374            } else {
375                self.is_finished = true;
376                break;
377            }
378        }
379        self.current_row = 0;
380        self.current_col = 0;
381        (self.rowbuf.len(), self.is_finished)
382    }
383}
384
385macro_rules! impl_produce {
386    ($($t: ty,)+) => {
387        $(
388            impl<'r, 'a> Produce<'r, $t> for MsSQLSourceParser<'a> {
389                type Error = MsSQLSourceError;
390
391                #[throws(MsSQLSourceError)]
392                fn produce(&'r mut self) -> $t {
393                    let (ridx, cidx) = self.next_loc()?;
394                    let res = self.rowbuf[ridx].get(cidx).ok_or_else(|| anyhow!("MsSQL get None at position: ({}, {})", ridx, cidx))?;
395                    res
396                }
397            }
398
399            impl<'r, 'a> Produce<'r, Option<$t>> for MsSQLSourceParser<'a> {
400                type Error = MsSQLSourceError;
401
402                #[throws(MsSQLSourceError)]
403                fn produce(&'r mut self) -> Option<$t> {
404                    let (ridx, cidx) = self.next_loc()?;
405                    let res = self.rowbuf[ridx].get(cidx);
406                    res
407                }
408            }
409        )+
410    };
411}
412
413impl_produce!(
414    u8,
415    i16,
416    i32,
417    i64,
418    IntN,
419    f32,
420    f64,
421    FloatN,
422    bool,
423    &'r str,
424    &'r [u8],
425    Uuid,
426    Decimal,
427    NaiveDateTime,
428    NaiveDate,
429    NaiveTime,
430    DateTime<Utc>,
431);