use std::future::Future;
use std::sync::OnceLock;
use std::sync::mpsc;
use tokio::runtime::Handle;
use crate::engine::{EngineConn, Target};
use crate::{DatabaseError, Row, Value};
fn bridge() -> Result<&'static Handle, DatabaseError> {
static BRIDGE: OnceLock<Result<Handle, String>> = OnceLock::new();
match BRIDGE.get_or_init(start_bridge) {
Ok(handle) => Ok(handle),
Err(message) => Err(DatabaseError::Storage(message.clone())),
}
}
fn start_bridge() -> Result<Handle, String> {
let (tx, rx) = mpsc::channel();
std::thread::Builder::new()
.name("frust-turso-db".to_string())
.spawn(move || {
let runtime = match tokio::runtime::Builder::new_current_thread().build() {
Ok(runtime) => runtime,
Err(e) => {
let _ = tx.send(Err(format!(
"could not start the turso bridge runtime: {e}"
)));
return;
}
};
let _ = tx.send(Ok(runtime.handle().clone()));
runtime.block_on(std::future::pending::<()>());
})
.map_err(|e| format!("could not start the turso bridge thread: {e}"))?;
rx.recv()
.map_err(|_| "the turso bridge thread stopped before reporting its runtime".to_string())?
}
fn in_async_context() -> bool {
Handle::try_current().is_ok() && tokio::task::try_id().is_none()
}
fn run<T>(fut: impl Future<Output = T> + Send + 'static) -> Result<T, DatabaseError>
where
T: Send + 'static,
{
if in_async_context() {
return Err(DatabaseError::AsyncContext);
}
let handle = bridge()?;
let (tx, rx) = mpsc::channel();
handle.spawn(async move {
let _ = tx.send(fut.await);
});
rx.recv().map_err(|_| {
DatabaseError::Storage(
"the turso bridge dropped a database call without producing a result".to_string(),
)
})
}
fn open_err(e: turso::Error) -> DatabaseError {
let is_storage_level = matches!(
e,
turso::Error::IoError(..) | turso::Error::NotAdb(_) | turso::Error::Corrupt(_)
);
if is_storage_level {
DatabaseError::Storage(e.to_string())
} else {
sql_err(e)
}
}
fn sql_err(e: turso::Error) -> DatabaseError {
DatabaseError::Sql {
message: e.to_string(),
}
}
fn to_turso_value(v: &Value) -> turso::Value {
match v {
Value::Null => turso::Value::Null,
Value::Integer(i) => turso::Value::Integer(*i),
Value::Real(f) => turso::Value::Real(*f),
Value::Text(s) => turso::Value::Text(s.clone()),
Value::Blob(b) => turso::Value::Blob(b.clone()),
}
}
fn from_turso_value(v: turso::Value) -> Value {
match v {
turso::Value::Null => Value::Null,
turso::Value::Integer(i) => Value::Integer(i),
turso::Value::Real(f) => Value::Real(f),
turso::Value::Text(s) => Value::Text(s),
turso::Value::Blob(b) => Value::Blob(b),
}
}
fn to_turso_params(params: &[Value]) -> Vec<turso::Value> {
params.iter().map(to_turso_value).collect()
}
pub(crate) struct TursoConn {
conn: turso::Connection,
}
pub(crate) fn open(target: Target) -> Result<TursoConn, DatabaseError> {
let (path, is_file) = match &target {
Target::Memory => (":memory:".to_string(), false),
Target::Path(path) => {
let path = path.to_str().ok_or_else(|| {
DatabaseError::Storage(format!("database path {path:?} is not valid UTF-8"))
})?;
(path.to_string(), true)
}
};
let conn = run(async move {
let db = turso::Builder::new_local(&path).build().await?;
db.connect()
})?
.map_err(open_err)?;
let mut conn = TursoConn { conn };
if is_file {
assert_wal(&mut conn)?;
}
conn.query("PRAGMA foreign_keys = ON", &[])?;
Ok(conn)
}
fn assert_wal(conn: &mut TursoConn) -> Result<(), DatabaseError> {
let rows = conn.query("PRAGMA journal_mode", &[])?;
let mode = match rows.first().and_then(|row| row.get(0)) {
Some(Value::Text(mode)) => mode.to_lowercase(),
other => {
return Err(DatabaseError::Storage(format!(
"turso reported an unreadable journal_mode: {other:?}"
)));
}
};
if mode == "wal" {
Ok(())
} else {
Err(DatabaseError::Storage(format!(
"database file is in {mode} journal mode; frust-database requires WAL"
)))
}
}
impl EngineConn for TursoConn {
fn execute(&mut self, sql: &str, params: &[Value]) -> Result<u64, DatabaseError> {
let conn = self.conn.clone();
let sql = sql.to_string();
let params = to_turso_params(params);
run(async move { conn.execute(sql, params).await })?.map_err(sql_err)
}
fn query(&mut self, sql: &str, params: &[Value]) -> Result<Vec<Row>, DatabaseError> {
let conn = self.conn.clone();
let sql = sql.to_string();
let params = to_turso_params(params);
run(async move {
let mut rows = conn.query(sql, params).await?;
let columns = rows.column_names();
let mut out = Vec::new();
while let Some(row) = rows.next().await? {
let mut values = Vec::with_capacity(columns.len());
for i in 0..columns.len() {
values.push(from_turso_value(row.get_value(i)?));
}
out.push(Row::new(columns.clone(), values));
}
Ok::<_, turso::Error>(out)
})?
.map_err(sql_err)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn open_memory() -> TursoConn {
open(Target::Memory).expect("in-memory open should succeed")
}
fn scratch_dir(case: &str) -> std::path::PathBuf {
let dir = std::env::temp_dir().join(format!(
"frust-database-turso-{case}-{}-{}",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::create_dir_all(&dir).unwrap();
dir
}
#[test]
fn open_in_memory_works() {
let mut conn = open_memory();
let affected = conn
.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, name TEXT)", &[])
.unwrap();
assert_eq!(affected, 0);
}
#[test]
fn file_open_round_trips_and_reports_wal() {
let dir = scratch_dir("file-open");
let path = dir.join("wal.db");
let mut conn = open(Target::Path(path)).expect("file open should succeed");
conn.execute("CREATE TABLE t (name TEXT)", &[]).unwrap();
conn.execute(
"INSERT INTO t (name) VALUES (?1)",
&[Value::Text("a".to_string())],
)
.unwrap();
let rows = conn.query("PRAGMA journal_mode", &[]).unwrap();
let mode = match rows[0].get(0).unwrap() {
Value::Text(s) => s.to_lowercase(),
other => panic!("expected a text journal_mode, got {other:?}"),
};
assert_eq!(mode, "wal");
let rows = conn.query("SELECT name FROM t", &[]).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].get(0), Some(&Value::Text("a".to_string())));
drop(conn);
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn open_in_missing_directory_is_storage() {
let dir = scratch_dir("missing-dir");
let path = dir.join("no-such-dir").join("x.db");
match open(Target::Path(path)) {
Err(DatabaseError::Storage(_)) => {}
Err(other) => panic!("expected DatabaseError::Storage, got {other:?}"),
Ok(_) => panic!("opening under a missing directory should fail"),
}
let _ = std::fs::remove_dir_all(&dir);
}
#[test]
fn execute_reports_rows_affected() {
let mut conn = open_memory();
conn.execute("CREATE TABLE t (id INTEGER PRIMARY KEY, name TEXT)", &[])
.unwrap();
let inserted = conn
.execute(
"INSERT INTO t (name) VALUES (?1), (?2), (?3)",
&[
Value::Text("a".to_string()),
Value::Text("b".to_string()),
Value::Text("c".to_string()),
],
)
.unwrap();
assert_eq!(inserted, 3);
let updated = conn
.execute(
"UPDATE t SET name = ?1 WHERE name = ?2",
&[Value::Text("z".to_string()), Value::Text("a".to_string())],
)
.unwrap();
assert_eq!(updated, 1);
let deleted = conn.execute("DELETE FROM t", &[]).unwrap();
assert_eq!(deleted, 3);
}
#[test]
fn concurrent_callers_share_one_bridge() {
let threads: Vec<_> = (0..4)
.map(|i| {
std::thread::spawn(move || {
let mut conn = open_memory();
conn.execute("CREATE TABLE t (v INTEGER)", &[]).unwrap();
conn.execute("INSERT INTO t (v) VALUES (?1)", &[Value::Integer(i)])
.unwrap();
let rows = conn.query("SELECT v FROM t", &[]).unwrap();
assert_eq!(rows[0].get(0), Some(&Value::Integer(i)));
})
})
.collect();
for thread in threads {
thread.join().expect("every bridged caller should finish");
}
}
#[test]
fn begin_commit_and_rollback_run_through_the_seam() {
let mut conn = open_memory();
conn.execute("CREATE TABLE t (name TEXT)", &[]).unwrap();
conn.execute("BEGIN", &[]).unwrap();
conn.execute(
"INSERT INTO t (name) VALUES (?1)",
&[Value::Text("kept".to_string())],
)
.unwrap();
conn.execute("COMMIT", &[]).unwrap();
conn.execute("BEGIN", &[]).unwrap();
conn.execute(
"INSERT INTO t (name) VALUES (?1)",
&[Value::Text("dropped".to_string())],
)
.unwrap();
conn.execute("ROLLBACK", &[]).unwrap();
let rows = conn.query("SELECT name FROM t", &[]).unwrap();
assert_eq!(rows.len(), 1);
assert_eq!(rows[0].get(0), Some(&Value::Text("kept".to_string())));
}
#[test]
fn query_round_trips_every_value_variant() {
let mut conn = open_memory();
conn.execute(
"CREATE TABLE t (i INTEGER, r REAL, t TEXT, b BLOB, n TEXT)",
&[],
)
.unwrap();
conn.execute(
"INSERT INTO t (i, r, t, b, n) VALUES (?1, ?2, ?3, ?4, ?5)",
&[
Value::Integer(42),
Value::Real(3.5),
Value::Text("hi".to_string()),
Value::Blob(vec![1, 2, 3]),
Value::Null,
],
)
.unwrap();
let rows = conn.query("SELECT i, r, t, b, n FROM t", &[]).unwrap();
assert_eq!(rows.len(), 1);
let row = &rows[0];
assert_eq!(row.get(0), Some(&Value::Integer(42)));
assert_eq!(row.get(1), Some(&Value::Real(3.5)));
assert_eq!(row.get(2), Some(&Value::Text("hi".to_string())));
assert_eq!(row.get(3), Some(&Value::Blob(vec![1, 2, 3])));
assert_eq!(row.get(4), Some(&Value::Null));
assert_eq!(row.get_named("t"), Some(&Value::Text("hi".to_string())));
}
#[test]
fn integer_bounds_and_empty_blob_round_trip() {
let mut conn = open_memory();
conn.execute("CREATE TABLE t (v)", &[]).unwrap();
for value in [
Value::Integer(i64::MIN),
Value::Integer(i64::MAX),
Value::Blob(Vec::new()),
Value::Blob(vec![0, 255, 128]),
Value::Text(String::new()),
Value::Real(f64::MIN_POSITIVE),
] {
conn.execute("DELETE FROM t", &[]).unwrap();
conn.execute(
"INSERT INTO t (v) VALUES (?1)",
std::slice::from_ref(&value),
)
.unwrap();
let rows = conn.query("SELECT v FROM t", &[]).unwrap();
assert_eq!(rows[0].get(0), Some(&value));
}
}
#[test]
fn sql_error_surfaces_message_text() {
let mut conn = open_memory();
let err = conn.execute("NOT VALID SQL", &[]).unwrap_err();
match err {
DatabaseError::Sql { message } => {
assert!(
!message.is_empty(),
"expected the engine's own message, got an empty string"
);
}
other => panic!("expected DatabaseError::Sql, got {other:?}"),
}
}
#[test]
fn query_error_surfaces_message_text() {
let mut conn = open_memory();
let err = conn.query("SELECT * FROM no_such_table", &[]).unwrap_err();
match err {
DatabaseError::Sql { message } => {
assert!(
message.contains("no_such_table"),
"expected the table name in {message:?}"
);
}
other => panic!("expected DatabaseError::Sql, got {other:?}"),
}
}
#[test]
fn async_context_is_reported_from_a_runtime_block_on() {
let mut conn = open_memory();
let runtime = tokio::runtime::Builder::new_current_thread()
.build()
.unwrap();
let err = runtime.block_on(async { conn.execute("SELECT 1", &[]).unwrap_err() });
assert!(
matches!(err, DatabaseError::AsyncContext),
"expected DatabaseError::AsyncContext, got {err:?}"
);
}
#[test]
fn spawn_blocking_is_not_treated_as_an_async_context() {
let runtime = tokio::runtime::Builder::new_current_thread()
.build()
.unwrap();
let rows = runtime.block_on(async {
tokio::task::spawn_blocking(|| {
let mut conn = open_memory();
conn.query("SELECT 1 AS one", &[])
})
.await
.unwrap()
});
let rows = rows.expect("a spawn_blocking caller must not be refused");
assert_eq!(rows[0].get(0), Some(&Value::Integer(1)));
}
#[test]
fn open_reports_async_context_too() {
let runtime = tokio::runtime::Builder::new_current_thread()
.build()
.unwrap();
runtime.block_on(async {
match open(Target::Memory) {
Err(DatabaseError::AsyncContext) => {}
Err(other) => panic!("expected DatabaseError::AsyncContext, got {other:?}"),
Ok(_) => panic!("an async-context open should be refused"),
}
});
}
}