use std::collections::HashMap;
use std::marker::PhantomData;
use anyhow::Result;
use serde::Serialize;
use sqlx::postgres::PgRow;
use sqlx::{FromRow, PgPool};
use uuid::Uuid;
use crate::repository::{
DatabaseOperations, PaginatedResult, PaginationInfo, PaginationParams, PostgresRepository,
};
pub trait EntityRepoMeta {
fn column_types() -> HashMap<String, String>;
fn search_fields() -> &'static [&'static str];
fn private_fields() -> &'static [&'static str] {
&[]
}
fn owner_field() -> Option<&'static str> {
None
}
fn company_field() -> Option<&'static str> {
None
}
fn relations() -> &'static [(&'static str, &'static str, &'static str)] {
&[]
}
}
#[derive(Debug, thiserror::Error)]
#[error("entity is company-scoped ({column}) but the request carries no company scope")]
pub struct MissingCompanyScope {
pub column: &'static str,
}
pub fn company_fence<T: EntityRepoMeta>(company: Option<Uuid>) -> Result<Option<String>, MissingCompanyScope> {
match (T::company_field(), company) {
(None, _) => Ok(None),
(Some(column), None) => Err(MissingCompanyScope { column }),
(Some(column), Some(id)) => Ok(Some(format!("{column} = '{id}'"))),
}
}
pub fn strip_client_company_filters<T: EntityRepoMeta>(filters: &mut HashMap<String, String>) {
let Some(column) = T::company_field() else {
return;
};
let camel = snake_to_camel(column);
filters.retain(|key, _| {
let base = key.split('[').next().unwrap_or(key);
!base.eq_ignore_ascii_case(column) && !base.eq_ignore_ascii_case(&camel)
});
}
pub fn and_conditions(a: Option<&str>, b: Option<String>) -> Option<String> {
match (a, b) {
(None, None) => None,
(Some(a), None) => Some(a.to_string()),
(None, Some(b)) => Some(b),
(Some(a), Some(b)) => Some(format!("{a} AND {b}")),
}
}
fn snake_to_camel(s: &str) -> String {
let mut out = String::with_capacity(s.len());
let mut upper = false;
for c in s.chars() {
if c == '_' {
upper = true;
} else if upper {
out.push(c.to_ascii_uppercase());
upper = false;
} else {
out.push(c);
}
}
out
}
pub struct SoftDelete;
pub struct HardDelete;
pub struct GenericCrudRepository<T, D = SoftDelete>
where
T: for<'r> FromRow<'r, PgRow> + Send + Unpin,
{
inner: PostgresRepository<T>,
_mode: PhantomData<D>,
}
impl<T, D> GenericCrudRepository<T, D>
where
T: for<'r> FromRow<'r, PgRow> + Send + Unpin,
{
pub fn new(pool: PgPool, table_name: &str) -> Self {
Self {
inner: PostgresRepository::new(pool, table_name),
_mode: PhantomData,
}
}
pub fn pool(&self) -> &PgPool {
self.inner.pool()
}
pub fn table_name(&self) -> &str {
self.inner.table_name()
}
pub async fn create(&self, entity: &T) -> Result<T>
where
T: Serialize + Send + Sync,
{
self.inner.create(entity).await
}
pub async fn bulk_create(&self, entities: &[T]) -> Result<Vec<T>>
where
T: Serialize + Send + Sync,
{
let tx = sqlx::pool::Pool::begin(self.pool()).await?;
let mut results = Vec::with_capacity(entities.len());
for entity in entities {
results.push(self.create(entity).await?);
}
tx.commit().await?;
Ok(results)
}
async fn find_by_text_field_with_cond(
&self,
field: &str,
value: &str,
extra: &str,
) -> Result<Option<T>> {
let query = format!(
"SELECT * FROM {} WHERE {} = $1{}",
self.table_name(), field, extra
);
let result = crate::company_scope::fetch_optional_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query).bind(value),
)
.await?;
Ok(result)
}
async fn exists_by_text_field_with_cond(
&self,
field: &str,
value: &str,
extra: &str,
) -> Result<bool> {
let query = format!(
"SELECT 1 FROM {} WHERE {} = $1{} LIMIT 1",
self.table_name(), field, extra
);
let result = crate::company_scope::fetch_optional_scalar_scoped(
self.pool(),
sqlx::query_scalar::<_, i32>(&query).bind(value),
)
.await?;
Ok(result.is_some())
}
async fn find_by_uuid_field_with_cond(
&self,
field: &str,
value: Uuid,
extra: &str,
) -> Result<Option<T>> {
let query = format!(
"SELECT * FROM {} WHERE {} = $1{}",
self.table_name(), field, extra
);
let result = crate::company_scope::fetch_optional_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query).bind(value),
)
.await?;
Ok(result)
}
async fn exists_by_uuid_field_with_cond(
&self,
field: &str,
value: Uuid,
extra: &str,
) -> Result<bool> {
let query = format!(
"SELECT 1 FROM {} WHERE {} = $1{} LIMIT 1",
self.table_name(), field, extra
);
let result = crate::company_scope::fetch_optional_scalar_scoped(
self.pool(),
sqlx::query_scalar::<_, i32>(&query).bind(value),
)
.await?;
Ok(result.is_some())
}
pub async fn run_filtered_query(
&self,
pagination: PaginationParams,
base_condition: Option<&str>,
filters: &HashMap<String, String>,
column_types: &HashMap<String, String>,
search_fields: &[&str],
) -> Result<PaginatedResult<T>>
where
T: Send + Sync,
{
let mut filters_map = filters.clone();
if let Some(cond) = base_condition {
filters_map.insert("__base_condition".to_string(), cond.to_string());
}
self.inner
.list_with_filters(pagination, &filters_map, column_types, search_fields)
.await
}
pub async fn run_aggregate_query(
&self,
spec: &crate::repository::AggregateSpec,
base_condition: Option<&str>,
filters: &HashMap<String, String>,
column_types: &HashMap<String, String>,
search_fields: &[&str],
) -> Result<crate::repository::AggregateResult>
where
T: crate::EntityRepoMeta + Send + Sync,
{
let mut spec = spec.clone();
if let (Some(group), Some(label)) = (&spec.group_by, &spec.label_field) {
let _ = label;
let camel = snake_to_camel(group);
if let Some((_, table, _)) = T::relations().iter().find(|(_, _, fk)| *fk == camel) {
let base_fk = camel_to_snake(&camel);
spec.label_relation = Some((table.to_string(), base_fk));
}
}
let mut filters_map = filters.clone();
if let Some(cond) = base_condition {
filters_map.insert("__base_condition".to_string(), cond.to_string());
}
self.inner
.aggregate_with_filters(&spec, &filters_map, column_types, search_fields)
.await
}
}
impl<T> GenericCrudRepository<T, SoftDelete>
where
T: for<'r> FromRow<'r, PgRow> + Send + Unpin,
{
pub async fn find_by_text_field(&self, field: &str, value: &str) -> Result<Option<T>> {
self.find_by_text_field_with_cond(field, value, " AND metadata->>'deleted_at' IS NULL").await
}
pub async fn exists_by_text_field(&self, field: &str, value: &str) -> Result<bool> {
self.exists_by_text_field_with_cond(field, value, " AND metadata->>'deleted_at' IS NULL").await
}
pub async fn find_by_uuid_field(&self, field: &str, value: Uuid) -> Result<Option<T>> {
self.find_by_uuid_field_with_cond(field, value, " AND metadata->>'deleted_at' IS NULL").await
}
pub async fn exists_by_uuid_field(&self, field: &str, value: Uuid) -> Result<bool> {
self.exists_by_uuid_field_with_cond(field, value, " AND metadata->>'deleted_at' IS NULL").await
}
pub async fn list_paginated_filtered(
&self,
pagination: PaginationParams,
filters: Option<&HashMap<String, String>>,
) -> Result<PaginatedResult<T>>
where
T: EntityRepoMeta + Send + Sync,
{
let filters_map = filters.cloned().unwrap_or_default();
let column_types = T::column_types();
let search_fields_owned: Vec<&'static str> = T::search_fields().iter().copied().collect();
self.run_filtered_query(
pagination,
Some("metadata->>'deleted_at' IS NULL"),
&filters_map,
&column_types,
&search_fields_owned,
).await
}
pub async fn aggregate_filtered(
&self,
spec: &crate::repository::AggregateSpec,
filters: Option<&HashMap<String, String>>,
) -> Result<crate::repository::AggregateResult>
where
T: EntityRepoMeta + Send + Sync,
{
let filters_map = filters.cloned().unwrap_or_default();
let column_types = T::column_types();
let search_fields_owned: Vec<&'static str> = T::search_fields().iter().copied().collect();
self.run_aggregate_query(spec, Some("metadata->>'deleted_at' IS NULL"), &filters_map, &column_types, &search_fields_owned).await
}
pub async fn list_paginated_filtered_scoped(
&self,
pagination: PaginationParams,
filters: Option<&HashMap<String, String>>,
company: Option<Uuid>,
) -> Result<PaginatedResult<T>>
where
T: EntityRepoMeta + Send + Sync,
{
let fence = company_fence::<T>(company)?;
let mut filters_map = filters.cloned().unwrap_or_default();
strip_client_company_filters::<T>(&mut filters_map);
let column_types = T::column_types();
let search_fields_owned: Vec<&'static str> = T::search_fields().iter().copied().collect();
self.run_filtered_query(
pagination,
and_conditions(Some("metadata->>'deleted_at' IS NULL"), fence).as_deref(),
&filters_map,
&column_types,
&search_fields_owned,
)
.await
}
pub async fn list_deleted_filtered_scoped(
&self,
pagination: PaginationParams,
filters: Option<&HashMap<String, String>>,
company: Option<Uuid>,
) -> Result<PaginatedResult<T>>
where
T: EntityRepoMeta + Send + Sync,
{
let fence = company_fence::<T>(company)?;
let mut filters_map = filters.cloned().unwrap_or_default();
strip_client_company_filters::<T>(&mut filters_map);
let column_types = T::column_types();
let empty: &[&str] = &[];
self.run_filtered_query(
pagination,
and_conditions(Some("metadata->>'deleted_at' IS NOT NULL"), fence).as_deref(),
&filters_map,
&column_types,
empty,
)
.await
}
pub async fn list_deleted_filtered(
&self,
pagination: PaginationParams,
filters: Option<&HashMap<String, String>>,
) -> Result<PaginatedResult<T>>
where
T: EntityRepoMeta + Send + Sync,
{
let filters_map = filters.cloned().unwrap_or_default();
let column_types = T::column_types();
let empty: &[&str] = &[];
self.run_filtered_query(
pagination,
Some("metadata->>'deleted_at' IS NOT NULL"),
&filters_map,
&column_types,
empty,
).await
}
pub async fn find_by_id(&self, id: &str) -> Result<Option<T>> {
let query = format!(
"SELECT * FROM {} WHERE id = $1::uuid AND metadata->>'deleted_at' IS NULL",
self.table_name()
);
let result = crate::company_scope::fetch_optional_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query).bind(id),
)
.await?;
Ok(result)
}
pub async fn find_all(&self) -> Result<Vec<T>> {
let query = format!(
"SELECT * FROM {} WHERE metadata->>'deleted_at' IS NULL",
self.table_name()
);
let results = crate::company_scope::fetch_all_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query),
)
.await?;
Ok(results)
}
pub async fn update(&self, id: &str, entity: &T) -> Result<Option<T>>
where
T: Serialize + Send + Sync,
{
if self.find_by_id(id).await?.is_none() {
return Ok(None);
}
self.inner.update(id, entity).await
}
pub async fn delete(&self, id: &str) -> Result<bool> {
self.soft_delete(id).await
}
pub async fn count(&self) -> Result<u64> {
self.count_active().await
}
pub async fn exists(&self, id: &str) -> Result<bool> {
let query = format!(
"SELECT 1 FROM {} WHERE id = $1::uuid AND metadata->>'deleted_at' IS NULL LIMIT 1",
self.table_name()
);
let result = crate::company_scope::fetch_optional_scalar_scoped(
self.pool(),
sqlx::query_scalar::<_, i32>(&query).bind(id),
)
.await?;
Ok(result.is_some())
}
pub async fn list_paginated(&self, pagination: PaginationParams) -> Result<PaginatedResult<T>> {
let offset = pagination.offset();
let limit = pagination.limit();
let query = format!(
"SELECT * FROM {} WHERE metadata->>'deleted_at' IS NULL \
ORDER BY id DESC LIMIT $1 OFFSET $2",
self.table_name()
);
let data = crate::company_scope::fetch_all_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query).bind(limit as i64).bind(offset as i64),
)
.await?;
let total = self.count_active().await?;
Ok(PaginatedResult {
data,
pagination: PaginationInfo::new(pagination.page, pagination.per_page, total),
})
}
pub async fn soft_delete(&self, id: &str) -> Result<bool> {
let query = format!(
"UPDATE {} SET metadata = jsonb_set(\
COALESCE(metadata, '{{}}'), \
'{{deleted_at}}', \
to_jsonb(NOW())\
) WHERE id = $1::uuid AND (metadata->>'deleted_at') IS NULL",
self.table_name()
);
let result = crate::company_scope::execute_scoped(
self.pool(),
sqlx::query(&query).bind(id),
)
.await?;
Ok(result.rows_affected() > 0)
}
pub async fn restore(&self, id: &str) -> Result<Option<T>> {
let query = format!(
"UPDATE {} SET metadata = metadata - 'deleted_at' \
WHERE id = $1::uuid AND (metadata->>'deleted_at') IS NOT NULL \
RETURNING *",
self.table_name()
);
let result = crate::company_scope::fetch_optional_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query).bind(id),
)
.await?;
Ok(result)
}
pub async fn list_deleted(&self, pagination: PaginationParams) -> Result<PaginatedResult<T>> {
let offset = pagination.offset();
let limit = pagination.limit();
let query = format!(
"SELECT * FROM {} WHERE (metadata->>'deleted_at') IS NOT NULL \
ORDER BY (metadata->>'deleted_at') DESC LIMIT $1 OFFSET $2",
self.table_name()
);
let data = crate::company_scope::fetch_all_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query).bind(limit as i64).bind(offset as i64),
)
.await?;
let count_query = format!(
"SELECT COUNT(*) FROM {} WHERE (metadata->>'deleted_at') IS NOT NULL",
self.table_name()
);
let total = crate::company_scope::fetch_one_scalar_scoped(
self.pool(),
sqlx::query_scalar::<_, i64>(&count_query),
)
.await? as u64;
Ok(PaginatedResult {
data,
pagination: PaginationInfo::new(pagination.page, pagination.per_page, total),
})
}
pub async fn empty_trash(&self) -> Result<u64> {
let query = format!(
"DELETE FROM {} WHERE (metadata->>'deleted_at') IS NOT NULL",
self.table_name()
);
let result = crate::company_scope::execute_scoped(self.pool(), sqlx::query(&query)).await?;
Ok(result.rows_affected())
}
pub async fn find_deleted_by_id(&self, id: &str) -> Result<Option<T>> {
let query = format!(
"SELECT * FROM {} WHERE id = $1::uuid AND (metadata->>'deleted_at') IS NOT NULL",
self.table_name()
);
let result = crate::company_scope::fetch_optional_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query).bind(id),
)
.await?;
Ok(result)
}
pub async fn permanent_delete(&self, id: &str) -> Result<bool> {
let query = format!(
"DELETE FROM {} WHERE id = $1::uuid AND (metadata->>'deleted_at') IS NOT NULL",
self.table_name()
);
let result = crate::company_scope::execute_scoped(
self.pool(),
sqlx::query(&query).bind(id),
)
.await?;
Ok(result.rows_affected() > 0)
}
pub async fn count_active(&self) -> Result<u64> {
let query = format!(
"SELECT COUNT(*) FROM {} WHERE (metadata->>'deleted_at') IS NULL",
self.table_name()
);
let count = crate::company_scope::fetch_one_scalar_scoped(
self.pool(),
sqlx::query_scalar::<_, i64>(&query),
)
.await? as u64;
Ok(count)
}
pub async fn count_deleted(&self) -> Result<u64> {
let query = format!(
"SELECT COUNT(*) FROM {} WHERE (metadata->>'deleted_at') IS NOT NULL",
self.table_name()
);
let count = crate::company_scope::fetch_one_scalar_scoped(
self.pool(),
sqlx::query_scalar::<_, i64>(&query),
)
.await? as u64;
Ok(count)
}
pub async fn bulk_soft_delete(&self, ids: &[String]) -> Result<u64> {
if ids.is_empty() {
return Ok(0);
}
let placeholders = id_in_placeholders(ids.len());
let query = format!(
"UPDATE {} SET metadata = jsonb_set(\
COALESCE(metadata, '{{}}'), '{{deleted_at}}', to_jsonb(NOW())\
) WHERE id IN ({placeholders}) AND (metadata->>'deleted_at') IS NULL",
self.table_name()
);
let mut tx = self.pool().begin().await?;
crate::company_scope::bind_current_company(&mut tx).await?;
let mut q = sqlx::query(&query);
for id in ids {
q = q.bind(id);
}
let affected = q.execute(&mut *tx).await?.rows_affected();
if affected != ids.len() as u64 {
return Err(anyhow::anyhow!(
"bulk_soft_delete: {} of {} ids were not active/deletable; rolled back",
ids.len() as u64 - affected,
ids.len()
));
}
tx.commit().await?;
Ok(affected)
}
pub async fn bulk_restore(&self, ids: &[String]) -> Result<Vec<T>> {
if ids.is_empty() {
return Ok(Vec::new());
}
let placeholders = id_in_placeholders(ids.len());
let query = format!(
"UPDATE {} SET metadata = metadata - 'deleted_at' \
WHERE id IN ({placeholders}) AND (metadata->>'deleted_at') IS NOT NULL \
RETURNING *",
self.table_name()
);
let mut tx = self.pool().begin().await?;
crate::company_scope::bind_current_company(&mut tx).await?;
let mut q = sqlx::query_as::<_, T>(&query);
for id in ids {
q = q.bind(id);
}
let rows = q.fetch_all(&mut *tx).await?;
if rows.len() != ids.len() {
return Err(anyhow::anyhow!(
"bulk_restore: {} of {} ids were not in trash; rolled back",
ids.len() - rows.len(),
ids.len()
));
}
tx.commit().await?;
Ok(rows)
}
pub async fn bulk_permanent_delete(&self, ids: &[String]) -> Result<u64> {
if ids.is_empty() {
return Ok(0);
}
let placeholders = id_in_placeholders(ids.len());
let query = format!(
"DELETE FROM {} WHERE id IN ({placeholders}) AND (metadata->>'deleted_at') IS NOT NULL",
self.table_name()
);
let mut tx = self.pool().begin().await?;
crate::company_scope::bind_current_company(&mut tx).await?;
let mut q = sqlx::query(&query);
for id in ids {
q = q.bind(id);
}
let affected = q.execute(&mut *tx).await?.rows_affected();
if affected != ids.len() as u64 {
return Err(anyhow::anyhow!(
"bulk_permanent_delete: {} of {} ids were not in trash; rolled back",
ids.len() as u64 - affected,
ids.len()
));
}
tx.commit().await?;
Ok(affected)
}
pub async fn restore_all(&self) -> Result<Vec<T>> {
let query = format!(
"UPDATE {} SET metadata = metadata - 'deleted_at' \
WHERE (metadata->>'deleted_at') IS NOT NULL \
RETURNING *",
self.table_name()
);
let rows = crate::company_scope::fetch_all_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query),
)
.await?;
Ok(rows)
}
pub async fn bulk_update(&self, entities: &[T]) -> Result<Vec<T>>
where
T: Serialize + Send + Sync,
{
bulk_update_rows(
self.pool(),
self.table_name(),
" AND t.metadata->>'deleted_at' IS NULL",
entities,
)
.await
}
}
fn id_in_placeholders(n: usize) -> String {
(1..=n)
.map(|i| format!("${i}::uuid"))
.collect::<Vec<_>>()
.join(", ")
}
fn build_update_parts<T: Serialize>(entity: &T) -> Result<(String, String, String)> {
let json_value = serde_json::to_value(entity)?;
let json_obj = match json_value {
serde_json::Value::Object(obj) => obj,
_ => return Err(anyhow::anyhow!("entity must serialize to a JSON object")),
};
let id = json_obj
.get("id")
.and_then(|v| v.as_str())
.ok_or_else(|| anyhow::anyhow!("entity missing string 'id' field"))?
.to_string();
let column_names = json_obj
.keys()
.filter(|k| *k != "id")
.map(|k| format!("\"{k}\""))
.collect::<Vec<_>>()
.join(", ");
let json_str = serde_json::to_string(&json_obj)?;
Ok((id, json_str, column_names))
}
impl<T> GenericCrudRepository<T, HardDelete>
where
T: for<'r> FromRow<'r, PgRow> + Send + Sync + Unpin + Serialize,
{
pub async fn find_by_text_field(&self, field: &str, value: &str) -> Result<Option<T>> {
self.find_by_text_field_with_cond(field, value, "").await
}
pub async fn exists_by_text_field(&self, field: &str, value: &str) -> Result<bool> {
self.exists_by_text_field_with_cond(field, value, "").await
}
pub async fn find_by_uuid_field(&self, field: &str, value: Uuid) -> Result<Option<T>> {
self.find_by_uuid_field_with_cond(field, value, "").await
}
pub async fn exists_by_uuid_field(&self, field: &str, value: Uuid) -> Result<bool> {
self.exists_by_uuid_field_with_cond(field, value, "").await
}
pub async fn list_paginated_filtered(
&self,
pagination: PaginationParams,
filters: Option<&HashMap<String, String>>,
) -> Result<PaginatedResult<T>>
where
T: EntityRepoMeta + Send + Sync,
{
let filters_map = filters.cloned().unwrap_or_default();
let column_types = T::column_types();
let search_fields_owned: Vec<&'static str> = T::search_fields().iter().copied().collect();
self.run_filtered_query(pagination, None, &filters_map, &column_types, &search_fields_owned).await
}
pub async fn aggregate_filtered(
&self,
spec: &crate::repository::AggregateSpec,
filters: Option<&HashMap<String, String>>,
) -> Result<crate::repository::AggregateResult>
where
T: EntityRepoMeta + Send + Sync,
{
let filters_map = filters.cloned().unwrap_or_default();
let column_types = T::column_types();
let search_fields_owned: Vec<&'static str> = T::search_fields().iter().copied().collect();
self.run_aggregate_query(spec, None, &filters_map, &column_types, &search_fields_owned).await
}
pub async fn find_by_id(&self, id: &str) -> Result<Option<T>> {
self.inner.find_by_id(id).await
}
pub async fn find_all(&self) -> Result<Vec<T>> {
let query = format!("SELECT * FROM {}", self.table_name());
let results = crate::company_scope::fetch_all_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query),
)
.await?;
Ok(results)
}
pub async fn update(&self, id: &str, entity: &T) -> Result<Option<T>> {
self.inner.update(id, entity).await
}
pub async fn delete(&self, id: &str) -> Result<bool> {
self.inner.delete(id).await
}
pub async fn count(&self) -> Result<u64> {
let query = format!("SELECT COUNT(*) FROM {}", self.table_name());
let count = crate::company_scope::fetch_one_scalar_scoped(
self.pool(),
sqlx::query_scalar::<_, i64>(&query),
)
.await? as u64;
Ok(count)
}
pub async fn exists(&self, id: &str) -> Result<bool> {
let query = format!(
"SELECT 1 FROM {} WHERE id = $1::uuid LIMIT 1",
self.table_name()
);
let result = crate::company_scope::fetch_optional_scalar_scoped(
self.pool(),
sqlx::query_scalar::<_, i32>(&query).bind(id),
)
.await?;
Ok(result.is_some())
}
pub async fn list_paginated(&self, pagination: PaginationParams) -> Result<PaginatedResult<T>> {
let offset = pagination.offset();
let limit = pagination.limit();
let query = format!(
"SELECT * FROM {} ORDER BY id DESC LIMIT $1 OFFSET $2",
self.table_name()
);
let data = crate::company_scope::fetch_all_scoped(
self.pool(),
sqlx::query_as::<_, T>(&query).bind(limit as i64).bind(offset as i64),
)
.await?;
let total = self.count().await?;
Ok(PaginatedResult {
data,
pagination: PaginationInfo::new(pagination.page, pagination.per_page, total),
})
}
pub async fn bulk_delete(&self, ids: &[String]) -> Result<u64> {
if ids.is_empty() {
return Ok(0);
}
let placeholders = id_in_placeholders(ids.len());
let query = format!(
"DELETE FROM {} WHERE id IN ({placeholders})",
self.table_name()
);
let mut tx = self.pool().begin().await?;
crate::company_scope::bind_current_company(&mut tx).await?;
let mut q = sqlx::query(&query);
for id in ids {
q = q.bind(id);
}
let affected = q.execute(&mut *tx).await?.rows_affected();
if affected != ids.len() as u64 {
return Err(anyhow::anyhow!(
"bulk_delete: {} of {} ids not found; rolled back",
ids.len() as u64 - affected,
ids.len()
));
}
tx.commit().await?;
Ok(affected)
}
pub async fn bulk_update(&self, entities: &[T]) -> Result<Vec<T>> {
bulk_update_rows(self.pool(), self.table_name(), "", entities).await
}
}
pub fn qualify_relation_table(caller_table: &str, target_table: &str) -> String {
if target_table.contains('.') || !caller_table.contains('.') {
target_table.to_string()
} else {
let schema = caller_table.split('.').next().unwrap_or(caller_table);
format!("{schema}.{target_table}")
}
}
fn is_undefined_table(err: &anyhow::Error) -> bool {
err.chain()
.filter_map(|cause| cause.downcast_ref::<sqlx::Error>())
.any(|sqlx_err| {
sqlx_err
.as_database_error()
.map(|db| db.code().as_deref() == Some("42P01"))
.unwrap_or(false)
})
}
pub async fn fetch_by_ids_as_json(
pool: &PgPool,
caller_table: &str,
table: &str,
ids: &[String],
) -> Result<Vec<serde_json::Value>> {
if ids.is_empty() {
return Ok(Vec::new());
}
let qualified = qualify_relation_table(caller_table, table);
match fetch_rows_as_json(pool, &qualified, ids).await {
Ok(rows) => Ok(rows),
Err(err) if qualified != table && is_undefined_table(&err) => {
fetch_rows_as_json(pool, table, ids).await
}
Err(err) => Err(err),
}
}
async fn fetch_rows_as_json(
pool: &PgPool,
table: &str,
ids: &[String],
) -> Result<Vec<serde_json::Value>> {
let query = format!("SELECT row_to_json(t) AS j FROM {table} t WHERE t.id = ANY($1::uuid[])");
let rows: Vec<(serde_json::Value,)> =
crate::company_scope::fetch_all_scoped(pool, sqlx::query_as(&query).bind(ids)).await?;
Ok(rows.into_iter().map(|(j,)| j).collect())
}
async fn bulk_update_rows<T>(
pool: &PgPool,
table: &str,
active_guard: &str,
entities: &[T],
) -> Result<Vec<T>>
where
T: for<'r> FromRow<'r, PgRow> + Send + Sync + Unpin + Serialize,
{
if entities.is_empty() {
return Ok(Vec::new());
}
let mut tx = pool.begin().await?;
crate::company_scope::bind_current_company(&mut tx).await?;
let mut out = Vec::with_capacity(entities.len());
for entity in entities {
let (id, json_str, column_names) = build_update_parts(entity)?;
let query = format!(
"WITH new_row AS (\
SELECT (jsonb_populate_record(NULL::{table}, $1::jsonb)).*\
) UPDATE {table} AS t \
SET ({columns}) = (SELECT {columns} FROM new_row) \
WHERE t.id = $2::uuid{guard} \
RETURNING t.*",
table = table,
columns = column_names,
guard = active_guard,
);
let updated = sqlx::query_as::<_, T>(&query)
.bind(&json_str)
.bind(&id)
.fetch_optional(&mut *tx)
.await?;
match updated {
Some(e) => out.push(e),
None => {
return Err(anyhow::anyhow!(
"bulk_update: id '{id}' not found or already deleted; rolled back"
));
}
}
}
tx.commit().await?;
Ok(out)
}
fn camel_to_snake(s: &str) -> String {
let mut out = String::with_capacity(s.len() + 4);
for ch in s.chars() {
if ch.is_uppercase() {
out.push('_');
out.extend(ch.to_lowercase());
} else {
out.push(ch);
}
}
out
}