use std::path::{Path, PathBuf};
use std::string::{String, ToString};
use std::vec::Vec;
use hashbrown::HashMap;
use rusqlite::{Connection, OpenFlags, OptionalExtension, TransactionBehavior, params};
use super::storage::{InsertSummary, Insertion, NamespaceSummary, Origin, Storage, replaces};
use crate::bytes::Bytes;
use crate::sync::{Arc, Lazy, Mutex};
pub fn db_file_name(name: &str) -> String {
crate::environment::file_name(name)
}
pub const SCHEMA_VERSION: u32 = 3;
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";
pub(crate) fn meta_get(conn: &Connection, key: &str) -> Result<Option<String>, rusqlite::Error> {
conn.query_row(META_GET, params![key], |row| row.get(0))
.optional()
}
pub(crate) fn meta_set(conn: &Connection, key: &str, value: &str) -> Result<(), rusqlite::Error> {
conn.execute(META_SET, params![key, value])?;
Ok(())
}
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 INSERT_SQL: &str = "INSERT INTO entries (namespace, key, value, origin) \
VALUES (?1, ?2, ?3, ?4) ON CONFLICT DO NOTHING";
const REPLACE_SQL: &str = "UPDATE entries SET value = ?3, origin = ?4 \
WHERE namespace = ?1 AND key = ?2";
const UPSERT_SQL: &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_SQL: &str = "SELECT value FROM entries WHERE namespace = ?1 AND key = ?2";
const SELECT_WITH_ORIGIN_SQL: &str =
"SELECT value, origin FROM entries WHERE namespace = ?1 AND key = ?2";
const SCAN_SQL: &str = "SELECT key, value FROM entries WHERE namespace = ?1";
const PURGE_SQL: &str = "DELETE FROM entries WHERE namespace = ?1";
const PURGE_KEY_SQL: &str = "DELETE FROM entries WHERE namespace = ?1 AND key = ?2";
fn origin_code(origin: Origin) -> i64 {
match origin {
Origin::Local => 0,
Origin::Imported => 1,
}
}
fn origin_from_code(code: i64) -> Origin {
match code {
1 => Origin::Imported,
_ => Origin::Local,
}
}
#[derive(Clone)]
pub struct Database {
conn: Arc<Mutex<Connection>>,
path: Arc<PathBuf>,
}
impl core::fmt::Debug for Database {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(f, "Database({:?})", self.path)
}
}
static OPENED: Lazy<Mutex<HashMap<PathBuf, Database>>> = Lazy::new(|| Mutex::new(HashMap::new()));
impl Database {
pub fn open_active() -> Option<Self> {
Self::open_path(crate::environment::path())
}
pub fn open_at<P: AsRef<Path>>(root: P, environment: &str) -> Option<Self> {
Self::open_path(root.as_ref().join(db_file_name(environment)))
}
fn open_path(path: PathBuf) -> Option<Self> {
let mut opened = OPENED.lock();
if let Some(database) = opened.get(&path) {
return Some(database.clone());
}
if let Some(root) = path.parent()
&& let Err(err) = std::fs::create_dir_all(root)
{
log::error!(
"cubecl cache: create_dir_all({root:?}) failed: {err}; \
persistence is disabled for this root"
);
return None;
}
match Self::open(&path, false) {
Ok(database) => {
opened.insert(path, database.clone());
Some(database)
}
Err(err) => {
log::error!(
"cubecl cache: can't open {path:?}: {err}; \
persistence is disabled for this root"
);
None
}
}
}
pub fn open(path: &Path, read_only: bool) -> Result<Self, rusqlite::Error> {
let conn = if read_only {
open_read_only(path)?
} else {
let flags = OpenFlags::SQLITE_OPEN_READ_WRITE
| OpenFlags::SQLITE_OPEN_CREATE
| OpenFlags::SQLITE_OPEN_NO_MUTEX;
Connection::open_with_flags(path, flags)?
};
conn.busy_timeout(core::time::Duration::from_secs(5))?;
if !read_only {
conn.pragma_update(None, "journal_mode", "WAL")?;
conn.pragma_update(None, "synchronous", "NORMAL")?;
conn.execute_batch(CREATE_META)?;
migrate(&conn)?;
conn.execute_batch(CREATE_ENTRIES)?;
}
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
path: Arc::new(path.to_path_buf()),
})
}
pub fn path(&self) -> &Path {
&self.path
}
pub fn with_connection<T>(&self, func: impl FnOnce(&mut Connection) -> T) -> T {
let mut guard = self.conn.lock();
func(&mut guard)
}
pub fn get(&self, namespace: &str, key: &[u8]) -> Option<Bytes> {
self.with_connection(|conn| {
conn.prepare_cached(SELECT_SQL)
.and_then(|mut stmt| {
stmt.query_row(params![namespace, key], |row| row.get(0))
.optional()
})
.unwrap_or_else(|err| {
self.warn("read", err);
None
})
.map(Bytes::from_bytes_vec)
})
}
pub fn insert(&self, namespace: &str, key: &[u8], value: &[u8], origin: Origin) -> Insertion {
let result = self.with_connection(|conn| {
let transaction = conn.transaction_with_behavior(TransactionBehavior::Immediate)?;
let outcome = insert_one(&transaction, namespace, key, value, origin)?;
transaction.commit()?;
Ok::<_, rusqlite::Error>(outcome)
});
self.report("write", result)
}
pub fn insert_many(
&self,
namespace: &str,
entries: &mut dyn Iterator<Item = (Bytes, Bytes)>,
origin: Origin,
) -> InsertSummary {
let result = self.with_connection(|conn| {
let transaction = conn.transaction_with_behavior(TransactionBehavior::Immediate)?;
let mut summary = InsertSummary::default();
for (key, value) in entries {
summary.record(&insert_one(&transaction, namespace, &key, &value, origin)?);
}
transaction.commit()?;
Ok::<_, rusqlite::Error>(summary)
});
match result {
Ok(summary) => summary,
Err(err) => {
self.warn("batch write", err);
InsertSummary::default()
}
}
}
pub fn replace(&self, namespace: &str, key: &[u8], value: &[u8], origin: Origin) -> Insertion {
let result = self.with_connection(|conn| {
conn.prepare_cached(UPSERT_SQL)?.execute(params![
namespace,
key,
value,
origin_code(origin)
])?;
Ok::<_, rusqlite::Error>(Insertion::Stored)
});
self.report("replace", result)
}
pub fn purge(&self, namespace: &str) {
let result = self.with_connection(|conn| {
conn.prepare_cached(PURGE_SQL)?
.execute(params![namespace])?;
Ok::<_, rusqlite::Error>(())
});
if let Err(err) = result {
log::warn!("cubecl cache: purge of '{namespace}' failed: {err}");
}
}
pub fn purge_key(&self, namespace: &str, key: &[u8]) {
let result = self.with_connection(|conn| {
conn.prepare_cached(PURGE_KEY_SQL)?
.execute(params![namespace, key])?;
Ok::<_, rusqlite::Error>(())
});
if let Err(err) = result {
log::warn!("cubecl cache: purge of a key in '{namespace}' failed: {err}");
}
}
pub fn finalize_for_shipping(&self) -> Result<(), rusqlite::Error> {
self.with_connection(|conn| {
conn.pragma_update_and_check(None, "journal_mode", "DELETE", |_row| Ok(()))
})
}
fn report(&self, operation: &str, result: Result<Insertion, rusqlite::Error>) -> Insertion {
match result {
Ok(insertion) => insertion,
Err(err) => {
let message = err.to_string();
self.warn(operation, err);
Insertion::Failed(message)
}
}
}
pub fn scan(&self, namespace: &str, visit: &mut dyn FnMut(&[u8], &[u8])) {
let result = self.with_connection(|conn| {
let mut stmt = conn.prepare_cached(SCAN_SQL)?;
let mut rows = stmt.query(params![namespace])?;
while let Some(row) = rows.next()? {
visit(row.get_ref(0)?.as_blob()?, row.get_ref(1)?.as_blob()?);
}
Ok::<_, rusqlite::Error>(())
});
if let Err(err) = result {
self.warn("scan", err);
}
}
pub fn namespaces(&self) -> Vec<String> {
self.summary()
.into_iter()
.map(|summary| summary.namespace)
.collect()
}
pub fn summary(&self) -> Vec<NamespaceSummary> {
let result = self.with_connection(|conn| {
let mut stmt = conn.prepare(
"SELECT namespace, count(*), sum(length(key) + length(value)) \
FROM entries GROUP BY namespace ORDER BY namespace",
)?;
let rows = stmt.query_map([], |row| {
Ok(NamespaceSummary {
namespace: row.get(0)?,
entries: row.get::<_, i64>(1)? as u64,
bytes: row.get::<_, Option<i64>>(2)?.unwrap_or(0) as u64,
})
})?;
rows.collect::<Result<Vec<_>, _>>()
});
result.unwrap_or_else(|err| {
self.warn("summary", err);
Vec::new()
})
}
fn warn(&self, operation: &str, err: rusqlite::Error) {
log::warn!("cubecl cache: {operation} on {:?} failed: {err}", self.path);
}
}
fn insert_one(
conn: &Connection,
namespace: &str,
key: &[u8],
value: &[u8],
origin: Origin,
) -> Result<Insertion, rusqlite::Error> {
let code = origin_code(origin);
let written = conn
.prepare_cached(INSERT_SQL)?
.execute(params![namespace, key, value, code])?;
if written != 0 {
return Ok(Insertion::Stored);
}
let found: Option<(Vec<u8>, i64)> = conn
.prepare_cached(SELECT_WITH_ORIGIN_SQL)?
.query_row(params![namespace, key], |row| {
Ok((row.get(0)?, row.get(1)?))
})
.optional()?;
match found {
Some((_, existing_code)) if replaces(origin, origin_from_code(existing_code)) => {
conn.prepare_cached(REPLACE_SQL)?
.execute(params![namespace, key, value, code])?;
Ok(Insertion::Stored)
}
Some((existing, _)) => Ok(Insertion::Conflict(Bytes::from_bytes_vec(existing))),
None => Ok(Insertion::Failed(String::from(
"the entry disappeared during the write",
))),
}
}
fn open_read_only(path: &Path) -> Result<Connection, rusqlite::Error> {
let flags = OpenFlags::SQLITE_OPEN_READ_ONLY | OpenFlags::SQLITE_OPEN_NO_MUTEX;
let attempt = Connection::open_with_flags(path, flags).and_then(|conn| {
conn.query_row("SELECT count(*) FROM sqlite_schema", [], |_| Ok(()))?;
Ok(conn)
});
let err = match attempt {
Ok(conn) => return Ok(conn),
Err(err) => err,
};
let Some(uri) = immutable_uri(path) else {
return Err(err);
};
log::debug!("cubecl cache: {path:?} is not readable in place ({err}); retrying immutable");
Connection::open_with_flags(uri, flags | OpenFlags::SQLITE_OPEN_URI)
}
fn immutable_uri(path: &Path) -> Option<String> {
let path = path.to_str()?;
let mut uri = String::from("file:");
for character in path.chars() {
match character {
'?' => uri.push_str("%3f"),
'#' => uri.push_str("%23"),
'%' => uri.push_str("%25"),
_ => uri.push(character),
}
}
uri.push_str("?immutable=1");
Some(uri)
}
fn migrate(conn: &Connection) -> Result<(), rusqlite::Error> {
let found = meta_get(conn, SCHEMA_VERSION_KEY)?;
let expected = SCHEMA_VERSION.to_string();
if found.as_deref() == Some(expected.as_str()) {
return Ok(());
}
if let Some(found) = found {
log::warn!(
"cubecl cache: database schema {found} is not {expected}, discarding cached entries"
);
}
conn.execute("DROP TABLE IF EXISTS entries", [])?;
meta_set(conn, SCHEMA_VERSION_KEY, &expected)?;
Ok(())
}
#[derive(Debug)]
pub struct SqliteStorage {
database: Database,
namespace: String,
}
impl SqliteStorage {
pub fn new(database: Database, namespace: String) -> Self {
Self {
database,
namespace,
}
}
}
impl Storage for SqliteStorage {
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 {
std::format!("{:?} [{}]", self.database.path(), self.namespace)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test_log::test]
#[cfg_attr(miri, ignore)]
fn concurrent_connections_agree_on_the_winner() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(db_file_name("test"));
let first = Database::open(&path, false).unwrap();
let second = Database::open(&path, false).unwrap();
assert_eq!(
first.insert("namespace", b"key", b"first", Origin::Local),
Insertion::Stored
);
assert_eq!(
second.insert("namespace", b"key", b"second", Origin::Local),
Insertion::Conflict(Bytes::from_bytes_vec(b"first".to_vec()))
);
assert_eq!(
second.get("namespace", b"key").as_deref(),
Some(&b"first"[..])
);
}
#[test_log::test]
#[cfg_attr(miri, ignore)]
fn concurrent_writers_do_not_lose_entries() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(db_file_name("test"));
let database = Database::open(&path, false).unwrap();
std::thread::scope(|scope| {
for writer in 0..4 {
let path = &path;
scope.spawn(move || {
let database = Database::open(path, false).unwrap();
for entry in 0..25 {
let key = std::format!("{writer}-{entry}");
assert_eq!(
database.insert("namespace", key.as_bytes(), b"value", Origin::Local),
Insertion::Stored
);
}
});
}
});
let mut count = 0;
database.scan("namespace", &mut |_key, _value| count += 1);
assert_eq!(count, 100);
}
#[test_log::test]
#[cfg(unix)]
#[cfg_attr(miri, ignore)]
fn a_database_in_a_read_only_directory_is_readable() {
use std::os::unix::fs::PermissionsExt;
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(db_file_name("test"));
let database = Database::open(&path, false).unwrap();
database.insert("namespace", b"key", b"value", Origin::Local);
drop(database);
let restore = std::fs::metadata(dir.path()).unwrap().permissions();
std::fs::set_permissions(dir.path(), std::fs::Permissions::from_mode(0o555)).unwrap();
let writable = std::fs::File::create(dir.path().join("probe")).is_ok();
let read = (!writable).then(|| {
let database = Database::open(&path, true).expect("read-only open");
database.get("namespace", b"key")
});
std::fs::set_permissions(dir.path(), restore).unwrap();
if let Some(read) = read {
assert_eq!(read.as_deref(), Some(&b"value"[..]));
}
}
#[test_log::test]
#[cfg_attr(miri, ignore)]
fn finalizing_clears_the_wal_header() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(db_file_name("test"));
let database = Database::open(&path, false).unwrap();
database.insert("namespace", b"key", b"value", Origin::Local);
database.finalize_for_shipping().unwrap();
let mode: String = database.with_connection(|conn| {
conn.query_row("PRAGMA journal_mode", [], |row| row.get(0))
.unwrap()
});
assert_eq!(mode, "delete");
assert!(!path.with_extension("db-wal").exists());
}
#[test_log::test]
#[cfg_attr(miri, ignore)]
fn replace_overwrites_an_existing_entry() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(db_file_name("test"));
let database = Database::open(&path, false).unwrap();
database.insert("namespace", b"key", b"corrupt", Origin::Local);
assert_eq!(
database.insert("namespace", b"key", b"value", Origin::Local),
Insertion::Conflict(Bytes::from_bytes_vec(b"corrupt".to_vec()))
);
assert_eq!(
database.replace("namespace", b"key", b"value", Origin::Local),
Insertion::Stored
);
assert_eq!(
database.get("namespace", b"key").as_deref(),
Some(&b"value"[..])
);
}
#[test_log::test]
#[cfg_attr(miri, ignore)]
fn insert_many_matches_insert() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(db_file_name("test"));
let database = Database::open(&path, false).unwrap();
database.insert("namespace", b"taken", b"local", Origin::Local);
let bytes = |value: &[u8]| Bytes::from_bytes_vec(value.to_vec());
let entries = std::vec![
(bytes(b"fresh"), bytes(b"imported")),
(bytes(b"taken"), bytes(b"imported")),
];
let summary = database.insert_many("namespace", &mut entries.into_iter(), Origin::Imported);
assert_eq!(summary.stored, 1);
assert_eq!(summary.conflict, 1);
assert_eq!(summary.failed, 0);
assert_eq!(
database.get("namespace", b"fresh").as_deref(),
Some(&b"imported"[..])
);
assert_eq!(
database.get("namespace", b"taken").as_deref(),
Some(&b"local"[..]),
"a bundle never overwrites a local value"
);
}
#[test_log::test]
#[cfg_attr(miri, ignore)]
fn an_incompatible_schema_is_rebuilt() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join(db_file_name("test"));
let database = Database::open(&path, false).unwrap();
database.with_connection(|conn| {
conn.execute("DROP TABLE entries", []).unwrap();
conn.execute(
"CREATE TABLE entries (store TEXT NOT NULL, key BLOB NOT NULL, \
value BLOB NOT NULL, PRIMARY KEY (store, key))",
[],
)
.unwrap();
conn.execute(
"INSERT INTO entries (store, key, value) VALUES ('old', x'00', x'00')",
[],
)
.unwrap();
conn.execute("UPDATE meta SET v = '999' WHERE k = 'schema_version'", [])
.unwrap();
});
drop(database);
let database = Database::open(&path, false).unwrap();
assert_eq!(database.get("old", b"\x00"), None, "stale rows are gone");
assert_eq!(
database.insert("namespace", b"key", b"value", Origin::Local),
Insertion::Stored
);
assert_eq!(
database.get("namespace", b"key").as_deref(),
Some(&b"value"[..]),
"the rebuilt table accepts the current column layout"
);
}
}