use crate::comm::ExecuteResult;
use crate::driver::non_blocking::mssql::MssqlAsyncConnection;
use crate::errors::{AkitaError, SmartBacktrace};
use crate::{database_err, mssql_err};
use akita_core::{AkitaValue, OperationType, Params, Row, Rows, SqlInjectionDetector};
use serde_json::Value;
use std::ops::{Deref, DerefMut};
use std::str::FromStr;
use std::sync::Arc;
use tokio::sync::Mutex;
pub struct MssqlAsyncAdapter {
client: Arc<Mutex<MssqlAsyncConnection>>,
}
impl MssqlAsyncAdapter {
pub fn new(client: MssqlAsyncConnection) -> Self {
Self {
client: Arc::new(Mutex::new(client)),
}
}
#[track_caller]
pub async fn start_transaction(&self) -> crate::prelude::Result<()> {
let mut client = self.client.lock().await;
client
.simple_query("BEGIN TRANSACTION")
.await
.map_err(|e| mssql_err!(e))?;
Ok(())
}
#[track_caller]
pub async fn commit_transaction(&self) -> crate::prelude::Result<()> {
let mut client = self.client.lock().await;
client
.simple_query("COMMIT")
.await
.map_err(|e| mssql_err!(e))?;
Ok(())
}
#[track_caller]
pub async fn rollback_transaction(&self) -> crate::prelude::Result<()> {
let mut client = self.client.lock().await;
client
.simple_query("ROLLBACK")
.await
.map_err(|e| mssql_err!(e))?;
Ok(())
}
#[track_caller]
pub async fn query(&self, sql: &str, params: Params) -> crate::prelude::Result<Rows> {
let mssql_params = convert_to_mssql_params(params);
let param_refs: Vec<&dyn tiberius::ToSql> = mssql_params
.iter()
.map(|p| &**p as &dyn tiberius::ToSql)
.collect();
self.inner_query(sql, ¶m_refs).await
}
#[track_caller]
async fn inner_query(
&self,
sql: &str,
param_refs: &[&dyn tiberius::ToSql],
) -> Result<Rows, AkitaError> {
let mut client = self.client.lock().await;
let stream = client
.query(sql, ¶m_refs)
.await
.map_err(|e| mssql_err!(e))?;
let rows: Vec<tiberius::Row> = stream
.into_first_result()
.await
.map_err(|e| database_err!(format!("Failed to get result: {}", e)))?;
if rows.is_empty() {
return Ok(Rows::new());
}
let first_row = &rows[0];
let column_names: Vec<String> = (0..first_row.columns().len())
.map(|i| first_row.columns()[i].name().to_string())
.collect();
let mut records = Rows::new();
for row in rows {
let mut record = Vec::new();
for i in 0..row.columns().len() {
let value = get_value_from_mssql_row(&row, i)?;
record.push(value);
}
records.push(Row {
columns: column_names.clone(),
data: record,
});
}
Ok(records)
}
#[track_caller]
pub async fn execute(
&self,
sql: &str,
params: Params,
) -> crate::prelude::Result<ExecuteResult> {
let mut client = self.client.lock().await;
let mssql_params = convert_to_mssql_params(params);
let param_refs: Vec<&dyn tiberius::ToSql> = mssql_params.iter().map(|p| &**p).collect();
let stmt_type = OperationType::detect_operation_type(sql);
match stmt_type {
OperationType::Select => {
let records = self.inner_query(sql, ¶m_refs).await?;
Ok(ExecuteResult::Rows(records))
}
_ => {
let result = client
.execute(sql, ¶m_refs)
.await
.map_err(|e| mssql_err!(e))?;
Ok(ExecuteResult::AffectedRows(result.total() as u64))
}
}
}
pub async fn affected_rows(&self) -> u64 {
0
}
}
fn convert_to_mssql_params(params: Params) -> Vec<Box<dyn tiberius::ToSql>> {
match params {
Params::None => vec![],
Params::Positional(param_values) => param_values
.into_iter()
.map(|v| convert_akita_value_to_mssql(v))
.collect::<Vec<_>>(),
Params::Named(named_params) => named_params
.values()
.cloned()
.into_iter()
.map(|v| convert_akita_value_to_mssql(v))
.collect::<Vec<_>>(),
}
}
fn convert_akita_value_to_mssql(val: AkitaValue) -> Box<dyn tiberius::ToSql> {
use tiberius::numeric::BigDecimal;
match val {
AkitaValue::Text(v) => Box::new(v),
AkitaValue::Bool(v) => Box::new(v),
AkitaValue::Tinyint(v) => {
let unsigned = if v >= 0 {
v as u8
} else {
(v as i16 + 256) as u8
};
Box::new(unsigned)
}
AkitaValue::Smallint(v) => Box::new(v),
AkitaValue::Int(v) => Box::new(v),
AkitaValue::Bigint(v) => Box::new(v),
AkitaValue::Float(v) => Box::new(v),
AkitaValue::Double(v) => Box::new(v),
AkitaValue::BigDecimal(v) => {
let bd = BigDecimal::from_str(v.to_string().as_str()).unwrap_or(BigDecimal::from(0));
Box::new(bd)
}
AkitaValue::Blob(v) => Box::new(v),
AkitaValue::Char(v) => Box::new(format!("{}", v)),
AkitaValue::Json(j) => match j {
Value::Bool(v) => Box::new(v),
Value::Number(v) => {
if let Some(n) = v.as_u64() {
Box::new(n as i64)
} else if let Some(n) = v.as_f64() {
Box::new(n)
} else if let Some(n) = v.as_i64() {
Box::new(n)
} else {
Box::new(v.to_string())
}
}
Value::String(v) => Box::new(v),
_ => Box::new(serde_json::to_string(&j).unwrap_or_default()),
},
AkitaValue::Uuid(v) => Box::new(v),
AkitaValue::Date(v) => Box::new(v),
AkitaValue::DateTime(v) => Box::new(v),
AkitaValue::Null => Box::new(Option::<String>::None),
_ => Box::new(val.to_string()),
}
}
fn get_value_from_mssql_row(row: &tiberius::Row, index: usize) -> Result<AkitaValue, AkitaError> {
use chrono::{NaiveDate, NaiveDateTime};
use tiberius::numeric::BigDecimal;
let column = &row.columns()[index];
match column.column_type() {
tiberius::ColumnType::Bit | tiberius::ColumnType::Bitn => {
let val: Option<bool> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Bool(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Int1 => {
let val: Option<u8> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Tinyint(v as i8),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Int2 => {
let val: Option<i16> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Smallint(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Int4 => {
let val: Option<i32> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Int(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Int8 => {
let val: Option<i64> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Bigint(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Intn => {
let val: Option<i64> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Bigint(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Float4 => {
let val: Option<f32> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Float(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Float8 => {
let val: Option<f64> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Double(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Floatn => {
let val: Option<f64> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Double(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Decimaln | tiberius::ColumnType::Numericn => {
let val: Option<BigDecimal> = row.get(index);
Ok(match val {
Some(v) => {
let decimal_str = v.to_string();
match decimal_str.parse() {
Ok(bd) => AkitaValue::BigDecimal(bd),
Err(_) => AkitaValue::Text(decimal_str),
}
}
None => AkitaValue::Null,
})
}
tiberius::ColumnType::BigVarChar
| tiberius::ColumnType::BigChar
| tiberius::ColumnType::NVarchar
| tiberius::ColumnType::NChar
| tiberius::ColumnType::Text
| tiberius::ColumnType::NText => {
let val: Option<&str> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Text(v.to_string()),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Xml => {
let val: Option<&str> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Text(v.to_string()), None => AkitaValue::Null,
})
}
tiberius::ColumnType::Daten => {
let val: Option<NaiveDate> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Date(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Datetime
| tiberius::ColumnType::Datetime4
| tiberius::ColumnType::Datetimen
| tiberius::ColumnType::Datetime2 => {
if let Ok(val) = row.try_get::<NaiveDateTime, _>(index) {
return Ok(match val {
Some(v) => AkitaValue::DateTime(v),
None => AkitaValue::Null,
});
}
let val: Option<&str> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::DateTime(
NaiveDateTime::parse_from_str(v, "%Y-%m-%d %H:%M:%S%.f").unwrap_or_else(|_| {
NaiveDate::from_ymd_opt(1970, 1, 1)
.unwrap()
.and_hms_opt(0, 0, 0)
.unwrap()
}),
),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Timen => {
let val: Option<&str> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Text(v.to_string()), None => AkitaValue::Null,
})
}
tiberius::ColumnType::DatetimeOffsetn => {
let val: Option<&str> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Text(v.to_string()),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::BigVarBin
| tiberius::ColumnType::BigBinary
| tiberius::ColumnType::Image => {
let val: Option<&[u8]> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Blob(v.to_vec()),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Guid => {
let val = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Uuid(v),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Money | tiberius::ColumnType::Money4 => {
if let Ok(val) = row.try_get::<f64, _>(index) {
return Ok(match val {
Some(v) => AkitaValue::Double(v),
None => AkitaValue::Null,
});
}
let val: Option<&str> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Text(v.to_string()),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::SSVariant => {
if let Ok(val) = row.try_get::<&str, _>(index) {
if let Some(v) = val {
return Ok(AkitaValue::Text(v.to_string()));
}
}
if let Ok(val) = row.try_get::<i64, _>(index) {
if let Some(v) = val {
return Ok(AkitaValue::Bigint(v));
}
}
if let Ok(val) = row.try_get::<f64, _>(index) {
if let Some(v) = val {
return Ok(AkitaValue::Double(v));
}
}
Ok(AkitaValue::Null)
}
tiberius::ColumnType::Udt => {
let val: Option<&[u8]> = row.get(index);
Ok(match val {
Some(v) => AkitaValue::Blob(v.to_vec()),
None => AkitaValue::Null,
})
}
tiberius::ColumnType::Null => Ok(AkitaValue::Null),
}
}