use std::{
collections::{BTreeMap, HashMap},
future::Future,
pin::Pin,
sync::Arc,
time::Duration,
};
use serde::{Deserialize, Serialize};
use super::{
cache::{CachedOutcome, IdentityCache},
failure::{DenyReason, IdentityResolution, ResolveError},
query::{MissingParam, prepare_enrichment_query},
};
pub(super) type BoxFuture<'a, T> = Pin<Box<dyn Future<Output = T> + Send + 'a>>;
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct EnrichmentQueryConfig {
#[serde(default)]
pub enabled: bool,
pub query: String,
#[serde(default)]
pub map: BTreeMap<String, String>,
#[serde(default = "default_cache_ttl_secs")]
pub cache_ttl_secs: u64,
#[serde(default = "default_negative_ttl_secs")]
pub negative_ttl_secs: u64,
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct IdentityConfig {
#[serde(default)]
pub enrichment: Option<EnrichmentQueryConfig>,
#[serde(default)]
pub sender: Option<EnrichmentQueryConfig>,
}
const fn default_cache_ttl_secs() -> u64 {
60
}
const fn default_negative_ttl_secs() -> u64 {
5
}
pub(super) trait IdentityStore: Send + Sync {
fn fetch_rows<'a>(
&'a self,
sql: &'a str,
binds: &'a [serde_json::Value],
) -> BoxFuture<'a, Result<Vec<serde_json::Map<String, serde_json::Value>>, ResolveError>>;
}
pub(super) struct PgIdentityStore {
pool: sqlx::PgPool,
}
impl PgIdentityStore {
pub(super) const fn new(pool: sqlx::PgPool) -> Self {
Self { pool }
}
}
impl IdentityStore for PgIdentityStore {
fn fetch_rows<'a>(
&'a self,
sql: &'a str,
binds: &'a [serde_json::Value],
) -> BoxFuture<'a, Result<Vec<serde_json::Map<String, serde_json::Value>>, ResolveError>> {
Box::pin(async move {
let wrapped = format!("SELECT row_to_json(t)::text FROM ({sql}) t LIMIT 2");
let mut query = sqlx::query_as::<_, (String,)>(&wrapped);
let string_binds: Vec<String> = binds
.iter()
.map(|v| match v {
serde_json::Value::String(s) => s.clone(),
other => other.to_string(),
})
.collect();
for (bind_value, string_val) in binds.iter().zip(&string_binds) {
query = match bind_value {
serde_json::Value::Number(n) => {
if let Some(i) = n.as_i64() {
query.bind(i)
} else if let Some(f) = n.as_f64() {
query.bind(f)
} else {
query.bind(string_val.as_str())
}
},
serde_json::Value::Bool(b) => query.bind(*b),
serde_json::Value::Null => query.bind(Option::<String>::None),
_ => query.bind(string_val.as_str()),
};
}
let rows = query
.fetch_all(&self.pool)
.await
.map_err(|e| ResolveError::new(format!("identity query failed: {e}")))?;
let mut out = Vec::with_capacity(rows.len());
for (json_text,) in rows {
let value: serde_json::Value = serde_json::from_str(&json_text).map_err(|e| {
ResolveError::new(format!("identity query returned invalid JSON: {e}"))
})?;
match value {
serde_json::Value::Object(map) => out.push(map),
_ => {
return Err(ResolveError::new(
"identity query did not return a JSON object",
));
},
}
}
Ok(out)
})
}
}
pub struct IdentityResolver {
config: EnrichmentQueryConfig,
store: Arc<dyn IdentityStore>,
cache: IdentityCache,
}
impl IdentityResolver {
pub(super) fn new(config: EnrichmentQueryConfig, store: Arc<dyn IdentityStore>) -> Self {
Self {
config,
store,
cache: IdentityCache::new(),
}
}
#[must_use]
pub fn postgres(config: EnrichmentQueryConfig, pool: sqlx::PgPool) -> Self {
Self::new(config, Arc::new(PgIdentityStore::new(pool)))
}
pub(super) async fn resolve(
&self,
sub: &str,
claims: &HashMap<String, serde_json::Value>,
) -> IdentityResolution {
let bound = match prepare_enrichment_query(&self.config.query, claims) {
Ok(bound) => bound,
Err(MissingParam(name)) => {
return self
.finalize(sub, IdentityResolution::Denied(DenyReason::MissingParam(name)));
},
};
let key = cache_key(&bound.binds);
if let Some(cached) = self.cache.get(&key) {
return self.finalize(sub, into_resolution(cached));
}
let rows = match self.store.fetch_rows(&bound.sql, &bound.binds).await {
Ok(rows) => rows,
Err(err) => return self.finalize(sub, IdentityResolution::Unavailable(err)),
};
let resolution = classify(rows, &self.config.map);
match &resolution {
IdentityResolution::Resolved(map) => self.cache.insert(
key,
sub.to_owned(),
CachedOutcome::Resolved(map.clone()),
Duration::from_secs(self.config.cache_ttl_secs),
),
IdentityResolution::Denied(reason) => self.cache.insert(
key,
sub.to_owned(),
CachedOutcome::Denied(reason.clone()),
Duration::from_secs(self.config.negative_ttl_secs),
),
IdentityResolution::Unavailable(_) => {},
}
self.finalize(sub, resolution)
}
pub(super) fn flush(&self, sub: &str) {
self.cache.flush(sub);
}
pub(super) fn flush_all(&self) {
self.cache.flush_all();
}
fn finalize(&self, sub: &str, resolution: IdentityResolution) -> IdentityResolution {
match &resolution {
IdentityResolution::Denied(reason) => tracing::warn!(
subject = %sub,
reason = %reason.log_label(),
"enriched-identity resolution denied",
),
IdentityResolution::Unavailable(err) => tracing::warn!(
subject = %sub,
error = %err,
"enriched-identity resolution unavailable",
),
IdentityResolution::Resolved(_) => {},
}
resolution
}
}
fn cache_key(binds: &[serde_json::Value]) -> String {
serde_json::Value::Array(binds.to_vec()).to_string()
}
fn into_resolution(cached: CachedOutcome) -> IdentityResolution {
match cached {
CachedOutcome::Resolved(map) => IdentityResolution::Resolved(map),
CachedOutcome::Denied(reason) => IdentityResolution::Denied(reason),
}
}
fn classify(
rows: Vec<serde_json::Map<String, serde_json::Value>>,
map: &BTreeMap<String, String>,
) -> IdentityResolution {
let mut rows = rows.into_iter();
let Some(row) = rows.next() else {
return IdentityResolution::Denied(DenyReason::ZeroRows);
};
if rows.next().is_some() {
return IdentityResolution::Denied(DenyReason::Ambiguous);
}
let mut resolved = serde_json::Map::with_capacity(map.len());
for (column, field) in map {
match row.get(column) {
Some(value) if !value.is_null() => {
resolved.insert(field.clone(), value.clone());
},
_ => return IdentityResolution::Denied(DenyReason::NullField(column.clone())),
}
}
IdentityResolution::Resolved(resolved)
}