use async_trait::async_trait;
use sqlx::{PgPool, FromRow, postgres::PgRow, Postgres};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use chrono::NaiveDateTime;
use std::collections::{HashMap, HashSet};
use crate::qualify_relation_table;
use crate::filter::{parse_filters as parse_query_filter};
use crate::filter::SortDirection as FilterSortDirection;
use sqlx::Row as _;
pub trait Entity {
fn id(&self) -> Option<&str>;
fn table_name() -> &'static str where Self: Sized;
fn is_deleted(&self) -> bool { false }
fn created_at(&self) -> Option<NaiveDateTime> { None }
fn updated_at(&self) -> Option<NaiveDateTime> { None }
}
#[derive(Debug, Clone, Default)]
pub struct PaginationParams {
pub page: u32,
pub per_page: u32,
}
impl PaginationParams {
pub fn new(page: u32, per_page: u32) -> Self {
Self {
page: page.max(1),
per_page: per_page.clamp(1, 100), }
}
pub fn offset(&self) -> u32 {
(self.page - 1) * self.per_page
}
pub fn limit(&self) -> u32 {
self.per_page
}
}
#[derive(Debug, Clone, Default)]
pub struct SortParams {
pub field: String,
pub direction: SortDirection,
}
#[derive(Debug, Clone, Default)]
pub enum SortDirection {
#[default]
Asc,
Desc,
}
#[derive(Debug, Clone, Default)]
pub struct FilterParams {
pub conditions: HashMap<String, FilterCondition>,
}
#[derive(Debug, Clone)]
pub enum FilterCondition {
Equals(String),
NotEquals(String),
GreaterThan(String),
LessThan(String),
Like(String),
In(Vec<String>),
IsNull,
IsNotNull,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PaginatedResult<T> {
pub data: Vec<T>,
pub pagination: PaginationInfo,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PaginationInfo {
pub page: u32,
pub per_page: u32,
pub total: u64,
pub total_pages: u32,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub next_cursor: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prev_cursor: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub has_more: Option<bool>,
#[serde(default = "default_count_mode")]
pub count_mode: String,
}
fn default_count_mode() -> String {
"exact".to_string()
}
impl PaginationInfo {
pub fn new(page: u32, per_page: u32, total: u64) -> Self {
let total_pages = ((total as f64) / (per_page as f64)).ceil() as u32;
Self {
page,
per_page,
total,
total_pages,
next_cursor: None,
prev_cursor: None,
has_more: None,
count_mode: default_count_mode(),
}
}
}
#[async_trait]
pub trait DatabaseOperations<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> {
async fn create(&self, entity: &T) -> anyhow::Result<T>;
async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>>;
async fn find_all(&self) -> anyhow::Result<Vec<T>>;
async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>>;
async fn delete(&self, id: &str) -> anyhow::Result<bool>;
async fn count(&self) -> anyhow::Result<u64>;
async fn exists(&self, id: &str) -> anyhow::Result<bool>;
async fn execute_query(&self, query: &str) -> anyhow::Result<u64>;
}
pub struct PostgresRepository<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> {
pool: PgPool,
table_name: String,
_phantom: std::marker::PhantomData<T>,
}
impl<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> PostgresRepository<T> {
pub fn new(pool: PgPool, table_name: &str) -> Self {
Self {
pool,
table_name: table_name.to_string(),
_phantom: std::marker::PhantomData,
}
}
pub fn pool(&self) -> &PgPool {
&self.pool
}
pub fn table_name(&self) -> &str {
&self.table_name
}
pub async fn list_with_filters(
&self,
pagination: PaginationParams,
filters: &HashMap<String, String>,
column_types: &HashMap<String, String>,
search_fields: &[&str],
) -> anyhow::Result<PaginatedResult<T>>
where
T: Send + Sync,
{
let mut query_filter = self.parse_typed_filters(filters, column_types, None).await?;
if !search_fields.is_empty() {
query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
}
self.execute_list(pagination, query_filter).await
}
pub async fn list_with_filters_whitelisted(
&self,
pagination: PaginationParams,
filters: &HashMap<String, String>,
column_types: &HashMap<String, String>,
search_fields: &[&str],
allowed_fields: Option<&HashSet<String>>,
) -> anyhow::Result<PaginatedResult<T>>
where
T: Send + Sync,
{
let mut query_filter =
self.parse_typed_filters(filters, column_types, allowed_fields).await?;
if !search_fields.is_empty() {
query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
}
self.execute_list(pagination, query_filter).await
}
async fn parse_typed_filters(
&self,
filters: &HashMap<String, String>,
column_types: &HashMap<String, String>,
allowed_fields: Option<&HashSet<String>>,
) -> anyhow::Result<crate::QueryFilter> {
let query_filter = parse_query_filter(filters, column_types, allowed_fields)?;
if !query_filter.has_uncast_value_conditions() {
return Ok(query_filter);
}
let catalog = catalog_filter_casts(&self.pool, &self.table_name).await?;
let merged = merge_filter_casts(column_types, catalog);
parse_query_filter(filters, &merged, allowed_fields)
}
#[allow(clippy::type_complexity)]
async fn execute_list(
&self,
pagination: PaginationParams,
mut query_filter: crate::QueryFilter,
) -> anyhow::Result<PaginatedResult<T>> {
let limit = pagination.limit() as i64;
let backwards =
query_filter.cursor_before.is_some() && query_filter.cursor_after.is_none();
let cursor_walk = query_filter.cursor_after.is_some() || backwards;
let (mut where_clause, mut filter_params) = query_filter.build_where_clause();
let order_clause;
let mut boundary_sorts: Vec<(String, FilterSortDirection)> = Vec::new();
let mut boundary_casts: Vec<Option<String>> = Vec::new();
let mut sorts: Vec<(String, FilterSortDirection)> = query_filter
.sorts
.iter()
.map(|s| (s.field.clone(), s.direction.clone()))
.collect();
if cursor_walk || !sorts.is_empty() {
if sorts.is_empty() {
sorts.push(("id".into(), FilterSortDirection::Asc));
} else if sorts.last().map(|(f, _)| f != "id").unwrap_or(true) {
sorts.push(("id".into(), FilterSortDirection::Asc));
}
}
if !sorts.is_empty() {
boundary_casts = self.sort_column_casts(&sorts).await?;
}
if cursor_walk {
let opaque = if backwards {
query_filter.cursor_before.clone().unwrap()
} else {
query_filter.cursor_after.clone().unwrap()
};
let payload = crate::filter::cursor::decode_cursor(&opaque, &sorts)
.map_err(|e| anyhow::anyhow!("cursor refused: {e}"))?;
let mut idx = filter_params.len() + 1;
let (keyset_sql, keyset_params) = crate::filter::cursor::build_keyset_predicate(
&payload,
&mut idx,
&boundary_casts,
backwards,
);
if where_clause.is_empty() {
where_clause = format!(" WHERE {}", keyset_sql);
} else {
where_clause = format!("{} AND ({})", where_clause, keyset_sql);
}
filter_params.extend(keyset_params);
let parts: Vec<String> = sorts
.iter()
.map(|(f, d)| {
let dir = if (*d == FilterSortDirection::Desc) != backwards {
"DESC"
} else {
"ASC"
};
format!("{} {}", f, dir)
})
.collect();
order_clause = format!(" ORDER BY {}", parts.join(", "));
} else {
if sorts.is_empty() {
order_clause = query_filter.build_order_by_clause();
} else {
let parts: Vec<String> = sorts
.iter()
.map(|(f, d)| {
let dir =
if *d == FilterSortDirection::Desc { "DESC" } else { "ASC" };
format!("{} {}", f, dir)
})
.collect();
order_clause = format!(" ORDER BY {}", parts.join(", "));
}
}
boundary_sorts = sorts;
let (total, count_mode) = if query_filter.estimate_total {
(
self.estimate_filtered_rows(&where_clause, &filter_params).await?,
"estimate",
)
} else if cursor_walk {
(0u64, "none")
} else {
let count_query = format!("SELECT COUNT(*) FROM {}{}", self.table_name, where_clause);
let mut count_query_builder = sqlx::query_scalar::<_, i64>(&count_query);
for param in &filter_params {
count_query_builder = count_query_builder.bind(param);
}
(
crate::company_scope::fetch_one_scalar_scoped(&self.pool, count_query_builder)
.await? as u64,
"exact",
)
};
let fetch = limit + 1;
let data_query = if cursor_walk {
format!(
"SELECT * FROM {}{}{} LIMIT {}",
self.table_name, where_clause, order_clause, fetch
)
} else {
format!(
"SELECT * FROM {}{}{} LIMIT {} OFFSET {}",
self.table_name,
where_clause,
order_clause,
fetch,
pagination.offset()
)
};
let mut pagination_info = PaginationInfo::new(pagination.page, pagination.per_page, total);
pagination_info.count_mode = count_mode.to_string();
let mut rows_query = sqlx::query(&data_query);
for param in &filter_params {
rows_query = rows_query.bind(param);
}
let rows: Vec<PgRow> =
crate::company_scope::fetch_all_rows_scoped(&self.pool, rows_query).await?;
let has_more = rows.len() as i64 > limit;
let mut page: Vec<PgRow> = rows.into_iter().take(limit as usize).collect();
if backwards {
page.reverse();
}
let data: anyhow::Result<Vec<T>> = page
.iter()
.map(|row| T::from_row(row).map_err(|e| anyhow::anyhow!("decode row: {e}")))
.collect();
let data = data?;
let deterministic = !boundary_sorts.is_empty();
let next_cursor = if (has_more || backwards) && deterministic {
page.last().and_then(|r| {
let casts = if boundary_casts.is_empty() {
&boundary_casts
} else {
&boundary_casts
};
self.row_cursor(r, &boundary_sorts, casts)
})
} else {
None
};
let prev_cursor = if !page.is_empty() && deterministic {
page.first().and_then(|r| self.row_cursor(r, &boundary_sorts, &boundary_casts))
} else {
None
};
pagination_info.has_more = Some(has_more);
pagination_info.next_cursor = next_cursor;
pagination_info.prev_cursor = prev_cursor;
Ok(PaginatedResult {
data,
pagination: pagination_info,
})
}
async fn estimate_filtered_rows(
&self,
where_clause: &str,
filter_params: &[String],
) -> anyhow::Result<u64> {
let explain = format!("EXPLAIN (FORMAT JSON) SELECT 1 FROM {}{}", self.table_name, where_clause);
let mut builder = sqlx::query_scalar::<_, serde_json::Value>(&explain);
for param in filter_params {
builder = builder.bind(param);
}
let plan: serde_json::Value =
crate::company_scope::fetch_one_scalar_scoped(&self.pool, builder).await?;
let rows = plan
.as_array()
.and_then(|a| a.first())
.and_then(|top| top.get("Plan"))
.and_then(|p| p.get("Plan Rows"))
.and_then(|r| r.as_i64())
.unwrap_or(0);
Ok(rows.max(0) as u64)
}
async fn sort_column_casts(
&self,
sorts: &[(String, FilterSortDirection)],
) -> anyhow::Result<Vec<Option<String>>> {
let (schema, table) = match self.table_name.rsplit_once('.') {
Some((s, t)) => (s.to_string(), t.to_string()),
None => ("public".to_string(), self.table_name.clone()),
};
let mut casts: Vec<Option<String>> = Vec::with_capacity(sorts.len());
for (field, _) in sorts {
let row: Option<(String, String)> = sqlx::query_as(
"SELECT data_type, coalesce(udt_name, '') FROM information_schema.columns \
WHERE table_schema = $1 AND table_name = $2 AND column_name = $3",
)
.bind(&schema)
.bind(&table)
.bind(field)
.fetch_optional(&self.pool)
.await?;
let cast = row.map(|(data_type, udt)| cast_suffix(&data_type, &udt)).flatten();
casts.push(cast);
}
Ok(casts)
}
fn row_cursor(
&self,
row: &PgRow,
sorts: &[(String, FilterSortDirection)],
casts: &[Option<String>],
) -> Option<String> {
let id: uuid::Uuid = row.try_get("id").ok()?;
let mut values: Vec<serde_json::Value> = Vec::with_capacity(sorts.len());
for (i, (field, _)) in sorts.iter().enumerate() {
let field = field.as_str();
let mut data_type = casts.get(i).and_then(|c| c.as_deref()).unwrap_or("");
if field == "id" && data_type.is_empty() {
data_type = "uuid";
}
let text: Option<String> = match data_type {
"numeric" => row
.try_get::<Option<sqlx::types::Decimal>, _>(field)
.ok()?
.map(|d| d.to_string()),
"uuid" => row
.try_get::<Option<uuid::Uuid>, _>(field)
.ok()?
.map(|u| u.to_string()),
"timestamptz" => row
.try_get::<Option<chrono::DateTime<chrono::Utc>>, _>(field)
.ok()?
.map(|t| t.to_rfc3339()),
"integer" | "smallint" => row
.try_get::<Option<i32>, _>(field)
.ok()?
.map(|n| n.to_string()),
"bigint" => row
.try_get::<Option<i64>, _>(field)
.ok()?
.map(|n| n.to_string()),
"boolean" => row
.try_get::<Option<bool>, _>(field)
.ok()?
.map(|b| b.to_string()),
"date" => row
.try_get::<Option<chrono::NaiveDate>, _>(field)
.ok()?
.map(|d| d.to_string()),
_ => row
.try_get::<Option<String>, _>(field)
.ok()?
.filter(|s| !s.is_empty() || data_type.is_empty()),
};
values.push(serde_json::Value::String(text?));
}
crate::filter::cursor::encode_cursor(sorts, &values, &id.to_string()).ok()
}
}
fn cast_suffix(data_type: &str, udt_name: &str) -> Option<String> {
match data_type {
"uuid" => Some("uuid".into()),
"numeric" => Some("numeric".into()),
"integer" => Some("integer".into()),
"smallint" => Some("smallint".into()),
"bigint" => Some("bigint".into()),
"boolean" => Some("boolean".into()),
"date" => Some("date".into()),
"timestamp with time zone" => Some("timestamptz".into()),
"timestamp without time zone" => Some("timestamp".into()),
"USER-DEFINED" if !udt_name.is_empty() => Some(udt_name.to_string()),
_ => None,
}
}
fn filter_cast_for(type_name: &str, typtype: &str) -> Option<String> {
match typtype {
"e" => Some(type_name.to_string()),
"b" => match type_name {
"uuid" | "boolean" | "smallint" | "integer" | "bigint" | "numeric" | "real"
| "double precision" | "date" | "interval" | "inet" | "cidr" | "macaddr" => {
Some(type_name.to_string())
}
"time without time zone" => Some("time".into()),
"time with time zone" => Some("timetz".into()),
"timestamp with time zone" => Some("timestamptz".into()),
"timestamp without time zone" => Some("timestamp".into()),
_ => None,
},
_ => None,
}
}
async fn catalog_filter_casts(
pool: &PgPool,
qualified_table: &str,
) -> anyhow::Result<HashMap<String, String>> {
let q = sqlx::query_as::<Postgres, (String, String, String)>(
"SELECT a.attname::text, format_type(a.atttypid, NULL), t.typtype::text
FROM pg_catalog.pg_attribute a
JOIN pg_catalog.pg_type t ON t.oid = a.atttypid
WHERE a.attrelid = to_regclass($1) AND a.attnum > 0 AND NOT a.attisdropped",
)
.bind(qualified_table.to_string());
let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
Ok(rows
.into_iter()
.filter_map(|(column, type_name, typtype)| {
filter_cast_for(&type_name, &typtype).map(|cast| (column, cast))
})
.collect())
}
fn merge_filter_casts(
hints: &HashMap<String, String>,
mut catalog: HashMap<String, String>,
) -> HashMap<String, String> {
for (column, cast) in hints {
catalog.insert(column.clone(), cast.clone());
}
catalog
}
fn explain_unknown_column(error: sqlx::Error, table: &str) -> anyhow::Error {
let text = error.to_string();
if text.contains("does not exist") && text.contains("column") {
return anyhow::Error::new(error).context(format!(
"insert into {table} named a column that does not exist: the entity serializes a field \
with no matching column. Every serialized field must be a column of the table (rename \
it, map it with #[serde(rename)], or skip it with #[serde(skip)])"
));
}
anyhow::Error::new(error)
}
fn quote_ident(name: &str) -> String {
format!("\"{}\"", name.replace('"', "\"\""))
}
#[async_trait]
impl<T> DatabaseOperations<T> for PostgresRepository<T>
where
T: for<'a> FromRow<'a, PgRow> + Send + Sync + Unpin + Serialize,
{
async fn create(&self, entity: &T) -> anyhow::Result<T> {
let json_value = serde_json::to_value(entity)?;
let json_obj = match json_value {
Value::Object(obj) => obj,
_ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
};
let json_str = serde_json::to_string(&json_obj)?;
let insert_columns: Vec<String> = json_obj.keys().map(|k| quote_ident(k)).collect();
let query = if insert_columns.is_empty() {
format!("INSERT INTO {table} DEFAULT VALUES RETURNING *", table = self.table_name)
} else {
let columns = insert_columns.join(", ");
format!(
r#"
INSERT INTO {table} ({columns})
SELECT {columns} FROM jsonb_populate_record(NULL::{table}, $1::jsonb)
RETURNING *
"#,
table = self.table_name,
columns = columns
)
};
let statement = if insert_columns.is_empty() {
sqlx::query_as::<_, T>(&query)
} else {
sqlx::query_as::<_, T>(&query).bind(&json_str)
};
let result = crate::company_scope::fetch_one_scoped(&self.pool, statement)
.await
.map_err(|e| explain_unknown_column(e, &self.table_name))?;
Ok(result)
}
async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>> {
let query = format!("SELECT * FROM {} WHERE id = $1::uuid", self.table_name);
let result = crate::company_scope::fetch_optional_scoped(
&self.pool,
sqlx::query_as::<Postgres, T>(&query).bind(id),
)
.await?;
Ok(result)
}
async fn find_all(&self) -> anyhow::Result<Vec<T>> {
let query = format!("SELECT * FROM {}", self.table_name);
let results = crate::company_scope::fetch_all_scoped(
&self.pool,
sqlx::query_as::<Postgres, T>(&query),
)
.await?;
Ok(results)
}
async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>> {
let json_value = serde_json::to_value(entity)?;
let json_obj = match json_value {
Value::Object(obj) => obj,
_ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
};
let update_columns: Vec<&String> = json_obj.keys()
.filter(|k| *k != "id")
.collect();
let column_names = update_columns.iter()
.map(|k| quote_ident(k))
.collect::<Vec<_>>()
.join(", ");
let json_str = serde_json::to_string(&json_obj)?;
let query = format!(
r#"
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
RETURNING t.*
"#,
table = self.table_name,
columns = column_names
);
let result = crate::company_scope::fetch_optional_scoped(
&self.pool,
sqlx::query_as::<_, T>(&query).bind(&json_str).bind(id),
)
.await?;
Ok(result)
}
async fn delete(&self, id: &str) -> anyhow::Result<bool> {
let query = format!("DELETE FROM {} WHERE id = $1::uuid", self.table_name);
let result = crate::company_scope::execute_scoped(
&self.pool,
sqlx::query(&query).bind(id),
)
.await?;
Ok(result.rows_affected() > 0)
}
async fn count(&self) -> anyhow::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)
}
async fn exists(&self, id: &str) -> anyhow::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())
}
async fn execute_query(&self, query: &str) -> anyhow::Result<u64> {
let result = crate::company_scope::execute_scoped(
&self.pool,
sqlx::query(query),
)
.await?;
Ok(result.rows_affected())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum AggregateFn {
Sum,
Avg,
Min,
Max,
}
impl AggregateFn {
fn sql(self) -> &'static str {
match self {
AggregateFn::Sum => "SUM",
AggregateFn::Avg => "AVG",
AggregateFn::Min => "MIN",
AggregateFn::Max => "MAX",
}
}
fn label(self) -> &'static str {
match self {
AggregateFn::Sum => "sum",
AggregateFn::Avg => "avg",
AggregateFn::Min => "min",
AggregateFn::Max => "max",
}
}
fn requires_numeric(self) -> bool {
matches!(self, AggregateFn::Sum | AggregateFn::Avg)
}
}
#[derive(Debug, Clone, Default)]
pub struct AggregateSpec {
pub group_by: Option<String>,
pub reductions: Vec<(AggregateFn, String)>,
pub group_limit: usize,
pub label_field: Option<String>,
pub label_relation: Option<(String, String)>,
}
pub const DEFAULT_GROUP_LIMIT: usize = 200;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AggregateGroup {
pub key: Option<String>,
pub label: Option<String>,
pub count: u64,
pub values: HashMap<String, Option<String>>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct AggregateResult {
pub groups: Vec<AggregateGroup>,
pub total: AggregateGroup,
pub truncated: bool,
}
#[derive(Debug, Clone)]
pub struct AggregateFieldError(pub String);
impl std::fmt::Display for AggregateFieldError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
impl std::error::Error for AggregateFieldError {}
fn is_numeric_pg_type(pg_type: &str) -> bool {
let t = pg_type.trim().to_ascii_lowercase();
let t = t.split('(').next().unwrap_or(&t).trim();
matches!(
t,
"numeric" | "decimal" | "money"
| "smallint" | "int2" | "integer" | "int" | "int4" | "bigint" | "int8"
| "real" | "float4" | "double precision" | "float8"
| "smallserial" | "serial" | "bigserial"
)
}
async fn catalog_columns(
pool: &PgPool,
qualified_table: &str,
) -> anyhow::Result<HashMap<String, String>> {
let (schema, table) = match qualified_table.split_once('.') {
Some((s, t)) => (s.to_string(), t.to_string()),
None => ("public".to_string(), qualified_table.to_string()),
};
let q = sqlx::query_as::<Postgres, (String, String)>(
"SELECT column_name, data_type FROM information_schema.columns
WHERE table_schema = $1 AND table_name = $2",
)
.bind(schema)
.bind(table);
let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
Ok(rows.into_iter().collect())
}
fn resolve_column<'a>(
name: &str,
column_types: &'a HashMap<String, String>,
) -> Result<(&'a str, &'a str), AggregateFieldError> {
column_types
.get_key_value(name)
.map(|(k, v)| (k.as_str(), v.as_str()))
.ok_or_else(|| {
AggregateFieldError(format!("unknown field `{name}` — not a column of this entity"))
})
}
impl<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> PostgresRepository<T> {
pub async fn aggregate_with_filters(
&self,
spec: &AggregateSpec,
filters: &HashMap<String, String>,
column_types: &HashMap<String, String>,
search_fields: &[&str],
) -> anyhow::Result<AggregateResult> {
let mut query_filter = self.parse_typed_filters(filters, column_types, None).await?;
if !search_fields.is_empty() {
query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
}
query_filter.limit = None;
query_filter.offset = None;
let (where_clause, filter_params) = query_filter.build_where_clause();
let columns = catalog_columns(&self.pool, &self.table_name).await?;
let mut selects: Vec<String> = Vec::new();
let mut value_keys: Vec<String> = Vec::new();
for (func, field) in &spec.reductions {
let (column, pg_type) = resolve_column(field, &columns)?;
if func.requires_numeric() && !is_numeric_pg_type(pg_type) {
return Err(AggregateFieldError(format!(
"cannot {} `{}`: its type is {} — {} needs a numeric column",
func.label(),
column,
pg_type,
func.label()
))
.into());
}
let key = format!("{}:{}", func.label(), column);
selects.push(format!("{}({})::text AS \"{}\"", func.sql(), column, key));
value_keys.push(key);
}
let group_limit = if spec.group_limit == 0 { DEFAULT_GROUP_LIMIT } else { spec.group_limit };
let reductions = if selects.is_empty() { String::new() } else { format!(", {}", selects.join(", ")) };
let mut label_select = String::new();
let mut label_join = String::new();
if let (Some(field), Some(label_field), Some((rel_table, base_fk))) = (
&spec.group_by,
&spec.label_field,
&spec.label_relation,
) {
let _ = field;
let qualified = qualify_relation_table(&self.table_name, rel_table);
let rel_columns = catalog_columns(&self.pool, &qualified).await?;
let (label_col, _) = resolve_column(label_field, &rel_columns).map_err(|_| {
AggregateFieldError(format!(
"cannot label groups by `{label_field}`: the related table `{qualified}` has no such column"
))
})?;
label_select = format!(", (label_rel.{label_col})::text AS __group_label");
label_join = format!(
" LEFT JOIN {qualified} AS label_rel ON label_rel.id IS NOT DISTINCT FROM {base_fk}"
);
}
let sql = match &spec.group_by {
Some(field) => {
let (column, _) = resolve_column(field, &columns)?;
format!(
"SELECT GROUPING({column}) AS __is_total, ({column})::text AS __group_key{label_select}, \
COUNT(*) AS __count{reductions} \
FROM {table}{label_join}{where_clause} \
GROUP BY GROUPING SETS (({column}), ()) \
ORDER BY __is_total DESC, __count DESC \
LIMIT {limit}",
column = column,
label_select = label_select,
label_join = label_join,
reductions = reductions,
table = self.table_name,
where_clause = where_clause,
limit = group_limit + 2,
)
}
None => format!(
"SELECT 1 AS __is_total, NULL::text AS __group_key, COUNT(*) AS __count{reductions} \
FROM {table}{where_clause}",
reductions = reductions,
table = self.table_name,
where_clause = where_clause,
),
};
let mut builder = sqlx::query(&sql);
for param in &filter_params {
builder = builder.bind(param);
}
let rows = crate::company_scope::fetch_all_rows_scoped(&self.pool, builder).await?;
let read_group = |row: &PgRow| -> AggregateGroup {
use sqlx::Row as _;
let mut values = HashMap::with_capacity(value_keys.len());
for key in &value_keys {
values.insert(key.clone(), row.try_get::<Option<String>, _>(key.as_str()).ok().flatten());
}
AggregateGroup {
key: row.try_get::<Option<String>, _>("__group_key").ok().flatten(),
label: row.try_get::<Option<String>, _>("__group_label").ok().flatten(),
count: row.try_get::<i64, _>("__count").unwrap_or(0).max(0) as u64,
values,
}
};
use sqlx::Row as _;
let mut total: Option<AggregateGroup> = None;
let mut groups: Vec<AggregateGroup> = Vec::new();
for row in &rows {
let is_total = row.try_get::<i32, _>("__is_total").unwrap_or(0) == 1;
if is_total {
total = Some(read_group(row));
} else {
groups.push(read_group(row));
}
}
let truncated = groups.len() > group_limit;
groups.truncate(group_limit);
let total = total.unwrap_or_else(|| AggregateGroup {
key: None,
label: None,
count: 0,
values: value_keys.iter().map(|k| (k.clone(), None)).collect(),
});
Ok(AggregateResult { groups, total, truncated })
}
}
#[cfg(test)]
mod aggregate_field_tests {
use super::*;
fn columns() -> HashMap<String, String> {
[
("status", "text"),
("total", "numeric"),
("qty", "integer"),
("notes", "text"),
]
.iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect()
}
#[test]
fn resolves_only_declared_columns() {
let cols = columns();
assert_eq!(resolve_column("total", &cols).unwrap().0, "total");
assert!(resolve_column("password_hash", &cols).is_err());
}
#[test]
fn rejects_injection_attempts_rather_than_escaping_them() {
let cols = columns();
for probe in [
"total) FROM selling.sales_orders; DROP TABLE users --",
"status\"",
"1=1",
"total, (SELECT password FROM users)",
"",
] {
assert!(
resolve_column(probe, &cols).is_err(),
"`{probe}` must be refused, never escaped into the query"
);
}
}
#[test]
fn returns_the_declared_key_not_the_callers_string() {
let cols = columns();
let (name, _) = resolve_column("total", &cols).unwrap();
assert!(std::ptr::eq(name, cols.get_key_value("total").unwrap().0.as_str()));
}
#[test]
fn sum_and_avg_require_a_numeric_type() {
assert!(AggregateFn::Sum.requires_numeric());
assert!(AggregateFn::Avg.requires_numeric());
assert!(!AggregateFn::Min.requires_numeric());
assert!(!AggregateFn::Max.requires_numeric());
}
#[test]
fn recognises_the_numeric_postgres_types() {
for t in ["numeric", "NUMERIC(14,2)", "integer", "bigint", "double precision", "money"] {
assert!(is_numeric_pg_type(t), "{t} should count as numeric");
}
for t in ["text", "uuid", "timestamptz", "boolean", "jsonb", "USER-DEFINED"] {
assert!(!is_numeric_pg_type(t), "{t} must not accept a SUM");
}
}
}
#[cfg(test)]
mod filter_cast_tests {
use super::*;
#[test]
fn typed_base_columns_get_their_own_cast() {
for (t, want) in [
("boolean", "boolean"),
("integer", "integer"),
("bigint", "bigint"),
("smallint", "smallint"),
("numeric", "numeric"),
("double precision", "double precision"),
("uuid", "uuid"),
("date", "date"),
("time without time zone", "time"),
("timestamp with time zone", "timestamptz"),
("timestamp without time zone", "timestamp"),
] {
assert_eq!(filter_cast_for(t, "b").as_deref(), Some(want), "{t}");
}
}
#[test]
fn text_like_and_composite_columns_keep_the_text_bind() {
for t in ["text", "character varying", "character", "jsonb", "json", "bytea", "text[]", "uuid[]"] {
assert_eq!(filter_cast_for(t, "b"), None, "{t}");
}
assert_eq!(filter_cast_for("approvals.money_amount", "d"), None);
assert_eq!(filter_cast_for("approvals.address", "c"), None);
}
#[test]
fn an_enum_casts_to_its_catalog_name() {
assert_eq!(filter_cast_for("approval_status", "e").as_deref(), Some("approval_status"));
assert_eq!(
filter_cast_for("recruitment.stage_kind", "e").as_deref(),
Some("recruitment.stage_kind")
);
}
#[test]
fn a_generated_hint_wins_over_the_catalog_for_its_column() {
let hints: HashMap<String, String> =
[("id", "uuid"), ("status", "approval_status")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
let catalog: HashMap<String, String> = [
("id", "uuid"),
("status", "approvals.approval_status"),
("folded", "boolean"),
("requested_by", "uuid"),
]
.iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect();
let merged = merge_filter_casts(&hints, catalog);
assert_eq!(merged["status"], "approval_status");
assert_eq!(merged["folded"], "boolean");
assert_eq!(merged["requested_by"], "uuid");
assert_eq!(merged.len(), 4);
}
#[test]
fn only_a_filter_with_an_uncast_comparison_reads_the_catalog() {
let hints: HashMap<String, String> =
[("id", "uuid")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
let parse = |pairs: &[(&str, &str)]| {
let f: HashMap<String, String> =
pairs.iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
parse_query_filter(&f, &hints, None).unwrap().has_uncast_value_conditions()
};
assert!(!parse(&[("id[in]", "a,b"), ("name[contain]", "x"), ("limit", "5")]));
assert!(!parse(&[("deleted_by[isnull]", "1")]));
assert!(parse(&[("folded[eq]", "false")]));
assert!(parse(&[("folded", "false")]));
assert!(parse(&[("scheduled_at[between]", "2026-10-01,2026-10-03")]));
assert!(parse(&[("sequence[or]", "3")]));
}
}