use std::io::Write;
use rusqlite::{Connection, OpenFlags};
use tempfile::NamedTempFile;
use crate::{KeyError, Result};
#[derive(Debug, thiserror::Error)]
pub enum KoboDbError {
#[error("temp file for Kobo DB: {0}")]
TempFile(#[source] std::io::Error),
#[error("writing Kobo DB copy: {0}")]
Write(#[source] std::io::Error),
#[error("Kobo DB: {0}")]
Open(#[source] rusqlite::Error),
}
pub fn open_kobo_db(
mut db_bytes: Vec<u8>,
) -> std::result::Result<(NamedTempFile, Connection), KoboDbError> {
if db_bytes.len() >= 20 {
db_bytes[18] = 0x01;
db_bytes[19] = 0x01;
}
let mut file = NamedTempFile::new().map_err(KoboDbError::TempFile)?;
file.write_all(&db_bytes)
.and_then(|()| file.as_file().sync_all())
.map_err(KoboDbError::Write)?;
let conn = Connection::open_with_flags(file.path(), OpenFlags::SQLITE_OPEN_READ_ONLY)
.map_err(KoboDbError::Open)?;
Ok((file, conn))
}
pub(super) fn read_userids(db_bytes: Vec<u8>) -> Result<Vec<String>> {
let (_tmp, conn) = open_kobo_db(db_bytes).map_err(|e| KeyError::Invalid(e.to_string()))?;
let mut stmt = conn
.prepare("SELECT UserID FROM user")
.map_err(sqlite_err)?;
let userids: Vec<String> = stmt
.query_map([], |row| row.get::<_, String>(0))
.map_err(sqlite_err)?
.filter_map(|r| r.ok())
.collect();
Ok(userids)
}
fn sqlite_err(e: rusqlite::Error) -> KeyError {
KeyError::Invalid(format!("Kobo DB: {e}"))
}
#[cfg(test)]
mod tests {
use super::*;
fn fixture_db(userids: &[&str]) -> Vec<u8> {
let file = tempfile::NamedTempFile::new().unwrap();
{
let conn = Connection::open(file.path()).unwrap();
conn.execute("CREATE TABLE user (UserID TEXT)", []).unwrap();
for id in userids {
conn.execute("INSERT INTO user (UserID) VALUES (?1)", [id])
.unwrap();
}
}
std::fs::read(file.path()).unwrap()
}
#[test]
fn reads_userids_from_fixture() {
let db = fixture_db(&[
"11111111-2222-3333-4444-555555555555",
"aaaaaaaa-bbbb-cccc-dddd-eeeeeeeeeeee",
]);
let ids = read_userids(db).unwrap();
assert_eq!(ids.len(), 2);
assert!(ids.contains(&"11111111-2222-3333-4444-555555555555".to_string()));
}
#[test]
fn skips_null_userids_but_keeps_the_rest() {
let file = tempfile::NamedTempFile::new().unwrap();
{
let conn = Connection::open(file.path()).unwrap();
conn.execute("CREATE TABLE user (UserID TEXT)", []).unwrap();
conn.execute("INSERT INTO user (UserID) VALUES (NULL)", [])
.unwrap();
conn.execute("INSERT INTO user (UserID) VALUES ('real-id')", [])
.unwrap();
}
let db = std::fs::read(file.path()).unwrap();
let ids = read_userids(db).unwrap();
assert_eq!(ids, vec!["real-id".to_string()]);
}
#[test]
fn missing_table_is_an_error_not_a_panic() {
let file = tempfile::NamedTempFile::new().unwrap();
{
let conn = Connection::open(file.path()).unwrap();
conn.execute("CREATE TABLE other (x TEXT)", []).unwrap();
}
let db = std::fs::read(file.path()).unwrap();
assert!(read_userids(db).is_err());
}
}