use sqlx::{PgPool, Row};
use tonic::Status;
use uuid::Uuid;
use crate::backend::BackendKind;
use crate::ir::compile::CompileContext;
use crate::ir::{
ComparisonOp, ConflictStrategy, LogicalAssignment, LogicalFilter, LogicalRecord, LogicalUpdate,
LogicalValue, LogicalWrite,
};
use crate::runtime::core::DataBrokerRuntime;
use crate::runtime::native_catalog::native_model;
fn seal_idp_secret(runtime: &DataBrokerRuntime, plaintext: &str) -> Result<String, Status> {
if plaintext.is_empty() {
return Ok(String::new());
}
runtime
.encrypt_secret_at_rest(plaintext)
.map_err(|err| Status::internal(format!("idp secret encryption-at-rest failed: {err}")))
}
pub const PROVIDER_MSG: &str = "udb.core.idp.entity.v1.IdentityProvider";
pub const EXTERNAL_IDENTITY_MSG: &str = "udb.core.idp.entity.v1.ExternalIdentity";
pub const SCIM_STATE_MSG: &str = "udb.core.idp.entity.v1.ScimDirectoryState";
pub const SAML_REPLAY_MSG: &str = "udb.core.idp.entity.v1.SamlReplayEntry";
pub const USER_MSG: &str = "udb.core.authn.entity.v1.User";
pub const SESSION_MSG: &str = "udb.core.authn.entity.v1.Session";
#[derive(Debug, Clone, Default)]
pub struct ProviderRow {
pub provider_id: String,
pub tenant_id: String,
pub kind: String,
pub display_name: String,
pub issuer: String,
pub entity_id: String,
pub jwks_url: String,
pub saml_metadata_url: String,
pub client_ids_json: String,
pub audiences_json: String,
pub claim_mapping_json: String,
pub group_mapping_json: String,
pub jit_policy_json: String,
pub account_linking_policy: String,
pub enabled: bool,
pub saml_idp_certs_json: String,
pub saml_sso_url: String,
pub health: String,
pub last_jwks_refresh_status: String,
pub created_by: String,
pub updated_by: String,
pub created_at_unix: i64,
pub updated_at_unix: i64,
pub last_jwks_refresh_at_unix: i64,
}
#[derive(Debug, Clone, Default)]
pub struct ExternalIdentityRow {
pub external_identity_id: String,
pub tenant_id: String,
pub provider_id: String,
pub subject: String,
pub user_id: String,
pub email: String,
pub email_verified: bool,
pub linked_at_unix: i64,
pub last_login_at_unix: i64,
}
fn map_err(context: &str) -> impl Fn(sqlx::Error) -> Status + '_ {
move |err| Status::internal(format!("{context}: {err}"))
}
fn native_compile_context() -> CompileContext<'static> {
CompileContext::new(crate::runtime::native_catalog::native_manifest())
}
fn uuid_value(value: &str, field: &str) -> Result<LogicalValue, Status> {
let uuid = Uuid::parse_str(value.trim())
.map_err(|_| Status::invalid_argument(format!("{field} must be a UUID")))?;
Ok(LogicalValue::String(uuid.to_string()))
}
fn string_or_null(value: &str) -> LogicalValue {
if value.is_empty() {
LogicalValue::Null
} else {
LogicalValue::String(value.to_string())
}
}
fn json_or_null(value: &str) -> Result<LogicalValue, Status> {
if value.is_empty() {
return Ok(LogicalValue::Null);
}
serde_json::from_str(value)
.map(LogicalValue::Json)
.map_err(|err| Status::invalid_argument(format!("invalid JSON field value: {err}")))
}
fn json_value(value: &str, field: &str) -> Result<LogicalValue, Status> {
serde_json::from_str(value)
.map(LogicalValue::Json)
.map_err(|err| Status::invalid_argument(format!("{field} must be valid JSON: {err}")))
}
fn eq(field: &str, value: LogicalValue) -> LogicalFilter {
LogicalFilter::Comparison {
field: field.to_string(),
op: ComparisonOp::Eq,
value,
}
}
fn active_provider_filter(provider_id: &str, tenant_id: &str) -> Result<LogicalFilter, Status> {
Ok(LogicalFilter::And(vec![
eq("provider_id", uuid_value(provider_id, "provider_id")?),
eq("tenant_id", LogicalValue::String(tenant_id.to_string())),
LogicalFilter::IsNull("deleted_at".to_string()),
]))
}
async fn execute_typed_write(
pool: &PgPool,
op: LogicalWrite,
) -> Result<(u64, Vec<serde_json::Value>), Status> {
let ctx = native_compile_context();
let compiled = crate::runtime::service::handlers_data::compile_logical_write_dispatch(
&BackendKind::Postgres,
&op,
&ctx,
)?;
execute_compiled_mutation(pool, &compiled.spec_json).await
}
async fn execute_typed_update(
pool: &PgPool,
op: LogicalUpdate,
) -> Result<(u64, Vec<serde_json::Value>), Status> {
let ctx = native_compile_context();
let compiled = crate::runtime::service::handlers_data::compile_logical_update_dispatch(
&BackendKind::Postgres,
&op,
&ctx,
)?;
execute_compiled_mutation(pool, &compiled.spec_json).await
}
async fn execute_compiled_mutation(
pool: &PgPool,
spec_json: &str,
) -> Result<(u64, Vec<serde_json::Value>), Status> {
let spec: serde_json::Value = serde_json::from_str(spec_json)
.map_err(|err| Status::internal(format!("typed native mutation JSON failed: {err}")))?;
let sql = spec
.get("sql")
.and_then(serde_json::Value::as_str)
.ok_or_else(|| Status::internal("typed native mutation did not compile to SQL"))?;
crate::runtime::core::validate_pg_mutation_sql(sql)?;
let params = crate::runtime::core::dispatch_params(&spec)?;
let param_types = crate::runtime::core::dispatch_param_types(&spec)?;
let return_rows = sql.to_ascii_lowercase().contains(" returning ");
if return_rows {
let rows = crate::runtime::core::bind_typed_generic_pg_params(
sqlx::query(sql),
¶ms,
param_types.as_deref(),
)?
.fetch_all(pool)
.await
.map_err(map_err("typed native mutation failed"))?;
let rows = crate::runtime::core::pg_rows_to_json(rows)?;
Ok((rows.len() as u64, rows))
} else {
let result = crate::runtime::core::bind_typed_generic_pg_params(
sqlx::query(sql),
¶ms,
param_types.as_deref(),
)?
.execute(pool)
.await
.map_err(map_err("typed native mutation failed"))?;
Ok((result.rows_affected(), Vec::new()))
}
}
fn provider_select_columns() -> Vec<&'static str> {
vec![
"provider_id",
"tenant_id",
"kind",
"display_name",
"issuer",
"entity_id",
"jwks_url",
"saml_metadata_url",
"client_ids_json",
"audiences_json",
"claim_mapping_json",
"group_mapping_json",
"jit_policy_json",
"account_linking_policy",
"enabled",
"saml_idp_certs_json",
"saml_sso_url",
"health",
"last_jwks_refresh_status",
"created_by",
"updated_by",
"created_at",
"updated_at",
"last_jwks_refresh_at",
"deleted_at",
]
}
fn provider_row_from(row: &sqlx::postgres::PgRow) -> ProviderRow {
ProviderRow {
provider_id: row.try_get("provider_id").unwrap_or_default(),
tenant_id: row.try_get("tenant_id").unwrap_or_default(),
kind: row.try_get("kind").unwrap_or_default(),
display_name: row.try_get("display_name").unwrap_or_default(),
issuer: row.try_get("issuer").unwrap_or_default(),
entity_id: row.try_get("entity_id").unwrap_or_default(),
jwks_url: row.try_get("jwks_url").unwrap_or_default(),
saml_metadata_url: row.try_get("saml_metadata_url").unwrap_or_default(),
client_ids_json: row
.try_get("client_ids_json")
.unwrap_or_else(|_| "[]".into()),
audiences_json: row
.try_get("audiences_json")
.unwrap_or_else(|_| "[]".into()),
claim_mapping_json: row
.try_get("claim_mapping_json")
.unwrap_or_else(|_| "{}".into()),
group_mapping_json: row
.try_get("group_mapping_json")
.unwrap_or_else(|_| "{}".into()),
jit_policy_json: row
.try_get("jit_policy_json")
.unwrap_or_else(|_| "{}".into()),
account_linking_policy: row.try_get("account_linking_policy").unwrap_or_default(),
enabled: row.try_get("enabled").unwrap_or(false),
saml_idp_certs_json: row
.try_get("saml_idp_certs_json")
.unwrap_or_else(|_| "[]".into()),
saml_sso_url: row.try_get("saml_sso_url").unwrap_or_default(),
health: row.try_get("health").unwrap_or_default(),
last_jwks_refresh_status: row.try_get("last_jwks_refresh_status").unwrap_or_default(),
created_by: row.try_get("created_by").unwrap_or_default(),
updated_by: row.try_get("updated_by").unwrap_or_default(),
created_at_unix: row.try_get("created_at_unix").unwrap_or_default(),
updated_at_unix: row.try_get("updated_at_unix").unwrap_or_default(),
last_jwks_refresh_at_unix: row.try_get("last_jwks_refresh_at_unix").unwrap_or_default(),
}
}
fn provider_select_clause() -> String {
let m = native_model(PROVIDER_MSG, &provider_select_columns());
let parts = vec![
m.text_or_empty_as("provider_id", "provider_id"),
m.text_or_empty_as("tenant_id", "tenant_id"),
m.text_or_empty_as("kind", "kind"),
m.text_or_empty_as("display_name", "display_name"),
m.text_or_empty_as("issuer", "issuer"),
m.text_or_empty_as("entity_id", "entity_id"),
m.text_or_empty_as("jwks_url", "jwks_url"),
m.text_or_empty_as("saml_metadata_url", "saml_metadata_url"),
m.json_text_as("client_ids_json", "client_ids_json"),
m.json_text_as("audiences_json", "audiences_json"),
m.json_text_as("claim_mapping_json", "claim_mapping_json"),
m.json_text_as("group_mapping_json", "group_mapping_json"),
m.json_text_as("jit_policy_json", "jit_policy_json"),
m.text_or_empty_as("account_linking_policy", "account_linking_policy"),
format!("{} AS enabled", m.q("enabled")),
m.json_text_as("saml_idp_certs_json", "saml_idp_certs_json"),
m.text_or_empty_as("saml_sso_url", "saml_sso_url"),
m.text_or_empty_as("health", "health"),
m.text_or_empty_as("last_jwks_refresh_status", "last_jwks_refresh_status"),
m.text_or_empty_as("created_by", "created_by"),
m.text_or_empty_as("updated_by", "updated_by"),
m.timestamp_unix_as("created_at", "created_at_unix"),
m.timestamp_unix_as("updated_at", "updated_at_unix"),
m.timestamp_unix_as("last_jwks_refresh_at", "last_jwks_refresh_at_unix"),
];
parts.join(", ")
}
#[allow(clippy::too_many_arguments)]
pub async fn insert_provider(
runtime: &DataBrokerRuntime,
pool: &PgPool,
row: &ProviderRow,
client_secret: &str,
saml_signing_key_pem: &str,
) -> Result<String, Status> {
let client_secret = seal_idp_secret(runtime, client_secret)?;
let saml_signing_key_pem = seal_idp_secret(runtime, saml_signing_key_pem)?;
let provider_id = Uuid::new_v4();
let mut record = LogicalRecord::new();
record.insert(
"provider_id".to_string(),
LogicalValue::String(provider_id.to_string()),
);
record.insert(
"tenant_id".to_string(),
LogicalValue::String(row.tenant_id.clone()),
);
record.insert("kind".to_string(), LogicalValue::String(row.kind.clone()));
record.insert(
"display_name".to_string(),
LogicalValue::String(row.display_name.clone()),
);
record.insert(
"issuer".to_string(),
LogicalValue::String(row.issuer.clone()),
);
record.insert(
"entity_id".to_string(),
LogicalValue::String(row.entity_id.clone()),
);
record.insert(
"jwks_url".to_string(),
LogicalValue::String(row.jwks_url.clone()),
);
record.insert(
"saml_metadata_url".to_string(),
LogicalValue::String(row.saml_metadata_url.clone()),
);
record.insert(
"client_ids_json".to_string(),
json_value(&row.client_ids_json, "client_ids_json")?,
);
record.insert(
"audiences_json".to_string(),
json_value(&row.audiences_json, "audiences_json")?,
);
record.insert(
"claim_mapping_json".to_string(),
json_value(&row.claim_mapping_json, "claim_mapping_json")?,
);
record.insert(
"group_mapping_json".to_string(),
json_value(&row.group_mapping_json, "group_mapping_json")?,
);
record.insert(
"jit_policy_json".to_string(),
json_value(&row.jit_policy_json, "jit_policy_json")?,
);
record.insert(
"account_linking_policy".to_string(),
LogicalValue::String(row.account_linking_policy.clone()),
);
record.insert("enabled".to_string(), LogicalValue::Bool(row.enabled));
record.insert("client_secret".to_string(), string_or_null(&client_secret));
record.insert(
"saml_signing_key_pem".to_string(),
string_or_null(&saml_signing_key_pem),
);
record.insert(
"saml_idp_certs_json".to_string(),
json_value(&row.saml_idp_certs_json, "saml_idp_certs_json")?,
);
record.insert(
"saml_sso_url".to_string(),
LogicalValue::String(row.saml_sso_url.clone()),
);
record.insert(
"health".to_string(),
LogicalValue::String(if row.health.is_empty() {
"UNSPECIFIED".to_string()
} else {
row.health.clone()
}),
);
record.insert(
"created_by".to_string(),
LogicalValue::String(row.created_by.clone()),
);
record.insert(
"updated_by".to_string(),
LogicalValue::String(row.updated_by.clone()),
);
let (_, rows) = execute_typed_write(
pool,
LogicalWrite {
message_type: PROVIDER_MSG.to_string(),
records: vec![record],
conflict: ConflictStrategy::Error,
return_fields: vec!["provider_id".to_string()],
},
)
.await?;
Ok(rows
.first()
.and_then(|row| row.get("provider_id"))
.and_then(serde_json::Value::as_str)
.unwrap_or_default()
.to_string())
}
pub async fn get_provider(
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
) -> Result<Option<ProviderRow>, Status> {
let m = native_model(PROVIDER_MSG, &["provider_id", "tenant_id", "deleted_at"]);
let sql = format!(
"SELECT {cols} FROM {rel} WHERE {pid} = $1::UUID AND {tenant} = $2 AND {del} IS NULL",
cols = provider_select_clause(),
rel = m.relation,
pid = m.q("provider_id"),
tenant = m.q("tenant_id"),
del = m.q("deleted_at"),
);
let pid = Uuid::parse_str(provider_id.trim())
.map_err(|_| Status::invalid_argument("provider_id must be a UUID"))?;
let row = sqlx::query(&sql)
.bind(pid)
.bind(tenant_id)
.fetch_optional(pool)
.await
.map_err(map_err("idp provider get failed"))?;
Ok(row.map(|r| provider_row_from(&r)))
}
pub async fn list_providers(
pool: &PgPool,
tenant_id: &str,
kind: &str,
enabled_only: bool,
limit: i64,
offset: i64,
) -> Result<(Vec<ProviderRow>, i64), Status> {
let m = native_model(
PROVIDER_MSG,
&["tenant_id", "kind", "enabled", "deleted_at", "created_at"],
);
let base_where = format!(
"{tenant} = $1 AND {del} IS NULL AND ($2 = '' OR {kind} = $2) \
AND ($3 = false OR {enabled} = true)",
tenant = m.q("tenant_id"),
del = m.q("deleted_at"),
kind = m.q("kind"),
enabled = m.q("enabled"),
);
let count_sql = format!(
"SELECT COUNT(*)::bigint AS cnt FROM {rel} WHERE {filter}",
rel = m.relation,
filter = base_where,
);
let total: i64 = sqlx::query(&count_sql)
.bind(tenant_id)
.bind(kind)
.bind(enabled_only)
.fetch_one(pool)
.await
.map_err(map_err("idp provider count failed"))?
.try_get("cnt")
.unwrap_or(0);
let sql = format!(
"SELECT {cols} FROM {rel} WHERE {filter} ORDER BY {created} DESC LIMIT $4 OFFSET $5",
cols = provider_select_clause(),
rel = m.relation,
filter = base_where,
created = m.q("created_at"),
);
let rows = sqlx::query(&sql)
.bind(tenant_id)
.bind(kind)
.bind(enabled_only)
.bind(limit)
.bind(offset)
.fetch_all(pool)
.await
.map_err(map_err("idp provider list failed"))?;
Ok((rows.iter().map(provider_row_from).collect(), total))
}
#[allow(clippy::too_many_arguments)]
pub async fn update_provider(
runtime: &DataBrokerRuntime,
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
fields: &ProviderRow,
client_secret: &str,
saml_signing_key_pem: &str,
) -> Result<Option<ProviderRow>, Status> {
let client_secret = seal_idp_secret(runtime, client_secret)?;
let saml_signing_key_pem = seal_idp_secret(runtime, saml_signing_key_pem)?;
let mut assignments = std::collections::BTreeMap::new();
for (field, value) in [
("display_name", fields.display_name.as_str()),
("issuer", fields.issuer.as_str()),
("entity_id", fields.entity_id.as_str()),
("jwks_url", fields.jwks_url.as_str()),
("saml_metadata_url", fields.saml_metadata_url.as_str()),
(
"account_linking_policy",
fields.account_linking_policy.as_str(),
),
("client_secret", client_secret.as_str()),
("saml_signing_key_pem", saml_signing_key_pem.as_str()),
("updated_by", fields.updated_by.as_str()),
] {
assignments.insert(
field.to_string(),
LogicalAssignment::Coalesce {
value: string_or_null(value),
},
);
}
for (field, value) in [
("client_ids_json", fields.client_ids_json.as_str()),
("audiences_json", fields.audiences_json.as_str()),
("claim_mapping_json", fields.claim_mapping_json.as_str()),
("group_mapping_json", fields.group_mapping_json.as_str()),
("jit_policy_json", fields.jit_policy_json.as_str()),
] {
assignments.insert(
field.to_string(),
LogicalAssignment::Coalesce {
value: json_or_null(value)?,
},
);
}
let (affected, _) = execute_typed_update(
pool,
LogicalUpdate {
message_type: PROVIDER_MSG.to_string(),
filter: active_provider_filter(provider_id, tenant_id)?,
assignments,
return_fields: Vec::new(),
require_affected: false,
},
)
.await?;
if affected == 0 {
return Ok(None);
}
get_provider(pool, tenant_id, provider_id).await
}
pub async fn disable_provider(
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
updated_by: &str,
) -> Result<Option<ProviderRow>, Status> {
let mut assignments = std::collections::BTreeMap::new();
assignments.insert(
"enabled".to_string(),
LogicalAssignment::Set {
value: LogicalValue::Bool(false),
},
);
assignments.insert(
"updated_by".to_string(),
LogicalAssignment::Coalesce {
value: string_or_null(updated_by),
},
);
let (affected, _) = execute_typed_update(
pool,
LogicalUpdate {
message_type: PROVIDER_MSG.to_string(),
filter: active_provider_filter(provider_id, tenant_id)?,
assignments,
return_fields: Vec::new(),
require_affected: false,
},
)
.await?;
if affected == 0 {
return Ok(None);
}
get_provider(pool, tenant_id, provider_id).await
}
pub async fn record_jwks_refresh(
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
health: &str,
status: &str,
) -> Result<(), Status> {
let mut assignments = std::collections::BTreeMap::new();
assignments.insert(
"health".to_string(),
LogicalAssignment::Set {
value: LogicalValue::String(health.to_string()),
},
);
assignments.insert(
"last_jwks_refresh_status".to_string(),
LogicalAssignment::Set {
value: LogicalValue::String(status.to_string()),
},
);
assignments.insert(
"last_jwks_refresh_at".to_string(),
LogicalAssignment::ServerNow,
);
execute_typed_update(
pool,
LogicalUpdate {
message_type: PROVIDER_MSG.to_string(),
filter: LogicalFilter::And(vec![
eq("provider_id", uuid_value(provider_id, "provider_id")?),
eq("tenant_id", LogicalValue::String(tenant_id.to_string())),
]),
assignments,
return_fields: Vec::new(),
require_affected: false,
},
)
.await?;
Ok(())
}
pub async fn record_scim_sync(
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
success: bool,
cursor: &str,
error: &str,
) -> Result<(), Status> {
let m = native_model(
SCIM_STATE_MSG,
&[
"scim_directory_state_id",
"tenant_id",
"provider_id",
"cursor",
"last_sync_at",
"failure_count",
"last_error",
],
);
let sql = format!(
"INSERT INTO {rel} ({tenant}, {pid}, {cursor}, {last}, {fc}, {err}) \
VALUES ($1, $2::UUID, NULLIF($4,''), \
CASE WHEN $3 THEN NOW() ELSE NULL END, \
CASE WHEN $3 THEN 0 ELSE 1 END, \
NULLIF($5,'')) \
ON CONFLICT ({tenant}, {pid}) DO UPDATE SET \
{cursor} = COALESCE(NULLIF($4,''), {rel}.{cursor}), \
{last} = CASE WHEN $3 THEN NOW() ELSE {rel}.{last} END, \
{fc} = CASE WHEN $3 THEN 0 ELSE {rel}.{fc} + 1 END, \
{err} = CASE WHEN $3 THEN NULL ELSE NULLIF($5,'') END",
rel = m.relation,
tenant = m.q("tenant_id"),
pid = m.q("provider_id"),
cursor = m.q("cursor"),
last = m.q("last_sync_at"),
fc = m.q("failure_count"),
err = m.q("last_error"),
);
let pid = Uuid::parse_str(provider_id.trim())
.map_err(|_| Status::invalid_argument("provider_id must be a UUID"))?;
sqlx::query(&sql)
.bind(tenant_id)
.bind(pid)
.bind(success)
.bind(cursor)
.bind(error)
.execute(pool)
.await
.map_err(map_err("scim directory-state update failed"))?;
Ok(())
}
pub async fn update_saml_metadata(
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
entity_id: &str,
sso_url: &str,
certs_json: &str,
updated_by: &str,
) -> Result<Option<ProviderRow>, Status> {
let mut assignments = std::collections::BTreeMap::new();
assignments.insert(
"entity_id".to_string(),
LogicalAssignment::Coalesce {
value: string_or_null(entity_id),
},
);
assignments.insert(
"saml_sso_url".to_string(),
LogicalAssignment::Coalesce {
value: string_or_null(sso_url),
},
);
assignments.insert(
"saml_idp_certs_json".to_string(),
LogicalAssignment::Set {
value: json_value(certs_json, "saml_idp_certs_json")?,
},
);
assignments.insert(
"updated_by".to_string(),
LogicalAssignment::Coalesce {
value: string_or_null(updated_by),
},
);
let (affected, _) = execute_typed_update(
pool,
LogicalUpdate {
message_type: PROVIDER_MSG.to_string(),
filter: active_provider_filter(provider_id, tenant_id)?,
assignments,
return_fields: Vec::new(),
require_affected: false,
},
)
.await?;
if affected == 0 {
return Ok(None);
}
get_provider(pool, tenant_id, provider_id).await
}
fn external_select_clause() -> String {
let m = native_model(
EXTERNAL_IDENTITY_MSG,
&[
"external_identity_id",
"tenant_id",
"provider_id",
"subject",
"user_id",
"email",
"email_verified",
"linked_at",
"last_login_at",
],
);
[
m.text_or_empty_as("external_identity_id", "external_identity_id"),
m.text_or_empty_as("tenant_id", "tenant_id"),
m.text_or_empty_as("provider_id", "provider_id"),
m.text_or_empty_as("subject", "subject"),
m.text_or_empty_as("user_id", "user_id"),
m.text_or_empty_as("email", "email"),
format!("{} AS email_verified", m.q("email_verified")),
m.timestamp_unix_as("linked_at", "linked_at_unix"),
m.timestamp_unix_as("last_login_at", "last_login_at_unix"),
]
.join(", ")
}
fn external_row_from(row: &sqlx::postgres::PgRow) -> ExternalIdentityRow {
ExternalIdentityRow {
external_identity_id: row.try_get("external_identity_id").unwrap_or_default(),
tenant_id: row.try_get("tenant_id").unwrap_or_default(),
provider_id: row.try_get("provider_id").unwrap_or_default(),
subject: row.try_get("subject").unwrap_or_default(),
user_id: row.try_get("user_id").unwrap_or_default(),
email: row.try_get("email").unwrap_or_default(),
email_verified: row.try_get("email_verified").unwrap_or(false),
linked_at_unix: row.try_get("linked_at_unix").unwrap_or_default(),
last_login_at_unix: row.try_get("last_login_at_unix").unwrap_or_default(),
}
}
pub async fn get_external_identity(
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
subject: &str,
) -> Result<Option<ExternalIdentityRow>, Status> {
let m = native_model(
EXTERNAL_IDENTITY_MSG,
&["tenant_id", "provider_id", "subject", "deleted_at"],
);
let sql = format!(
"SELECT {cols} FROM {rel} WHERE {tenant} = $1 AND {pid} = $2::UUID AND {sub} = $3 \
AND {del} IS NULL",
cols = external_select_clause(),
rel = m.relation,
tenant = m.q("tenant_id"),
pid = m.q("provider_id"),
sub = m.q("subject"),
del = m.q("deleted_at"),
);
let pid = Uuid::parse_str(provider_id.trim())
.map_err(|_| Status::invalid_argument("provider_id must be a UUID"))?;
let row = sqlx::query(&sql)
.bind(tenant_id)
.bind(pid)
.bind(subject)
.fetch_optional(pool)
.await
.map_err(map_err("external identity lookup failed"))?;
Ok(row.map(|r| external_row_from(&r)))
}
#[allow(clippy::too_many_arguments)]
pub async fn upsert_external_identity(
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
subject: &str,
user_id: &str,
email: &str,
email_verified: bool,
) -> Result<ExternalIdentityRow, Status> {
let m = native_model(
EXTERNAL_IDENTITY_MSG,
&[
"external_identity_id",
"tenant_id",
"provider_id",
"subject",
"user_id",
"email",
"email_verified",
"last_login_at",
],
);
let sql = format!(
"INSERT INTO {rel} ({id}, {tenant}, {pid}, {sub}, {uid}, {email}, {ev}, {lla}) \
VALUES (gen_random_uuid(), $1, $2::UUID, $3, $4::UUID, $5, $6, NOW()) \
ON CONFLICT ({tenant}, {pid}, {sub}) DO UPDATE SET \
{uid} = EXCLUDED.{uid}, {email} = EXCLUDED.{email}, \
{ev} = EXCLUDED.{ev}, {lla} = NOW()",
rel = m.relation,
id = m.q("external_identity_id"),
tenant = m.q("tenant_id"),
pid = m.q("provider_id"),
sub = m.q("subject"),
uid = m.q("user_id"),
email = m.q("email"),
ev = m.q("email_verified"),
lla = m.q("last_login_at"),
);
let pid = Uuid::parse_str(provider_id.trim())
.map_err(|_| Status::invalid_argument("provider_id must be a UUID"))?;
let uid = Uuid::parse_str(user_id.trim())
.map_err(|_| Status::invalid_argument("user_id must be a UUID"))?;
sqlx::query(&sql)
.bind(tenant_id)
.bind(pid)
.bind(subject)
.bind(uid)
.bind(email)
.bind(email_verified)
.execute(pool)
.await
.map_err(map_err("external identity upsert failed"))?;
get_external_identity(pool, tenant_id, provider_id, subject)
.await?
.ok_or_else(|| Status::internal("external identity vanished after upsert"))
}
pub async fn list_external_identities(
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
user_id: &str,
limit: i64,
offset: i64,
) -> Result<(Vec<ExternalIdentityRow>, i64), Status> {
let m = native_model(
EXTERNAL_IDENTITY_MSG,
&[
"tenant_id",
"provider_id",
"user_id",
"deleted_at",
"linked_at",
],
);
let base_where = format!(
"{tenant} = $1 AND {del} IS NULL \
AND ($2 = '' OR {pid}::text = $2) \
AND ($3 = '' OR {uid}::text = $3)",
tenant = m.q("tenant_id"),
del = m.q("deleted_at"),
pid = m.q("provider_id"),
uid = m.q("user_id"),
);
let count_sql = format!(
"SELECT COUNT(*)::bigint AS cnt FROM {rel} WHERE {filter}",
rel = m.relation,
filter = base_where,
);
let total: i64 = sqlx::query(&count_sql)
.bind(tenant_id)
.bind(provider_id.trim())
.bind(user_id.trim())
.fetch_one(pool)
.await
.map_err(map_err("external identity count failed"))?
.try_get("cnt")
.unwrap_or(0);
let sql = format!(
"SELECT {cols} FROM {rel} WHERE {filter} ORDER BY {linked} DESC LIMIT $4 OFFSET $5",
cols = external_select_clause(),
rel = m.relation,
filter = base_where,
linked = m.q("linked_at"),
);
let rows = sqlx::query(&sql)
.bind(tenant_id)
.bind(provider_id.trim())
.bind(user_id.trim())
.bind(limit)
.bind(offset)
.fetch_all(pool)
.await
.map_err(map_err("external identity list failed"))?;
Ok((rows.iter().map(external_row_from).collect(), total))
}
pub async fn unlink_external_identity(
pool: &PgPool,
tenant_id: &str,
external_identity_id: &str,
) -> Result<bool, Status> {
let mut assignments = std::collections::BTreeMap::new();
assignments.insert("deleted_at".to_string(), LogicalAssignment::ServerNow);
let (affected, _) = execute_typed_update(
pool,
LogicalUpdate {
message_type: EXTERNAL_IDENTITY_MSG.to_string(),
filter: LogicalFilter::And(vec![
eq(
"external_identity_id",
uuid_value(external_identity_id, "external_identity_id")?,
),
eq("tenant_id", LogicalValue::String(tenant_id.to_string())),
LogicalFilter::IsNull("deleted_at".to_string()),
]),
assignments,
return_fields: Vec::new(),
require_affected: false,
},
)
.await?;
Ok(affected > 0)
}
pub async fn record_saml_assertion(
pool: &PgPool,
tenant_id: &str,
provider_id: &str,
assertion_id: &str,
not_on_or_after_unix: i64,
) -> Result<bool, Status> {
let not_on_or_after = chrono::DateTime::from_timestamp(not_on_or_after_unix, 0)
.ok_or_else(|| Status::invalid_argument("not_on_or_after_unix is out of range"))?;
let mut record = LogicalRecord::new();
record.insert(
"saml_replay_entry_id".to_string(),
LogicalValue::String(Uuid::new_v4().to_string()),
);
record.insert(
"tenant_id".to_string(),
LogicalValue::String(tenant_id.to_string()),
);
record.insert(
"provider_id".to_string(),
uuid_value(provider_id, "provider_id")?,
);
record.insert(
"assertion_id".to_string(),
LogicalValue::String(assertion_id.to_string()),
);
record.insert(
"not_on_or_after".to_string(),
LogicalValue::Timestamp(not_on_or_after),
);
let (affected, _) = execute_typed_write(
pool,
LogicalWrite {
message_type: SAML_REPLAY_MSG.to_string(),
records: vec![record],
conflict: ConflictStrategy::Ignore,
return_fields: Vec::new(),
},
)
.await?;
Ok(affected > 0)
}
pub async fn find_user_by_email(
pool: &PgPool,
tenant_id: &str,
email: &str,
) -> Result<Option<(String, bool)>, Status> {
if email.trim().is_empty() {
return Ok(None);
}
let m = native_model(
USER_MSG,
&[
"user_id",
"tenant_id",
"email",
"email_verified_at",
"deleted_at",
],
);
let sql = format!(
"SELECT {uid}::TEXT AS user_id, ({ev} IS NOT NULL) AS email_verified \
FROM {rel} WHERE {tenant} = $1 AND lower({email}) = lower($2) AND {del} IS NULL LIMIT 1",
uid = m.q("user_id"),
ev = m.q("email_verified_at"),
rel = m.relation,
tenant = m.q("tenant_id"),
email = m.q("email"),
del = m.q("deleted_at"),
);
let row = sqlx::query(&sql)
.bind(tenant_id)
.bind(email)
.fetch_optional(pool)
.await
.map_err(map_err("user lookup by email failed"))?;
Ok(row.map(|r| {
(
r.try_get::<String, _>("user_id").unwrap_or_default(),
r.try_get::<bool, _>("email_verified").unwrap_or(false),
)
}))
}
pub async fn create_external_user(
pool: &PgPool,
tenant_id: &str,
project_id: &str,
provider_id: &str,
subject: &str,
email: &str,
full_name: &str,
email_verified: bool,
created_by: &str,
) -> Result<String, Status> {
let m = native_model(
USER_MSG,
&[
"user_id",
"username",
"email",
"password_hash",
"account_kind",
"status",
"tenant_id",
"project_id",
"full_name",
"email_verified_at",
"external_provider_id",
"external_subject",
"created_by",
],
);
let username = external_username(provider_id, subject);
let sql = format!(
"INSERT INTO {rel} ({uid}, {uname}, {email}, {pwd}, {ak}, {status}, {tenant}, \
{project}, {fname}, {eva}, {epid}, {esub}, {cby}) \
VALUES (gen_random_uuid(), $1, $2, '', 'EXTERNAL_IDENTITY', \
'ACTIVE', $3, $4, $5, \
CASE WHEN $6 THEN NOW() ELSE NULL END, $7, $8, $9) \
ON CONFLICT ({uname}) DO UPDATE SET {email} = EXCLUDED.{email} \
RETURNING {uid}::TEXT AS user_id",
rel = m.relation,
uid = m.q("user_id"),
uname = m.q("username"),
email = m.q("email"),
pwd = m.q("password_hash"),
ak = m.q("account_kind"),
status = m.q("status"),
tenant = m.q("tenant_id"),
project = m.q("project_id"),
fname = m.q("full_name"),
eva = m.q("email_verified_at"),
epid = m.q("external_provider_id"),
esub = m.q("external_subject"),
cby = m.q("created_by"),
);
let row = sqlx::query(&sql)
.bind(&username)
.bind(email)
.bind(tenant_id)
.bind(project_id)
.bind(full_name)
.bind(email_verified)
.bind(provider_id)
.bind(subject)
.bind(created_by)
.fetch_one(pool)
.await
.map_err(map_err("external user JIT provision failed"))?;
Ok(row.try_get::<String, _>("user_id").unwrap_or_default())
}
pub async fn deactivate_user(
pool: &PgPool,
tenant_id: &str,
user_id: &str,
) -> Result<bool, Status> {
let mut user_assignments = std::collections::BTreeMap::new();
user_assignments.insert(
"status".to_string(),
LogicalAssignment::Set {
value: LogicalValue::String("SUSPENDED".to_string()),
},
);
let (affected, _) = execute_typed_update(
pool,
LogicalUpdate {
message_type: USER_MSG.to_string(),
filter: LogicalFilter::And(vec![
eq("user_id", uuid_value(user_id, "user_id")?),
eq("tenant_id", LogicalValue::String(tenant_id.to_string())),
LogicalFilter::IsNull("deleted_at".to_string()),
]),
assignments: user_assignments,
return_fields: Vec::new(),
require_affected: false,
},
)
.await?;
let mut session_assignments = std::collections::BTreeMap::new();
session_assignments.insert(
"is_active".to_string(),
LogicalAssignment::Set {
value: LogicalValue::Bool(false),
},
);
session_assignments.insert(
"revoke_reason".to_string(),
LogicalAssignment::Set {
value: LogicalValue::String("scim_deprovision".to_string()),
},
);
let _ = execute_typed_update(
pool,
LogicalUpdate {
message_type: SESSION_MSG.to_string(),
filter: LogicalFilter::And(vec![
eq("user_id", uuid_value(user_id, "user_id")?),
eq("tenant_id", LogicalValue::String(tenant_id.to_string())),
eq("is_active", LogicalValue::Bool(true)),
]),
assignments: session_assignments,
return_fields: Vec::new(),
require_affected: false,
},
)
.await; Ok(affected > 0)
}
fn external_username(provider_id: &str, subject: &str) -> String {
use sha2::{Digest, Sha256};
let mut hasher = Sha256::new();
hasher.update(b"udb.idp.external-username.v1:");
hasher.update(provider_id.as_bytes());
hasher.update(b"\x1f");
hasher.update(subject.as_bytes());
let digest = hasher.finalize();
let hex: String = digest.iter().take(20).map(|b| format!("{b:02x}")).collect();
format!("ext-{hex}")
}