connectorx/sources/mssql/
tiberius_impl.rs1use 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 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 config.database(decode(&url.path()[1..])?.to_owned());
65 #[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 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)?; 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)?; 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);