use alloc::boxed::Box;
use alloc::format;
use alloc::string::{String, ToString};
use alloc::sync::Arc;
use alloc::vec::Vec;
use core::future::Future;
use hashbrown::HashMap;
use turso::transaction::TransactionBehavior;
use super::{InsertSummary, Insertion, NamespaceSummary, Origin, Storage};
use crate::bytes::Bytes;
use crate::sync::{LazyLock, Mutex};
pub const SCHEMA_VERSION: u32 = 3;
pub(crate) const SCHEMA_VERSION_KEY: &str = "schema_version";
const CREATE_META: &str = "
CREATE TABLE IF NOT EXISTS meta (
k TEXT PRIMARY KEY,
v TEXT NOT NULL
)
";
const META_GET: &str = "SELECT v FROM meta WHERE k = ?1";
const META_SET: &str = "INSERT INTO meta (k, v) VALUES (?1, ?2) \
ON CONFLICT(k) DO UPDATE SET v = excluded.v";
const CREATE_ENTRIES: &str = "
CREATE TABLE IF NOT EXISTS entries (
namespace TEXT NOT NULL,
key BLOB NOT NULL,
value BLOB NOT NULL,
origin INTEGER NOT NULL,
PRIMARY KEY (namespace, key)
)
";
const DROP_ENTRIES: &str = "DROP TABLE IF EXISTS entries";
const INSERT: &str = "INSERT INTO entries (namespace, key, value, origin) \
VALUES (?1, ?2, ?3, ?4) \
ON CONFLICT(namespace, key) DO UPDATE \
SET value = excluded.value, origin = excluded.origin \
WHERE entries.origin = 1 AND excluded.origin = 0";
const REPLACE: &str = "INSERT INTO entries (namespace, key, value, origin) \
VALUES (?1, ?2, ?3, ?4) \
ON CONFLICT(namespace, key) DO UPDATE \
SET value = excluded.value, origin = excluded.origin";
const SELECT: &str = "SELECT value FROM entries WHERE namespace = ?1 AND key = ?2";
const SCAN: &str = "SELECT key, value FROM entries WHERE namespace = ?1";
const PURGE: &str = "DELETE FROM entries WHERE namespace = ?1";
const PURGE_KEY: &str = "DELETE FROM entries WHERE namespace = ?1 AND key = ?2";
const SUMMARY: &str = "SELECT namespace, COUNT(*), SUM(length(key) + length(value)) \
FROM entries GROUP BY namespace ORDER BY namespace";
type DatabaseResult = Result<Arc<turso::Database>, String>;
enum DatabaseState {
Opening(Vec<async_channel::Sender<DatabaseResult>>),
Ready(Arc<turso::Database>),
}
static DATABASES: LazyLock<Mutex<HashMap<String, DatabaseState>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
pub struct Database {
connection: Mutex<turso::Connection>,
location: String,
}
impl core::fmt::Debug for Database {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
formatter
.debug_struct("Database")
.field("location", &self.location)
.finish()
}
}
impl Database {
fn open() -> Result<Self, String> {
let location = location()?;
let database = database()?;
let connection = connect(&database).map_err(error)?;
Ok(Self {
connection: Mutex::new(connection),
location,
})
}
pub fn open_active() -> Option<Self> {
match Self::open() {
Ok(database) => Some(database),
Err(error) => {
log::warn!("Unable to open the Turso cache: {error}");
None
}
}
}
fn run<T>(
&self,
name: &str,
namespace: &str,
operation: impl FnOnce(&mut turso::Connection) -> Result<T, turso::Error>,
) -> Result<T, String> {
operation(&mut self.connection.lock()).map_err(|err| {
log::warn!(
"Unable to {name} {}: {err}",
describe(&self.location, namespace)
);
err.to_string()
})
}
pub fn get(&self, namespace: &str, key: &[u8]) -> Option<Bytes> {
self.run("read", namespace, |connection| {
let mut rows = drive(connection.query(SELECT, (namespace, key.to_vec())))?;
let Some(row) = drive(rows.next())? else {
return Ok(None);
};
let value: Vec<u8> = row.get(0)?;
Ok(Some(Bytes::from_bytes_vec(value)))
})
.unwrap_or_default()
}
pub fn insert(&self, namespace: &str, key: &[u8], value: &[u8], origin: Origin) -> Insertion {
self.run("write", namespace, |connection| {
insert_on(connection, namespace, key, value, origin)
})
.unwrap_or_else(Insertion::Failed)
}
pub fn replace(&self, namespace: &str, key: &[u8], value: &[u8], origin: Origin) -> Insertion {
self.run("replace", namespace, |connection| {
let params = (namespace, key.to_vec(), value.to_vec(), origin_code(origin));
drive(connection.execute(REPLACE, params))?;
Ok(Insertion::Stored)
})
.unwrap_or_else(Insertion::Failed)
}
pub fn insert_many(
&self,
namespace: &str,
entries: &mut dyn Iterator<Item = (Bytes, Bytes)>,
origin: Origin,
) -> InsertSummary {
let mut connection = self.connection.lock();
let transaction = match drive(connection.transaction_with_behavior(write_transaction())) {
Ok(transaction) => transaction,
Err(err) => {
log::warn!(
"Unable to batch write {}: {err}",
describe(&self.location, namespace)
);
return InsertSummary {
failed: entries.count(),
..InsertSummary::default()
};
}
};
let mut summary = InsertSummary::default();
for (key, value) in entries {
match insert_on(&transaction, namespace, &key, &value, origin) {
Ok(insertion) => summary.record(&insertion),
Err(_) => summary.failed += 1,
}
}
if let Err(err) = drive(transaction.commit()) {
log::warn!(
"Unable to commit a batch write to {}: {err}",
describe(&self.location, namespace)
);
summary.failed += summary.stored;
summary.stored = 0;
}
summary
}
pub fn scan(&self, namespace: &str, visit: &mut dyn FnMut(&[u8], &[u8])) {
let _ = self.run("scan", namespace, |connection| {
let mut rows = drive(connection.query(SCAN, (namespace,)))?;
while let Some(row) = drive(rows.next())? {
let key: Vec<u8> = row.get(0)?;
let value: Vec<u8> = row.get(1)?;
visit(&key, &value);
}
Ok(())
});
}
pub fn purge(&self, namespace: &str) {
let _ = self.run("purge", namespace, |connection| {
drive(connection.execute(PURGE, (namespace,))).map(|_| ())
});
}
pub fn purge_key(&self, namespace: &str, key: &[u8]) {
let _ = self.run("purge a key from", namespace, |connection| {
drive(connection.execute(PURGE_KEY, (namespace, key.to_vec()))).map(|_| ())
});
}
pub fn namespaces(&self) -> Vec<String> {
self.run("summarize", "", |connection| summarize(connection))
.unwrap_or_default()
.into_iter()
.map(|summary| summary.namespace)
.collect()
}
}
fn insert_on(
connection: &turso::Connection,
namespace: &str,
key: &[u8],
value: &[u8],
origin: Origin,
) -> Result<Insertion, turso::Error> {
let params = (namespace, key.to_vec(), value.to_vec(), origin_code(origin));
if drive(connection.execute(INSERT, params))? == 1 {
return Ok(Insertion::Stored);
}
let mut rows = drive(connection.query(SELECT, (namespace, key.to_vec())))?;
match drive(rows.next())? {
Some(row) => Ok(Insertion::Conflict(Bytes::from_bytes_vec(row.get(0)?))),
None => Ok(Insertion::Failed(
"the entry was removed while it was being written".to_string(),
)),
}
}
fn describe(location: &str, namespace: &str) -> String {
format!("Turso {location} ({namespace})")
}
pub struct TursoStorage {
database: Database,
namespace: String,
}
impl core::fmt::Debug for TursoStorage {
fn fmt(&self, formatter: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
formatter
.debug_struct("TursoStorage")
.field("namespace", &self.namespace)
.field("location", &self.database.location)
.finish()
}
}
impl TursoStorage {
pub fn open(namespace: String) -> Result<Self, String> {
Ok(Self {
database: Database::open()?,
namespace,
})
}
}
impl Storage for TursoStorage {
fn get(&self, key: &[u8]) -> Option<Bytes> {
self.database.get(&self.namespace, key)
}
fn insert(&self, key: &[u8], value: Bytes, origin: Origin) -> Insertion {
self.database.insert(&self.namespace, key, &value, origin)
}
fn replace(&self, key: &[u8], value: Bytes, origin: Origin) -> Insertion {
self.database.replace(&self.namespace, key, &value, origin)
}
fn insert_many(
&self,
entries: &mut dyn Iterator<Item = (Bytes, Bytes)>,
origin: Origin,
) -> InsertSummary {
self.database.insert_many(&self.namespace, entries, origin)
}
fn scan(&self, visit: &mut dyn FnMut(&[u8], &[u8])) {
self.database.scan(&self.namespace, visit);
}
fn purge(&self) {
self.database.purge(&self.namespace);
}
fn purge_key(&self, key: &[u8]) {
self.database.purge_key(&self.namespace, key);
}
fn describe(&self) -> String {
describe(&self.database.location, &self.namespace)
}
}
fn drive<T>(future: impl Future<Output = T>) -> T {
#[cfg(not(browser_cache))]
{
crate::future::block_on(future)
}
#[cfg(browser_cache)]
{
use core::pin::pin;
use core::task::{Context, Poll, Waker};
let mut future = pin!(future);
let mut context = Context::from_waker(Waker::noop());
loop {
if let Poll::Ready(output) = future.as_mut().poll(&mut context) {
return output;
}
}
}
}
pub(crate) fn database() -> DatabaseResult {
let location = location()?;
if let Some(DatabaseState::Ready(database)) = DATABASES.lock().get(&location) {
return Ok(database.clone());
}
#[cfg(native_cache)]
{
crate::future::block_on(shared_database(&location))
}
#[cfg(browser_cache)]
{
Err(format!(
"the database of environment '{}' is not open; await `environment::open()` \
before the first use of its caches",
crate::environment::active()
))
}
}
pub(crate) async fn open_ahead() -> Result<(), String> {
shared_database(&location()?).await.map(|_| ())
}
async fn shared_database(location: &str) -> DatabaseResult {
let receiver = {
let mut databases = DATABASES.lock();
match databases.get_mut(location) {
Some(DatabaseState::Ready(database)) => return Ok(database.clone()),
Some(DatabaseState::Opening(waiters)) => {
let (sender, receiver) = async_channel::bounded(1);
waiters.push(sender);
Some(receiver)
}
None => {
databases.insert(location.to_string(), DatabaseState::Opening(Vec::new()));
None
}
}
};
if let Some(receiver) = receiver {
return receiver
.recv()
.await
.map_err(|error| format!("database initialization was cancelled: {error}"))?;
}
let mut opener = Opener {
location,
result: Err("database initialization was cancelled".to_string()),
};
let writable = match open_database(location).await {
Ok(database) => migrate(&database).map(|()| database),
Err(err) => Err(err),
};
opener.result = match writable {
Ok(database) => Ok(Arc::new(database)),
Err(err @ (turso::Error::Busy(_) | turso::Error::BusySnapshot(_))) => Err(error(err)),
#[cfg(native_cache)]
Err(err) => open_read_only(location, &err.to_string())
.await
.map(Arc::new),
#[cfg(not(native_cache))]
Err(err) => Err(error(err)),
};
opener.result.clone()
}
struct Opener<'a> {
location: &'a str,
result: DatabaseResult,
}
impl Drop for Opener<'_> {
fn drop(&mut self) {
let waiters = {
let mut databases = DATABASES.lock();
let waiters = match databases.remove(self.location) {
Some(DatabaseState::Opening(waiters)) => waiters,
_ => Vec::new(),
};
if let Ok(database) = &self.result {
databases.insert(
self.location.to_string(),
DatabaseState::Ready(database.clone()),
);
}
waiters
};
for waiter in waiters {
let _ = waiter.try_send(self.result.clone());
}
}
}
#[cfg(native_cache)]
async fn open_database(location: &str) -> Result<turso::Database, turso::Error> {
if let Some(parent) = std::path::Path::new(location).parent() {
std::fs::create_dir_all(parent)
.map_err(|error| turso::Error::IoError(error.kind(), "creating the cache directory"))?;
}
turso::Builder::new_local(location)
.experimental_multiprocess_wal(true)
.build()
.await
}
#[cfg(native_cache)]
async fn open_read_only(location: &str, err: &str) -> Result<turso::Database, String> {
log::debug!("cubecl cache: {location} is not writable ({err}); opening read-only");
let database = turso::Builder::new_local(location)
.read_only(true)
.build()
.await
.map_err(error)?;
let connection = connect(&database).map_err(error)?;
let expected = SCHEMA_VERSION.to_string();
match meta_get(&connection, SCHEMA_VERSION_KEY).map_err(error)? {
Some(found) if found == expected => Ok(database),
found => Err(format!(
"read-only database at {location} has schema {found:?}, expected {expected}"
)),
}
}
#[cfg(browser_cache)]
async fn open_database(location: &str) -> Result<turso::Database, turso::Error> {
let wal = format!("{location}-wal");
let io = super::turso_browser::BrowserIo::new(&[location, &wal])
.await
.map_err(turso::Error::Error)?;
turso::Builder::new_local(location)
.with_io_impl(Arc::new(io))
.build()
.await
}
#[cfg(native_cache)]
fn location() -> Result<String, String> {
crate::environment::path()
.to_str()
.map(ToString::to_string)
.ok_or_else(|| "cache path is not valid UTF-8".to_string())
}
#[cfg(browser_cache)]
fn location() -> Result<String, String> {
Ok(format!("cubecl-{}.db", crate::environment::active()))
}
fn write_transaction() -> TransactionBehavior {
#[cfg(browser_cache)]
{
TransactionBehavior::Deferred
}
#[cfg(not(browser_cache))]
{
TransactionBehavior::Immediate
}
}
pub(crate) fn migrate(database: &turso::Database) -> Result<(), turso::Error> {
let mut connection = connect(database)?;
drive(connection.execute(CREATE_META, ()))?;
let transaction = drive(connection.transaction_with_behavior(write_transaction()))?;
let expected = SCHEMA_VERSION.to_string();
let found = meta_get(&transaction, SCHEMA_VERSION_KEY)?;
if found.as_deref() != Some(expected.as_str()) {
match &found {
Some(found) => log::warn!(
"cubecl cache: database schema {found} is not {expected}, discarding cached entries"
),
None => log::debug!("cubecl cache: initializing database schema {expected}"),
}
drive(transaction.execute(DROP_ENTRIES, ()))?;
meta_set(&transaction, SCHEMA_VERSION_KEY, &expected)?;
}
drive(transaction.execute(CREATE_ENTRIES, ()))?;
drive(transaction.commit())
}
pub(crate) fn meta_get(
connection: &turso::Connection,
key: &str,
) -> Result<Option<String>, turso::Error> {
let mut rows = drive(connection.query(META_GET, (key,)))?;
match drive(rows.next())? {
Some(row) => Ok(Some(row.get(0)?)),
None => Ok(None),
}
}
pub(crate) fn meta_set(
connection: &turso::Connection,
key: &str,
value: &str,
) -> Result<(), turso::Error> {
drive(connection.execute(META_SET, (key, value)))?;
Ok(())
}
pub(crate) fn summary() -> Vec<NamespaceSummary> {
let result = database().and_then(|database| {
let connection = connect(&database).map_err(error)?;
summarize(&connection).map_err(error)
});
result.unwrap_or_else(|err| {
log::warn!("Unable to summarize the cache: {err}");
Vec::new()
})
}
pub(crate) fn summarize(
connection: &turso::Connection,
) -> Result<Vec<NamespaceSummary>, turso::Error> {
let mut rows = drive(connection.query(SUMMARY, ()))?;
let mut summaries = Vec::new();
while let Some(row) = drive(rows.next())? {
summaries.push(NamespaceSummary {
namespace: row.get(0)?,
entries: row.get::<i64>(1)? as u64,
bytes: row.get::<i64>(2)? as u64,
});
}
Ok(summaries)
}
#[cfg(native_cache)]
pub(crate) fn checkpoint(connection: &turso::Connection) -> Result<bool, turso::Error> {
let mut incomplete = false;
drive(connection.pragma_query("wal_checkpoint(TRUNCATE)", |row| {
incomplete |= row.get::<i64>(0).map_or(true, |busy| busy != 0);
Ok(())
}))?;
Ok(!incomplete)
}
const BUSY_TIMEOUT: core::time::Duration = core::time::Duration::from_secs(5);
pub(crate) fn connect(database: &turso::Database) -> Result<turso::Connection, turso::Error> {
let connection = database.connect()?;
connection.busy_timeout(BUSY_TIMEOUT)?;
drive(connection.pragma_query("synchronous = NORMAL", |_| Ok(())))?;
Ok(connection)
}
fn origin_code(origin: Origin) -> i64 {
match origin {
Origin::Local => 0,
Origin::Imported => 1,
}
}
fn error(error: turso::Error) -> String {
error.to_string()
}
pub fn open(namespace: &str) -> Result<Box<dyn Storage>, String> {
TursoStorage::open(namespace.to_string()).map(|storage| Box::new(storage) as Box<dyn Storage>)
}
#[cfg(all(test, native_cache))]
mod tests {
use super::*;
use crate::future::block_on;
use alloc::vec;
fn tables(location: &str) -> Vec<String> {
let database = block_on(open_database(location)).unwrap();
let connection = connect(&database).unwrap();
let mut rows = drive(connection.query(
"SELECT name FROM sqlite_schema WHERE type = 'table' ORDER BY name",
(),
))
.unwrap();
let mut names = Vec::new();
while let Some(row) = drive(rows.next()).unwrap() {
names.push(row.get::<String>(0).unwrap());
}
names
}
fn active_location(root: &std::path::Path) -> String {
crate::environment::set_root(root);
location().unwrap()
}
#[test_log::test]
#[serial_test::serial]
#[cfg_attr(miri, ignore)]
fn an_incompatible_schema_is_rebuilt() {
let dir = tempfile::tempdir().unwrap();
let location = active_location(dir.path());
{
let database = block_on(open_database(&location)).unwrap();
let connection = connect(&database).unwrap();
drive(connection.execute(CREATE_META, ())).unwrap();
drive(connection.execute(META_SET, (SCHEMA_VERSION_KEY, "999"))).unwrap();
drive(connection.execute(
"CREATE TABLE entries (store TEXT NOT NULL, key BLOB NOT NULL, \
value BLOB NOT NULL, PRIMARY KEY (store, key))",
(),
))
.unwrap();
drive(connection.execute("INSERT INTO entries VALUES ('old', X'01', X'02')", ()))
.unwrap();
}
let storage = TursoStorage::open("old".to_string()).unwrap();
assert_eq!(storage.get(b"\x01"), None, "stale rows are gone");
assert_eq!(
storage.insert(b"key", Bytes::from_bytes_vec(vec![1]), Origin::Local),
Insertion::Stored,
"the rebuilt table accepts the current column layout"
);
assert_eq!(tables(&location), vec!["entries", "meta"]);
let database = block_on(open_database(&location)).unwrap();
let connection = connect(&database).unwrap();
assert_eq!(
meta_get(&connection, SCHEMA_VERSION_KEY).unwrap(),
Some(SCHEMA_VERSION.to_string())
);
}
#[test_log::test]
#[serial_test::serial]
fn a_cancelled_open_releases_the_location() {
let location = "cancelled.db";
let (sender, receiver) = async_channel::bounded(1);
DATABASES
.lock()
.insert(location.to_string(), DatabaseState::Opening(vec![sender]));
drop(Opener {
location,
result: Err("database initialization was cancelled".to_string()),
});
assert!(!DATABASES.lock().contains_key(location));
assert!(block_on(receiver.recv()).unwrap().is_err());
}
#[test_log::test]
#[serial_test::serial]
#[cfg_attr(miri, ignore)]
fn a_current_file_keeps_its_entries() {
let dir = tempfile::tempdir().unwrap();
let location = active_location(dir.path());
let storage = TursoStorage::open("kept".to_string()).unwrap();
assert_eq!(
storage.insert(b"key", Bytes::from_bytes_vec(vec![7]), Origin::Local),
Insertion::Stored
);
let database = block_on(open_database(&location)).unwrap();
migrate(&database).unwrap();
let reopened = TursoStorage::open("kept".to_string()).unwrap();
assert_eq!(reopened.get(b"key"), Some(Bytes::from_bytes_vec(vec![7])));
}
#[test_log::test]
#[serial_test::serial]
#[cfg_attr(miri, ignore)]
fn concurrent_connections_agree_on_the_winner() {
let dir = tempfile::tempdir().unwrap();
active_location(dir.path());
let first = TursoStorage::open("namespace".to_string()).unwrap();
let second = TursoStorage::open("namespace".to_string()).unwrap();
let bytes = |value: &[u8]| Bytes::from_bytes_vec(value.to_vec());
assert_eq!(
first.insert(b"key", bytes(b"first"), Origin::Local),
Insertion::Stored
);
assert_eq!(
second.insert(b"key", bytes(b"second"), Origin::Local),
Insertion::Conflict(bytes(b"first"))
);
assert_eq!(second.get(b"key"), Some(bytes(b"first")));
}
}