Skip to main content

sa_token_storage_database/
lib.rs

1// Author: 金书记
2//
3//! # sa-token-storage-database
4//!
5//! 基于 sqlx 的关系型数据库存储实现(默认 PostgreSQL,可选 MySQL)。
6//!
7//! ## DDL
8//!
9//! 见 [`migrations/001_sa_token_storage.sql`](../../migrations/001_sa_token_storage.sql)。
10
11// postgres 与 mysql 后端互斥,避免 sqlx 同时编译两套驱动
12#[cfg(all(feature = "postgres", feature = "mysql"))]
13compile_error!(
14    "Features `postgres` and `mysql` are mutually exclusive; enable only one database backend"
15);
16
17use std::time::Duration;
18
19use async_trait::async_trait;
20use chrono::{DateTime, Utc};
21use sa_token_adapter::storage::{SaStorage, ScanPage, StorageError, StorageResult};
22use sqlx::{Pool, Postgres};
23
24/// PostgreSQL 存储实现
25#[derive(Clone)]
26pub struct DatabaseStorage {
27    pool: Pool<Postgres>,
28}
29
30impl std::fmt::Debug for DatabaseStorage {
31    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
32        f.write_str("DatabaseStorage { .. }")
33    }
34}
35
36impl DatabaseStorage {
37    /// 连接数据库并确保表结构存在
38    pub async fn new(database_url: &str) -> StorageResult<Self> {
39        let pool = Pool::<Postgres>::connect(database_url)
40            .await
41            .map_err(|e| StorageError::ConnectionError(e.to_string()))?;
42
43        let storage = Self { pool };
44        storage.migrate().await?;
45        Ok(storage)
46    }
47
48    /// 使用已有连接池
49    pub fn from_pool(pool: Pool<Postgres>) -> Self {
50        Self { pool }
51    }
52
53    /// 执行内嵌 DDL(幂等)
54    pub async fn migrate(&self) -> StorageResult<()> {
55        let ddl = include_str!("../migrations/001_sa_token_storage.sql");
56        for statement in ddl.split(';').map(str::trim).filter(|s| !s.is_empty()) {
57            sqlx::query(statement)
58                .execute(&self.pool)
59                .await
60                .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
61        }
62        Ok(())
63    }
64
65    async fn delete_expired(&self, key: &str) -> StorageResult<()> {
66        sqlx::query("DELETE FROM sa_token_storage WHERE key = $1")
67            .bind(key)
68            .execute(&self.pool)
69            .await
70            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
71        Ok(())
72    }
73
74    fn is_expired(expire_at: Option<DateTime<Utc>>) -> bool {
75        expire_at.is_some_and(|t| Utc::now() > t)
76    }
77}
78
79/// 将 `*` 通配符转为 SQL LIKE 模式,并转义 `%` / `_`
80pub fn pattern_to_like(pattern: &str) -> String {
81    let mut out = String::with_capacity(pattern.len());
82    for ch in pattern.chars() {
83        match ch {
84            '*' => out.push('%'),
85            '%' | '_' | '\\' => {
86                out.push('\\');
87                out.push(ch);
88            }
89            other => out.push(other),
90        }
91    }
92    out
93}
94
95#[async_trait]
96impl SaStorage for DatabaseStorage {
97    async fn get(&self, key: &str) -> StorageResult<Option<String>> {
98        let row: Option<(String, Option<DateTime<Utc>>)> =
99            sqlx::query_as("SELECT value, expire_at FROM sa_token_storage WHERE key = $1")
100                .bind(key)
101                .fetch_optional(&self.pool)
102                .await
103                .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
104
105        match row {
106            Some((_value, expire_at)) if Self::is_expired(expire_at) => {
107                self.delete_expired(key).await?;
108                Ok(None)
109            }
110            Some((value, _)) => Ok(Some(value)),
111            None => Ok(None),
112        }
113    }
114
115    async fn set(&self, key: &str, value: &str, ttl: Option<Duration>) -> StorageResult<()> {
116        let expire_at: Option<DateTime<Utc>> = ttl
117            .and_then(|d| chrono::Duration::from_std(d).ok())
118            .map(|d| Utc::now() + d);
119
120        sqlx::query(
121            r#"
122            INSERT INTO sa_token_storage (key, value, expire_at, updated_at)
123            VALUES ($1, $2, $3, NOW())
124            ON CONFLICT (key) DO UPDATE
125            SET value = EXCLUDED.value,
126                expire_at = EXCLUDED.expire_at,
127                updated_at = NOW()
128            "#,
129        )
130        .bind(key)
131        .bind(value)
132        .bind(expire_at)
133        .execute(&self.pool)
134        .await
135        .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
136
137        Ok(())
138    }
139
140    async fn delete(&self, key: &str) -> StorageResult<()> {
141        sqlx::query("DELETE FROM sa_token_storage WHERE key = $1")
142            .bind(key)
143            .execute(&self.pool)
144            .await
145            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
146        Ok(())
147    }
148
149    async fn exists(&self, key: &str) -> StorageResult<bool> {
150        Ok(self.get(key).await?.is_some())
151    }
152
153    async fn expire(&self, key: &str, ttl: Duration) -> StorageResult<()> {
154        let Some(delta) = chrono::Duration::from_std(ttl).ok() else {
155            return Ok(());
156        };
157        let expire_at = Utc::now() + delta;
158        let updated = sqlx::query(
159            "UPDATE sa_token_storage SET expire_at = $1, updated_at = NOW() WHERE key = $2",
160        )
161        .bind(expire_at)
162        .bind(key)
163        .execute(&self.pool)
164        .await
165        .map_err(|e| StorageError::OperationFailed(e.to_string()))?
166        .rows_affected();
167
168        if updated == 0 {
169            return Err(StorageError::KeyNotFound(key.to_string()));
170        }
171        Ok(())
172    }
173
174    async fn ttl(&self, key: &str) -> StorageResult<Option<Duration>> {
175        let row: Option<Option<DateTime<Utc>>> =
176            sqlx::query_scalar("SELECT expire_at FROM sa_token_storage WHERE key = $1")
177                .bind(key)
178                .fetch_optional(&self.pool)
179                .await
180                .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
181
182        match row {
183            None => Ok(None),
184            Some(None) => Ok(None),
185            Some(Some(expire_at)) if Self::is_expired(Some(expire_at)) => {
186                self.delete_expired(key).await?;
187                Ok(None)
188            }
189            Some(Some(expire_at)) => {
190                let remaining = expire_at.signed_duration_since(Utc::now());
191                if remaining.num_milliseconds() <= 0 {
192                    Ok(Some(Duration::ZERO))
193                } else {
194                    Ok(Some(remaining.to_std().unwrap_or(Duration::ZERO)))
195                }
196            }
197        }
198    }
199
200    async fn clear(&self) -> StorageResult<()> {
201        sqlx::query("TRUNCATE TABLE sa_token_storage")
202            .execute(&self.pool)
203            .await
204            .map_err(|e| StorageError::OperationFailed(e.to_string()))?;
205        Ok(())
206    }
207
208    async fn set_if_absent(
209        &self,
210        _key: &str,
211        _value: &str,
212        _ttl: Option<Duration>,
213    ) -> StorageResult<bool> {
214        Err(StorageError::Unsupported("set_if_absent"))
215    }
216
217    async fn get_del(&self, _key: &str) -> StorageResult<Option<String>> {
218        Err(StorageError::Unsupported("get_del"))
219    }
220
221    async fn compare_and_swap(
222        &self,
223        _key: &str,
224        _expected: Option<&str>,
225        _new_value: &str,
226        _ttl: Option<Duration>,
227    ) -> StorageResult<bool> {
228        Err(StorageError::Unsupported("compare_and_swap"))
229    }
230
231    async fn compare_and_delete(&self, _key: &str, _expected: &str) -> StorageResult<bool> {
232        Err(StorageError::Unsupported("compare_and_delete"))
233    }
234
235    async fn list_push(
236        &self,
237        _key: &str,
238        _member: &str,
239        _unique: bool,
240        _ttl: Option<Duration>,
241    ) -> StorageResult<usize> {
242        Err(StorageError::Unsupported("list_push"))
243    }
244
245    async fn list_remove(&self, _key: &str, _member: &str) -> StorageResult<bool> {
246        Err(StorageError::Unsupported("list_remove"))
247    }
248
249    async fn list_range(
250        &self,
251        _key: &str,
252        _start: usize,
253        _limit: Option<usize>,
254    ) -> StorageResult<Vec<String>> {
255        Err(StorageError::Unsupported("list_range"))
256    }
257
258    async fn list_len(&self, _key: &str) -> StorageResult<usize> {
259        Err(StorageError::Unsupported("list_len"))
260    }
261
262    async fn scan(&self, _pattern: &str, _cursor: u64, _limit: usize) -> StorageResult<ScanPage> {
263        Err(StorageError::Unsupported("scan"))
264    }
265}
266
267#[cfg(test)]
268mod tests {
269    use super::*;
270
271    #[test]
272    fn like_pattern_escapes_wildcards() {
273        assert_eq!(pattern_to_like("sa:token:*"), "sa:token:%");
274        assert_eq!(pattern_to_like("a%b_c"), "a\\%b\\_c");
275    }
276}
277
278#[cfg(all(test, feature = "postgres"))]
279mod postgres_tests {
280    use super::*;
281
282    fn database_url() -> Option<String> {
283        std::env::var("DATABASE_URL").ok()
284    }
285
286    #[tokio::test]
287    #[ignore = "requires PostgreSQL (set DATABASE_URL)"]
288    async fn database_storage_roundtrip() {
289        let Some(url) = database_url() else {
290            return;
291        };
292        let storage = DatabaseStorage::new(&url).await.expect("connect");
293        storage
294            .set("sa:test:1", "v1", Some(Duration::from_secs(60)))
295            .await
296            .unwrap();
297        assert_eq!(storage.get("sa:test:1").await.unwrap(), Some("v1".into()));
298        assert!(storage.exists("sa:test:1").await.unwrap());
299        let ttl = storage.ttl("sa:test:1").await.unwrap();
300        assert!(ttl.is_some());
301        storage.delete("sa:test:1").await.unwrap();
302        assert!(!storage.exists("sa:test:1").await.unwrap());
303
304        let err = storage.scan("sa:test:*", 0, 100).await.unwrap_err();
305        assert!(matches!(err, StorageError::Unsupported(_)));
306    }
307}