sa_token_storage_database/
lib.rs1#[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#[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 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 pub fn from_pool(pool: Pool<Postgres>) -> Self {
50 Self { pool }
51 }
52
53 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
79pub 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}