use std::collections::HashSet;
use std::sync::Arc;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use sqlx::PgPool;
use uuid::Uuid;
#[derive(Debug, Clone)]
pub struct RecipientContext {
pub user_id: Uuid,
pub email: String,
pub company_id: Option<Uuid>,
pub is_internal: bool,
pub last_login: Option<DateTime<Utc>>,
}
#[derive(Debug, Clone, thiserror::Error)]
pub enum RecipientPortError {
#[error("recipient context port not wired: {detail}")]
NotWired { detail: String },
#[error("recipient context source failed: {0}")]
Unavailable(String),
}
impl RecipientPortError {
pub fn is_not_wired(&self) -> bool {
matches!(self, RecipientPortError::NotWired { .. })
}
}
#[async_trait]
pub trait RecipientContextPort: Send + Sync {
async fn resolve(&self, user_id: Uuid) -> Result<RecipientContext, RecipientPortError>;
async fn count_connected(
&self,
company_id: Option<Uuid>,
since: DateTime<Utc>,
until: DateTime<Utc>,
) -> Result<i64, RecipientPortError>;
async fn any_logged_in_since(
&self,
user_ids: &[Uuid],
since: DateTime<Utc>,
) -> Result<bool, RecipientPortError>;
async fn group_keys(&self, user_id: Uuid) -> Result<HashSet<String>, RecipientPortError>;
}
pub struct RefusingRecipientContext;
fn not_wired(what: &str) -> RecipientPortError {
RecipientPortError::NotWired {
detail: format!(
"no RecipientContextPort is wired; refusing {what} — compose one via \
DigestModule::set_recipient_port (SqlRecipientContext is the standard adapter)"
),
}
}
#[async_trait]
impl RecipientContextPort for RefusingRecipientContext {
async fn resolve(&self, _user_id: Uuid) -> Result<RecipientContext, RecipientPortError> {
Err(not_wired("recipient resolution"))
}
async fn count_connected(
&self,
_company_id: Option<Uuid>,
_since: DateTime<Utc>,
_until: DateTime<Utc>,
) -> Result<i64, RecipientPortError> {
Err(not_wired("the connected-users count"))
}
async fn any_logged_in_since(
&self,
_user_ids: &[Uuid],
_since: DateTime<Utc>,
) -> Result<bool, RecipientPortError> {
Err(not_wired("the slowdown login signal"))
}
async fn group_keys(&self, _user_id: Uuid) -> Result<HashSet<String>, RecipientPortError> {
Err(not_wired("the tip group-key resolution"))
}
}
pub struct CannedRecipientContext {
pub resolved: std::sync::Mutex<Vec<Uuid>>,
pub answer: std::sync::Mutex<CannedAnswers>,
}
#[derive(Default, Clone)]
pub struct CannedAnswers {
pub context: Option<RecipientContext>,
pub connected: i64,
pub any_logged_in: bool,
pub keys: HashSet<String>,
}
impl CannedRecipientContext {
pub fn with(answer: CannedAnswers) -> Self {
Self {
resolved: std::sync::Mutex::new(Vec::new()),
answer: std::sync::Mutex::new(answer),
}
}
}
#[async_trait]
impl RecipientContextPort for CannedRecipientContext {
async fn resolve(&self, user_id: Uuid) -> Result<RecipientContext, RecipientPortError> {
self.resolved.lock().unwrap_or_else(|e| e.into_inner()).push(user_id);
self.answer
.lock()
.unwrap_or_else(|e| e.into_inner())
.context
.clone()
.ok_or_else(|| RecipientPortError::Unavailable("no canned context".into()))
}
async fn count_connected(
&self,
_company_id: Option<Uuid>,
_since: DateTime<Utc>,
_until: DateTime<Utc>,
) -> Result<i64, RecipientPortError> {
Ok(self.answer.lock().unwrap_or_else(|e| e.into_inner()).connected)
}
async fn any_logged_in_since(
&self,
_user_ids: &[Uuid],
_since: DateTime<Utc>,
) -> Result<bool, RecipientPortError> {
Ok(self.answer.lock().unwrap_or_else(|e| e.into_inner()).any_logged_in)
}
async fn group_keys(&self, _user_id: Uuid) -> Result<HashSet<String>, RecipientPortError> {
Ok(self.answer.lock().unwrap_or_else(|e| e.into_inner()).keys.clone())
}
}
#[derive(Clone)]
pub struct RecipientContextSlot {
inner: Arc<std::sync::RwLock<Arc<dyn RecipientContextPort>>>,
}
impl RecipientContextSlot {
pub fn install(&self, port: Arc<dyn RecipientContextPort>) {
*self.inner.write().unwrap_or_else(|e| e.into_inner()) = port;
}
pub fn current(&self) -> Arc<dyn RecipientContextPort> {
self.inner.read().unwrap_or_else(|e| e.into_inner()).clone()
}
}
impl Default for RecipientContextSlot {
fn default() -> Self {
Self {
inner: Arc::new(std::sync::RwLock::new(Arc::new(RefusingRecipientContext))),
}
}
}
#[async_trait]
impl RecipientContextPort for RecipientContextSlot {
async fn resolve(&self, user_id: Uuid) -> Result<RecipientContext, RecipientPortError> {
self.current().resolve(user_id).await
}
async fn count_connected(
&self,
company_id: Option<Uuid>,
since: DateTime<Utc>,
until: DateTime<Utc>,
) -> Result<i64, RecipientPortError> {
self.current().count_connected(company_id, since, until).await
}
async fn any_logged_in_since(
&self,
user_ids: &[Uuid],
since: DateTime<Utc>,
) -> Result<bool, RecipientPortError> {
self.current().any_logged_in_since(user_ids, since).await
}
async fn group_keys(&self, user_id: Uuid) -> Result<HashSet<String>, RecipientPortError> {
self.current().group_keys(user_id).await
}
}
pub struct SqlRecipientContext {
pool: PgPool,
}
impl SqlRecipientContext {
pub fn new(pool: PgPool) -> Self {
Self { pool }
}
}
#[async_trait]
impl RecipientContextPort for SqlRecipientContext {
async fn resolve(&self, user_id: Uuid) -> Result<RecipientContext, RecipientPortError> {
let row = sqlx::query_as::<_, (String, Option<DateTime<Utc>>, Option<Uuid>)>(
r#"SELECT u.email, u.last_login,
(SELECT ou.org_unit_id
FROM sapiens.organization_users ou
WHERE ou.user_id = u.id AND ou.status = 'active'
AND (ou.metadata->>'deleted_at') IS NULL
ORDER BY ou.joined_at
LIMIT 1) AS company_id
FROM users u
WHERE u.id = $1 AND (u.metadata->>'deleted_at') IS NULL"#,
)
.bind(user_id)
.fetch_optional(&self.pool)
.await
.map_err(|e| RecipientPortError::Unavailable(e.to_string()))?;
let Some((email, last_login, company_id)) = row else {
return Err(RecipientPortError::Unavailable(format!(
"user {user_id} not found (or soft-deleted) in sapiens"
)));
};
Ok(RecipientContext {
user_id,
email,
company_id,
is_internal: company_id.is_some(),
last_login,
})
}
async fn count_connected(
&self,
company_id: Option<Uuid>,
since: DateTime<Utc>,
until: DateTime<Utc>,
) -> Result<i64, RecipientPortError> {
let (n,) = sqlx::query_as::<_, (i64,)>(
r#"SELECT count(DISTINCT u.id)
FROM users u
JOIN sapiens.organization_users ou
ON ou.user_id = u.id
AND ou.status = 'active'
AND (ou.metadata->>'deleted_at') IS NULL
AND ou.org_unit_id = $1
WHERE u.last_login >= $2 AND u.last_login < $3
AND u.status = 'active'
AND (u.metadata->>'deleted_at') IS NULL
AND u.status <> 'suspended'
AND u.status <> 'pending_verification'"#,
)
.bind(company_id)
.bind(since)
.bind(until)
.fetch_one(&self.pool)
.await
.map_err(|e| RecipientPortError::Unavailable(e.to_string()))?;
Ok(n)
}
async fn any_logged_in_since(
&self,
user_ids: &[Uuid],
since: DateTime<Utc>,
) -> Result<bool, RecipientPortError> {
if user_ids.is_empty() {
return Ok(false);
}
let (n,) = sqlx::query_as::<_, (i64,)>(
r#"SELECT count(*) FROM users
WHERE id = ANY($1)
AND last_login IS NOT NULL
AND last_login >= $2"#,
)
.bind(user_ids)
.bind(since)
.fetch_one(&self.pool)
.await
.map_err(|e| RecipientPortError::Unavailable(e.to_string()))?;
Ok(n > 0)
}
async fn group_keys(&self, user_id: Uuid) -> Result<HashSet<String>, RecipientPortError> {
let rows = sqlx::query_as::<_, (String,)>(
r#"SELECT r.name
FROM roles r
JOIN user_roles ur ON ur.role_id = r.id
WHERE ur.user_id = $1
AND (r.metadata->>'deleted_at') IS NULL"#,
)
.bind(user_id)
.fetch_all(&self.pool)
.await
.map_err(|e| RecipientPortError::Unavailable(e.to_string()))?;
Ok(rows.into_iter().map(|(name,)| name).collect())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn unwired_port_refuses_loudly_everywhere() {
let slot = RecipientContextSlot::default();
let u = Uuid::new_v4();
assert!(slot.resolve(u).await.unwrap_err().is_not_wired());
assert!(slot
.count_connected(Some(u), Utc::now(), Utc::now())
.await
.unwrap_err()
.is_not_wired());
assert!(slot
.any_logged_in_since(&[u], Utc::now())
.await
.unwrap_err()
.is_not_wired());
assert!(slot.group_keys(u).await.unwrap_err().is_not_wired());
}
#[tokio::test]
async fn install_replaces_the_refusal() {
let slot = RecipientContextSlot::default();
slot.install(Arc::new(CannedRecipientContext::with(CannedAnswers {
context: Some(RecipientContext {
user_id: Uuid::new_v4(),
email: "someone@example.test".into(),
company_id: Some(Uuid::new_v4()),
is_internal: true,
last_login: None,
}),
..Default::default()
})));
assert!(slot.resolve(Uuid::new_v4()).await.is_ok());
}
}