use async_trait::async_trait;
use chrono::{DateTime, Utc};
use dashmap::{DashMap, mapref::entry::Entry};
use sqlx::PgPool;
use super::SamlError;
#[async_trait]
pub trait SamlReplayStore: Send + Sync + std::fmt::Debug {
async fn check_and_record(
&self,
assertion_id: &str,
expires_at: DateTime<Utc>,
now: DateTime<Utc>,
) -> Result<bool, SamlError>;
fn is_distributed(&self) -> bool;
}
#[derive(Debug, Default)]
pub struct SamlReplayCache {
seen: DashMap<String, DateTime<Utc>>,
}
impl SamlReplayCache {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn len(&self) -> usize {
self.seen.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.seen.is_empty()
}
}
#[async_trait]
impl SamlReplayStore for SamlReplayCache {
async fn check_and_record(
&self,
assertion_id: &str,
expires_at: DateTime<Utc>,
now: DateTime<Utc>,
) -> Result<bool, SamlError> {
self.seen.retain(|_, expiry| *expiry > now);
Ok(match self.seen.entry(assertion_id.to_owned()) {
Entry::Occupied(_) => false,
Entry::Vacant(slot) => {
slot.insert(expires_at);
true
},
})
}
fn is_distributed(&self) -> bool {
false
}
}
pub const PG_SAML_REPLAY_SCHEMA_SQL: &str = r"
CREATE SCHEMA IF NOT EXISTS core;
CREATE TABLE IF NOT EXISTS core.tb_saml_replay (
assertion_id TEXT PRIMARY KEY,
expires_at TIMESTAMPTZ NOT NULL
);
CREATE INDEX IF NOT EXISTS idx_saml_replay_expires ON core.tb_saml_replay(expires_at);
REVOKE ALL ON core.tb_saml_replay FROM PUBLIC;
";
#[derive(Debug, Clone)]
pub struct PgSamlReplayStore {
db: PgPool,
}
impl PgSamlReplayStore {
#[must_use]
pub const fn new(db: PgPool) -> Self {
Self { db }
}
pub async fn init(&self) -> Result<(), SamlError> {
for stmt in PG_SAML_REPLAY_SCHEMA_SQL.split(';').map(str::trim).filter(|s| !s.is_empty()) {
sqlx::query(stmt)
.execute(&self.db)
.await
.map_err(|e| SamlError::Config(format!("saml replay table: {e}")))?;
}
Ok(())
}
pub async fn sweep_expired(&self, now: DateTime<Utc>) -> Result<u64, SamlError> {
sqlx::query("DELETE FROM core.tb_saml_replay WHERE expires_at <= $1")
.bind(now)
.execute(&self.db)
.await
.map(|r| r.rows_affected())
.map_err(|e| SamlError::Verification(format!("saml replay sweep: {e}")))
}
}
#[async_trait]
impl SamlReplayStore for PgSamlReplayStore {
async fn check_and_record(
&self,
assertion_id: &str,
expires_at: DateTime<Utc>,
now: DateTime<Utc>,
) -> Result<bool, SamlError> {
let inserted = sqlx::query(
"INSERT INTO core.tb_saml_replay (assertion_id, expires_at)
VALUES ($1, $2)
ON CONFLICT (assertion_id) DO UPDATE
SET expires_at = EXCLUDED.expires_at
WHERE core.tb_saml_replay.expires_at <= $3",
)
.bind(assertion_id)
.bind(expires_at)
.bind(now)
.execute(&self.db)
.await
.map_err(|e| SamlError::Verification(format!("saml replay store: {e}")))?
.rows_affected();
Ok(inserted > 0)
}
fn is_distributed(&self) -> bool {
true
}
}