use crate::error::{Error, Result};
use crate::storage::StorageBackend;
use crate::utils::security::{ensure_secure_dir, set_secure_file_permissions};
use rusqlite::Connection;
use serde::{Serialize, de::DeserializeOwned};
use std::path::Path;
pub const DEFAULT_TABLE: &str = "rcman_settings";
pub const DEFAULT_KEY: &str = "main";
#[derive(Clone)]
pub struct SqliteStorage {
table_name: String,
key: String,
}
impl Default for SqliteStorage {
fn default() -> Self {
Self {
table_name: DEFAULT_TABLE.to_string(),
key: DEFAULT_KEY.to_string(),
}
}
}
impl SqliteStorage {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn with_table(mut self, table: impl Into<String>) -> Self {
let table = table.into();
if is_valid_identifier(&table) {
self.table_name = table;
} else {
log::warn!("rejected invalid SQLite table name {table:?}; keeping default");
}
self
}
#[must_use]
pub fn with_key(mut self, key: impl Into<String>) -> Self {
self.key = key.into();
self
}
fn connect(&self, path: &Path) -> Result<Connection> {
if !is_valid_identifier(&self.table_name) {
return Err(Error::Config(format!(
"invalid SQLite table name: {:?}",
self.table_name
)));
}
if let Some(parent) = path.parent()
&& !parent.as_os_str().is_empty()
&& !parent.exists()
{
ensure_secure_dir(parent)?;
}
Connection::open(path)
.map_err(|e| Error::Config(format!("sqlite open {}: {e}", path.display())))
}
fn ensure_schema(&self, conn: &Connection) -> Result<()> {
let sql = format!(
"CREATE TABLE IF NOT EXISTS {table} (
key TEXT PRIMARY KEY NOT NULL,
data TEXT NOT NULL
)",
table = self.table_name
);
conn.execute(&sql, [])
.map_err(|e| Error::Config(format!("sqlite create table: {e}")))?;
Ok(())
}
}
impl StorageBackend for SqliteStorage {
fn extension(&self) -> &'static str {
"db"
}
fn serialize<T: Serialize>(&self, data: &T) -> Result<String> {
serde_json::to_string(data).map_err(Error::from)
}
fn deserialize<T: DeserializeOwned>(&self, content: &str) -> Result<T> {
serde_json::from_str(content).map_err(Error::from)
}
fn read<T: DeserializeOwned>(&self, path: &Path) -> Result<T> {
let conn = self.connect(path)?;
self.ensure_schema(&conn)?;
let sql = format!(
"SELECT data FROM {table} WHERE key = ?1",
table = self.table_name
);
let row_data: Option<String> = conn
.query_row(&sql, rusqlite::params![self.key], |row| row.get(0))
.or_else(|e| match e {
rusqlite::Error::QueryReturnedNoRows => Ok(None),
_ => Err(e),
})
.map_err(|e| Error::Config(format!("sqlite query: {e}")))?;
match row_data {
Some(content) => self.deserialize(&content),
None => Err(Error::FileRead {
path: path.to_path_buf(),
source: std::io::Error::new(
std::io::ErrorKind::NotFound,
format!("no settings row for key {:?}", self.key),
),
}),
}
}
fn write<T: Serialize>(&self, path: &Path, data: &T) -> Result<()> {
let content = self.serialize(data)?;
let conn = self.connect(path)?;
self.ensure_schema(&conn)?;
let sql = format!(
"INSERT INTO {table} (key, data) VALUES (?1, ?2)
ON CONFLICT(key) DO UPDATE SET data = excluded.data",
table = self.table_name
);
conn.execute(&sql, rusqlite::params![self.key, content])
.map_err(|e| Error::Config(format!("sqlite upsert: {e}")))?;
let _ = set_secure_file_permissions(path);
Ok(())
}
}
fn is_valid_identifier(name: &str) -> bool {
let mut bytes = name.bytes();
let Some(first) = bytes.next() else {
return false;
};
if !(first.is_ascii_alphabetic() || first == b'_') {
return false;
}
bytes.all(|b| b.is_ascii_alphanumeric() || b == b'_')
}
#[cfg(test)]
mod tests {
use super::*;
use serde::{Deserialize, Serialize};
use tempfile::tempdir;
#[derive(Debug, Serialize, Deserialize, PartialEq)]
struct TestData {
name: String,
value: i32,
nested: Nested,
}
#[derive(Debug, Serialize, Deserialize, PartialEq)]
struct Nested {
flag: bool,
items: Vec<String>,
}
fn sample() -> TestData {
TestData {
name: "alice".into(),
value: 42,
nested: Nested {
flag: true,
items: vec!["a".into(), "b".into()],
},
}
}
#[test]
fn extension_is_db() {
assert_eq!(SqliteStorage::new().extension(), "db");
}
#[test]
fn roundtrip_default_settings() {
let storage = SqliteStorage::new();
let dir = tempdir().unwrap();
let path = dir.path().join("settings.db");
let data = sample();
storage.write(&path, &data).unwrap();
let loaded: TestData = storage.read(&path).unwrap();
assert_eq!(data, loaded);
}
#[test]
fn roundtrip_creates_parent_dirs() {
let storage = SqliteStorage::new();
let dir = tempdir().unwrap();
let path = dir.path().join("nested").join("dir").join("settings.db");
let data = sample();
storage.write(&path, &data).unwrap();
let loaded: TestData = storage.read(&path).unwrap();
assert_eq!(data, loaded);
}
#[test]
fn read_missing_path_errors() {
let storage = SqliteStorage::new();
let dir = tempdir().unwrap();
let path = dir.path().join("missing.db");
let result: Result<TestData> = storage.read(&path);
assert!(result.is_err());
match result.unwrap_err() {
Error::FileRead { .. } => {}
other => panic!("expected FileRead, got {other:?}"),
}
}
#[test]
fn write_overwrites_existing_row() {
let storage = SqliteStorage::new();
let dir = tempdir().unwrap();
let path = dir.path().join("settings.db");
let first = sample();
storage.write(&path, &first).unwrap();
let second = TestData {
name: "bob".into(),
value: 7,
nested: Nested {
flag: false,
items: vec![],
},
};
storage.write(&path, &second).unwrap();
let loaded: TestData = storage.read(&path).unwrap();
assert_eq!(loaded, second);
}
#[test]
fn custom_table_and_key_share_database() {
let dir = tempdir().unwrap();
let path = dir.path().join("multi.db");
let alpha = SqliteStorage::new().with_key("alpha");
let beta = SqliteStorage::new().with_key("beta");
alpha.write(&path, &sample()).unwrap();
beta.write(
&path,
&TestData {
name: "beta".into(),
value: 99,
nested: Nested {
flag: false,
items: vec!["z".into()],
},
},
)
.unwrap();
let a: TestData = alpha.read(&path).unwrap();
let b: TestData = beta.read(&path).unwrap();
assert_eq!(a.name, "alice");
assert_eq!(b.name, "beta");
}
#[test]
fn invalid_table_name_is_rejected_at_use() {
let storage = SqliteStorage::new().with_table("valid name with spaces");
let dir = tempdir().unwrap();
let path = dir.path().join("settings.db");
storage.write(&path, &sample()).unwrap();
let loaded: TestData = storage.read(&path).unwrap();
assert_eq!(loaded, sample());
}
#[test]
fn identifier_validation() {
assert!(is_valid_identifier("rcman_settings"));
assert!(is_valid_identifier("_private"));
assert!(is_valid_identifier("abc_123"));
assert!(!is_valid_identifier(""));
assert!(!is_valid_identifier("1leads_with_digit"));
assert!(!is_valid_identifier("has space"));
assert!(!is_valid_identifier("has;sql;injection"));
assert!(!is_valid_identifier("quoted\"name"));
}
}