use std::{
collections::HashSet,
fs,
marker::PhantomData,
path::Path,
sync::{Arc, RwLock},
};
use gix_tempfile::{AutoRemove, ContainingDirectory};
use r2d2::PooledConnection;
use r2d2_sqlite::SqliteConnectionManager;
use serde::{Serialize, de::DeserializeOwned};
use super::{Result, SqliteDatasetError, SqliteDatasetWriter, connection::create_conn_pool};
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,
splits: Arc::new(RwLock::new(HashSet::new())),
overwrite,
conn_pool: 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));
}
}
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 mut db_file_tmp = self.db_file.clone();
db_file_tmp.set_extension("db.tmp");
if db_file_tmp.exists() {
fs::remove_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 conn_pool = create_conn_pool(db_file_tmp, true)?;
self.conn_pool = Some(conn_pool);
Ok(self)
}
pub fn write(&self, split: &str, item: &I) -> Result<usize> {
let is_completed = self.is_completed.read().unwrap();
if *is_completed {
return Err(SqliteDatasetError::Other(
"Cannot save to a completed dataset writer",
));
}
if !self.splits.read().unwrap().contains(split) {
self.create_table(split)?;
}
let conn_pool = self.conn_pool.as_ref().unwrap();
let conn = conn_pool.get()?;
let serialized_item = rmp_serde::to_vec(item)?;
pragma_update_with_error_handling(&conn, "synchronous", "OFF")?;
pragma_update_with_error_handling(&conn, "journal_mode", "OFF")?;
let insert_statement = format!("insert into {split} (item) values (?)");
conn.execute(insert_statement.as_str(), [serialized_item])?;
let index = (conn.last_insert_rowid() - 1) as usize;
Ok(index)
}
pub fn set_completed(&mut self) -> Result<()> {
let mut is_completed = self.is_completed.write().unwrap();
if let Some(pool) = self.conn_pool.take() {
std::mem::drop(pool);
}
let _file_result = self
.db_file_tmp
.take() .unwrap() .persist(&self.db_file)?
.ok_or("Unable to persist the database file")?;
*is_completed = true;
Ok(())
}
fn create_table(&self, split: &str) -> Result<()> {
if self.splits.read().unwrap().contains(split) {
return Ok(());
}
let conn_pool = self.conn_pool.as_ref().unwrap();
let connection = conn_pool.get()?;
let create_table_statement = format!(
"create table if not exists {split} (row_id integer primary key autoincrement not \
null, item blob not null)"
);
connection.execute(create_table_statement.as_str(), [])?;
self.splits.write().unwrap().insert(split.to_string());
Ok(())
}
}
fn pragma_update_with_error_handling(
conn: &PooledConnection<SqliteConnectionManager>,
setting: &str,
value: &str,
) -> Result<()> {
let result = conn.pragma_update(None, setting, value);
if let Err(error) = result
&& error != rusqlite::Error::ExecuteReturnedResults
{
return Err(SqliteDatasetError::Sql(error));
}
Ok(())
}