use async_trait::async_trait;
use chrono::{DateTime, Utc};
use sqlx::{Row, postgres::PgPool};
use uuid::Uuid;
use super::SamlError;
pub const PG_SAML_IDP_SCHEMA_SQL: &str = r"
CREATE SCHEMA IF NOT EXISTS core;
CREATE TABLE IF NOT EXISTS core.tb_saml_idp (
pk_saml_idp BIGINT GENERATED ALWAYS AS IDENTITY PRIMARY KEY,
id UUID NOT NULL DEFAULT gen_random_uuid(),
idp_name TEXT NOT NULL,
tenant_id UUID,
sp_entity_id TEXT NOT NULL,
acs_url TEXT NOT NULL,
metadata_xml TEXT NOT NULL,
idp_entity_id TEXT NOT NULL,
trust_asserted_email BOOLEAN NOT NULL DEFAULT FALSE,
certificate_expires_at TIMESTAMPTZ,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT now(),
deleted_at TIMESTAMPTZ
);
-- Spans deleted rows on purpose: a retired name must never be reissued to another
-- tenant, or the new IdP inherits the old one's `saml:<name>` identity rows.
CREATE UNIQUE INDEX IF NOT EXISTS uq_saml_idp_name ON core.tb_saml_idp (idp_name);
CREATE INDEX IF NOT EXISTS idx_saml_idp_tenant ON core.tb_saml_idp (tenant_id)
WHERE deleted_at IS NULL;
-- RLS deny-by-default, mirroring core.tb_user / core.tb_auth_identity. ENABLE not FORCE:
-- the owner (this store, running the trusted admin path) operates freely, while any other
-- role reads a row only once it has set fraiseql.tenant_id to that row's tenant.
ALTER TABLE core.tb_saml_idp ENABLE ROW LEVEL SECURITY;
DROP POLICY IF EXISTS p_saml_idp_tenant_read ON core.tb_saml_idp;
CREATE POLICY p_saml_idp_tenant_read ON core.tb_saml_idp
FOR SELECT USING (tenant_id = NULLIF(current_setting('fraiseql.tenant_id', true), '')::uuid);
DROP POLICY IF EXISTS p_saml_idp_insert ON core.tb_saml_idp;
CREATE POLICY p_saml_idp_insert ON core.tb_saml_idp FOR INSERT WITH CHECK (true);
REVOKE ALL ON core.tb_saml_idp FROM PUBLIC;
";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SamlIdpRecord {
pub id: Uuid,
pub idp_name: String,
pub tenant_id: Option<Uuid>,
pub sp_entity_id: String,
pub acs_url: String,
pub metadata_xml: String,
pub idp_entity_id: String,
pub trust_asserted_email: bool,
pub certificate_expires_at: Option<DateTime<Utc>>,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct SamlIdpSpec {
pub idp_name: String,
pub tenant_id: Option<Uuid>,
pub sp_entity_id: String,
pub acs_url: String,
pub metadata_xml: String,
pub trust_asserted_email: bool,
}
#[async_trait]
pub trait SamlIdpStore: Send + Sync {
async fn list(&self) -> Result<Vec<SamlIdpRecord>, SamlError>;
async fn get(&self, idp_name: &str) -> Result<Option<SamlIdpRecord>, SamlError>;
async fn create(&self, spec: &SamlIdpSpec) -> Result<SamlIdpRecord, SamlError>;
async fn update(&self, spec: &SamlIdpSpec) -> Result<SamlIdpRecord, SamlError>;
async fn delete(&self, idp_name: &str) -> Result<(), SamlError>;
}
#[derive(Debug, Clone)]
pub struct PgSamlIdpStore {
db: PgPool,
}
const COLUMNS: &str = "id, idp_name, tenant_id, sp_entity_id, acs_url, metadata_xml, \
idp_entity_id, trust_asserted_email, certificate_expires_at, \
created_at, updated_at";
impl PgSamlIdpStore {
#[must_use]
pub const fn new(db: PgPool) -> Self {
Self { db }
}
pub async fn init(&self) -> Result<(), SamlError> {
sqlx::raw_sql(PG_SAML_IDP_SCHEMA_SQL)
.execute(&self.db)
.await
.map_err(|e| SamlError::Store(format!("initialize SAML IdP store: {e}")))?;
Ok(())
}
fn decode(row: &sqlx::postgres::PgRow) -> SamlIdpRecord {
SamlIdpRecord {
id: row.get("id"),
idp_name: row.get("idp_name"),
tenant_id: row.get("tenant_id"),
sp_entity_id: row.get("sp_entity_id"),
acs_url: row.get("acs_url"),
metadata_xml: row.get("metadata_xml"),
idp_entity_id: row.get("idp_entity_id"),
trust_asserted_email: row.get("trust_asserted_email"),
certificate_expires_at: row.get("certificate_expires_at"),
created_at: row.get("created_at"),
updated_at: row.get("updated_at"),
}
}
}
fn derive(spec: &SamlIdpSpec) -> Result<(String, Option<DateTime<Utc>>), SamlError> {
let config = super::SamlIdpConfig::builder(
spec.idp_name.clone(),
spec.sp_entity_id.clone(),
spec.acs_url.clone(),
)
.idp_metadata_xml(&spec.metadata_xml)?
.tenant_id(spec.tenant_id)
.trust_asserted_email(spec.trust_asserted_email)
.build()?;
Ok((config.idp_entity_id().to_string(), config.signing_certificate_expiry()))
}
#[async_trait]
impl SamlIdpStore for PgSamlIdpStore {
async fn list(&self) -> Result<Vec<SamlIdpRecord>, SamlError> {
let rows = sqlx::query(&format!(
"SELECT {COLUMNS} FROM core.tb_saml_idp WHERE deleted_at IS NULL ORDER BY idp_name"
))
.fetch_all(&self.db)
.await
.map_err(|e| SamlError::Store(format!("list SAML IdPs: {e}")))?;
Ok(rows.iter().map(Self::decode).collect())
}
async fn get(&self, idp_name: &str) -> Result<Option<SamlIdpRecord>, SamlError> {
let row = sqlx::query(&format!(
"SELECT {COLUMNS} FROM core.tb_saml_idp WHERE idp_name = $1 AND deleted_at IS NULL"
))
.bind(idp_name)
.fetch_optional(&self.db)
.await
.map_err(|e| SamlError::Store(format!("get SAML IdP: {e}")))?;
Ok(row.as_ref().map(Self::decode))
}
async fn create(&self, spec: &SamlIdpSpec) -> Result<SamlIdpRecord, SamlError> {
let (idp_entity_id, expires_at) = derive(spec)?;
let row = sqlx::query(&format!(
"INSERT INTO core.tb_saml_idp \
(idp_name, tenant_id, sp_entity_id, acs_url, metadata_xml, idp_entity_id, \
trust_asserted_email, certificate_expires_at) \
VALUES ($1, $2, $3, $4, $5, $6, $7, $8) RETURNING {COLUMNS}"
))
.bind(&spec.idp_name)
.bind(spec.tenant_id)
.bind(&spec.sp_entity_id)
.bind(&spec.acs_url)
.bind(&spec.metadata_xml)
.bind(&idp_entity_id)
.bind(spec.trust_asserted_email)
.bind(expires_at)
.fetch_one(&self.db)
.await
.map_err(|e| match &e {
sqlx::Error::Database(db) if db.code().as_deref() == Some("23505") => {
SamlError::NameTaken(spec.idp_name.clone())
},
_ => SamlError::Store(format!("create SAML IdP: {e}")),
})?;
Ok(Self::decode(&row))
}
async fn update(&self, spec: &SamlIdpSpec) -> Result<SamlIdpRecord, SamlError> {
let (idp_entity_id, expires_at) = derive(spec)?;
let row = sqlx::query(&format!(
"UPDATE core.tb_saml_idp SET sp_entity_id = $2, acs_url = $3, metadata_xml = $4, \
idp_entity_id = $5, trust_asserted_email = $6, certificate_expires_at = $7, \
updated_at = now() \
WHERE idp_name = $1 AND deleted_at IS NULL RETURNING {COLUMNS}"
))
.bind(&spec.idp_name)
.bind(&spec.sp_entity_id)
.bind(&spec.acs_url)
.bind(&spec.metadata_xml)
.bind(&idp_entity_id)
.bind(spec.trust_asserted_email)
.bind(expires_at)
.fetch_optional(&self.db)
.await
.map_err(|e| SamlError::Store(format!("update SAML IdP: {e}")))?
.ok_or_else(|| SamlError::NotFound(spec.idp_name.clone()))?;
Ok(Self::decode(&row))
}
async fn delete(&self, idp_name: &str) -> Result<(), SamlError> {
let affected = sqlx::query(
"UPDATE core.tb_saml_idp SET deleted_at = now(), updated_at = now() \
WHERE idp_name = $1 AND deleted_at IS NULL",
)
.bind(idp_name)
.execute(&self.db)
.await
.map_err(|e| SamlError::Store(format!("delete SAML IdP: {e}")))?
.rows_affected();
if affected == 0 {
return Err(SamlError::NotFound(idp_name.to_string()));
}
Ok(())
}
}