use adbc_core::options::{OptionConnection, OptionValue};
use adbc_core::{Connection, Database};
use std::error::Error as StdError;
use std::{
sync::{Arc, Mutex},
fmt,
};
pub struct AdbcConnectionManager<D>
where
D: Database + Send,
D::ConnectionType: Send + Sync,
{
database: Arc<Mutex<D>>,
connection_options: Vec<(String, String)>,
}
impl<D> AdbcConnectionManager<D>
where
D: Database + Send,
D::ConnectionType: Send + Sync,
{
pub fn new(database: D) -> Self {
Self {
database: Arc::new(Mutex::new(database)),
connection_options: Vec::new(),
}
}
pub fn with_options<I>(database: D, options: I) -> Self
where
I: IntoIterator<Item = (String, String)>,
{
Self {
database: Arc::new(Mutex::new(database)),
connection_options: options.into_iter().collect(),
}
}
pub fn add_option(&mut self, key: impl Into<String>, value: impl Into<String>) {
self.connection_options.push((key.into(), value.into()));
}
pub fn clear_options(&mut self) {
self.connection_options.clear();
}
pub fn options(&self) -> &[(String, String)] {
&self.connection_options
}
}
#[derive(Debug)]
pub struct AdbcError(pub adbc_core::error::Error);
impl fmt::Display for AdbcError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "ADBC error: {}", self.0)
}
}
impl StdError for AdbcError {
fn source(&self) -> Option<&(dyn StdError + 'static)> {
Some(&self.0)
}
}
impl From<adbc_core::error::Error> for AdbcError {
fn from(err: adbc_core::error::Error) -> Self {
AdbcError(err)
}
}
impl<D> From<std::sync::PoisonError<std::sync::MutexGuard<'_, D>>> for AdbcError
where
D: Database,
{
fn from(err: std::sync::PoisonError<std::sync::MutexGuard<'_, D>>) -> Self {
AdbcError(adbc_core::error::Error::with_message_and_status(
format!("Failed to acquire database lock: {}", err),
adbc_core::error::Status::Internal,
))
}
}
unsafe impl<D> Send for AdbcConnectionManager<D>
where
D: Database + Send,
D::ConnectionType: Send + Sync,
{}
impl<D> r2d2::ManageConnection for AdbcConnectionManager<D>
where
D: Database + Send + 'static,
D::ConnectionType: Send + Sync + 'static,
{
type Connection = D::ConnectionType;
type Error = AdbcError;
fn connect(&self) -> Result<Self::Connection, Self::Error> {
let database = self.database.lock()?;
if self.connection_options.is_empty() {
database.new_connection().map_err(AdbcError::from)
} else {
database
.new_connection_with_opts(
self.connection_options
.iter()
.map(|(k, v)| (OptionConnection::from(k.as_str()), OptionValue::from(v.as_str()))),
)
.map_err(AdbcError::from)
}
}
fn is_valid(&self, conn: &mut Self::Connection) -> Result<(), Self::Error> {
conn.new_statement().map(|_| ()).map_err(AdbcError::from)
}
fn has_broken(&self, _conn: &mut Self::Connection) -> bool {
false
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_error_display() {
use adbc_core::error::{Error, Status};
let adbc_err = Error::with_message_and_status("test error", Status::Internal);
let wrapped_err = AdbcError(adbc_err);
let display = format!("{}", wrapped_err);
assert!(display.contains("ADBC error"));
}
#[test]
fn test_error_source() {
use adbc_core::error::{Error, Status};
let adbc_err = Error::with_message_and_status("test error", Status::Internal);
let wrapped_err = AdbcError(adbc_err);
assert!(wrapped_err.source().is_some());
}
}