use crate::Result;
use elefant_client::tokio_connection::TokioPostgresPool;
use elefant_client::{
CollectBatch, ElefantClientError, FlattenTuple, FromSqlOwned, FromSqlRowOwned,
PostgresConnectionSettings, PostgresDataRow,
};
use tracing::instrument;
pub struct PostgresClientWrapper {
pool: TokioPostgresPool,
version: i32,
}
impl PostgresClientWrapper {
#[instrument(skip_all)]
pub async fn new(settings: PostgresConnectionSettings) -> Result<Self> {
let pool = TokioPostgresPool::new(
elefant_client::tokio_connection::TokioConnectionFactory,
settings,
)
.await?;
let client = pool.get_client().await?;
let version_str = client
.get_parameter("server_version")
.ok_or(crate::ElefantToolsError::InvalidPostgresVersionResponse)?;
let major_version: i32 = version_str
.split('.')
.next()
.and_then(|s| s.parse().ok())
.ok_or(crate::ElefantToolsError::InvalidPostgresVersionResponse)?;
if major_version < 12 {
return Err(crate::ElefantToolsError::UnsupportedPostgresVersion(
version_str.to_string(),
));
}
let version = major_version * 10;
Ok(PostgresClientWrapper { pool, version })
}
pub fn version(&self) -> i32 {
self.version
}
pub fn pool(&self) -> &TokioPostgresPool {
&self.pool
}
pub async fn execute_non_query(&self, sql: &str) -> Result {
let mut client = self.pool.get_client().await?;
client.execute_non_query_simple(sql).await.map_err(|e| {
crate::ElefantToolsError::PostgresErrorWithQuery {
source: e,
query: sql.to_string(),
}
})?;
Ok(())
}
pub async fn get_results<T: FromSqlRowOwned>(&self, sql: &str) -> Result<Vec<T>> {
let mut client = self.pool.get_client().await?;
let query_result = client.query_simple(sql).await.map_err(|e| {
crate::ElefantToolsError::PostgresErrorWithQuery {
source: e,
query: sql.to_string(),
}
})?;
let rows = query_result.collect_to_vec::<T>().await.map_err(|e| {
crate::ElefantToolsError::PostgresErrorWithQuery {
source: e,
query: sql.to_string(),
}
})?;
Ok(rows)
}
pub async fn get_result<T: FromSqlRowOwned>(&self, sql: &str) -> Result<T> {
let results = self.get_results(sql).await?;
if results.len() != 1 {
return Err(crate::ElefantToolsError::InvalidNumberOfResults {
actual: results.len(),
expected: 1,
});
}
let r = results.into_iter().next().unwrap();
Ok(r)
}
pub async fn get_single_results<T: FromSqlOwned>(&self, sql: &str) -> Result<Vec<T>> {
let r = self
.get_results::<(T,)>(sql)
.await?
.into_iter()
.map(|t| t.0)
.collect();
Ok(r)
}
pub async fn get_single_result<T: FromSqlOwned>(&self, sql: &str) -> Result<T> {
let result = self.get_result::<(T,)>(sql).await?;
Ok(result.0)
}
}
pub(crate) trait FromPgChar: Sized {
fn from_pg_char(c: char) -> std::result::Result<Self, crate::ElefantToolsError>;
}
pub(crate) trait RowEnumExt {
fn try_get_enum_value<T: FromPgChar>(
&self,
idx: usize,
) -> std::result::Result<T, ElefantClientError>;
fn try_get_opt_enum_value<T: FromPgChar>(
&self,
idx: usize,
) -> std::result::Result<Option<T>, ElefantClientError>;
}
impl RowEnumExt for PostgresDataRow<'_, '_> {
fn try_get_enum_value<T: FromPgChar>(
&self,
idx: usize,
) -> std::result::Result<T, ElefantClientError> {
let c: char = self.get(idx)?;
T::from_pg_char(c).map_err(|e| ElefantClientError::PostgresError(e.to_string()))
}
fn try_get_opt_enum_value<T: FromPgChar>(
&self,
idx: usize,
) -> std::result::Result<Option<T>, ElefantClientError> {
let c: Option<char> = self.get(idx)?;
match c {
Some('\0') => Ok(None),
Some(c) => {
Ok(Some(T::from_pg_char(c).map_err(|e| {
ElefantClientError::PostgresError(e.to_string())
})?))
}
None => Ok(None),
}
}
}
pub(crate) trait QueryResult: FromSqlRowOwned {
fn query(version: i32) -> &'static str;
}
pub(crate) struct BatchQueryBuilder<Batch> {
query: String,
version: i32,
_batch: std::marker::PhantomData<Batch>,
}
impl BatchQueryBuilder<()> {
pub(crate) fn new(connection: &PostgresClientWrapper) -> Self {
Self {
query: String::new(),
version: connection.version(),
_batch: std::marker::PhantomData,
}
}
}
impl<Batch> BatchQueryBuilder<Batch> {
pub(crate) fn add<T: QueryResult>(mut self) -> BatchQueryBuilder<(Batch, Vec<T>)> {
self.query.push_str(T::query(self.version));
BatchQueryBuilder {
query: self.query,
version: self.version,
_batch: std::marker::PhantomData,
}
}
}
impl<Batch: CollectBatch + FlattenTuple> BatchQueryBuilder<Batch> {
pub(crate) async fn execute(self, connection: &PostgresClientWrapper) -> Result<Batch::Output> {
let mut client = connection.pool().get_client().await?;
let mut result = client.query_simple(&self.query).await?;
Ok(Batch::collect(&mut result).await?.flatten())
}
}