1use std::sync::Arc;
14
15use async_trait::async_trait;
16use chrono::{DateTime, Utc};
17use dashmap::DashMap;
18use serde::Deserialize;
19use tracing::{debug, info, warn};
20
21#[derive(Debug, Clone, Deserialize)]
27pub struct TokenRevocationConfig {
28 #[serde(default)]
30 pub enabled: bool,
31
32 #[serde(default = "default_backend")]
34 pub backend: String,
35
36 #[serde(default = "default_true")]
38 pub require_jti: bool,
39
40 #[serde(default)]
44 pub fail_open: bool,
45
46 pub redis_url: Option<String>,
48}
49
50fn default_backend() -> String {
51 "memory".into()
52}
53const fn default_true() -> bool {
54 true
55}
56
57#[async_trait]
65pub trait RevocationStore: Send + Sync {
66 async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError>;
68
69 async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError>;
72
73 async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError>;
76}
77
78#[derive(Debug, thiserror::Error)]
80#[non_exhaustive]
81pub enum RevocationError {
82 #[error("revocation store error: {0}")]
84 Backend(String),
85}
86
87pub struct InMemoryRevocationStore {
93 pub(crate) entries: DashMap<String, (String, DateTime<Utc>)>,
95}
96
97impl InMemoryRevocationStore {
98 #[must_use]
100 pub fn new() -> Self {
101 Self {
102 entries: DashMap::new(),
103 }
104 }
105
106 pub fn cleanup_expired(&self) {
108 let now = Utc::now();
109 self.entries.retain(|_, (_, exp)| *exp > now);
110 }
111}
112
113impl Default for InMemoryRevocationStore {
114 fn default() -> Self {
115 Self::new()
116 }
117}
118
119#[async_trait]
123impl RevocationStore for InMemoryRevocationStore {
124 async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError> {
125 if let Some(entry) = self.entries.get(jti) {
126 let (_, expires_at) = entry.value();
127 if *expires_at > Utc::now() {
128 return Ok(true);
129 }
130 drop(entry);
132 self.entries.remove(jti);
133 }
134 Ok(false)
135 }
136
137 async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
138 let expires_at = Utc::now() + chrono::Duration::seconds(ttl_secs.cast_signed());
139 self.entries.insert(jti.to_string(), (String::new(), expires_at));
141 Ok(())
142 }
143
144 async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
145 let keys_to_remove: Vec<String> = self
149 .entries
150 .iter()
151 .filter(|entry| {
152 let (s, _) = entry.value();
153 s == sub
154 })
155 .map(|entry| entry.key().clone())
156 .collect();
157
158 let count = keys_to_remove.len() as u64;
159 for key in &keys_to_remove {
160 self.entries.remove(key);
161 }
162 Ok(count)
163 }
164}
165
166#[cfg(feature = "redis-rate-limiting")]
175pub struct RedisRevocationStore {
176 client: redis::Client,
177 key_prefix: String,
178}
179
180#[cfg(feature = "redis-rate-limiting")]
181impl RedisRevocationStore {
182 pub fn new(redis_url: &str) -> Result<Self, RevocationError> {
188 let client = redis::Client::open(redis_url)
189 .map_err(|e| RevocationError::Backend(format!("Redis connection error: {e}")))?;
190 Ok(Self {
191 client,
192 key_prefix: "fraiseql:revoked:".into(),
193 })
194 }
195}
196
197#[cfg(feature = "redis-rate-limiting")]
198#[async_trait]
202impl RevocationStore for RedisRevocationStore {
203 async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError> {
204 use redis::AsyncCommands;
205 let mut conn = self
206 .client
207 .get_multiplexed_async_connection()
208 .await
209 .map_err(|e| RevocationError::Backend(format!("Redis: {e}")))?;
210 let key = format!("{}{jti}", self.key_prefix);
211 let exists: bool = conn
212 .exists(&key)
213 .await
214 .map_err(|e| RevocationError::Backend(format!("Redis EXISTS: {e}")))?;
215 Ok(exists)
216 }
217
218 async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
219 use redis::AsyncCommands;
220 let mut conn = self
221 .client
222 .get_multiplexed_async_connection()
223 .await
224 .map_err(|e| RevocationError::Backend(format!("Redis: {e}")))?;
225 let key = format!("{}{jti}", self.key_prefix);
226 let _: () = conn
227 .set_ex(&key, "1", ttl_secs)
228 .await
229 .map_err(|e| RevocationError::Backend(format!("Redis SET EX: {e}")))?;
230 Ok(())
231 }
232
233 async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
234 let mut conn = self
235 .client
236 .get_multiplexed_async_connection()
237 .await
238 .map_err(|e| RevocationError::Backend(format!("Redis: {e}")))?;
239 let pattern = format!("{}user:{sub}:*", self.key_prefix);
244 let mut cursor: u64 = 0;
245 let mut all_keys: Vec<String> = Vec::new();
246 loop {
247 let (next_cursor, batch): (u64, Vec<String>) = redis::cmd("SCAN")
248 .arg(cursor)
249 .arg("MATCH")
250 .arg(&pattern)
251 .arg("COUNT")
252 .arg(100u32)
253 .query_async(&mut conn)
254 .await
255 .map_err(|e| RevocationError::Backend(format!("Redis SCAN: {e}")))?;
256 all_keys.extend(batch);
257 cursor = next_cursor;
258 if cursor == 0 {
259 break;
260 }
261 }
262 let count = all_keys.len() as u64;
263 if !all_keys.is_empty() {
264 let _: () = redis::cmd("DEL")
265 .arg(&all_keys)
266 .query_async(&mut conn)
267 .await
268 .map_err(|e| RevocationError::Backend(format!("Redis DEL: {e}")))?;
269 }
270 Ok(count)
271 }
272}
273
274pub struct TokenRevocationManager {
280 store: Arc<dyn RevocationStore>,
281 require_jti: bool,
282 fail_open: bool,
283}
284
285impl TokenRevocationManager {
286 #[must_use]
288 pub fn new(store: Arc<dyn RevocationStore>, require_jti: bool, fail_open: bool) -> Self {
289 Self {
290 store,
291 require_jti,
292 fail_open,
293 }
294 }
295
296 pub async fn check_token(&self, jti: Option<&str>) -> Result<(), TokenRejection> {
307 let jti = match jti {
308 Some(j) if !j.is_empty() => j,
309 _ => {
310 if self.require_jti {
311 return Err(TokenRejection::MissingJti);
312 }
313 return Ok(());
315 },
316 };
317
318 match self.store.is_revoked(jti).await {
319 Ok(true) => Err(TokenRejection::Revoked),
320 Ok(false) => Ok(()),
321 Err(e) => {
322 warn!(error = %e, jti = %jti, "Revocation store check failed");
323 if self.fail_open {
324 debug!("fail_open=true — allowing request despite store error");
325 Ok(())
326 } else {
327 Err(TokenRejection::StoreUnavailable)
328 }
329 },
330 }
331 }
332
333 pub async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
339 self.store.revoke(jti, ttl_secs).await
340 }
341
342 pub async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
348 self.store.revoke_all_for_user(sub).await
349 }
350
351 #[must_use]
353 pub const fn require_jti(&self) -> bool {
354 self.require_jti
355 }
356}
357
358impl std::fmt::Debug for TokenRevocationManager {
359 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
360 f.debug_struct("TokenRevocationManager")
361 .field("require_jti", &self.require_jti)
362 .field("fail_open", &self.fail_open)
363 .finish_non_exhaustive()
364 }
365}
366
367#[derive(Debug, Clone, PartialEq, Eq)]
369#[non_exhaustive]
370pub enum TokenRejection {
371 Revoked,
373 MissingJti,
375 StoreUnavailable,
377}
378
379pub fn revocation_manager_from_schema(
385 schema: &fraiseql_core::schema::CompiledSchema,
386) -> Option<Arc<TokenRevocationManager>> {
387 let security = schema.security.as_ref()?;
388 let revocation_val = security.additional.get("token_revocation")?;
389 let config: TokenRevocationConfig = serde_json::from_value(revocation_val.clone())
390 .map_err(|e| {
391 warn!(error = %e, "Failed to parse security.token_revocation config");
392 })
393 .ok()?;
394
395 if !config.enabled {
396 return None;
397 }
398
399 let store: Arc<dyn RevocationStore> = match config.backend.as_str() {
400 #[cfg(feature = "redis-rate-limiting")]
401 "redis" => {
402 let url = config.redis_url.as_deref().unwrap_or("redis://localhost:6379");
403 match RedisRevocationStore::new(url) {
404 Ok(s) => {
405 info!(backend = "redis", "Token revocation store initialized");
406 Arc::new(s)
407 },
408 Err(e) => {
409 warn!(error = %e, "Failed to init Redis revocation store — falling back to in-memory");
410 Arc::new(InMemoryRevocationStore::new())
411 },
412 }
413 },
414 #[cfg(not(feature = "redis-rate-limiting"))]
415 "redis" => {
416 warn!(
417 "token_revocation.backend = \"redis\" but the `redis-rate-limiting` feature is \
418 not compiled in. Falling back to in-memory."
419 );
420 Arc::new(InMemoryRevocationStore::new())
421 },
422 "memory" | "env" => {
423 info!(backend = "memory", "Token revocation store initialized (in-memory)");
424 Arc::new(InMemoryRevocationStore::new())
425 },
426 other => {
427 warn!(backend = %other, "Unknown revocation backend — falling back to in-memory");
428 Arc::new(InMemoryRevocationStore::new())
429 },
430 };
431
432 Some(Arc::new(TokenRevocationManager::new(
433 store,
434 config.require_jti,
435 config.fail_open,
436 )))
437}