use crate::config::AkitaConfig;
use crate::database_err;
use crate::driver::non_blocking::get_tokio_context;
use crate::driver::DriverType;
use crate::errors::AkitaError;
use async_trait::async_trait;
use deadpool::managed::{Metrics, Object, Pool, RecycleResult};
use deadpool::Runtime;
use tiberius::error::Error;
use tokio_util::compat::{Compat, TokioAsyncReadCompatExt};
use tracing::log::trace;
pub type MssqlAsyncPool = Pool<MssqlAsyncConnectionManager>;
pub type MssqlAsyncConnection = Object<MssqlAsyncConnectionManager>;
#[derive(Clone)]
pub struct MssqlAsyncConnectionManager {
config: tiberius::Config,
}
impl MssqlAsyncConnectionManager {
pub fn new(cfg: &AkitaConfig) -> Result<Self, AkitaError> {
if cfg.get_platform()? != DriverType::Mssql {
return Err(database_err!(
"Database type mismatch: expected SQL Server".to_string()
));
}
let connection_string = cfg.get_connection_string()?;
let config = tiberius::Config::from_ado_string(&connection_string)
.map_err(|e| database_err!(format!("Invalid SQL Server connection string: {}", e)))?;
Ok(Self { config })
}
}
#[async_trait]
impl deadpool::managed::Manager for MssqlAsyncConnectionManager {
type Type = tiberius::Client<Compat<tokio::net::TcpStream>>;
type Error = Error;
async fn create(&self) -> Result<Self::Type, Self::Error> {
let tcp = tokio::net::TcpStream::connect(self.config.get_addr()).await?;
tcp.set_nodelay(true)?;
let compat_tcp = tcp.compat();
tiberius::Client::connect(self.config.clone(), compat_tcp).await
}
async fn recycle(
&self,
conn: &mut Self::Type,
_metrics: &Metrics,
) -> RecycleResult<Self::Error> {
conn.simple_query("SELECT 1").await?;
Ok(())
}
}
pub async fn init_mssql_async_pool(
config: crate::config::AkitaConfig,
) -> Result<MssqlAsyncPool, AkitaError> {
use tokio::runtime::Handle;
let _handle = get_tokio_context()?;
let manager = MssqlAsyncConnectionManager::new(&config)?;
let pool_config = deadpool::managed::PoolConfig {
max_size: config.get_max_size() as usize,
timeouts: deadpool::managed::Timeouts {
wait: Some(config.get_connection_timeout()),
create: Some(config.get_connection_timeout()),
recycle: Some(config.get_idle_timeout()),
},
..Default::default()
};
let pool = Pool::builder(manager)
.runtime(Runtime::Tokio1)
.config(pool_config)
.build()?;
let mut client: MssqlAsyncConnection = pool
.get()
.await
.map_err(|e| database_err!(format!("Failed to get connection from pool: {}", e)))?;
client
.simple_query("SELECT 1")
.await
.map_err(|e| database_err!(format!("SQL Server async connection test failed: {}", e)))?;
tracing::info!("SQL Server async connection pool initialized successfully");
Ok(pool)
}