use std::{
collections::HashSet,
ffi::OsString,
fs, io,
marker::PhantomData,
path::{Path, PathBuf},
sync::{Arc, Mutex, PoisonError, RwLock},
};
use crate::Dataset;
use futures_lite::future::block_on;
use gix_tempfile::{
AutoRemove, ContainingDirectory, Handle,
handle::{Writable, persist},
};
use r2d2::Pool;
use sanitize_filename::sanitize;
use serde::{Serialize, de::DeserializeOwned};
use turso::{Builder, Connection, Database, Value};
mod de;
use de::from_row_with_columns;
pub use de::RowError;
pub type Result<T> = core::result::Result<T, SqliteDatasetError>;
#[derive(thiserror::Error, Debug)]
pub enum SqliteDatasetError {
#[error("IO error: {0}")]
Io(#[from] io::Error),
#[error("Sql error: {0}")]
Sql(#[from] turso::Error),
#[error("Serde error: {0}")]
Serde(#[from] rmp_serde::encode::Error),
#[error("Deserialize error: {0}")]
Deserialize(#[from] rmp_serde::decode::Error),
#[error("Row error: {0}")]
Row(#[from] RowError),
#[error("Overwrite flag is set to false and the database file already exists: {0}")]
FileExists(PathBuf),
#[error("Failed to create connection pool: {0}")]
ConnectionPool(#[from] r2d2::Error),
#[error("Could not persist the temporary database file: {0}")]
PersistDbFile(#[from] persist::Error<Writable>),
#[error("{0}")]
Other(&'static str),
}
impl From<&'static str> for SqliteDatasetError {
fn from(s: &'static str) -> Self {
SqliteDatasetError::Other(s)
}
}
#[derive(Debug)]
pub struct SqliteDataset<I> {
db_file: PathBuf,
split: String,
conn_pool: Pool<TursoConnectionManager>,
columns: Vec<String>,
len: usize,
select_statement: String,
row_serialized: bool,
phantom: PhantomData<I>,
}
impl<I> SqliteDataset<I> {
pub fn from_db_file<P: AsRef<Path>>(db_file: P, split: &str) -> Result<Self> {
let database = open_database(&db_file, false)?;
let conn_pool = Pool::new(TursoConnectionManager { database })?;
let row_serialized = Self::check_if_row_serialized(&conn_pool, split)?;
let select_statement = if row_serialized {
format!("select item from {split} where row_id = ?")
} else {
format!("select * from {split} where row_id = ?")
};
let (columns, len) = fetch_columns_and_len(&conn_pool, &select_statement, split)?;
Ok(SqliteDataset {
db_file: db_file.as_ref().to_path_buf(),
split: split.to_string(),
conn_pool,
columns,
len,
select_statement,
row_serialized,
phantom: PhantomData,
})
}
fn check_if_row_serialized(
conn_pool: &Pool<TursoConnectionManager>,
split: &str,
) -> Result<bool> {
let conn = conn_pool.get()?;
let columns = block_on(conn.prepare(&format!("select * from {split}")))?.columns();
let matches = |index: usize, name: &str, ty: &str| {
columns[index].name().eq_ignore_ascii_case(name)
&& columns[index]
.decl_type()
.is_some_and(|declared| declared.eq_ignore_ascii_case(ty))
};
Ok(columns.len() == 2 && matches(0, "row_id", "integer") && matches(1, "item", "blob"))
}
pub fn db_file(&self) -> PathBuf {
self.db_file.clone()
}
pub fn split(&self) -> &str {
self.split.as_str()
}
}
impl<I: DeserializeOwned> SqliteDataset<I> {
fn item_from_row(&self, row: &turso::Row) -> Result<I> {
if self.row_serialized {
match row.get_value(0)? {
Value::Blob(blob) => Ok(rmp_serde::from_slice::<I>(&blob)?),
_ => Err(SqliteDatasetError::Other("expected a blob column")),
}
} else {
Ok(from_row_with_columns::<I>(row, &self.columns)?)
}
}
}
impl<I> Dataset<I, SqliteDatasetError> for SqliteDataset<I>
where
I: Clone + Send + Sync + DeserializeOwned,
{
fn get(&self, index: usize) -> Result<I> {
assert!(
index < self.len,
"Index out of bounds for SqliteDataset: {} >= {}",
index,
self.len,
);
let row_id = (index + 1) as i64;
let connection = self.conn_pool.get()?;
let row = block_on(async {
let mut statement = connection.prepare_cached(&self.select_statement).await?;
statement.query_row([row_id]).await
})
.map_err(|error| match error {
turso::Error::QueryReturnedNoRows => SqliteDatasetError::Other(
"no row for this index; row_id values must be contiguous and start at 1",
),
error => error.into(),
})?;
self.item_from_row(&row)
}
fn get_many(&self, indexes: Vec<usize>) -> Result<Vec<I>> {
if indexes.is_empty() {
return Ok(Vec::new());
}
for &index in &indexes {
assert!(
index < self.len,
"Index out of bounds for SqliteDataset: {} >= {}",
index,
self.len,
);
}
let values_clause = vec!["(?, ?)"; indexes.len()].join(", ");
let params: Vec<i64> = indexes
.iter()
.enumerate()
.flat_map(|(ord, &index)| [(index + 1) as i64, ord as i64])
.collect();
let split = &self.split;
let connection = self.conn_pool.get()?;
let selection = if self.row_serialized { "t.item" } else { "t.*" };
let query = format!(
"WITH req(row_id, ord) AS (VALUES {values_clause}) \
SELECT {selection} FROM req JOIN {split} t ON t.row_id = req.row_id ORDER BY req.ord"
);
block_on(async {
let mut statement = connection.prepare(&query).await?;
let mut rows = statement.query(params).await?;
let mut items = Vec::with_capacity(indexes.len());
while let Some(row) = rows.next().await? {
items.push(self.item_from_row(&row)?);
}
if items.len() != indexes.len() {
return Err(SqliteDatasetError::Other(
"fewer rows than indexes requested; row_id values must be contiguous and \
start at 1",
));
}
Ok(items)
})
}
fn len(&self) -> usize {
self.len
}
}
fn fetch_columns_and_len(
conn_pool: &Pool<TursoConnectionManager>,
select_statement: &str,
split: &str,
) -> Result<(Vec<String>, usize)> {
let connection = conn_pool.get()?;
let (columns, max_row_id) = block_on(async {
let statement = connection.prepare(select_statement).await?;
let columns = statement.column_names();
let mut statement = connection
.prepare(format!("select max(row_id) from {split}").as_str())
.await?;
let max_row_id = statement.query_row(()).await?.get_value(0)?;
Ok::<_, turso::Error>((columns, max_row_id))
})?;
let len = match max_row_id {
Value::Null => 0,
Value::Integer(max_row_id) => usize::try_from(max_row_id).map_err(|_| {
SqliteDatasetError::Other("row_id is negative, so it cannot index the dataset")
})?,
_ => {
return Err(SqliteDatasetError::Other(
"row_id is not an integer, so it cannot index the dataset",
));
}
};
Ok((columns, len))
}
#[derive(Debug)]
struct TursoConnectionManager {
database: Database,
}
impl r2d2::ManageConnection for TursoConnectionManager {
type Connection = Connection;
type Error = turso::Error;
fn connect(&self) -> core::result::Result<Connection, turso::Error> {
self.database.connect()
}
fn is_valid(&self, _conn: &mut Connection) -> core::result::Result<(), turso::Error> {
Ok(())
}
fn has_broken(&self, _conn: &mut Connection) -> bool {
false
}
}
fn open_database<P: AsRef<Path>>(db_file: P, write: bool) -> Result<Database> {
let db_file = db_file.as_ref().to_str().ok_or(SqliteDatasetError::Other(
"database path is not valid UTF-8, which the turso engine requires",
))?;
Ok(block_on(
Builder::new_local(db_file).read_only(!write).build(),
)?)
}
#[derive(Clone, Debug)]
pub struct SqliteDatasetStorage {
name: Option<String>,
db_file: Option<PathBuf>,
base_dir: Option<PathBuf>,
}
impl SqliteDatasetStorage {
pub fn from_name(name: &str) -> Self {
SqliteDatasetStorage {
name: Some(name.to_string()),
db_file: None,
base_dir: None,
}
}
pub fn from_file<P: AsRef<Path>>(db_file: P) -> Self {
SqliteDatasetStorage {
name: None,
db_file: Some(db_file.as_ref().to_path_buf()),
base_dir: None,
}
}
pub fn with_base_dir<P: AsRef<Path>>(mut self, base_dir: P) -> Self {
self.base_dir = Some(base_dir.as_ref().to_path_buf());
self
}
pub fn exists(&self) -> bool {
self.db_file().exists()
}
pub fn db_file(&self) -> PathBuf {
match &self.db_file {
Some(db_file) => db_file.clone(),
None => {
let name = sanitize(self.name.as_ref().expect("Name is not set"));
Self::base_dir(self.base_dir.to_owned()).join(format!("{name}.db"))
}
}
}
pub fn base_dir(base_dir: Option<PathBuf>) -> PathBuf {
match base_dir {
Some(base_dir) => base_dir,
None => dirs::cache_dir()
.expect("Could not get cache directory")
.join("burn-dataset"),
}
}
pub fn writer<I>(&self, overwrite: bool) -> Result<SqliteDatasetWriter<I>>
where
I: Clone + Send + Sync + Serialize + DeserializeOwned,
{
SqliteDatasetWriter::new(self.db_file(), overwrite)
}
pub fn reader<I>(&self, split: &str) -> Result<SqliteDataset<I>>
where
I: Clone + Send + Sync + Serialize + DeserializeOwned,
{
if !self.exists() {
panic!("The database file does not exist");
}
SqliteDataset::from_db_file(self.db_file(), split)
}
}
#[derive(Debug)]
pub struct SqliteDatasetWriter<I> {
db_file: PathBuf,
db_file_tmp: Option<Handle<Writable>>,
overwrite: bool,
state: Option<Mutex<WriteState>>,
is_completed: Arc<RwLock<bool>>,
phantom: PhantomData<I>,
}
const WRITE_BATCH_SIZE: usize = 256;
#[derive(Debug)]
struct WriteState {
connection: Connection,
pending: usize,
lost_a_batch: bool,
splits: HashSet<String>,
}
impl WriteState {
fn insert(&mut self, split: &str, item: Vec<u8>) -> Result<usize> {
block_on(async {
if self.pending == 0 {
self.connection.execute("BEGIN", ()).await?;
}
self.pending += 1;
let insert_statement = format!("insert into {split} (item) values (?)");
let mut statement = self.connection.prepare_cached(&insert_statement).await?;
statement.execute([item]).await
})?;
let index = (self.connection.last_insert_rowid() - 1) as usize;
if self.pending >= WRITE_BATCH_SIZE {
self.commit()?;
}
Ok(index)
}
fn commit(&mut self) -> Result<()> {
if self.pending > 0 {
let committed = block_on(self.connection.execute("COMMIT", ()));
if committed.is_err() {
let _ = block_on(self.connection.execute("ROLLBACK", ()));
self.lost_a_batch = true;
}
self.pending = 0;
committed?;
}
Ok(())
}
fn checkpoint(&self) -> Result<()> {
if self.lost_a_batch {
return Err(SqliteDatasetError::Other(
"A batch was rolled back earlier, so this dataset is missing rows and will not be \
published. Rebuild it from scratch",
));
}
let mut failed = false;
block_on(
self.connection
.pragma_query("wal_checkpoint(TRUNCATE)", |row| {
failed |= !matches!(row.get_value(0), Ok(Value::Integer(0)));
Ok(())
}),
)?;
if failed {
return Err(SqliteDatasetError::Other(
"Could not checkpoint the write-ahead log, so the database was left unpublished \
rather than published incomplete. Turso does not report why; enable the \
`tracing` feature and log at debug level to see the underlying cause",
));
}
Ok(())
}
fn create_table(&mut self, split: &str) -> Result<()> {
if self.splits.contains(split) {
return Ok(());
}
self.commit()?;
let create_table_statement = format!(
"create table if not exists {split} (row_id integer primary key autoincrement not \
null, item blob not null)"
);
block_on(self.connection.execute(&create_table_statement, ()))?;
self.splits.insert(split.to_string());
Ok(())
}
}
impl<I> SqliteDatasetWriter<I>
where
I: Clone + Send + Sync + Serialize + DeserializeOwned,
{
pub fn new<P: AsRef<Path>>(db_file: P, overwrite: bool) -> Result<Self> {
let writer = Self {
db_file: db_file.as_ref().to_path_buf(),
db_file_tmp: None,
overwrite,
state: None,
is_completed: Arc::new(RwLock::new(false)),
phantom: PhantomData,
};
writer.init()
}
fn init(mut self) -> Result<Self> {
if self.db_file.exists() {
if self.overwrite {
fs::remove_file(&self.db_file)?;
} else {
return Err(SqliteDatasetError::FileExists(self.db_file.clone()));
}
}
remove_wal_file(&self.db_file)?;
let db_file_dir = self
.db_file
.parent()
.ok_or("Unable to get parent directory")?;
if !db_file_dir.exists() {
fs::create_dir_all(db_file_dir)?;
}
let db_file_tmp = tmp_db_file(&self.db_file);
if db_file_tmp.exists() {
fs::remove_file(&db_file_tmp)?;
}
remove_wal_file(&db_file_tmp)?;
gix_tempfile::signal::setup(Default::default());
self.db_file_tmp = Some(gix_tempfile::writable_at(
&db_file_tmp,
ContainingDirectory::Exists,
AutoRemove::Tempfile,
)?);
let connection = open_database(&db_file_tmp, true)?.connect()?;
self.state = Some(Mutex::new(WriteState {
connection,
pending: 0,
lost_a_batch: false,
splits: HashSet::new(),
}));
Ok(self)
}
pub fn write(&self, split: &str, item: &I) -> Result<usize> {
let is_completed = self
.is_completed
.read()
.unwrap_or_else(PoisonError::into_inner);
if *is_completed {
return Err(SqliteDatasetError::Other(
"Cannot save to a completed dataset writer",
));
}
let serialized_item = rmp_serde::to_vec(item)?;
let state = self.state.as_ref().ok_or(SqliteDatasetError::Other(
"Cannot save to a completed dataset writer",
))?;
let mut state = state.lock().unwrap_or_else(PoisonError::into_inner);
state.create_table(split)?;
state.insert(split, serialized_item)
}
pub fn set_completed(&mut self) -> Result<()> {
let mut is_completed = self
.is_completed
.write()
.unwrap_or_else(PoisonError::into_inner);
{
let state = self
.state
.as_mut()
.ok_or(SqliteDatasetError::Other(
"Cannot complete a dataset writer twice",
))?
.get_mut()
.unwrap_or_else(PoisonError::into_inner);
state.commit()?;
state.checkpoint()?;
}
self.state = None;
let _file_result = self
.db_file_tmp
.take() .unwrap() .persist(&self.db_file)?
.ok_or("Unable to persist the database file")?;
let _ = fs::remove_file(wal_file(&tmp_db_file(&self.db_file)));
*is_completed = true;
Ok(())
}
}
impl<I> Drop for SqliteDatasetWriter<I> {
fn drop(&mut self) {
self.state = None;
let _ = fs::remove_file(wal_file(&tmp_db_file(&self.db_file)));
}
}
fn tmp_db_file(db_file: &Path) -> PathBuf {
let mut db_file_tmp = db_file.to_path_buf();
db_file_tmp.set_extension("db.tmp");
db_file_tmp
}
fn wal_file(db_file: &Path) -> PathBuf {
let mut wal_file = OsString::from(db_file);
wal_file.push("-wal");
PathBuf::from(wal_file)
}
fn remove_wal_file(db_file: &Path) -> Result<()> {
let wal_file = wal_file(db_file);
if wal_file.exists() {
fs::remove_file(wal_file)?;
}
Ok(())
}
#[cfg(test)]
mod tests {
use rayon::prelude::*;
use rstest::{fixture, rstest};
use serde::{Deserialize, Serialize};
use tempfile::{NamedTempFile, TempDir, tempdir};
use super::*;
type SqlDs = SqliteDataset<Sample>;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct Sample {
column_str: String,
column_bytes: Vec<u8>,
column_int: i64,
column_bool: bool,
column_float: f64,
}
#[fixture]
fn train_dataset() -> SqlDs {
SqliteDataset::<Sample>::from_db_file("tests/data/sqlite-dataset.db", "train").unwrap()
}
#[rstest]
pub fn len(train_dataset: SqlDs) {
assert_eq!(train_dataset.len(), 2);
}
#[rstest]
pub fn get_some(train_dataset: SqlDs) {
let item = train_dataset.get(0).unwrap();
assert_eq!(item.column_str, "HI1");
assert_eq!(item.column_bytes, vec![55, 231, 159]);
assert_eq!(item.column_int, 1);
assert!(item.column_bool);
assert_eq!(item.column_float, 1.0);
}
#[rstest]
#[should_panic(expected = "Index out of bounds")]
pub fn get_none(train_dataset: SqlDs) {
train_dataset.get(10).unwrap();
}
#[rstest]
pub fn get_many_out_of_order_with_duplicates(train_dataset: SqlDs) {
let items = train_dataset.get_many(vec![1, 0, 1]).unwrap();
assert_eq!(items.len(), 3);
assert_eq!(items[0], train_dataset.get(1).unwrap());
assert_eq!(items[1], train_dataset.get(0).unwrap());
assert_eq!(items[2], train_dataset.get(1).unwrap());
}
#[rstest]
pub fn get_many_empty(train_dataset: SqlDs) {
assert_eq!(train_dataset.get_many(vec![]).unwrap(), Vec::new());
}
#[rstest]
#[should_panic(expected = "Index out of bounds")]
pub fn get_many_out_of_bounds(train_dataset: SqlDs) {
train_dataset.get_many(vec![0, 10]).unwrap();
}
#[rstest]
pub fn multi_thread(train_dataset: SqlDs) {
let dataset_len = train_dataset.len();
let indices: Vec<usize> = vec![0, 1, 1, 3, 4, 5, 6, 0, 8, 1];
let valid_indices: Vec<usize> = indices.into_iter().filter(|&i| i < dataset_len).collect();
let results: Vec<Sample> = valid_indices
.par_iter()
.map(|&i| train_dataset.get(i).unwrap())
.collect();
assert_eq!(results.len(), 5);
}
#[rstest]
pub fn reading_leaves_the_database_file_untouched(tmp_dir: TempDir) {
let db_file = tmp_dir.path().join("sqlite-dataset.db");
fs::copy("tests/data/sqlite-dataset.db", &db_file).unwrap();
let before = fs::read(&db_file).unwrap();
let dataset = SqliteDataset::<Sample>::from_db_file(&db_file, "train").unwrap();
assert_eq!(dataset.len(), 2);
dataset.get(0).unwrap();
dataset.get_many(vec![0, 1]).unwrap();
drop(dataset);
assert_eq!(fs::read(&db_file).unwrap(), before);
assert!(!wal_file(&db_file).exists());
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct Typed {
#[serde(rename = "MedInc")]
median_income: f32,
count: usize,
ratio: f64,
label: Option<String>,
missing: Option<f64>,
}
#[derive(Debug, Clone, Deserialize, PartialEq)]
enum Status {
Ready,
Pending,
}
#[derive(Debug, Clone, Deserialize, PartialEq)]
struct UserId(i64);
#[derive(Debug, Clone, Deserialize, PartialEq)]
struct Marker;
#[derive(Debug, Clone, Deserialize, PartialEq)]
struct SerdeShapes {
status: Status,
user_id: UserId,
marker: Marker,
}
#[rstest]
fn get_maps_columns_onto_serde_shapes(tmp_dir: TempDir) {
let db_file = tmp_dir.path().join("serde-shapes.db");
{
let connection = open_database(&db_file, true).unwrap().connect().unwrap();
block_on(async {
connection
.execute(
"create table train (status TEXT, user_id INTEGER, marker TEXT, \
row_id INTEGER NOT NULL, PRIMARY KEY (row_id))",
(),
)
.await?;
connection
.execute("insert into train values ('Ready', 42, 'Marker', 1)", ())
.await?;
connection
.execute(
"create table statuses (status TEXT, row_id INTEGER NOT NULL, \
PRIMARY KEY (row_id))",
(),
)
.await?;
connection
.execute("insert into statuses values ('Pending', 1)", ())
.await?;
connection
.execute(
"create table users (user_id INTEGER, row_id INTEGER NOT NULL, \
PRIMARY KEY (row_id))",
(),
)
.await?;
connection
.execute("insert into users values (7, 1)", ())
.await?;
connection
.pragma_query("wal_checkpoint(TRUNCATE)", |_| Ok(()))
.await
})
.unwrap();
}
fs::remove_file(wal_file(&db_file)).unwrap();
let dataset = SqliteDataset::<SerdeShapes>::from_db_file(&db_file, "train").unwrap();
assert_eq!(
dataset.get(0).unwrap(),
SerdeShapes {
status: Status::Ready,
user_id: UserId(42),
marker: Marker,
}
);
let statuses = SqliteDataset::<Status>::from_db_file(&db_file, "statuses").unwrap();
assert_eq!(statuses.get(0).unwrap(), Status::Pending);
let users = SqliteDataset::<UserId>::from_db_file(&db_file, "users").unwrap();
assert_eq!(users.get(0).unwrap(), UserId(7));
}
#[rstest]
pub fn get_maps_columns_onto_typed_fields(tmp_dir: TempDir) {
let db_file = tmp_dir.path().join("typed.db");
{
let connection = open_database(&db_file, true).unwrap().connect().unwrap();
block_on(async {
connection
.execute(
"create table train (\"MedInc\" REAL, count INTEGER, ratio REAL, \
label TEXT, missing REAL, row_id INTEGER NOT NULL, PRIMARY KEY (row_id))",
(),
)
.await?;
connection
.execute(
"insert into train values (8.3252, 41, 2, 'first', NULL, 1)",
(),
)
.await?;
connection
.execute(
"insert into train values (8.3014, 21, 0.5, NULL, 1.5, 2)",
(),
)
.await?;
connection
.execute(
"insert into train values (NULL, 7, 1.0, 'third', NULL, 3)",
(),
)
.await?;
connection
.pragma_query("wal_checkpoint(TRUNCATE)", |_| Ok(()))
.await
})
.unwrap();
}
fs::remove_file(wal_file(&db_file)).unwrap();
let dataset = SqliteDataset::<Typed>::from_db_file(&db_file, "train").unwrap();
assert_eq!(dataset.len(), 3);
assert_eq!(
dataset.get(0).unwrap(),
Typed {
median_income: 8.3252,
count: 41,
ratio: 2.0,
label: Some("first".to_string()),
missing: None,
}
);
assert_eq!(
dataset.get(1).unwrap(),
Typed {
median_income: 8.3014,
count: 21,
ratio: 0.5,
label: None,
missing: Some(1.5),
}
);
let third = dataset.get(2).unwrap();
assert!(third.median_income.is_nan());
assert_eq!(third.count, 7);
assert_eq!(third.label, Some("third".to_string()));
}
#[test]
fn sqlite_dataset_storage() {
let storage = SqliteDatasetStorage::from_file("non-existing.db");
assert!(!storage.exists());
let storage = SqliteDatasetStorage::from_name("non-existing.db");
assert!(!storage.exists());
let storage = SqliteDatasetStorage::from_file("tests/data/sqlite-dataset.db");
assert!(storage.exists());
let result = storage.reader::<Sample>("train");
assert!(result.is_ok());
let train = result.unwrap();
assert_eq!(train.len(), 2);
let temp_file = NamedTempFile::new().unwrap();
let storage = SqliteDatasetStorage::from_file(temp_file.path());
assert!(storage.exists());
let result = storage.writer::<Sample>(true);
assert!(result.is_ok());
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
pub struct Complex {
column_str: String,
column_bytes: Vec<u8>,
column_int: i64,
column_bool: bool,
column_float: f64,
column_complex: Vec<Vec<Vec<[u8; 3]>>>,
}
#[fixture]
fn tmp_dir() -> TempDir {
tempdir().unwrap()
}
type Writer = SqliteDatasetWriter<Complex>;
#[fixture]
fn writer_fixture(tmp_dir: TempDir) -> (Writer, TempDir) {
let temp_dir_str = tmp_dir.path();
let storage = SqliteDatasetStorage::from_name("preprocessed").with_base_dir(temp_dir_str);
let overwrite = true;
let result = storage.writer::<Complex>(overwrite);
assert!(result.is_ok());
let writer = result.unwrap();
(writer, tmp_dir)
}
#[test]
fn test_new() {
let test_path = NamedTempFile::new().unwrap();
let _writer = SqliteDatasetWriter::<Complex>::new(&test_path, true).unwrap();
assert!(!test_path.path().exists());
let test_path = NamedTempFile::new().unwrap();
let result = SqliteDatasetWriter::<Complex>::new(&test_path, false);
assert!(result.is_err());
let temp = NamedTempFile::new().unwrap();
let test_path = temp.path().to_path_buf();
assert!(temp.close().is_ok());
assert!(!test_path.exists());
let _writer = SqliteDatasetWriter::<Complex>::new(&test_path, true).unwrap();
assert!(!test_path.exists());
}
#[rstest]
pub fn sqlite_writer_write(writer_fixture: (Writer, TempDir)) {
let (writer, _tmp_dir) = writer_fixture;
assert!(writer.overwrite);
assert!(!writer.db_file.exists());
let new_item = Complex {
column_str: "HI1".to_string(),
column_bytes: vec![1_u8, 2, 3],
column_int: 0,
column_bool: true,
column_float: 1.0,
column_complex: vec![vec![vec![[1, 23_u8, 3]]]],
};
let index = writer.write("train", &new_item).unwrap();
assert_eq!(index, 0);
let mut writer = writer;
writer.set_completed().expect("Failed to set completed");
assert!(writer.db_file.exists());
assert!(writer.db_file_tmp.is_none());
let result = writer.write("train", &new_item);
assert!(result.is_err());
let dataset = SqliteDataset::<Complex>::from_db_file(&writer.db_file, "train").unwrap();
let fetched_item = dataset.get(0).unwrap();
assert_eq!(fetched_item, new_item);
assert_eq!(dataset.len(), 1);
}
#[rstest]
pub fn sqlite_writer_write_multi_thread(writer_fixture: (Writer, TempDir)) {
let (writer, _tmp_dir) = writer_fixture;
let writer = Arc::new(writer);
let record_count = 20;
let splits = ["train", "test"];
(0..record_count).into_par_iter().for_each(|index: i64| {
let thread_id: std::thread::ThreadId = std::thread::current().id();
let sample = Complex {
column_str: format!("test_{thread_id:?}_{index}"),
column_bytes: vec![index as u8, 2, 3],
column_int: index,
column_bool: true,
column_float: 1.0,
column_complex: vec![vec![vec![[1, index as u8, 3]]]],
};
let split = splits[index as usize % 2];
let _index = writer.write(split, &sample).unwrap();
});
let mut writer = Arc::try_unwrap(writer).unwrap();
writer
.set_completed()
.expect("Should set completed successfully");
let train =
SqliteDataset::<Complex>::from_db_file(writer.db_file.clone(), "train").unwrap();
let test = SqliteDataset::<Complex>::from_db_file(&writer.db_file, "test").unwrap();
assert_eq!(train.len(), record_count as usize / 2);
assert_eq!(test.len(), record_count as usize / 2);
}
#[rstest]
pub fn a_lost_batch_can_never_be_published(writer_fixture: (Writer, TempDir)) {
let (writer, _tmp_dir) = writer_fixture;
let item = Complex {
column_str: "HI".to_string(),
column_bytes: vec![1, 2, 3],
column_int: 0,
column_bool: true,
column_float: 1.0,
column_complex: vec![vec![vec![[1, 2, 3]]]],
};
writer.write("train", &item).unwrap();
writer.state.as_ref().unwrap().lock().unwrap().lost_a_batch = true;
let mut writer = writer;
assert!(writer.set_completed().is_err(), "published a lost batch");
assert!(!writer.db_file.exists());
writer.write("train", &item).unwrap();
assert!(writer.set_completed().is_err(), "published a lost batch");
assert!(!writer.db_file.exists());
}
#[rstest]
pub fn get_many_handles_batches_past_the_sqlite_variable_limit(tmp_dir: TempDir) {
let db_file = tmp_dir.path().join("wide.db");
let item = |index: usize| Complex {
column_str: format!("item_{index}"),
column_bytes: vec![index as u8, 2, 3],
column_int: index as i64,
column_bool: true,
column_float: 1.0,
column_complex: vec![vec![vec![[1, index as u8, 3]]]],
};
let count = 1200;
let mut writer = SqliteDatasetWriter::<Complex>::new(&db_file, true).unwrap();
for index in 0..count {
writer.write("train", &item(index)).unwrap();
}
writer.set_completed().unwrap();
let dataset = SqliteDataset::<Complex>::from_db_file(&db_file, "train").unwrap();
let items = dataset.get_many((0..count).collect()).unwrap();
assert_eq!(items.len(), count);
for (index, got) in items.iter().enumerate() {
assert_eq!(got.column_int, index as i64);
}
}
#[rstest]
pub fn non_integer_row_id_is_an_error(tmp_dir: TempDir) {
let db_file = tmp_dir.path().join("textual.db");
{
let connection = open_database(&db_file, true).unwrap().connect().unwrap();
block_on(async {
connection
.execute("create table train (row_id TEXT, name TEXT)", ())
.await?;
connection
.execute("insert into train values ('a', 'first')", ())
.await?;
connection
.pragma_query("wal_checkpoint(TRUNCATE)", |_| Ok(()))
.await
})
.unwrap();
}
fs::remove_file(wal_file(&db_file)).unwrap();
#[derive(Debug, Clone, Serialize, Deserialize)]
struct Named {
name: String,
}
let result = SqliteDataset::<Named>::from_db_file(&db_file, "train");
assert!(
result.is_err(),
"a textual row_id should be rejected, not reported as an empty dataset"
);
}
#[rstest]
pub fn sqlite_writer_abandoned_leaves_nothing_behind(tmp_dir: TempDir) {
let db_file = tmp_dir.path().join("abandoned.db");
let item = Complex {
column_str: "HI".to_string(),
column_bytes: vec![1, 2, 3],
column_int: 0,
column_bool: true,
column_float: 1.0,
column_complex: vec![vec![vec![[1, 2, 3]]]],
};
{
let writer = SqliteDatasetWriter::<Complex>::new(&db_file, true).unwrap();
for _ in 0..1000 {
writer.write("train", &item).unwrap();
}
}
let db_file_tmp = tmp_db_file(&db_file);
assert!(!db_file_tmp.exists(), "temporary database left behind");
assert!(
!wal_file(&db_file_tmp).exists(),
"abandoned writer leaked its write-ahead log"
);
assert!(!db_file.exists(), "nothing should have been published");
}
#[rstest]
pub fn concurrent_reads_scale_across_the_pool(tmp_dir: TempDir) {
let db_file = tmp_dir.path().join("concurrent.db");
let item = |index: usize| Complex {
column_str: format!("item_{index}"),
column_bytes: vec![index as u8, 2, 3],
column_int: index as i64,
column_bool: true,
column_float: 1.0,
column_complex: vec![vec![vec![[1, index as u8, 3]]]],
};
let mut writer = SqliteDatasetWriter::<Complex>::new(&db_file, true).unwrap();
for index in 0..1000 {
writer.write("train", &item(index)).unwrap();
}
writer.set_completed().unwrap();
let dataset = SqliteDataset::<Complex>::from_db_file(&db_file, "train").unwrap();
let requested: Vec<usize> = (0..8000).map(|i| i % 1000).collect();
let got: Vec<i64> = requested
.par_iter()
.map(|&index| dataset.get(index).unwrap().column_int)
.collect();
for (position, &index) in requested.iter().enumerate() {
assert_eq!(got[position], index as i64);
}
(0..256usize).into_par_iter().for_each(|batch| {
let indexes: Vec<usize> = (0..32).map(|offset| (batch * 32 + offset) % 1000).collect();
let items = dataset.get_many(indexes.clone()).unwrap();
assert_eq!(items.len(), indexes.len());
for (item, index) in items.iter().zip(&indexes) {
assert_eq!(item.column_int, *index as i64);
}
});
}
#[rstest]
pub fn sqlite_writer_set_completed_is_retryable(writer_fixture: (Writer, TempDir)) {
let (writer, _tmp_dir) = writer_fixture;
let item = Complex {
column_str: "HI".to_string(),
column_bytes: vec![1, 2, 3],
column_int: 0,
column_bool: true,
column_float: 1.0,
column_complex: vec![vec![vec![[1, 2, 3]]]],
};
for _ in 0..500 {
writer.write("train", &item).unwrap();
}
let mut writer = writer;
let blocker = open_database(tmp_db_file(&writer.db_file), true)
.unwrap()
.connect()
.unwrap();
block_on(async {
blocker.execute("BEGIN", ()).await?;
blocker
.prepare("select count(*) from train")
.await?
.query_row(())
.await
})
.unwrap();
let first = writer.set_completed();
assert!(first.is_err(), "checkpoint should have been refused");
assert!(!writer.db_file.exists(), "nothing may be published yet");
assert!(
writer.set_completed().is_err(),
"retry published while the checkpoint was still refused"
);
assert!(!writer.db_file.exists(), "nothing may be published yet");
writer.write("train", &item).unwrap();
drop(blocker);
writer
.set_completed()
.expect("should complete once the database is no longer busy");
let dataset = SqliteDataset::<Complex>::from_db_file(&writer.db_file, "train").unwrap();
assert_eq!(dataset.len(), 501);
assert!(writer.write("train", &item).is_err());
}
#[rstest]
pub fn sqlite_writer_write_across_batches(writer_fixture: (Writer, TempDir)) {
let (writer, _tmp_dir) = writer_fixture;
let record_count = WRITE_BATCH_SIZE * 2 + 5;
let item = |index: usize| Complex {
column_str: format!("item_{index}"),
column_bytes: vec![index as u8, 2, 3],
column_int: index as i64,
column_bool: true,
column_float: 1.0,
column_complex: vec![vec![vec![[1, index as u8, 3]]]],
};
for index in 0..record_count {
assert_eq!(writer.write("train", &item(index)).unwrap(), index);
}
let mut writer = writer;
writer.set_completed().expect("Failed to set completed");
let train = SqliteDataset::<Complex>::from_db_file(&writer.db_file, "train").unwrap();
assert_eq!(train.len(), record_count);
assert_eq!(train.get(0).unwrap(), item(0));
assert_eq!(train.get(WRITE_BATCH_SIZE).unwrap(), item(WRITE_BATCH_SIZE));
assert_eq!(train.get(record_count - 1).unwrap(), item(record_count - 1));
assert!(!wal_file(&tmp_db_file(&writer.db_file)).exists());
}
}