use std::fmt;
use std::future::Future;
use std::ops::Deref;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use crate::error::{Error, Result};
use crate::executor::Conn;
use crate::options::{ConnectOptions, Source};
use turso_sql::Statement;
enum Engine {
Local(turso::Database),
#[cfg(feature = "sync")]
Sync(turso::sync::Database),
#[cfg(feature = "serverless")]
Remote(turso_serverless::Database),
}
pub(crate) struct Inner {
engine: Engine,
pub(crate) options: ConnectOptions,
idle: Mutex<Vec<Conn>>,
permits: Arc<Semaphore>,
}
#[derive(Clone)]
pub struct Database {
pub(crate) inner: Arc<Inner>,
}
impl fmt::Debug for Database {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Database")
.field("source", &self.inner.options.source)
.field("max_connections", &self.inner.options.max_connections)
.finish_non_exhaustive()
}
}
impl Database {
pub async fn connect(options: impl Into<ConnectOptions>) -> Result<Self> {
let options = options.into();
let engine = match &options.source {
Source::Memory => Engine::Local(options.local_builder(":memory:").build().await?),
Source::File(path) => Engine::Local(
options
.local_builder(&path.to_string_lossy())
.build()
.await?,
),
#[cfg(feature = "sync")]
Source::Sync(sync) => {
if options.encryption.is_some() {
return Err(Error::InvalidOptions(
"local encryption is not supported together with sync".into(),
));
}
let mut builder = turso::sync::Builder::new_remote(&sync.path.to_string_lossy())
.with_remote_url(&sync.remote_url)
.bootstrap_if_empty(sync.bootstrap_if_empty);
if let Some(token) = &sync.auth_token {
builder = builder.with_auth_token(token);
}
Engine::Sync(builder.build().await?)
}
#[cfg(feature = "serverless")]
Source::Remote(remote) => {
if options.encryption.is_some() {
return Err(Error::InvalidOptions(
"local encryption does not apply to a remote database".into(),
));
}
let mut builder = turso_serverless::Builder::new_remote(remote.url.clone());
if let Some(token) = &remote.auth_token {
builder = builder.with_auth_token(token.clone());
}
if let Some(key) = &remote.remote_encryption_key {
builder = builder.with_remote_encryption_key(key.clone());
}
Engine::Remote(builder.build().await?)
}
};
let db = Self {
inner: Arc::new(Inner {
engine,
permits: Arc::new(Semaphore::new(options.max_connections)),
idle: Mutex::new(Vec::with_capacity(options.max_connections)),
options,
}),
};
drop(db.acquire().await?);
Ok(db)
}
pub fn options(&self) -> &ConnectOptions {
&self.inner.options
}
pub async fn ping(&self) -> Result<()> {
let conn = self.acquire().await?;
crate::executor::query_one(&conn, &Statement::from_string("SELECT 1")).await?;
Ok(())
}
#[cfg(feature = "sync")]
#[cfg_attr(docsrs, doc(cfg(feature = "sync")))]
pub async fn push(&self) -> Result<()> {
match &self.inner.engine {
Engine::Sync(db) => Ok(db.push().await?),
_ => Err(Error::InvalidOptions("not an embedded replica".into())),
}
}
#[cfg(feature = "sync")]
#[cfg_attr(docsrs, doc(cfg(feature = "sync")))]
pub async fn pull(&self) -> Result<bool> {
match &self.inner.engine {
Engine::Sync(db) => Ok(db.pull().await?),
_ => Err(Error::InvalidOptions("not an embedded replica".into())),
}
}
pub(crate) async fn acquire(&self) -> Result<PooledConnection> {
let permit = tokio::time::timeout(
self.inner.options.acquire_timeout,
Arc::clone(&self.inner.permits).acquire_owned(),
)
.await
.map_err(|_| Error::PoolTimeout)?
.map_err(|_| Error::Misuse("connection pool closed".into()))?;
let idle = self
.inner
.idle
.lock()
.map_err(|_| Error::Misuse("pool mutex poisoned".into()))?
.pop();
let conn = match idle {
Some(conn) => conn,
None => self.open_connection().await?,
};
Ok(PooledConnection {
conn: Some(conn),
pool: Arc::clone(&self.inner),
_permit: permit,
discard: AtomicBool::new(false),
})
}
async fn open_connection(&self) -> Result<Conn> {
let conn = match &self.inner.engine {
Engine::Local(db) => Conn::Embedded(db.connect()?),
#[cfg(feature = "sync")]
Engine::Sync(db) => Conn::Embedded(db.connect().await?),
#[cfg(feature = "serverless")]
Engine::Remote(db) => Conn::Remote(db.connect()?),
};
let options = &self.inner.options;
if let Some(timeout) = options.busy_timeout {
conn.busy_timeout(timeout)?;
}
if options.foreign_keys {
conn.pragma_update("foreign_keys", "ON").await?;
}
if options.mvcc && matches!(conn, Conn::Embedded(_)) {
conn.pragma_update("journal_mode", "'mvcc'").await?;
}
for (name, value) in &options.pragmas {
conn.pragma_update(name, value).await?;
}
Ok(conn)
}
}
pub(crate) struct PooledConnection {
conn: Option<Conn>,
pool: Arc<Inner>,
_permit: OwnedSemaphorePermit,
discard: AtomicBool,
}
impl PooledConnection {
pub(crate) fn discard(&self) {
self.discard.store(true, Ordering::Release);
}
pub(crate) fn options(&self) -> &ConnectOptions {
&self.pool.options
}
}
impl Deref for PooledConnection {
type Target = Conn;
fn deref(&self) -> &Self::Target {
self.conn.as_ref().expect("connection present until drop")
}
}
impl Drop for PooledConnection {
fn drop(&mut self) {
if let Some(conn) = self.conn.take()
&& !self.discard.load(Ordering::Acquire)
&& conn.is_autocommit().unwrap_or(false)
&& let Ok(mut idle) = self.pool.idle.lock()
{
idle.push(conn);
}
}
}
impl fmt::Debug for PooledConnection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("PooledConnection").finish_non_exhaustive()
}
}
pub(crate) async fn retry_busy<T, F, Fut>(budget: Duration, mut op: F) -> Result<T>
where
F: FnMut() -> Fut,
Fut: Future<Output = Result<T>>,
{
let start = std::time::Instant::now();
let mut delay = Duration::from_millis(5);
loop {
match op().await {
Err(err) if err.is_busy() && start.elapsed() < budget => {
tracing::debug!(%err, ?delay, "busy, retrying");
tokio::time::sleep(delay).await;
delay = (delay * 2).min(Duration::from_millis(250));
}
other => return other,
}
}
}