use chrono::{DateTime, Utc};
use sqlx::{Row, postgres::PgPool};
use subtle::ConstantTimeEq as _;
use uuid::Uuid;
use crate::error::{AuthError, Result};
const TOKEN_BYTES: usize = 32;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ScimTokenRecord {
pub id: Uuid,
pub idp_name: String,
pub tenant_id: Option<Uuid>,
pub description: Option<String>,
pub created_at: DateTime<Utc>,
pub last_used_at: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone)]
pub struct MintedScimToken {
pub record: ScimTokenRecord,
pub token: String,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ScimPrincipal {
pub idp_name: String,
pub tenant_id: Option<Uuid>,
}
#[derive(Debug, Clone)]
pub struct PgScimTokenStore {
db: PgPool,
}
fn hash_token(token: &str) -> String {
use sha2::{Digest as _, Sha256};
let mut hasher = Sha256::new();
hasher.update(token.as_bytes());
hex::encode(hasher.finalize())
}
fn db_error(context: &str, e: &sqlx::Error) -> AuthError {
AuthError::DatabaseError {
message: format!("{context}: {e}"),
}
}
impl PgScimTokenStore {
#[must_use]
pub const fn new(db: PgPool) -> Self {
Self { db }
}
pub async fn mint(
&self,
idp_name: &str,
tenant_id: Option<Uuid>,
description: Option<&str>,
) -> Result<MintedScimToken> {
use base64::Engine as _;
use rand::RngCore as _;
let mut bytes = [0u8; TOKEN_BYTES];
rand::rng().fill_bytes(&mut bytes);
let token = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(bytes);
let row = sqlx::query(
"INSERT INTO core.tb_scim_token (token_hash, idp_name, tenant_id, description) \
VALUES ($1, $2, $3, $4) \
RETURNING id, idp_name, tenant_id, description, created_at, last_used_at",
)
.bind(hash_token(&token))
.bind(idp_name)
.bind(tenant_id)
.bind(description)
.fetch_one(&self.db)
.await
.map_err(|e| db_error("mint SCIM token", &e))?;
Ok(MintedScimToken {
record: decode(&row),
token,
})
}
pub async fn authenticate(&self, token: &str) -> Result<ScimPrincipal> {
let presented = hash_token(token);
let row = sqlx::query(
"SELECT id, token_hash, idp_name, tenant_id, description, created_at, last_used_at \
FROM core.tb_scim_token WHERE token_hash = $1 AND revoked_at IS NULL",
)
.bind(&presented)
.fetch_optional(&self.db)
.await
.map_err(|e| db_error("authenticate SCIM token", &e))?
.ok_or_else(|| AuthError::InvalidToken {
reason: "unknown or revoked SCIM provisioning token".to_string(),
})?;
let stored: String = row.get("token_hash");
if stored.as_bytes().ct_eq(presented.as_bytes()).unwrap_u8() != 1 {
return Err(AuthError::InvalidToken {
reason: "SCIM provisioning token hash mismatch".to_string(),
});
}
let id: Uuid = row.get("id");
let _ = sqlx::query("UPDATE core.tb_scim_token SET last_used_at = now() WHERE id = $1")
.bind(id)
.execute(&self.db)
.await;
Ok(ScimPrincipal {
idp_name: row.get("idp_name"),
tenant_id: row.get("tenant_id"),
})
}
pub async fn list(&self) -> Result<Vec<ScimTokenRecord>> {
let rows = sqlx::query(
"SELECT id, idp_name, tenant_id, description, created_at, last_used_at \
FROM core.tb_scim_token WHERE revoked_at IS NULL ORDER BY created_at",
)
.fetch_all(&self.db)
.await
.map_err(|e| db_error("list SCIM tokens", &e))?;
Ok(rows.iter().map(decode).collect())
}
pub async fn revoke(&self, id: Uuid) -> Result<()> {
let affected = sqlx::query(
"UPDATE core.tb_scim_token SET revoked_at = now() \
WHERE id = $1 AND revoked_at IS NULL",
)
.bind(id)
.execute(&self.db)
.await
.map_err(|e| db_error("revoke SCIM token", &e))?
.rows_affected();
if affected == 0 {
return Err(AuthError::TokenNotFound);
}
Ok(())
}
}
fn decode(row: &sqlx::postgres::PgRow) -> ScimTokenRecord {
ScimTokenRecord {
id: row.get("id"),
idp_name: row.get("idp_name"),
tenant_id: row.get("tenant_id"),
description: row.get("description"),
created_at: row.get("created_at"),
last_used_at: row.get("last_used_at"),
}
}