use std::collections::HashMap;
use std::future::Future;
use std::path::PathBuf;
use std::pin::Pin;
use std::sync::{Arc, RwLock};
use std::time::Duration;
use pylon_core::ir::SessionConfig;
use pylon_core::schema::SchemaDescriptor;
use pylon_value::DecodedValue;
use tokio::sync::OnceCell;
use crate::error::{Error, Result};
use crate::exec;
use crate::query_arg::QueryArgs;
use crate::queryable::{Queryable, decode_optional_row, decode_row, decode_rows};
use crate::schema;
use crate::transaction::{Isolation, Transaction};
pub type TxFuture<'a, T> = Pin<Box<dyn Future<Output = Result<T>> + Send + 'a>>;
enum CacheSource {
Open { path: PathBuf, max_size_mb: usize },
Shared(Arc<pylon_cache::Cache>),
}
pub struct Builder {
dsn: String,
max_pool_size: usize,
cache: Option<CacheSource>,
}
impl Builder {
pub fn new(dsn: impl Into<String>) -> Self {
Self {
dsn: dsn.into(),
max_pool_size: 10,
cache: None,
}
}
pub fn max_pool_size(mut self, max_pool_size: usize) -> Self {
self.max_pool_size = max_pool_size;
self
}
pub fn cache(mut self, path: impl Into<PathBuf>, max_size_mb: usize) -> Self {
self.cache = Some(CacheSource::Open {
path: path.into(),
max_size_mb,
});
self
}
pub fn cache_handle(mut self, cache: Arc<pylon_cache::Cache>) -> Self {
self.cache = Some(CacheSource::Shared(cache));
self
}
pub fn build(self) -> Result<Client> {
let cache = match self.cache {
None => None,
Some(CacheSource::Open { path, max_size_mb }) => Some(Arc::new(
pylon_cache::Cache::open(&path, max_size_mb).map_err(|e| Error::Cache(e.to_string()))?,
)),
Some(CacheSource::Shared(cache)) => Some(cache),
};
Ok(Client {
dsn: Arc::new(self.dsn),
max_pool_size: self.max_pool_size,
connected: Arc::new(OnceCell::new()),
globals: Arc::new(HashMap::new()),
config: SessionConfig::default(),
cache,
})
}
}
struct Connected {
pool: pylon_pgcon::PgPool,
schema: Arc<RwLock<SchemaDescriptor>>,
}
#[derive(Clone)]
pub struct Client {
dsn: Arc<String>,
max_pool_size: usize,
connected: Arc<OnceCell<Connected>>,
globals: Arc<HashMap<String, DecodedValue>>,
config: SessionConfig,
cache: Option<Arc<pylon_cache::Cache>>,
}
impl Client {
pub fn builder(dsn: impl Into<String>) -> Builder {
Builder::new(dsn)
}
async fn connected(&self) -> Result<&Connected> {
self.connected
.get_or_try_init(|| async {
let pool = pylon_pgcon::PgPool::connect(&self.dsn, self.max_pool_size)
.await
.map_err(Error::Db)?;
let schema = schema::fetch(&pool).await?;
Ok(Connected {
pool,
schema: Arc::new(RwLock::new(schema)),
})
})
.await
}
pub async fn ensure_connected(&self) -> Result<()> {
self.connected().await.map(|_| ())
}
pub async fn reload_schema(&self) -> Result<()> {
let conn = self.connected().await?;
let fresh = schema::fetch(&conn.pool).await?;
*conn.schema.write().unwrap() = fresh;
pylon_core::query::clear_query_cache();
Ok(())
}
pub fn with_globals(&self, globals: impl IntoIterator<Item = (String, DecodedValue)>) -> Client {
let mut merged = (*self.globals).clone();
merged.extend(globals);
Client {
globals: Arc::new(merged),
..self.clone()
}
}
pub fn with_config(&self, config: SessionConfig) -> Client {
Client { config, ..self.clone() }
}
pub async fn raw_connection(&self) -> Result<&pylon_pgcon::PgPool> {
Ok(&self.connected().await?.pool)
}
pub fn pool_if_connected(&self) -> Option<&pylon_pgcon::PgPool> {
self.connected.get().map(|conn| &conn.pool)
}
pub async fn schema(&self) -> Result<SchemaDescriptor> {
Ok(self.connected().await?.schema.read().unwrap().clone())
}
pub fn cache_handle(&self) -> Option<Arc<pylon_cache::Cache>> {
self.cache.clone()
}
pub async fn query<R: Queryable, A: QueryArgs + ?Sized>(&self, pyql: &str, args: &A) -> Result<Vec<R>> {
let params = args.to_params();
let conn = self.connected().await?;
let schema = conn.schema.read().unwrap().clone();
let values = exec::query(
&conn.pool,
pyql,
¶ms,
&schema,
&self.config,
&self.globals,
crate::cache::CacheAccess::read_write(self.cache.as_deref()),
)
.await?;
decode_rows(values)
}
pub async fn query_single<R: Queryable, A: QueryArgs + ?Sized>(&self, pyql: &str, args: &A) -> Result<Option<R>> {
let params = args.to_params();
let conn = self.connected().await?;
let schema = conn.schema.read().unwrap().clone();
let values = exec::query_single(
&conn.pool,
pyql,
¶ms,
&schema,
&self.config,
&self.globals,
crate::cache::CacheAccess::read_write(self.cache.as_deref()),
)
.await?;
decode_optional_row(values)
}
pub async fn query_required_single<R: Queryable, A: QueryArgs + ?Sized>(&self, pyql: &str, args: &A) -> Result<R> {
let params = args.to_params();
let conn = self.connected().await?;
let schema = conn.schema.read().unwrap().clone();
let values = exec::query_required_single(
&conn.pool,
pyql,
¶ms,
&schema,
&self.config,
&self.globals,
crate::cache::CacheAccess::read_write(self.cache.as_deref()),
)
.await?;
decode_row(values)
}
pub async fn execute<A: QueryArgs + ?Sized>(&self, pyql: &str, args: &A) -> Result<()> {
let params = args.to_params();
let conn = self.connected().await?;
let schema = conn.schema.read().unwrap().clone();
exec::execute(
&conn.pool,
pyql,
¶ms,
&schema,
&self.config,
&self.globals,
crate::cache::CacheAccess::read_write(self.cache.as_deref()),
)
.await
}
pub async fn listen(&self, channel: &str) -> Result<crate::ChannelListener> {
let conn = self.connected().await?;
let schema = conn.schema.read().unwrap().clone();
crate::listen::listen(&self.dsn, &schema, channel).await
}
pub async fn query_json<A: QueryArgs + ?Sized>(&self, pyql: &str, args: &A) -> Result<String> {
let params = args.to_params();
let conn = self.connected().await?;
let schema = conn.schema.read().unwrap().clone();
exec::query_json(
&conn.pool,
pyql,
¶ms,
&schema,
&self.config,
&self.globals,
crate::cache::CacheAccess::read_write(self.cache.as_deref()),
)
.await
}
pub async fn query_single_json<A: QueryArgs + ?Sized>(&self, pyql: &str, args: &A) -> Result<Option<String>> {
let params = args.to_params();
let conn = self.connected().await?;
let schema = conn.schema.read().unwrap().clone();
exec::query_single_json(
&conn.pool,
pyql,
¶ms,
&schema,
&self.config,
&self.globals,
crate::cache::CacheAccess::read_write(self.cache.as_deref()),
)
.await
}
pub async fn query_required_single_json<A: QueryArgs + ?Sized>(&self, pyql: &str, args: &A) -> Result<String> {
let params = args.to_params();
let conn = self.connected().await?;
let schema = conn.schema.read().unwrap().clone();
exec::query_required_single_json(
&conn.pool,
pyql,
¶ms,
&schema,
&self.config,
&self.globals,
crate::cache::CacheAccess::read_write(self.cache.as_deref()),
)
.await
}
pub fn cache_stat(&self) -> Result<Option<pylon_cache::CacheStats>> {
match &self.cache {
None => Ok(None),
Some(cache) => cache.stat().map(Some).map_err(|e| Error::Cache(e.to_string())),
}
}
pub fn cache_clear(&self) -> Result<()> {
match &self.cache {
None => Ok(()),
Some(cache) => cache.clear().map_err(|e| Error::Cache(e.to_string())),
}
}
pub async fn analyze<A: QueryArgs + ?Sized>(&self, pyql: &str, args: &A) -> Result<String> {
let params = args.to_params();
let conn = self.connected().await?;
let schema = conn.schema.read().unwrap().clone();
exec::analyze(&conn.pool, pyql, ¶ms, &schema, &self.config, &self.globals).await
}
pub async fn transaction<T, F>(&self, isolation: Isolation, body: F) -> Result<T>
where
F: for<'a> FnMut(&'a Transaction) -> TxFuture<'a, T>,
{
self.transaction_with_attempts(isolation, 3, body).await
}
pub async fn transaction_opt<T, F>(&self, isolation: Isolation, body: F) -> Result<Option<T>>
where
F: for<'a> FnMut(&'a Transaction) -> TxFuture<'a, T>,
{
self.transaction_opt_with_attempts(isolation, 3, body).await
}
pub async fn transaction_opt_with_attempts<T, F>(
&self,
isolation: Isolation,
max_attempts: u32,
body: F,
) -> Result<Option<T>>
where
F: for<'a> FnMut(&'a Transaction) -> TxFuture<'a, T>,
{
match self.transaction_with_attempts(isolation, max_attempts, body).await {
Ok(value) => Ok(Some(value)),
Err(Error::Rollback) => Ok(None),
Err(e) => Err(e),
}
}
pub async fn transaction_with_attempts<T, F>(
&self,
isolation: Isolation,
max_attempts: u32,
mut body: F,
) -> Result<T>
where
F: for<'a> FnMut(&'a Transaction) -> TxFuture<'a, T>,
{
let conn = self.connected().await?;
let mut attempt = 0u32;
loop {
attempt += 1;
if attempt > 1 {
tokio::time::sleep(Duration::from_millis(100 * u64::from(attempt - 1))).await;
}
let pg_tx = conn.pool.begin(isolation.as_str()).await.map_err(Error::Db)?;
let tx = Transaction {
inner: pg_tx,
schema: conn.schema.clone(),
config: self.config.clone(),
globals: self.globals.clone(),
cache: self.cache.clone(),
};
let result = body(&tx).await;
match result {
Ok(value) => {
tx.inner.commit().await.map_err(Error::Db)?;
return Ok(value);
}
Err(e) if e.is_retriable() && attempt < max_attempts => {
let _ = tx.inner.rollback().await;
}
Err(e) => {
let _ = tx.inner.rollback().await;
return Err(e);
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const UNREACHABLE_DSN: &str = "postgresql://nobody@127.0.0.1:1/nothing";
#[test]
fn build_does_not_connect() {
let client = Client::builder(UNREACHABLE_DSN).build().unwrap();
assert!(client.pool_if_connected().is_none());
}
#[tokio::test]
async fn a_failed_connect_is_retried_rather_than_remembered() {
let client = Client::builder(UNREACHABLE_DSN).build().unwrap();
assert!(client.ensure_connected().await.is_err());
assert!(client.pool_if_connected().is_none());
assert!(client.ensure_connected().await.is_err());
}
#[test]
fn views_share_the_connection_slot() {
let client = Client::builder(UNREACHABLE_DSN).build().unwrap();
let view = client.with_globals([]);
assert!(Arc::ptr_eq(&client.connected, &view.connected));
}
}