#[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>,
#[serde(default = "default_revoke_all_ttl")]
pub revoke_all_ttl_secs: u64,
}
fn default_backend() -> String {
"memory".into()
}
const fn default_true() -> bool {
true
}
const fn default_revoke_all_ttl() -> u64 {
86_400
}
#[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, ttl_secs: u64) -> Result<(), RevocationError>;
async fn user_revoked_after(&self, sub: &str) -> Result<Option<i64>, 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>)>,
pub(crate) user_epochs: DashMap<String, (i64, DateTime<Utc>)>,
}
impl InMemoryRevocationStore {
#[must_use]
pub fn new() -> Self {
Self {
entries: DashMap::new(),
user_epochs: DashMap::new(),
}
}
pub fn cleanup_expired(&self) {
let now = Utc::now();
self.entries.retain(|_, (_, exp)| *exp > now);
self.user_epochs.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, ttl_secs: u64) -> Result<(), RevocationError> {
let now = Utc::now();
let expires_at = now + chrono::Duration::seconds(ttl_secs.cast_signed());
self.user_epochs.insert(sub.to_string(), (now.timestamp(), expires_at));
Ok(())
}
async fn user_revoked_after(&self, sub: &str) -> Result<Option<i64>, RevocationError> {
if let Some(entry) = self.user_epochs.get(sub) {
let (revoked_after, expires_at) = *entry.value();
if expires_at > Utc::now() {
return Ok(Some(revoked_after));
}
drop(entry);
self.user_epochs.remove(sub);
}
Ok(None)
}
}
#[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, 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!("{}user:{sub}", self.key_prefix);
let now = Utc::now().timestamp();
let _: () = conn
.set_ex(&key, now, ttl_secs)
.await
.map_err(|e| RevocationError::Backend(format!("Redis SET EX: {e}")))?;
Ok(())
}
async fn user_revoked_after(&self, sub: &str) -> Result<Option<i64>, 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!("{}user:{sub}", self.key_prefix);
let epoch: Option<i64> = conn
.get(&key)
.await
.map_err(|e| RevocationError::Backend(format!("Redis GET: {e}")))?;
Ok(epoch)
}
}
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);
CREATE TABLE IF NOT EXISTS fraiseql_revoked_users (
sub TEXT PRIMARY KEY,
revoked_after BIGINT NOT NULL,
expires_at TIMESTAMPTZ NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_fraiseql_revoked_users_expires
ON fraiseql_revoked_users (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 tokens = sqlx::query("DELETE FROM fraiseql_revoked_tokens WHERE expires_at <= NOW()")
.execute(&self.pool)
.await
.map_err(|e| RevocationError::Backend(format!("cleanup failed: {e}")))?;
let users = sqlx::query("DELETE FROM fraiseql_revoked_users WHERE expires_at <= NOW()")
.execute(&self.pool)
.await
.map_err(|e| RevocationError::Backend(format!("cleanup failed: {e}")))?;
Ok(tokens.rows_affected() + users.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, ttl_secs: u64) -> Result<(), RevocationError> {
let expires_at = Utc::now() + chrono::Duration::seconds(ttl_secs.cast_signed());
sqlx::query(
"INSERT INTO fraiseql_revoked_users (sub, revoked_after, expires_at)
VALUES ($1, EXTRACT(EPOCH FROM NOW())::BIGINT, $2)
ON CONFLICT (sub) DO UPDATE
SET revoked_after = EXTRACT(EPOCH FROM NOW())::BIGINT,
expires_at = EXCLUDED.expires_at",
)
.bind(sub)
.bind(expires_at)
.execute(&self.pool)
.await
.map_err(|e| RevocationError::Backend(format!("revoke_all_for_user failed: {e}")))?;
Ok(())
}
async fn user_revoked_after(&self, sub: &str) -> Result<Option<i64>, RevocationError> {
let epoch: Option<i64> = sqlx::query_scalar(
"SELECT revoked_after FROM fraiseql_revoked_users
WHERE sub = $1 AND expires_at > NOW()",
)
.bind(sub)
.fetch_optional(&self.pool)
.await
.map_err(|e| RevocationError::Backend(format!("user_revoked_after query failed: {e}")))?;
Ok(epoch)
}
}
pub struct TokenRevocationManager {
store: Arc<dyn RevocationStore>,
require_jti: bool,
fail_open: bool,
revoke_all_ttl_secs: u64,
}
impl TokenRevocationManager {
#[must_use]
pub fn new(
store: Arc<dyn RevocationStore>,
require_jti: bool,
fail_open: bool,
revoke_all_ttl_secs: u64,
) -> Self {
Self {
store,
require_jti,
fail_open,
revoke_all_ttl_secs,
}
}
pub async fn check_token(
&self,
jti: Option<&str>,
sub: &str,
iat: Option<i64>,
) -> Result<(), TokenRejection> {
match jti {
Some(j) if !j.is_empty() => match self.store.is_revoked(j).await {
Ok(true) => return Err(TokenRejection::Revoked),
Ok(false) => {},
Err(e) => {
warn!(error = %e, jti = %j, "Revocation store check failed");
if self.fail_open {
debug!("fail_open=true — allowing request despite store error");
return Ok(());
}
return Err(TokenRejection::StoreUnavailable);
},
},
_ => {
if self.require_jti {
return Err(TokenRejection::MissingJti);
}
},
}
match self.store.user_revoked_after(sub).await {
Ok(Some(epoch)) => {
if iat.is_some_and(|issued| issued <= epoch) {
return Err(TokenRejection::Revoked);
}
Ok(())
},
Ok(None) => Ok(()),
Err(e) => {
warn!(error = %e, sub = %sub, "Revoke-all epoch 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<(), RevocationError> {
self.store.revoke_all_for_user(sub, self.revoke_all_ttl_secs).await
}
pub async fn user_revoked_after(&self, sub: &str) -> Result<Option<i64>, RevocationError> {
self.store.user_revoked_after(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,
config.revoke_all_ttl_secs,
))))
}
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,
config.revoke_all_ttl_secs,
))))
}
#[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")
}