Skip to main content

bone_persistence/backends/
sqlite.rs

1use chrono::Utc;
2use serde::{Deserialize, Serialize};
3use sqlx::sqlite::SqliteConnectOptions;
4use sqlx::{Row, SqlitePool};
5use std::collections::HashMap;
6use std::path::Path;
7
8use crate::error::{PersistenceError, Result};
9use crate::persistence::{Persistence, PersistentState, UpdateStrategy};
10
11/// SQLite 后端实现
12#[derive(Clone)]
13pub struct SqliteBackend {
14    pool: SqlitePool,
15}
16
17impl SqliteBackend {
18    /// 创建新的 SQLite 后端
19    pub async fn new(db_path: &Path) -> Result<Self> {
20        std::fs::create_dir_all(db_path.parent().unwrap())?;
21        let database_filename = format!("{}", db_path.display());
22        let pool = SqlitePool::connect_with(
23            SqliteConnectOptions::new()
24                .filename(&database_filename)
25                .create_if_missing(true),
26        )
27        .await?;
28
29        Ok(Self { pool })
30    }
31}
32
33#[async_trait::async_trait]
34impl<V, M> Persistence<V, M> for SqliteBackend
35where
36    V: Clone + Serialize + for<'de> Deserialize<'de> + Send + Sync + 'static,
37    M: Clone + Default + Serialize + for<'de> Deserialize<'de> + Send + Sync + 'static,
38{
39    /// 初始化
40    async fn init(&self) -> Result<()> {
41        // 创建工具数据项表
42        sqlx::query(
43            "CREATE TABLE IF NOT EXISTS data (
44                id TEXT PRIMARY KEY,
45                value_json TEXT NOT NULL,
46                created_at TEXT NOT NULL,
47                updated_at TEXT NOT NULL
48            )",
49        )
50        .execute(&self.pool)
51        .await?;
52
53        // 创建索引
54        sqlx::query("CREATE INDEX IF NOT EXISTS idx_values_updated_at ON data(updated_at)")
55            .execute(&self.pool)
56            .await?;
57
58        Ok(())
59    }
60
61    /// 更新整个持久化状态
62    async fn update(&self, state: PersistentState<V, M>) -> Result<()> {
63        let now = Utc::now().to_rfc3339();
64        let metadata_json = serde_json::to_string(&state.metadata)?;
65
66        // metadata
67        sqlx::query(
68            "INSERT OR REPLACE INTO data (id, value_json, created_at, updated_at)
69             VALUES (?1, ?2, COALESCE((SELECT created_at FROM data WHERE id = ?1), ?3), ?3)",
70        )
71        .bind("metadata")
72        .bind(&metadata_json)
73        .bind(&now)
74        .execute(&self.pool)
75        .await?;
76
77        // data
78        for item in &state.data {
79            let key = format!("value_{}", item.0);
80            let data_json = serde_json::to_string(&item.1)?;
81
82            sqlx::query(
83                "INSERT OR REPLACE INTO data (id, value_json, created_at, updated_at)
84                 VALUES (?1, ?2, COALESCE((SELECT created_at FROM data WHERE id = ?1), ?3), ?3)",
85            )
86            .bind(&key)
87            .bind(&data_json)
88            .bind(&now)
89            .execute(&self.pool)
90            .await?;
91        }
92
93        Ok(())
94    }
95
96    /// 加载整个持久化状态
97    async fn load(&self) -> Result<PersistentState<V, M>> {
98        let rows = sqlx::query("SELECT id, value_json FROM data")
99            .fetch_all(&self.pool)
100            .await?;
101
102        let mut data: HashMap<String, V> = HashMap::new();
103        let mut metadata: Option<M> = None;
104
105        for row in rows {
106            let id: String = row.get("id");
107            let value_json: String = row.get("value_json");
108
109            if id == "metadata" {
110                let _metadata: M = serde_json::from_str(&value_json)?;
111                metadata = Some(_metadata);
112            } else if id.starts_with("value_") {
113                let key = id.strip_prefix("value_").unwrap().to_string();
114                let _value: V = serde_json::from_str(&value_json)?;
115                data.insert(key, _value);
116            }
117        }
118
119        match metadata {
120            None => {
121                let metadata_json = serde_json::to_string(&M::default())?;
122                let now = Utc::now().to_rfc3339();
123                sqlx::query(
124                    "INSERT OR REPLACE INTO data (id, value_json, created_at, updated_at)
125                     VALUES (?1, ?2, COALESCE((SELECT created_at FROM data WHERE id = ?1), ?3), ?3)"
126                )
127                .bind("metadata")
128                .bind(&metadata_json)
129                .bind(&now)
130                .execute(&self.pool)
131                .await?;
132                Ok(PersistentState {
133                    metadata: M::default(),
134                    data,
135                })
136            }
137            Some(metadata) => Ok(PersistentState { metadata, data }),
138        }
139    }
140
141    /// 添加新数据项
142    async fn add_item(&self, key: &str, value: V, strategy: Option<UpdateStrategy>) -> Result<()> {
143        match strategy {
144            Some(UpdateStrategy::Overwrite) | None => {
145                let now = Utc::now().to_rfc3339();
146                let value_json = serde_json::to_string(&value)?;
147
148                sqlx::query(
149                    "INSERT OR REPLACE INTO data (id, value_json, created_at, updated_at)
150                     VALUES (?1, ?2, COALESCE((SELECT created_at FROM data WHERE id = ?1), ?3), ?3)"
151                )
152                .bind(format!("value_{}", key))
153                .bind(&value_json)
154                .bind(&now)
155                .execute(&self.pool)
156                .await?;
157
158                Ok(())
159            }
160        }
161    }
162
163    /// 加载指定数据项
164    async fn get_item(&self, key: &str) -> Result<V> {
165        let row = sqlx::query("SELECT value_json FROM data WHERE id = ?1")
166            .bind(format!("value_{}", key))
167            .fetch_optional(&self.pool)
168            .await?;
169
170        let row = row.ok_or_else(|| PersistenceError::NotFound {
171            message: "未找到指定的数据项".to_string(),
172        })?;
173
174        let value_json: String = row.get("value_json");
175        let value: V = serde_json::from_str(&value_json)?;
176
177        Ok(value)
178    }
179
180    /// 删除数据项
181    async fn remove_item(&self, key: &str) -> Result<()> {
182        sqlx::query("DELETE FROM data WHERE id = ?1")
183            .bind(format!("value_{}", key))
184            .execute(&self.pool)
185            .await?;
186        Ok(())
187    }
188
189    /// 列出所有数据项键
190    async fn list_keys(&self) -> Result<Vec<String>> {
191        let rows = sqlx::query("SELECT id FROM data ORDER BY updated_at DESC")
192            .fetch_all(&self.pool)
193            .await?;
194
195        let mut keys = Vec::new();
196        for row in rows {
197            let id: String = row.get("id");
198            if let Some(key) = id.strip_prefix("value_") {
199                keys.push(key.to_string());
200            }
201        }
202
203        Ok(keys)
204    }
205
206    /// 列出所有数据项
207    async fn list_items(&self) -> Result<Vec<V>> {
208        let rows = sqlx::query("SELECT id, value_json FROM data ORDER BY updated_at DESC")
209            .fetch_all(&self.pool)
210            .await?;
211
212        let mut items = Vec::new();
213        for row in rows {
214            let id: String = row.get("id");
215            if id.starts_with("value_") {
216                let value_json: String = row.get("value_json");
217                let value: V = serde_json::from_str(&value_json)?;
218                items.push(value);
219            }
220        }
221
222        Ok(items)
223    }
224}