use async_trait::async_trait;
use chrono::{DateTime, Utc};
use sqlx::{Row, postgres::PgPool};
use uuid::Uuid;
use crate::error::{AuthError, Result};
pub const PG_SCIM_SCHEMA_SQL: &str = r"
CREATE SCHEMA IF NOT EXISTS core;
CREATE TABLE IF NOT EXISTS core.tb_scim_group (
pk_scim_group BIGINT GENERATED ALWAYS AS IDENTITY PRIMARY KEY,
id UUID NOT NULL DEFAULT gen_random_uuid(),
display_name TEXT NOT NULL,
external_id TEXT,
tenant_id UUID,
version BIGINT NOT NULL DEFAULT 1,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
updated_at TIMESTAMPTZ NOT NULL DEFAULT now()
);
CREATE UNIQUE INDEX IF NOT EXISTS uq_scim_group_id ON core.tb_scim_group (id);
CREATE UNIQUE INDEX IF NOT EXISTS uq_scim_group_name
ON core.tb_scim_group (tenant_id, display_name) NULLS NOT DISTINCT;
CREATE TABLE IF NOT EXISTS core.tb_scim_group_member (
pk_scim_group_member BIGINT GENERATED ALWAYS AS IDENTITY PRIMARY KEY,
fk_scim_group BIGINT NOT NULL REFERENCES core.tb_scim_group (pk_scim_group)
ON DELETE CASCADE,
user_id TEXT NOT NULL,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
UNIQUE (fk_scim_group, user_id)
);
CREATE INDEX IF NOT EXISTS idx_scim_group_member_user
ON core.tb_scim_group_member (user_id);
-- Provisioning credentials. Distinct from the admin token by construction: nothing reads
-- this table except the SCIM router, so holding one grants provisioning and nothing else.
CREATE TABLE IF NOT EXISTS core.tb_scim_token (
pk_scim_token BIGINT GENERATED ALWAYS AS IDENTITY PRIMARY KEY,
id UUID NOT NULL DEFAULT gen_random_uuid(),
token_hash TEXT NOT NULL UNIQUE,
idp_name TEXT NOT NULL,
tenant_id UUID,
description TEXT,
created_at TIMESTAMPTZ NOT NULL DEFAULT now(),
last_used_at TIMESTAMPTZ,
revoked_at TIMESTAMPTZ
);
CREATE INDEX IF NOT EXISTS idx_scim_token_live
ON core.tb_scim_token (token_hash) WHERE revoked_at IS NULL;
ALTER TABLE core.tb_scim_group ENABLE ROW LEVEL SECURITY;
ALTER TABLE core.tb_scim_group_member ENABLE ROW LEVEL SECURITY;
ALTER TABLE core.tb_scim_token ENABLE ROW LEVEL SECURITY;
DROP POLICY IF EXISTS p_scim_group_tenant_read ON core.tb_scim_group;
CREATE POLICY p_scim_group_tenant_read ON core.tb_scim_group
FOR SELECT USING (tenant_id = NULLIF(current_setting('fraiseql.tenant_id', true), '')::uuid);
DROP POLICY IF EXISTS p_scim_group_insert ON core.tb_scim_group;
CREATE POLICY p_scim_group_insert ON core.tb_scim_group FOR INSERT WITH CHECK (true);
REVOKE ALL ON core.tb_scim_group FROM PUBLIC;
REVOKE ALL ON core.tb_scim_group_member FROM PUBLIC;
REVOKE ALL ON core.tb_scim_token FROM PUBLIC;
";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ScimUser {
pub id: String,
pub user_name: String,
pub external_id: Option<String>,
pub email: Option<String>,
pub given_name: Option<String>,
pub family_name: Option<String>,
pub display_name: Option<String>,
pub active: bool,
pub version: i64,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ScimGroup {
pub id: Uuid,
pub display_name: String,
pub external_id: Option<String>,
pub members: Vec<String>,
pub version: i64,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct ScimUserWrite {
pub user_name: String,
pub external_id: Option<String>,
pub email: Option<String>,
pub given_name: Option<String>,
pub family_name: Option<String>,
pub display_name: Option<String>,
pub active: bool,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ScimPage<T> {
pub resources: Vec<T>,
pub total_results: i64,
}
#[async_trait]
pub trait ScimStore: Send + Sync {
async fn init(&self) -> Result<()>;
async fn list_users(
&self,
user_name: Option<&str>,
start_index: i64,
count: i64,
) -> Result<ScimPage<ScimUser>>;
async fn get_user(&self, id: &str) -> Result<Option<ScimUser>>;
async fn create_user(&self, write: &ScimUserWrite) -> Result<ScimUser>;
async fn replace_user(&self, id: &str, write: &ScimUserWrite) -> Result<ScimUser>;
async fn set_user_active(&self, id: &str, active: bool) -> Result<ScimUser>;
async fn delete_user(&self, id: &str) -> Result<()>;
async fn list_groups(
&self,
display_name: Option<&str>,
start_index: i64,
count: i64,
) -> Result<ScimPage<ScimGroup>>;
async fn get_group(&self, id: Uuid) -> Result<Option<ScimGroup>>;
async fn create_group(
&self,
display_name: &str,
external_id: Option<&str>,
members: &[String],
) -> Result<ScimGroup>;
async fn replace_group(
&self,
id: Uuid,
display_name: &str,
external_id: Option<&str>,
members: &[String],
) -> Result<ScimGroup>;
async fn patch_group_members(
&self,
id: Uuid,
add: &[String],
remove: &[String],
) -> Result<ScimGroup>;
async fn delete_group(&self, id: Uuid) -> Result<()>;
async fn groups_of_user(&self, user_id: &str) -> Result<Vec<String>>;
}
#[derive(Debug, Clone)]
pub struct PgScimStore {
db: PgPool,
tenant_id: Option<Uuid>,
}
impl PgScimStore {
#[must_use]
pub const fn new(db: PgPool, tenant_id: Option<Uuid>) -> Self {
Self { db, tenant_id }
}
#[must_use]
pub const fn pool(&self) -> &PgPool {
&self.db
}
}
fn db_error(context: &str, e: &sqlx::Error) -> AuthError {
AuthError::DatabaseError {
message: format!("{context}: {e}"),
}
}
fn is_unique_violation(e: &sqlx::Error) -> bool {
matches!(e, sqlx::Error::Database(db) if db.code().as_deref() == Some("23505"))
}
const fn user_columns() -> &'static str {
"user_id, user_name, external_id, email, given_name, family_name, display_name, \
active, version, created_at, updated_at"
}
fn decode_user(row: &sqlx::postgres::PgRow) -> ScimUser {
ScimUser {
id: row.get("user_id"),
user_name: row.get::<Option<String>, _>("user_name").unwrap_or_default(),
external_id: row.get("external_id"),
email: row.get("email"),
given_name: row.get("given_name"),
family_name: row.get("family_name"),
display_name: row.get("display_name"),
active: row.get("active"),
version: row.get("version"),
created_at: row.get("created_at"),
updated_at: row.get("updated_at"),
}
}
fn new_user_id() -> String {
format!("user_{}", Uuid::new_v4().as_simple())
}
#[async_trait]
impl ScimStore for PgScimStore {
async fn init(&self) -> Result<()> {
sqlx::raw_sql(PG_SCIM_SCHEMA_SQL)
.execute(&self.db)
.await
.map_err(|e| db_error("initialize SCIM schema", &e))?;
Ok(())
}
async fn list_users(
&self,
user_name: Option<&str>,
start_index: i64,
count: i64,
) -> Result<ScimPage<ScimUser>> {
let offset = (start_index - 1).max(0);
let total: i64 = sqlx::query_scalar(
"SELECT count(*) FROM core.tb_user WHERE ($1::text IS NULL OR user_name = $1) \
AND tenant_id IS NOT DISTINCT FROM $2",
)
.bind(user_name)
.bind(self.tenant_id)
.fetch_one(&self.db)
.await
.map_err(|e| db_error("count SCIM users", &e))?;
let rows = sqlx::query(&format!(
"SELECT {} FROM core.tb_user WHERE ($1::text IS NULL OR user_name = $1) \
AND tenant_id IS NOT DISTINCT FROM $4 \
ORDER BY pk_user LIMIT $2 OFFSET $3",
user_columns()
))
.bind(user_name)
.bind(count)
.bind(offset)
.bind(self.tenant_id)
.fetch_all(&self.db)
.await
.map_err(|e| db_error("list SCIM users", &e))?;
Ok(ScimPage {
resources: rows.iter().map(decode_user).collect(),
total_results: total,
})
}
async fn get_user(&self, id: &str) -> Result<Option<ScimUser>> {
let row = sqlx::query(&format!(
"SELECT {} FROM core.tb_user WHERE user_id = $1 AND tenant_id IS NOT DISTINCT FROM $2",
user_columns()
))
.bind(id)
.bind(self.tenant_id)
.fetch_optional(&self.db)
.await
.map_err(|e| db_error("get SCIM user", &e))?;
Ok(row.as_ref().map(decode_user))
}
async fn create_user(&self, write: &ScimUserWrite) -> Result<ScimUser> {
let user_id = new_user_id();
let row = sqlx::query(&format!(
"INSERT INTO core.tb_user \
(user_id, user_name, external_id, email, given_name, family_name, display_name, \
active, tenant_id) \
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9) RETURNING {}",
user_columns()
))
.bind(&user_id)
.bind(&write.user_name)
.bind(write.external_id.as_deref())
.bind(write.email.as_deref())
.bind(write.given_name.as_deref())
.bind(write.family_name.as_deref())
.bind(write.display_name.as_deref())
.bind(write.active)
.bind(self.tenant_id)
.fetch_one(&self.db)
.await
.map_err(|e| {
if is_unique_violation(&e) {
AuthError::EmailAlreadyRegistered
} else {
db_error("create SCIM user", &e)
}
})?;
Ok(decode_user(&row))
}
async fn replace_user(&self, id: &str, write: &ScimUserWrite) -> Result<ScimUser> {
let row = sqlx::query(&format!(
"UPDATE core.tb_user SET user_name = $2, external_id = $3, email = $4, \
given_name = $5, family_name = $6, display_name = $7, active = $8, \
version = version + 1, updated_at = now() \
WHERE user_id = $1 AND tenant_id IS NOT DISTINCT FROM $9 RETURNING {}",
user_columns()
))
.bind(id)
.bind(&write.user_name)
.bind(write.external_id.as_deref())
.bind(write.email.as_deref())
.bind(write.given_name.as_deref())
.bind(write.family_name.as_deref())
.bind(write.display_name.as_deref())
.bind(write.active)
.bind(self.tenant_id)
.fetch_optional(&self.db)
.await
.map_err(|e| {
if is_unique_violation(&e) {
AuthError::EmailAlreadyRegistered
} else {
db_error("replace SCIM user", &e)
}
})?
.ok_or(AuthError::TokenNotFound)?;
Ok(decode_user(&row))
}
async fn set_user_active(&self, id: &str, active: bool) -> Result<ScimUser> {
let row = sqlx::query(&format!(
"UPDATE core.tb_user SET active = $2, version = version + 1, updated_at = now() \
WHERE user_id = $1 AND tenant_id IS NOT DISTINCT FROM $3 RETURNING {}",
user_columns()
))
.bind(id)
.bind(active)
.bind(self.tenant_id)
.fetch_optional(&self.db)
.await
.map_err(|e| db_error("set SCIM user active", &e))?
.ok_or(AuthError::TokenNotFound)?;
Ok(decode_user(&row))
}
async fn delete_user(&self, id: &str) -> Result<()> {
let affected = sqlx::query(
"DELETE FROM core.tb_user WHERE user_id = $1 AND tenant_id IS NOT DISTINCT FROM $2",
)
.bind(id)
.bind(self.tenant_id)
.execute(&self.db)
.await
.map_err(|e| db_error("delete SCIM user", &e))?
.rows_affected();
if affected == 0 {
return Err(AuthError::TokenNotFound);
}
Ok(())
}
async fn list_groups(
&self,
display_name: Option<&str>,
start_index: i64,
count: i64,
) -> Result<ScimPage<ScimGroup>> {
let offset = (start_index - 1).max(0);
let total: i64 = sqlx::query_scalar(
"SELECT count(*) FROM core.tb_scim_group \
WHERE tenant_id IS NOT DISTINCT FROM $1 \
AND ($2::text IS NULL OR display_name = $2)",
)
.bind(self.tenant_id)
.bind(display_name)
.fetch_one(&self.db)
.await
.map_err(|e| db_error("count SCIM groups", &e))?;
let rows = sqlx::query(
"SELECT pk_scim_group, id, display_name, external_id, version, created_at, \
updated_at FROM core.tb_scim_group \
WHERE tenant_id IS NOT DISTINCT FROM $1 \
AND ($2::text IS NULL OR display_name = $2) \
ORDER BY pk_scim_group LIMIT $3 OFFSET $4",
)
.bind(self.tenant_id)
.bind(display_name)
.bind(count)
.bind(offset)
.fetch_all(&self.db)
.await
.map_err(|e| db_error("list SCIM groups", &e))?;
let mut resources = Vec::with_capacity(rows.len());
for row in &rows {
resources.push(self.hydrate_group(row).await?);
}
Ok(ScimPage {
resources,
total_results: total,
})
}
async fn get_group(&self, id: Uuid) -> Result<Option<ScimGroup>> {
let row = sqlx::query(
"SELECT pk_scim_group, id, display_name, external_id, version, created_at, \
updated_at FROM core.tb_scim_group \
WHERE id = $1 AND tenant_id IS NOT DISTINCT FROM $2",
)
.bind(id)
.bind(self.tenant_id)
.fetch_optional(&self.db)
.await
.map_err(|e| db_error("get SCIM group", &e))?;
match row {
Some(row) => Ok(Some(self.hydrate_group(&row).await?)),
None => Ok(None),
}
}
async fn create_group(
&self,
display_name: &str,
external_id: Option<&str>,
members: &[String],
) -> Result<ScimGroup> {
self.ensure_members_in_tenant(members).await?;
let row = sqlx::query(
"INSERT INTO core.tb_scim_group (display_name, external_id, tenant_id) \
VALUES ($1, $2, $3) \
RETURNING pk_scim_group, id, display_name, external_id, version, created_at, \
updated_at",
)
.bind(display_name)
.bind(external_id)
.bind(self.tenant_id)
.fetch_one(&self.db)
.await
.map_err(|e| {
if is_unique_violation(&e) {
AuthError::EmailAlreadyRegistered
} else {
db_error("create SCIM group", &e)
}
})?;
let pk: i64 = row.get("pk_scim_group");
self.set_members(pk, members).await?;
self.hydrate_group(&row).await
}
async fn replace_group(
&self,
id: Uuid,
display_name: &str,
external_id: Option<&str>,
members: &[String],
) -> Result<ScimGroup> {
self.ensure_members_in_tenant(members).await?;
let row = sqlx::query(
"UPDATE core.tb_scim_group SET display_name = $3, external_id = $4, \
version = version + 1, updated_at = now() \
WHERE id = $1 AND tenant_id IS NOT DISTINCT FROM $2 \
RETURNING pk_scim_group, id, display_name, external_id, version, created_at, \
updated_at",
)
.bind(id)
.bind(self.tenant_id)
.bind(display_name)
.bind(external_id)
.fetch_optional(&self.db)
.await
.map_err(|e| {
if is_unique_violation(&e) {
AuthError::EmailAlreadyRegistered
} else {
db_error("replace SCIM group", &e)
}
})?
.ok_or(AuthError::TokenNotFound)?;
let pk: i64 = row.get("pk_scim_group");
sqlx::query("DELETE FROM core.tb_scim_group_member WHERE fk_scim_group = $1")
.bind(pk)
.execute(&self.db)
.await
.map_err(|e| db_error("clear SCIM group members", &e))?;
self.set_members(pk, members).await?;
self.hydrate_group(&row).await
}
async fn patch_group_members(
&self,
id: Uuid,
add: &[String],
remove: &[String],
) -> Result<ScimGroup> {
self.ensure_members_in_tenant(add).await?;
let row = sqlx::query(
"UPDATE core.tb_scim_group SET version = version + 1, updated_at = now() \
WHERE id = $1 AND tenant_id IS NOT DISTINCT FROM $2 \
RETURNING pk_scim_group, id, display_name, external_id, version, created_at, \
updated_at",
)
.bind(id)
.bind(self.tenant_id)
.fetch_optional(&self.db)
.await
.map_err(|e| db_error("patch SCIM group", &e))?
.ok_or(AuthError::TokenNotFound)?;
let pk: i64 = row.get("pk_scim_group");
if !remove.is_empty() {
sqlx::query(
"DELETE FROM core.tb_scim_group_member \
WHERE fk_scim_group = $1 AND user_id = ANY($2)",
)
.bind(pk)
.bind(remove)
.execute(&self.db)
.await
.map_err(|e| db_error("remove SCIM group members", &e))?;
}
self.set_members(pk, add).await?;
self.hydrate_group(&row).await
}
async fn delete_group(&self, id: Uuid) -> Result<()> {
let affected = sqlx::query(
"DELETE FROM core.tb_scim_group WHERE id = $1 AND tenant_id IS NOT DISTINCT FROM $2",
)
.bind(id)
.bind(self.tenant_id)
.execute(&self.db)
.await
.map_err(|e| db_error("delete SCIM group", &e))?
.rows_affected();
if affected == 0 {
return Err(AuthError::TokenNotFound);
}
Ok(())
}
async fn groups_of_user(&self, user_id: &str) -> Result<Vec<String>> {
let rows = sqlx::query(
"SELECT g.display_name FROM core.tb_scim_group_member m \
JOIN core.tb_scim_group g ON g.pk_scim_group = m.fk_scim_group \
WHERE m.user_id = $1 AND g.tenant_id IS NOT DISTINCT FROM $2 \
ORDER BY g.display_name",
)
.bind(user_id)
.bind(self.tenant_id)
.fetch_all(&self.db)
.await
.map_err(|e| db_error("list groups of user", &e))?;
Ok(rows.iter().map(|r| r.get("display_name")).collect())
}
}
impl PgScimStore {
async fn ensure_members_in_tenant(&self, members: &[String]) -> Result<()> {
if members.is_empty() {
return Ok(());
}
let found: i64 = sqlx::query_scalar(
"SELECT count(DISTINCT user_id) FROM core.tb_user \
WHERE user_id = ANY($1) AND tenant_id IS NOT DISTINCT FROM $2",
)
.bind(members)
.bind(self.tenant_id)
.fetch_one(&self.db)
.await
.map_err(|e| db_error("check SCIM group members", &e))?;
let wanted = members.iter().collect::<std::collections::HashSet<_>>().len();
if usize::try_from(found).ok() != Some(wanted) {
return Err(AuthError::InvalidRegistration {
reason: "a group member is not a user of this provisioning tenant".to_string(),
});
}
Ok(())
}
async fn set_members(&self, pk_group: i64, members: &[String]) -> Result<()> {
for user_id in members {
sqlx::query(
"INSERT INTO core.tb_scim_group_member (fk_scim_group, user_id) \
VALUES ($1, $2) ON CONFLICT DO NOTHING",
)
.bind(pk_group)
.bind(user_id)
.execute(&self.db)
.await
.map_err(|e| db_error("add SCIM group member", &e))?;
}
Ok(())
}
async fn hydrate_group(&self, row: &sqlx::postgres::PgRow) -> Result<ScimGroup> {
let pk: i64 = row.get("pk_scim_group");
let members = sqlx::query(
"SELECT user_id FROM core.tb_scim_group_member \
WHERE fk_scim_group = $1 ORDER BY pk_scim_group_member",
)
.bind(pk)
.fetch_all(&self.db)
.await
.map_err(|e| db_error("list SCIM group members", &e))?;
Ok(ScimGroup {
id: row.get("id"),
display_name: row.get("display_name"),
external_id: row.get("external_id"),
members: members.iter().map(|r| r.get("user_id")).collect(),
version: row.get("version"),
created_at: row.get("created_at"),
updated_at: row.get("updated_at"),
})
}
}