bone-persistence 0.1.0

Persistence layer for BoneTools
Documentation
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};

/// SQLite 后端实现
#[derive(Clone)]
pub struct SqliteBackend {
    pool: SqlitePool,
}

impl SqliteBackend {
    /// 创建新的 SQLite 后端
    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)?;

        // 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?;

        // data
        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)
    }
}