Skip to main content

fraiseql_server/
token_revocation.rs

1//! Token revocation — reject JWTs whose `jti` claim has been revoked.
2//!
3//! After JWT signature verification succeeds, the server checks the token's
4//! `jti` (JWT ID) claim against a revocation store.  If the `jti` is present,
5//! the token is rejected with 401.
6//!
7//! Two production backends: Redis (recommended) and PostgreSQL (fallback).
8//! An in-memory backend is provided for testing and single-instance dev.
9//!
10//! Revoked JTIs expire automatically when the JWT's `exp` claim passes, keeping
11//! the store bounded.
12
13use 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// ───────────────────────────────────────────────────────────────
22// Configuration
23// ───────────────────────────────────────────────────────────────
24
25/// Token revocation configuration embedded in the compiled schema.
26#[derive(Debug, Clone, Deserialize)]
27pub struct TokenRevocationConfig {
28    /// Whether token revocation is enabled.
29    #[serde(default)]
30    pub enabled: bool,
31
32    /// Storage backend: `"redis"` or `"postgres"` or `"memory"`.
33    #[serde(default = "default_backend")]
34    pub backend: String,
35
36    /// Reject JWTs that lack a `jti` claim when revocation is enabled.
37    #[serde(default = "default_true")]
38    pub require_jti: bool,
39
40    /// If the revocation store is unreachable:
41    /// - `false` (default): reject the request (fail-closed)
42    /// - `true`: allow the request (fail-open)
43    #[serde(default)]
44    pub fail_open: bool,
45
46    /// Redis URL (inherited from `[fraiseql.redis]` if not set here).
47    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// ───────────────────────────────────────────────────────────────
58// Trait
59// ───────────────────────────────────────────────────────────────
60
61/// Revocation store abstraction.
62// Reason: used as dyn Trait (Arc<dyn RevocationStore>); async_trait ensures Send bounds and
63// dyn-compatibility async_trait: dyn-dispatch required; remove when RTN + Send is stable (RFC 3425)
64#[async_trait]
65pub trait RevocationStore: Send + Sync {
66    /// Check if a JTI has been revoked.
67    async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError>;
68
69    /// Revoke a single JTI.  `ttl_secs` is the remaining JWT lifetime —
70    /// the store should auto-expire the entry after this duration.
71    async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError>;
72
73    /// Revoke all tokens for a user (by `sub` claim).
74    /// Returns the number of tokens revoked.
75    async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError>;
76}
77
78/// Revocation store error.
79#[derive(Debug, thiserror::Error)]
80#[non_exhaustive]
81pub enum RevocationError {
82    /// Backend is unreachable or returned an error.
83    #[error("revocation store error: {0}")]
84    Backend(String),
85}
86
87// ───────────────────────────────────────────────────────────────
88// In-memory backend
89// ───────────────────────────────────────────────────────────────
90
91/// In-memory revocation store for testing and single-instance dev.
92pub struct InMemoryRevocationStore {
93    /// Map of JTI → (sub, `expires_at`).
94    pub(crate) entries: DashMap<String, (String, DateTime<Utc>)>,
95}
96
97impl InMemoryRevocationStore {
98    /// Create a new, empty in-memory revocation store.
99    #[must_use]
100    pub fn new() -> Self {
101        Self {
102            entries: DashMap::new(),
103        }
104    }
105
106    /// Remove expired entries.
107    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// Reason: RevocationStore is defined with #[async_trait]; all implementations must match
120// its transformed method signatures to satisfy the trait contract
121// async_trait: dyn-dispatch required; remove when RTN + Send is stable (RFC 3425)
122#[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            // Expired — remove lazily.
131            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        // We store an empty sub — single-JTI revocation doesn't need sub.
140        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        // Collect all JTIs belonging to this user and remove them from the store.
146        // Two-pass approach (collect keys, then remove) avoids holding a mutable
147        // reference to DashMap while iterating, which would deadlock.
148        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// ───────────────────────────────────────────────────────────────
167// Redis backend (optional)
168// ───────────────────────────────────────────────────────────────
169
170/// Redis-backed JWT revocation store.
171///
172/// Stores revoked JTI claims in Redis with automatic TTL-based expiry.
173/// Requires the `redis-rate-limiting` feature.
174#[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    /// Create a new Redis-backed revocation store.
183    ///
184    /// # Errors
185    ///
186    /// Returns error if the Redis URL is invalid.
187    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// Reason: RevocationStore is defined with #[async_trait]; all implementations must match
199// its transformed method signatures to satisfy the trait contract
200// async_trait: dyn-dispatch required; remove when RTN + Send is stable (RFC 3425)
201#[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        // SECURITY: Use SCAN cursor iteration instead of KEYS to avoid O(N) blocking.
240        // KEYS blocks Redis for the entire scan duration; SCAN is non-blocking and
241        // yields results in small batches, making it safe for production use.
242        // User-keyed entries use prefix: fraiseql:revoked:user:{sub}:*
243        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
274// ───────────────────────────────────────────────────────────────
275// Token Revocation Manager
276// ───────────────────────────────────────────────────────────────
277
278/// High-level token revocation manager wrapping a backend store.
279pub struct TokenRevocationManager {
280    store:       Arc<dyn RevocationStore>,
281    require_jti: bool,
282    fail_open:   bool,
283}
284
285impl TokenRevocationManager {
286    /// Create a new revocation manager.
287    #[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    /// Check if a token should be rejected.
297    ///
298    /// Returns `Ok(())` if the token is allowed, or an error reason if rejected.
299    ///
300    /// # Errors
301    ///
302    /// Returns `TokenRejection::MissingJti` if JTI is required but absent.
303    /// Returns `TokenRejection::Revoked` if the token has been revoked.
304    /// Returns `TokenRejection::StoreUnavailable` if the revocation store is unreachable and
305    /// `fail_open` is false.
306    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                // No JTI and not required — allow through.
314                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    /// Revoke a single token by JTI.
334    ///
335    /// # Errors
336    ///
337    /// Returns `RevocationError` if the underlying revocation store operation fails.
338    pub async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
339        self.store.revoke(jti, ttl_secs).await
340    }
341
342    /// Revoke all tokens for a user.
343    ///
344    /// # Errors
345    ///
346    /// Returns `RevocationError` if the underlying revocation store operation fails.
347    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    /// Whether JTI is required.
352    #[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/// Why a token was rejected.
368#[derive(Debug, Clone, PartialEq, Eq)]
369#[non_exhaustive]
370pub enum TokenRejection {
371    /// Token has been revoked.
372    Revoked,
373    /// Token lacks a `jti` claim and `require_jti` is enabled.
374    MissingJti,
375    /// Revocation store is unavailable and `fail_open` is false.
376    StoreUnavailable,
377}
378
379// ───────────────────────────────────────────────────────────────
380// Builder from compiled schema
381// ───────────────────────────────────────────────────────────────
382
383/// Build a `TokenRevocationManager` from the compiled schema's `security.token_revocation` JSON.
384pub 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}