use chrono::Utc;
use serde::{Deserialize, Serialize};
use sqlx::sqlite::SqliteConnectOptions;
use sqlx::{Row, SqlitePool};
use std::collections::HashMap;
use std::path::Path;
use crate::error::{PersistenceError, Result};
use crate::persistence::{Persistence, PersistentState, UpdateStrategy};
#[derive(Clone)]
pub struct SqliteBackend {
pool: SqlitePool,
}
impl SqliteBackend {
pub async fn new(db_path: &Path) -> Result<Self> {
std::fs::create_dir_all(db_path.parent().unwrap())?;
let database_filename = format!("{}", db_path.display());
let pool = SqlitePool::connect_with(
SqliteConnectOptions::new()
.filename(&database_filename)
.create_if_missing(true),
)
.await?;
Ok(Self { pool })
}
}
#[async_trait::async_trait]
impl<V, M> Persistence<V, M> for SqliteBackend
where
V: Clone + Serialize + for<'de> Deserialize<'de> + Send + Sync + 'static,
M: Clone + Default + Serialize + for<'de> Deserialize<'de> + Send + Sync + 'static,
{
async fn init(&self) -> Result<()> {
sqlx::query(
"CREATE TABLE IF NOT EXISTS data (
id TEXT PRIMARY KEY,
value_json TEXT NOT NULL,
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
)",
)
.execute(&self.pool)
.await?;
sqlx::query("CREATE INDEX IF NOT EXISTS idx_values_updated_at ON data(updated_at)")
.execute(&self.pool)
.await?;
Ok(())
}
async fn update(&self, state: PersistentState<V, M>) -> Result<()> {
let now = Utc::now().to_rfc3339();
let metadata_json = serde_json::to_string(&state.metadata)?;
sqlx::query(
"INSERT OR REPLACE INTO data (id, value_json, created_at, updated_at)
VALUES (?1, ?2, COALESCE((SELECT created_at FROM data WHERE id = ?1), ?3), ?3)",
)
.bind("metadata")
.bind(&metadata_json)
.bind(&now)
.execute(&self.pool)
.await?;
for item in &state.data {
let key = format!("value_{}", item.0);
let data_json = serde_json::to_string(&item.1)?;
sqlx::query(
"INSERT OR REPLACE INTO data (id, value_json, created_at, updated_at)
VALUES (?1, ?2, COALESCE((SELECT created_at FROM data WHERE id = ?1), ?3), ?3)",
)
.bind(&key)
.bind(&data_json)
.bind(&now)
.execute(&self.pool)
.await?;
}
Ok(())
}
async fn load(&self) -> Result<PersistentState<V, M>> {
let rows = sqlx::query("SELECT id, value_json FROM data")
.fetch_all(&self.pool)
.await?;
let mut data: HashMap<String, V> = HashMap::new();
let mut metadata: Option<M> = None;
for row in rows {
let id: String = row.get("id");
let value_json: String = row.get("value_json");
if id == "metadata" {
let _metadata: M = serde_json::from_str(&value_json)?;
metadata = Some(_metadata);
} else if id.starts_with("value_") {
let key = id.strip_prefix("value_").unwrap().to_string();
let _value: V = serde_json::from_str(&value_json)?;
data.insert(key, _value);
}
}
match metadata {
None => {
let metadata_json = serde_json::to_string(&M::default())?;
let now = Utc::now().to_rfc3339();
sqlx::query(
"INSERT OR REPLACE INTO data (id, value_json, created_at, updated_at)
VALUES (?1, ?2, COALESCE((SELECT created_at FROM data WHERE id = ?1), ?3), ?3)"
)
.bind("metadata")
.bind(&metadata_json)
.bind(&now)
.execute(&self.pool)
.await?;
Ok(PersistentState {
metadata: M::default(),
data,
})
}
Some(metadata) => Ok(PersistentState { metadata, data }),
}
}
async fn add_item(&self, key: &str, value: V, strategy: Option<UpdateStrategy>) -> Result<()> {
match strategy {
Some(UpdateStrategy::Overwrite) | None => {
let now = Utc::now().to_rfc3339();
let value_json = serde_json::to_string(&value)?;
sqlx::query(
"INSERT OR REPLACE INTO data (id, value_json, created_at, updated_at)
VALUES (?1, ?2, COALESCE((SELECT created_at FROM data WHERE id = ?1), ?3), ?3)"
)
.bind(format!("value_{}", key))
.bind(&value_json)
.bind(&now)
.execute(&self.pool)
.await?;
Ok(())
}
}
}
async fn get_item(&self, key: &str) -> Result<V> {
let row = sqlx::query("SELECT value_json FROM data WHERE id = ?1")
.bind(format!("value_{}", key))
.fetch_optional(&self.pool)
.await?;
let row = row.ok_or_else(|| PersistenceError::NotFound {
message: "未找到指定的数据项".to_string(),
})?;
let value_json: String = row.get("value_json");
let value: V = serde_json::from_str(&value_json)?;
Ok(value)
}
async fn remove_item(&self, key: &str) -> Result<()> {
sqlx::query("DELETE FROM data WHERE id = ?1")
.bind(format!("value_{}", key))
.execute(&self.pool)
.await?;
Ok(())
}
async fn list_keys(&self) -> Result<Vec<String>> {
let rows = sqlx::query("SELECT id FROM data ORDER BY updated_at DESC")
.fetch_all(&self.pool)
.await?;
let mut keys = Vec::new();
for row in rows {
let id: String = row.get("id");
if let Some(key) = id.strip_prefix("value_") {
keys.push(key.to_string());
}
}
Ok(keys)
}
async fn list_items(&self) -> Result<Vec<V>> {
let rows = sqlx::query("SELECT id, value_json FROM data ORDER BY updated_at DESC")
.fetch_all(&self.pool)
.await?;
let mut items = Vec::new();
for row in rows {
let id: String = row.get("id");
if id.starts_with("value_") {
let value_json: String = row.get("value_json");
let value: V = serde_json::from_str(&value_json)?;
items.push(value);
}
}
Ok(items)
}
}