#[cfg(test)]
mod tests;
use std::sync::Arc;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use dashmap::DashMap;
use serde::Deserialize;
use tracing::{debug, info, warn};
#[derive(Debug, Clone, Deserialize)]
pub struct TokenRevocationConfig {
#[serde(default)]
pub enabled: bool,
#[serde(default = "default_backend")]
pub backend: String,
#[serde(default = "default_true")]
pub require_jti: bool,
#[serde(default)]
pub fail_open: bool,
pub redis_url: Option<String>,
}
fn default_backend() -> String {
"memory".into()
}
const fn default_true() -> bool {
true
}
#[async_trait]
pub trait RevocationStore: Send + Sync {
async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError>;
async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError>;
async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError>;
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum RevocationError {
#[error("revocation store error: {0}")]
Backend(String),
}
pub struct InMemoryRevocationStore {
pub(crate) entries: DashMap<String, (String, DateTime<Utc>)>,
}
impl InMemoryRevocationStore {
#[must_use]
pub fn new() -> Self {
Self {
entries: DashMap::new(),
}
}
pub fn cleanup_expired(&self) {
let now = Utc::now();
self.entries.retain(|_, (_, exp)| *exp > now);
}
}
impl Default for InMemoryRevocationStore {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl RevocationStore for InMemoryRevocationStore {
async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError> {
if let Some(entry) = self.entries.get(jti) {
let (_, expires_at) = entry.value();
if *expires_at > Utc::now() {
return Ok(true);
}
drop(entry);
self.entries.remove(jti);
}
Ok(false)
}
async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
let expires_at = Utc::now() + chrono::Duration::seconds(ttl_secs.cast_signed());
self.entries.insert(jti.to_string(), (String::new(), expires_at));
Ok(())
}
async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
let keys_to_remove: Vec<String> = self
.entries
.iter()
.filter(|entry| {
let (s, _) = entry.value();
s == sub
})
.map(|entry| entry.key().clone())
.collect();
let count = keys_to_remove.len() as u64;
for key in &keys_to_remove {
self.entries.remove(key);
}
Ok(count)
}
}
#[cfg(feature = "redis-rate-limiting")]
pub struct RedisRevocationStore {
client: redis::Client,
key_prefix: String,
}
#[cfg(feature = "redis-rate-limiting")]
impl RedisRevocationStore {
pub fn new(redis_url: &str) -> Result<Self, RevocationError> {
let client = redis::Client::open(redis_url)
.map_err(|e| RevocationError::Backend(format!("Redis connection error: {e}")))?;
Ok(Self {
client,
key_prefix: "fraiseql:revoked:".into(),
})
}
}
#[cfg(feature = "redis-rate-limiting")]
#[async_trait]
impl RevocationStore for RedisRevocationStore {
async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError> {
use redis::AsyncCommands;
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| RevocationError::Backend(format!("Redis: {e}")))?;
let key = format!("{}{jti}", self.key_prefix);
let exists: bool = conn
.exists(&key)
.await
.map_err(|e| RevocationError::Backend(format!("Redis EXISTS: {e}")))?;
Ok(exists)
}
async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
use redis::AsyncCommands;
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| RevocationError::Backend(format!("Redis: {e}")))?;
let key = format!("{}{jti}", self.key_prefix);
let _: () = conn
.set_ex(&key, "1", ttl_secs)
.await
.map_err(|e| RevocationError::Backend(format!("Redis SET EX: {e}")))?;
Ok(())
}
async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
let mut conn = self
.client
.get_multiplexed_async_connection()
.await
.map_err(|e| RevocationError::Backend(format!("Redis: {e}")))?;
let pattern = format!("{}user:{sub}:*", self.key_prefix);
let mut cursor: u64 = 0;
let mut all_keys: Vec<String> = Vec::new();
loop {
let (next_cursor, batch): (u64, Vec<String>) = redis::cmd("SCAN")
.arg(cursor)
.arg("MATCH")
.arg(&pattern)
.arg("COUNT")
.arg(100u32)
.query_async(&mut conn)
.await
.map_err(|e| RevocationError::Backend(format!("Redis SCAN: {e}")))?;
all_keys.extend(batch);
cursor = next_cursor;
if cursor == 0 {
break;
}
}
let count = all_keys.len() as u64;
if !all_keys.is_empty() {
let _: () = redis::cmd("DEL")
.arg(&all_keys)
.query_async(&mut conn)
.await
.map_err(|e| RevocationError::Backend(format!("Redis DEL: {e}")))?;
}
Ok(count)
}
}
const REVOCATION_POOL_MAX: u32 = 5;
const REVOKED_TOKENS_SCHEMA_SQL: &str = "\
CREATE TABLE IF NOT EXISTS fraiseql_revoked_tokens (
jti TEXT PRIMARY KEY,
sub TEXT,
expires_at TIMESTAMPTZ NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_fraiseql_revoked_tokens_sub
ON fraiseql_revoked_tokens (sub);
CREATE INDEX IF NOT EXISTS idx_fraiseql_revoked_tokens_expires
ON fraiseql_revoked_tokens (expires_at);";
pub struct PostgresRevocationStore {
pool: sqlx::PgPool,
}
impl PostgresRevocationStore {
pub async fn new(pool: sqlx::PgPool) -> Result<Self, RevocationError> {
sqlx::raw_sql(REVOKED_TOKENS_SCHEMA_SQL)
.execute(&pool)
.await
.map_err(|e| RevocationError::Backend(format!("schema creation failed: {e}")))?;
Ok(Self { pool })
}
pub async fn cleanup_expired(&self) -> Result<u64, RevocationError> {
let result = sqlx::query("DELETE FROM fraiseql_revoked_tokens WHERE expires_at <= NOW()")
.execute(&self.pool)
.await
.map_err(|e| RevocationError::Backend(format!("cleanup failed: {e}")))?;
Ok(result.rows_affected())
}
}
#[async_trait]
impl RevocationStore for PostgresRevocationStore {
async fn is_revoked(&self, jti: &str) -> Result<bool, RevocationError> {
let revoked: bool = sqlx::query_scalar(
"SELECT EXISTS (
SELECT 1 FROM fraiseql_revoked_tokens WHERE jti = $1 AND expires_at > NOW()
)",
)
.bind(jti)
.fetch_one(&self.pool)
.await
.map_err(|e| RevocationError::Backend(format!("is_revoked query failed: {e}")))?;
Ok(revoked)
}
async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
let expires_at = Utc::now() + chrono::Duration::seconds(ttl_secs.cast_signed());
sqlx::query(
"INSERT INTO fraiseql_revoked_tokens (jti, sub, expires_at)
VALUES ($1, NULL, $2)
ON CONFLICT (jti) DO UPDATE SET expires_at = EXCLUDED.expires_at",
)
.bind(jti)
.bind(expires_at)
.execute(&self.pool)
.await
.map_err(|e| RevocationError::Backend(format!("revoke insert failed: {e}")))?;
Ok(())
}
async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
let result = sqlx::query("DELETE FROM fraiseql_revoked_tokens WHERE sub = $1")
.bind(sub)
.execute(&self.pool)
.await
.map_err(|e| RevocationError::Backend(format!("revoke_all_for_user failed: {e}")))?;
Ok(result.rows_affected())
}
}
pub struct TokenRevocationManager {
store: Arc<dyn RevocationStore>,
require_jti: bool,
fail_open: bool,
}
impl TokenRevocationManager {
#[must_use]
pub fn new(store: Arc<dyn RevocationStore>, require_jti: bool, fail_open: bool) -> Self {
Self {
store,
require_jti,
fail_open,
}
}
pub async fn check_token(&self, jti: Option<&str>) -> Result<(), TokenRejection> {
let jti = match jti {
Some(j) if !j.is_empty() => j,
_ => {
if self.require_jti {
return Err(TokenRejection::MissingJti);
}
return Ok(());
},
};
match self.store.is_revoked(jti).await {
Ok(true) => Err(TokenRejection::Revoked),
Ok(false) => Ok(()),
Err(e) => {
warn!(error = %e, jti = %jti, "Revocation store check failed");
if self.fail_open {
debug!("fail_open=true — allowing request despite store error");
Ok(())
} else {
Err(TokenRejection::StoreUnavailable)
}
},
}
}
pub async fn revoke(&self, jti: &str, ttl_secs: u64) -> Result<(), RevocationError> {
self.store.revoke(jti, ttl_secs).await
}
pub async fn revoke_all_for_user(&self, sub: &str) -> Result<u64, RevocationError> {
self.store.revoke_all_for_user(sub).await
}
#[must_use]
pub const fn require_jti(&self) -> bool {
self.require_jti
}
}
impl std::fmt::Debug for TokenRevocationManager {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TokenRevocationManager")
.field("require_jti", &self.require_jti)
.field("fail_open", &self.fail_open)
.finish_non_exhaustive()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum TokenRejection {
Revoked,
MissingJti,
StoreUnavailable,
}
pub fn revocation_manager_from_schema(
schema: &fraiseql_core::schema::CompiledSchema,
) -> crate::Result<Option<Arc<TokenRevocationManager>>> {
let Some(security) = schema.security.as_ref() else {
return Ok(None);
};
let Some(revocation_val) = security.additional.get("token_revocation") else {
return Ok(None);
};
if revocation_val.is_null() {
return Ok(None);
}
let config: TokenRevocationConfig =
serde_json::from_value(revocation_val.clone()).map_err(|e| {
crate::ServerError::ConfigError(format!(
"invalid security.token_revocation config: {e}"
))
})?;
if !config.enabled {
return Ok(None);
}
let store: Arc<dyn RevocationStore> = match config.backend.as_str() {
#[cfg(feature = "redis-rate-limiting")]
"redis" => {
let url = config.redis_url.as_deref().unwrap_or("redis://localhost:6379");
match RedisRevocationStore::new(url) {
Ok(s) => {
info!(backend = "redis", "Token revocation store initialized");
Arc::new(s)
},
Err(e) => {
warn!(error = %e, "Failed to init Redis revocation store — falling back to in-memory");
Arc::new(InMemoryRevocationStore::new())
},
}
},
#[cfg(not(feature = "redis-rate-limiting"))]
"redis" => {
warn!(
"token_revocation.backend = \"redis\" but the `redis-rate-limiting` feature is \
not compiled in. Falling back to in-memory."
);
Arc::new(InMemoryRevocationStore::new())
},
"memory" | "env" => {
info!(backend = "memory", "Token revocation store initialized (in-memory)");
Arc::new(InMemoryRevocationStore::new())
},
"postgres" => {
info!(
backend = "postgres",
"Token revocation backend = postgres; provisioned by the PostgreSQL runtime"
);
return Ok(None);
},
other => {
return Err(crate::ServerError::ConfigError(format!(
"unknown token_revocation backend {other:?}; \
expected \"memory\", \"redis\", or \"postgres\""
)));
},
};
Ok(Some(Arc::new(TokenRevocationManager::new(
store,
config.require_jti,
config.fail_open,
))))
}
pub async fn build_postgres_revocation_manager(
database_url: &str,
schema: &fraiseql_core::schema::CompiledSchema,
) -> std::result::Result<Option<Arc<TokenRevocationManager>>, String> {
let Some(security) = schema.security.as_ref() else {
return Ok(None);
};
let Some(revocation_val) = security.additional.get("token_revocation") else {
return Ok(None);
};
if revocation_val.is_null() {
return Ok(None);
}
let config: TokenRevocationConfig = serde_json::from_value(revocation_val.clone())
.map_err(|e| format!("invalid security.token_revocation config: {e}"))?;
if !config.enabled || config.backend != "postgres" {
return Ok(None);
}
let pool = sqlx::postgres::PgPoolOptions::new()
.max_connections(REVOCATION_POOL_MAX)
.connect(database_url)
.await
.map_err(|e| format!("token revocation: failed to connect to PostgreSQL: {e}"))?;
let store = PostgresRevocationStore::new(pool)
.await
.map_err(|e| format!("token revocation: {e}"))?;
info!(backend = "postgres", "Token revocation store initialized (PostgreSQL)");
Ok(Some(Arc::new(TokenRevocationManager::new(
Arc::new(store),
config.require_jti,
config.fail_open,
))))
}
#[must_use]
pub fn revocation_backend_is_postgres(schema: &fraiseql_core::schema::CompiledSchema) -> bool {
schema
.security
.as_ref()
.and_then(|s| s.additional.get("token_revocation"))
.and_then(|v| serde_json::from_value::<TokenRevocationConfig>(v.clone()).ok())
.is_some_and(|c| c.enabled && c.backend == "postgres")
}