1#[cfg(test)]
14mod tests;
15
16use std::sync::Arc;
17
18use async_trait::async_trait;
19use chrono::{DateTime, Utc};
20use dashmap::DashMap;
21use serde::Deserialize;
22use tracing::{debug, info, warn};
23
24#[derive(Debug, Clone, Deserialize)]
30pub struct TokenRevocationConfig {
31 #[serde(default)]
33 pub enabled: bool,
34
35 #[serde(default = "default_backend")]
37 pub backend: String,
38
39 #[serde(default = "default_true")]
41 pub require_jti: bool,
42
43 #[serde(default)]
47 pub fail_open: bool,
48
49 pub redis_url: Option<String>,
51}
52
53fn default_backend() -> String {
54 "memory".into()
55}
56const fn default_true() -> bool {
57 true
58}
59
60#[async_trait]
68pub trait RevocationStore: Send + Sync {
69 async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError>;
71
72 async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError>;
75
76 async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError>;
79}
80
81#[derive(Debug, thiserror::Error)]
83#[non_exhaustive]
84pub enum RevocationError {
85 #[error("revocation store error: {0}")]
87 Backend(String),
88}
89
90pub struct InMemoryRevocationStore {
96 pub(crate) entries: DashMap<String, (String, DateTime<Utc>)>,
98}
99
100impl InMemoryRevocationStore {
101 #[must_use]
103 pub fn new() -> Self {
104 Self {
105 entries: DashMap::new(),
106 }
107 }
108
109 pub fn cleanup_expired(&self) {
111 let now = Utc::now();
112 self.entries.retain(|_, (_, exp)| *exp > now);
113 }
114}
115
116impl Default for InMemoryRevocationStore {
117 fn default() -> Self {
118 Self::new()
119 }
120}
121
122#[async_trait]
126impl RevocationStore for InMemoryRevocationStore {
127 async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError> {
128 if let Some(entry) = self.entries.get(jti) {
129 let (_, expires_at) = entry.value();
130 if *expires_at > Utc::now() {
131 return Ok(true);
132 }
133 drop(entry);
135 self.entries.remove(jti);
136 }
137 Ok(false)
138 }
139
140 async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
141 let expires_at = Utc::now() + chrono::Duration::seconds(ttl_secs.cast_signed());
142 self.entries.insert(jti.to_string(), (String::new(), expires_at));
144 Ok(())
145 }
146
147 async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
148 let keys_to_remove: Vec<String> = self
152 .entries
153 .iter()
154 .filter(|entry| {
155 let (s, _) = entry.value();
156 s == sub
157 })
158 .map(|entry| entry.key().clone())
159 .collect();
160
161 let count = keys_to_remove.len() as u64;
162 for key in &keys_to_remove {
163 self.entries.remove(key);
164 }
165 Ok(count)
166 }
167}
168
169#[cfg(feature = "redis-rate-limiting")]
178pub struct RedisRevocationStore {
179 client: redis::Client,
180 key_prefix: String,
181}
182
183#[cfg(feature = "redis-rate-limiting")]
184impl RedisRevocationStore {
185 pub fn new(redis_url: &str) -> Result<Self, RevocationError> {
191 let client = redis::Client::open(redis_url)
192 .map_err(|e| RevocationError::Backend(format!("Redis connection error: {e}")))?;
193 Ok(Self {
194 client,
195 key_prefix: "fraiseql:revoked:".into(),
196 })
197 }
198}
199
200#[cfg(feature = "redis-rate-limiting")]
201#[async_trait]
205impl RevocationStore for RedisRevocationStore {
206 async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError> {
207 use redis::AsyncCommands;
208 let mut conn = self
209 .client
210 .get_multiplexed_async_connection()
211 .await
212 .map_err(|e| RevocationError::Backend(format!("Redis: {e}")))?;
213 let key = format!("{}{jti}", self.key_prefix);
214 let exists: bool = conn
215 .exists(&key)
216 .await
217 .map_err(|e| RevocationError::Backend(format!("Redis EXISTS: {e}")))?;
218 Ok(exists)
219 }
220
221 async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
222 use redis::AsyncCommands;
223 let mut conn = self
224 .client
225 .get_multiplexed_async_connection()
226 .await
227 .map_err(|e| RevocationError::Backend(format!("Redis: {e}")))?;
228 let key = format!("{}{jti}", self.key_prefix);
229 let _: () = conn
230 .set_ex(&key, "1", ttl_secs)
231 .await
232 .map_err(|e| RevocationError::Backend(format!("Redis SET EX: {e}")))?;
233 Ok(())
234 }
235
236 async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
237 let mut conn = self
238 .client
239 .get_multiplexed_async_connection()
240 .await
241 .map_err(|e| RevocationError::Backend(format!("Redis: {e}")))?;
242 let pattern = format!("{}user:{sub}:*", self.key_prefix);
247 let mut cursor: u64 = 0;
248 let mut all_keys: Vec<String> = Vec::new();
249 loop {
250 let (next_cursor, batch): (u64, Vec<String>) = redis::cmd("SCAN")
251 .arg(cursor)
252 .arg("MATCH")
253 .arg(&pattern)
254 .arg("COUNT")
255 .arg(100u32)
256 .query_async(&mut conn)
257 .await
258 .map_err(|e| RevocationError::Backend(format!("Redis SCAN: {e}")))?;
259 all_keys.extend(batch);
260 cursor = next_cursor;
261 if cursor == 0 {
262 break;
263 }
264 }
265 let count = all_keys.len() as u64;
266 if !all_keys.is_empty() {
267 let _: () = redis::cmd("DEL")
268 .arg(&all_keys)
269 .query_async(&mut conn)
270 .await
271 .map_err(|e| RevocationError::Backend(format!("Redis DEL: {e}")))?;
272 }
273 Ok(count)
274 }
275}
276
277const REVOCATION_POOL_MAX: u32 = 5;
285
286const REVOKED_TOKENS_SCHEMA_SQL: &str = "\
288CREATE TABLE IF NOT EXISTS fraiseql_revoked_tokens (
289 jti TEXT PRIMARY KEY,
290 sub TEXT,
291 expires_at TIMESTAMPTZ NOT NULL
292);
293CREATE INDEX IF NOT EXISTS idx_fraiseql_revoked_tokens_sub
294 ON fraiseql_revoked_tokens (sub);
295CREATE INDEX IF NOT EXISTS idx_fraiseql_revoked_tokens_expires
296 ON fraiseql_revoked_tokens (expires_at);";
297
298pub struct PostgresRevocationStore {
307 pool: sqlx::PgPool,
308}
309
310impl PostgresRevocationStore {
311 pub async fn new(pool: sqlx::PgPool) -> Result<Self, RevocationError> {
318 sqlx::raw_sql(REVOKED_TOKENS_SCHEMA_SQL)
319 .execute(&pool)
320 .await
321 .map_err(|e| RevocationError::Backend(format!("schema creation failed: {e}")))?;
322 Ok(Self { pool })
323 }
324
325 pub async fn cleanup_expired(&self) -> Result<u64, RevocationError> {
332 let result = sqlx::query("DELETE FROM fraiseql_revoked_tokens WHERE expires_at <= NOW()")
333 .execute(&self.pool)
334 .await
335 .map_err(|e| RevocationError::Backend(format!("cleanup failed: {e}")))?;
336 Ok(result.rows_affected())
337 }
338}
339
340#[async_trait]
344impl RevocationStore for PostgresRevocationStore {
345 async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError> {
346 let revoked: bool = sqlx::query_scalar(
347 "SELECT EXISTS (
348 SELECT 1 FROM fraiseql_revoked_tokens WHERE jti = $1 AND expires_at > NOW()
349 )",
350 )
351 .bind(jti)
352 .fetch_one(&self.pool)
353 .await
354 .map_err(|e| RevocationError::Backend(format!("is_revoked query failed: {e}")))?;
355 Ok(revoked)
356 }
357
358 async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
359 let expires_at = Utc::now() + chrono::Duration::seconds(ttl_secs.cast_signed());
360 sqlx::query(
363 "INSERT INTO fraiseql_revoked_tokens (jti, sub, expires_at)
364 VALUES ($1, NULL, $2)
365 ON CONFLICT (jti) DO UPDATE SET expires_at = EXCLUDED.expires_at",
366 )
367 .bind(jti)
368 .bind(expires_at)
369 .execute(&self.pool)
370 .await
371 .map_err(|e| RevocationError::Backend(format!("revoke insert failed: {e}")))?;
372 Ok(())
373 }
374
375 async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
376 let result = sqlx::query("DELETE FROM fraiseql_revoked_tokens WHERE sub = $1")
377 .bind(sub)
378 .execute(&self.pool)
379 .await
380 .map_err(|e| RevocationError::Backend(format!("revoke_all_for_user failed: {e}")))?;
381 Ok(result.rows_affected())
382 }
383}
384
385pub struct TokenRevocationManager {
391 store: Arc<dyn RevocationStore>,
392 require_jti: bool,
393 fail_open: bool,
394}
395
396impl TokenRevocationManager {
397 #[must_use]
399 pub fn new(store: Arc<dyn RevocationStore>, require_jti: bool, fail_open: bool) -> Self {
400 Self {
401 store,
402 require_jti,
403 fail_open,
404 }
405 }
406
407 pub async fn check_token(&self, jti: Option<&str>) -> Result<(), TokenRejection> {
418 let jti = match jti {
419 Some(j) if !j.is_empty() => j,
420 _ => {
421 if self.require_jti {
422 return Err(TokenRejection::MissingJti);
423 }
424 return Ok(());
426 },
427 };
428
429 match self.store.is_revoked(jti).await {
430 Ok(true) => Err(TokenRejection::Revoked),
431 Ok(false) => Ok(()),
432 Err(e) => {
433 warn!(error = %e, jti = %jti, "Revocation store check failed");
434 if self.fail_open {
435 debug!("fail_open=true — allowing request despite store error");
436 Ok(())
437 } else {
438 Err(TokenRejection::StoreUnavailable)
439 }
440 },
441 }
442 }
443
444 pub async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
450 self.store.revoke(jti, ttl_secs).await
451 }
452
453 pub async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
459 self.store.revoke_all_for_user(sub).await
460 }
461
462 #[must_use]
464 pub const fn require_jti(&self) -> bool {
465 self.require_jti
466 }
467}
468
469impl std::fmt::Debug for TokenRevocationManager {
470 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
471 f.debug_struct("TokenRevocationManager")
472 .field("require_jti", &self.require_jti)
473 .field("fail_open", &self.fail_open)
474 .finish_non_exhaustive()
475 }
476}
477
478#[derive(Debug, Clone, PartialEq, Eq)]
480#[non_exhaustive]
481pub enum TokenRejection {
482 Revoked,
484 MissingJti,
486 StoreUnavailable,
488}
489
490pub fn revocation_manager_from_schema(
507 schema: &fraiseql_core::schema::CompiledSchema,
508) -> crate::Result<Option<Arc<TokenRevocationManager>>> {
509 let Some(security) = schema.security.as_ref() else {
510 return Ok(None);
511 };
512 let Some(revocation_val) = security.additional.get("token_revocation") else {
513 return Ok(None);
514 };
515 if revocation_val.is_null() {
519 return Ok(None);
520 }
521 let config: TokenRevocationConfig =
522 serde_json::from_value(revocation_val.clone()).map_err(|e| {
523 crate::ServerError::ConfigError(format!(
524 "invalid security.token_revocation config: {e}"
525 ))
526 })?;
527
528 if !config.enabled {
529 return Ok(None);
530 }
531
532 let store: Arc<dyn RevocationStore> = match config.backend.as_str() {
533 #[cfg(feature = "redis-rate-limiting")]
534 "redis" => {
535 let url = config.redis_url.as_deref().unwrap_or("redis://localhost:6379");
536 match RedisRevocationStore::new(url) {
537 Ok(s) => {
538 info!(backend = "redis", "Token revocation store initialized");
539 Arc::new(s)
540 },
541 Err(e) => {
542 warn!(error = %e, "Failed to init Redis revocation store — falling back to in-memory");
543 Arc::new(InMemoryRevocationStore::new())
544 },
545 }
546 },
547 #[cfg(not(feature = "redis-rate-limiting"))]
548 "redis" => {
549 warn!(
550 "token_revocation.backend = \"redis\" but the `redis-rate-limiting` feature is \
551 not compiled in. Falling back to in-memory."
552 );
553 Arc::new(InMemoryRevocationStore::new())
554 },
555 "memory" | "env" => {
556 info!(backend = "memory", "Token revocation store initialized (in-memory)");
557 Arc::new(InMemoryRevocationStore::new())
558 },
559 "postgres" => {
560 info!(
563 backend = "postgres",
564 "Token revocation backend = postgres; provisioned by the PostgreSQL runtime"
565 );
566 return Ok(None);
567 },
568 other => {
569 return Err(crate::ServerError::ConfigError(format!(
570 "unknown token_revocation backend {other:?}; \
571 expected \"memory\", \"redis\", or \"postgres\""
572 )));
573 },
574 };
575
576 Ok(Some(Arc::new(TokenRevocationManager::new(
577 store,
578 config.require_jti,
579 config.fail_open,
580 ))))
581}
582
583pub async fn build_postgres_revocation_manager(
597 database_url: &str,
598 schema: &fraiseql_core::schema::CompiledSchema,
599) -> std::result::Result<Option<Arc<TokenRevocationManager>>, String> {
600 let Some(security) = schema.security.as_ref() else {
601 return Ok(None);
602 };
603 let Some(revocation_val) = security.additional.get("token_revocation") else {
604 return Ok(None);
605 };
606 if revocation_val.is_null() {
608 return Ok(None);
609 }
610 let config: TokenRevocationConfig = serde_json::from_value(revocation_val.clone())
611 .map_err(|e| format!("invalid security.token_revocation config: {e}"))?;
612
613 if !config.enabled || config.backend != "postgres" {
614 return Ok(None);
615 }
616
617 let pool = sqlx::postgres::PgPoolOptions::new()
618 .max_connections(REVOCATION_POOL_MAX)
619 .connect(database_url)
620 .await
621 .map_err(|e| format!("token revocation: failed to connect to PostgreSQL: {e}"))?;
622
623 let store = PostgresRevocationStore::new(pool)
624 .await
625 .map_err(|e| format!("token revocation: {e}"))?;
626
627 info!(backend = "postgres", "Token revocation store initialized (PostgreSQL)");
628 Ok(Some(Arc::new(TokenRevocationManager::new(
629 Arc::new(store),
630 config.require_jti,
631 config.fail_open,
632 ))))
633}
634
635#[must_use]
641pub fn revocation_backend_is_postgres(schema: &fraiseql_core::schema::CompiledSchema) -> bool {
642 schema
643 .security
644 .as_ref()
645 .and_then(|s| s.additional.get("token_revocation"))
646 .and_then(|v| serde_json::from_value::<TokenRevocationConfig>(v.clone()).ok())
647 .is_some_and(|c| c.enabled && c.backend == "postgres")
648}