1use async_trait::async_trait;
4use sqlx::{PgPool, FromRow, postgres::PgRow, Postgres};
5use serde::{Deserialize, Serialize};
6use serde_json::Value;
7use chrono::NaiveDateTime;
8use std::collections::{HashMap, HashSet};
9
10use crate::qualify_relation_table;
11use crate::filter::{parse_filters as parse_query_filter};
12use crate::filter::SortDirection as FilterSortDirection;
13use sqlx::Row as _;
14
15pub trait Entity {
17 fn id(&self) -> Option<&str>;
19
20 fn table_name() -> &'static str where Self: Sized;
22
23 fn is_deleted(&self) -> bool { false }
25
26 fn created_at(&self) -> Option<NaiveDateTime> { None }
28
29 fn updated_at(&self) -> Option<NaiveDateTime> { None }
31}
32
33#[derive(Debug, Clone, Default)]
35pub struct PaginationParams {
36 pub page: u32,
37 pub per_page: u32,
38}
39
40impl PaginationParams {
41 pub fn new(page: u32, per_page: u32) -> Self {
42 Self {
43 page: page.max(1),
44 per_page: per_page.clamp(1, 100), }
46 }
47
48 pub fn offset(&self) -> u32 {
49 (self.page - 1) * self.per_page
50 }
51
52 pub fn limit(&self) -> u32 {
53 self.per_page
54 }
55}
56
57#[derive(Debug, Clone, Default)]
59pub struct SortParams {
60 pub field: String,
61 pub direction: SortDirection,
62}
63
64#[derive(Debug, Clone, Default)]
65pub enum SortDirection {
66 #[default]
67 Asc,
68 Desc,
69}
70
71#[derive(Debug, Clone, Default)]
73pub struct FilterParams {
74 pub conditions: HashMap<String, FilterCondition>,
75}
76
77#[derive(Debug, Clone)]
78pub enum FilterCondition {
79 Equals(String),
80 NotEquals(String),
81 GreaterThan(String),
82 LessThan(String),
83 Like(String),
84 In(Vec<String>),
85 IsNull,
86 IsNotNull,
87}
88
89#[derive(Debug, Clone, Serialize, Deserialize)]
91pub struct PaginatedResult<T> {
92 pub data: Vec<T>,
93 pub pagination: PaginationInfo,
94}
95
96#[derive(Debug, Clone, Serialize, Deserialize)]
97pub struct PaginationInfo {
98 pub page: u32,
99 pub per_page: u32,
100 pub total: u64,
101 pub total_pages: u32,
102 #[serde(default, skip_serializing_if = "Option::is_none")]
105 pub next_cursor: Option<String>,
106 #[serde(default, skip_serializing_if = "Option::is_none")]
109 pub prev_cursor: Option<String>,
110 #[serde(default, skip_serializing_if = "Option::is_none")]
113 pub has_more: Option<bool>,
114 #[serde(default = "default_count_mode")]
118 pub count_mode: String,
119}
120
121fn default_count_mode() -> String {
122 "exact".to_string()
123}
124
125impl PaginationInfo {
126 pub fn new(page: u32, per_page: u32, total: u64) -> Self {
127 let total_pages = ((total as f64) / (per_page as f64)).ceil() as u32;
128 Self {
129 page,
130 per_page,
131 total,
132 total_pages,
133 next_cursor: None,
134 prev_cursor: None,
135 has_more: None,
136 count_mode: default_count_mode(),
137 }
138 }
139}
140
141#[async_trait]
143pub trait DatabaseOperations<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> {
144 async fn create(&self, entity: &T) -> anyhow::Result<T>;
146
147 async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>>;
149
150 async fn find_all(&self) -> anyhow::Result<Vec<T>>;
152
153 async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>>;
155
156 async fn delete(&self, id: &str) -> anyhow::Result<bool>;
158
159 async fn count(&self) -> anyhow::Result<u64>;
161
162 async fn exists(&self, id: &str) -> anyhow::Result<bool>;
164
165 async fn execute_query(&self, query: &str) -> anyhow::Result<u64>;
167}
168
169pub struct PostgresRepository<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> {
171 pool: PgPool,
172 table_name: String,
173 _phantom: std::marker::PhantomData<T>,
174}
175
176impl<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> PostgresRepository<T> {
177 pub fn new(pool: PgPool, table_name: &str) -> Self {
178 Self {
179 pool,
180 table_name: table_name.to_string(),
181 _phantom: std::marker::PhantomData,
182 }
183 }
184
185 pub fn pool(&self) -> &PgPool {
186 &self.pool
187 }
188
189 pub fn table_name(&self) -> &str {
190 &self.table_name
191 }
192
193 pub async fn list_with_filters(
249 &self,
250 pagination: PaginationParams,
251 filters: &HashMap<String, String>,
252 column_types: &HashMap<String, String>,
253 search_fields: &[&str],
254 ) -> anyhow::Result<PaginatedResult<T>>
255 where
256 T: Send + Sync,
257 {
258 let mut query_filter = self.parse_typed_filters(filters, column_types, None).await?;
260
261 if !search_fields.is_empty() {
263 query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
264 }
265
266 self.execute_list(pagination, query_filter).await
267 }
268
269 pub async fn list_with_filters_whitelisted(
301 &self,
302 pagination: PaginationParams,
303 filters: &HashMap<String, String>,
304 column_types: &HashMap<String, String>,
305 search_fields: &[&str],
306 allowed_fields: Option<&HashSet<String>>,
307 ) -> anyhow::Result<PaginatedResult<T>>
308 where
309 T: Send + Sync,
310 {
311 let mut query_filter =
313 self.parse_typed_filters(filters, column_types, allowed_fields).await?;
314
315 if !search_fields.is_empty() {
317 query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
318 }
319
320 self.execute_list(pagination, query_filter).await
321 }
322
323 async fn parse_typed_filters(
338 &self,
339 filters: &HashMap<String, String>,
340 column_types: &HashMap<String, String>,
341 allowed_fields: Option<&HashSet<String>>,
342 ) -> anyhow::Result<crate::QueryFilter> {
343 let query_filter = parse_query_filter(filters, column_types, allowed_fields)?;
344 if !query_filter.has_uncast_value_conditions() {
345 return Ok(query_filter);
346 }
347 let catalog = catalog_filter_casts(&self.pool, &self.table_name).await?;
348 let merged = merge_filter_casts(column_types, catalog);
349 parse_query_filter(filters, &merged, allowed_fields)
350 }
351
352 #[allow(clippy::type_complexity)]
366 async fn execute_list(
367 &self,
368 pagination: PaginationParams,
369 query_filter: crate::QueryFilter,
370 ) -> anyhow::Result<PaginatedResult<T>> {
371 let limit = pagination.limit() as i64;
372 let backwards =
373 query_filter.cursor_before.is_some() && query_filter.cursor_after.is_none();
374 let cursor_walk = query_filter.cursor_after.is_some() || backwards;
375
376 let (mut where_clause, mut filter_params) = query_filter.build_where_clause();
377 let order_clause;
378 let mut boundary_casts: Vec<Option<String>> = Vec::new();
380
381 let mut sorts: Vec<(String, FilterSortDirection)> = query_filter
386 .sorts
387 .iter()
388 .map(|s| (s.field.clone(), s.direction.clone()))
389 .collect();
390 if cursor_walk || !sorts.is_empty() {
391 if sorts.is_empty() {
392 sorts.push(("id".into(), FilterSortDirection::Asc));
393 } else if sorts.last().map(|(f, _)| f != "id").unwrap_or(true) {
394 sorts.push(("id".into(), FilterSortDirection::Asc));
395 }
396 }
397
398 if !sorts.is_empty() {
401 boundary_casts = self.sort_column_casts(&sorts).await?;
402 }
403
404 if cursor_walk {
405 let opaque = if backwards {
406 query_filter.cursor_before.clone().unwrap()
407 } else {
408 query_filter.cursor_after.clone().unwrap()
409 };
410 let payload = crate::filter::cursor::decode_cursor(&opaque, &sorts)
411 .map_err(|e| anyhow::anyhow!("cursor refused: {e}"))?;
412 let mut idx = filter_params.len() + 1;
413 let (keyset_sql, keyset_params) = crate::filter::cursor::build_keyset_predicate(
414 &payload,
415 &mut idx,
416 &boundary_casts,
417 backwards,
418 );
419 if where_clause.is_empty() {
420 where_clause = format!(" WHERE {}", keyset_sql);
421 } else {
422 where_clause = format!("{} AND ({})", where_clause, keyset_sql);
423 }
424 filter_params.extend(keyset_params);
425 let parts: Vec<String> = sorts
426 .iter()
427 .map(|(f, d)| {
428 let dir = if (*d == FilterSortDirection::Desc) != backwards {
429 "DESC"
430 } else {
431 "ASC"
432 };
433 format!("{} {}", f, dir)
434 })
435 .collect();
436 order_clause = format!(" ORDER BY {}", parts.join(", "));
437 } else {
438 if sorts.is_empty() {
441 order_clause = query_filter.build_order_by_clause();
442 } else {
443 let parts: Vec<String> = sorts
444 .iter()
445 .map(|(f, d)| {
446 let dir =
447 if *d == FilterSortDirection::Desc { "DESC" } else { "ASC" };
448 format!("{} {}", f, dir)
449 })
450 .collect();
451 order_clause = format!(" ORDER BY {}", parts.join(", "));
452 }
453 }
454 let boundary_sorts: Vec<(String, FilterSortDirection)> = sorts;
456
457 let (total, count_mode) = if query_filter.estimate_total {
460 (
461 self.estimate_filtered_rows(&where_clause, &filter_params).await?,
462 "estimate",
463 )
464 } else if cursor_walk {
465 (0u64, "none")
466 } else {
467 let count_query = format!("SELECT COUNT(*) FROM {}{}", self.table_name, where_clause);
468 let mut count_query_builder = sqlx::query_scalar::<_, i64>(&count_query);
469 for param in &filter_params {
470 count_query_builder = count_query_builder.bind(param);
471 }
472 (
473 crate::company_scope::fetch_one_scalar_scoped(&self.pool, count_query_builder)
474 .await? as u64,
475 "exact",
476 )
477 };
478
479 let fetch = limit + 1;
482 let data_query = if cursor_walk {
483 format!(
484 "SELECT * FROM {}{}{} LIMIT {}",
485 self.table_name, where_clause, order_clause, fetch
486 )
487 } else {
488 format!(
489 "SELECT * FROM {}{}{} LIMIT {} OFFSET {}",
490 self.table_name,
491 where_clause,
492 order_clause,
493 fetch,
494 pagination.offset()
495 )
496 };
497
498 let mut pagination_info = PaginationInfo::new(pagination.page, pagination.per_page, total);
499 pagination_info.count_mode = count_mode.to_string();
500
501 let mut rows_query = sqlx::query(&data_query);
506 for param in &filter_params {
507 rows_query = rows_query.bind(param);
508 }
509 let rows: Vec<PgRow> =
510 crate::company_scope::fetch_all_rows_scoped(&self.pool, rows_query).await?;
511 let has_more = rows.len() as i64 > limit;
512 let mut page: Vec<PgRow> = rows.into_iter().take(limit as usize).collect();
513 if backwards {
514 page.reverse();
515 }
516 let data: anyhow::Result<Vec<T>> = page
517 .iter()
518 .map(|row| T::from_row(row).map_err(|e| anyhow::anyhow!("decode row: {e}")))
519 .collect();
520 let data = data?;
521
522 let deterministic = !boundary_sorts.is_empty();
526 let next_cursor = if (has_more || backwards) && deterministic {
527 page.last().and_then(|r| {
528 let casts = if boundary_casts.is_empty() {
529 &boundary_casts
532 } else {
533 &boundary_casts
534 };
535 self.row_cursor(r, &boundary_sorts, casts)
536 })
537 } else {
538 None
539 };
540 let prev_cursor = if !page.is_empty() && deterministic {
541 page.first().and_then(|r| self.row_cursor(r, &boundary_sorts, &boundary_casts))
542 } else {
543 None
544 };
545
546 pagination_info.has_more = Some(has_more);
547 pagination_info.next_cursor = next_cursor;
548 pagination_info.prev_cursor = prev_cursor;
549
550 Ok(PaginatedResult {
551 data,
552 pagination: pagination_info,
553 })
554 }
555
556 async fn estimate_filtered_rows(
560 &self,
561 where_clause: &str,
562 filter_params: &[String],
563 ) -> anyhow::Result<u64> {
564 let explain = format!("EXPLAIN (FORMAT JSON) SELECT 1 FROM {}{}", self.table_name, where_clause);
565 let mut builder = sqlx::query_scalar::<_, serde_json::Value>(&explain);
566 for param in filter_params {
567 builder = builder.bind(param);
568 }
569 let plan: serde_json::Value =
570 crate::company_scope::fetch_one_scalar_scoped(&self.pool, builder).await?;
571 let rows = plan
572 .as_array()
573 .and_then(|a| a.first())
574 .and_then(|top| top.get("Plan"))
575 .and_then(|p| p.get("Plan Rows"))
576 .and_then(|r| r.as_i64())
577 .unwrap_or(0);
578 Ok(rows.max(0) as u64)
579 }
580
581 async fn sort_column_casts(
586 &self,
587 sorts: &[(String, FilterSortDirection)],
588 ) -> anyhow::Result<Vec<Option<String>>> {
589 let (schema, table) = match self.table_name.rsplit_once('.') {
590 Some((s, t)) => (s.to_string(), t.to_string()),
591 None => ("public".to_string(), self.table_name.clone()),
592 };
593 let mut casts: Vec<Option<String>> = Vec::with_capacity(sorts.len());
594 for (field, _) in sorts {
595 let row: Option<(String, String)> = sqlx::query_as(
596 "SELECT data_type, coalesce(udt_name, '') FROM information_schema.columns \
597 WHERE table_schema = $1 AND table_name = $2 AND column_name = $3",
598 )
599 .bind(&schema)
600 .bind(&table)
601 .bind(field)
602 .fetch_optional(&self.pool)
603 .await?;
604 let cast = row.map(|(data_type, udt)| cast_suffix(&data_type, &udt)).flatten();
605 casts.push(cast);
606 }
607 Ok(casts)
608 }
609
610 fn row_cursor(
614 &self,
615 row: &PgRow,
616 sorts: &[(String, FilterSortDirection)],
617 casts: &[Option<String>],
618 ) -> Option<String> {
619 let id: uuid::Uuid = row.try_get("id").ok()?;
620 let mut values: Vec<serde_json::Value> = Vec::with_capacity(sorts.len());
621 for (i, (field, _)) in sorts.iter().enumerate() {
622 let field = field.as_str();
623 let mut data_type = casts.get(i).and_then(|c| c.as_deref()).unwrap_or("");
624 if field == "id" && data_type.is_empty() {
628 data_type = "uuid";
629 }
630 let text: Option<String> = match data_type {
631 "numeric" => row
632 .try_get::<Option<sqlx::types::Decimal>, _>(field)
633 .ok()?
634 .map(|d| d.to_string()),
635 "uuid" => row
636 .try_get::<Option<uuid::Uuid>, _>(field)
637 .ok()?
638 .map(|u| u.to_string()),
639 "timestamptz" => row
640 .try_get::<Option<chrono::DateTime<chrono::Utc>>, _>(field)
641 .ok()?
642 .map(|t| t.to_rfc3339()),
643 "integer" | "smallint" => row
644 .try_get::<Option<i32>, _>(field)
645 .ok()?
646 .map(|n| n.to_string()),
647 "bigint" => row
648 .try_get::<Option<i64>, _>(field)
649 .ok()?
650 .map(|n| n.to_string()),
651 "boolean" => row
652 .try_get::<Option<bool>, _>(field)
653 .ok()?
654 .map(|b| b.to_string()),
655 "date" => row
656 .try_get::<Option<chrono::NaiveDate>, _>(field)
657 .ok()?
658 .map(|d| d.to_string()),
659 _ => row
660 .try_get::<Option<String>, _>(field)
661 .ok()?
662 .filter(|s| !s.is_empty() || data_type.is_empty()),
663 };
664 values.push(serde_json::Value::String(text?));
665 }
666 crate::filter::cursor::encode_cursor(sorts, &values, &id.to_string()).ok()
667 }
668}
669
670fn cast_suffix(data_type: &str, udt_name: &str) -> Option<String> {
673 match data_type {
674 "uuid" => Some("uuid".into()),
675 "numeric" => Some("numeric".into()),
676 "integer" => Some("integer".into()),
677 "smallint" => Some("smallint".into()),
678 "bigint" => Some("bigint".into()),
679 "boolean" => Some("boolean".into()),
680 "date" => Some("date".into()),
681 "timestamp with time zone" => Some("timestamptz".into()),
682 "timestamp without time zone" => Some("timestamp".into()),
683 "USER-DEFINED" if !udt_name.is_empty() => Some(udt_name.to_string()),
686 _ => None,
687 }
688}
689
690fn filter_cast_for(type_name: &str, typtype: &str) -> Option<String> {
695 match typtype {
696 "e" => Some(type_name.to_string()),
699 "b" => match type_name {
700 "uuid" | "boolean" | "smallint" | "integer" | "bigint" | "numeric" | "real"
701 | "double precision" | "date" | "interval" | "inet" | "cidr" | "macaddr" => {
702 Some(type_name.to_string())
703 }
704 "time without time zone" => Some("time".into()),
705 "time with time zone" => Some("timetz".into()),
706 "timestamp with time zone" => Some("timestamptz".into()),
707 "timestamp without time zone" => Some("timestamp".into()),
708 _ => None,
709 },
710 _ => None,
711 }
712}
713
714async fn catalog_filter_casts(
721 pool: &PgPool,
722 qualified_table: &str,
723) -> anyhow::Result<HashMap<String, String>> {
724 let q = sqlx::query_as::<Postgres, (String, String, String)>(
725 "SELECT a.attname::text, format_type(a.atttypid, NULL), t.typtype::text
726 FROM pg_catalog.pg_attribute a
727 JOIN pg_catalog.pg_type t ON t.oid = a.atttypid
728 WHERE a.attrelid = to_regclass($1) AND a.attnum > 0 AND NOT a.attisdropped",
729 )
730 .bind(qualified_table.to_string());
731 let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
732 Ok(rows
733 .into_iter()
734 .filter_map(|(column, type_name, typtype)| {
735 filter_cast_for(&type_name, &typtype).map(|cast| (column, cast))
736 })
737 .collect())
738}
739
740fn merge_filter_casts(
743 hints: &HashMap<String, String>,
744 mut catalog: HashMap<String, String>,
745) -> HashMap<String, String> {
746 for (column, cast) in hints {
747 catalog.insert(column.clone(), cast.clone());
748 }
749 catalog
750}
751
752fn explain_unknown_column(error: sqlx::Error, table: &str) -> anyhow::Error {
760 let text = error.to_string();
761 if text.contains("does not exist") && text.contains("column") {
762 return anyhow::Error::new(error).context(format!(
763 "insert into {table} named a column that does not exist: the entity serializes a field \
764 with no matching column. Every serialized field must be a column of the table (rename \
765 it, map it with #[serde(rename)], or skip it with #[serde(skip)])"
766 ));
767 }
768 anyhow::Error::new(error)
769}
770
771fn quote_ident(name: &str) -> String {
778 format!("\"{}\"", name.replace('"', "\"\""))
779}
780
781#[async_trait]
782impl<T> DatabaseOperations<T> for PostgresRepository<T>
783where
784 T: for<'a> FromRow<'a, PgRow> + Send + Sync + Unpin + Serialize,
785{
786 async fn create(&self, entity: &T) -> anyhow::Result<T> {
787 let json_value = serde_json::to_value(entity)?;
789
790 let json_obj = match json_value {
791 Value::Object(obj) => obj,
792 _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
793 };
794
795 let json_str = serde_json::to_string(&json_obj)?;
798
799 let insert_columns: Vec<String> = json_obj.keys().map(|k| quote_ident(k)).collect();
813
814 let query = if insert_columns.is_empty() {
815 format!("INSERT INTO {table} DEFAULT VALUES RETURNING *", table = self.table_name)
818 } else {
819 let columns = insert_columns.join(", ");
820 format!(
821 r#"
822 INSERT INTO {table} ({columns})
823 SELECT {columns} FROM jsonb_populate_record(NULL::{table}, $1::jsonb)
824 RETURNING *
825 "#,
826 table = self.table_name,
827 columns = columns
828 )
829 };
830
831 let statement = if insert_columns.is_empty() {
833 sqlx::query_as::<_, T>(&query)
834 } else {
835 sqlx::query_as::<_, T>(&query).bind(&json_str)
836 };
837 let result = crate::company_scope::fetch_one_scoped(&self.pool, statement)
838 .await
839 .map_err(|e| explain_unknown_column(e, &self.table_name))?;
840
841 Ok(result)
842 }
843
844 async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>> {
845 let query = format!("SELECT * FROM {} WHERE id = $1::uuid", self.table_name);
847 let result = crate::company_scope::fetch_optional_scoped(
848 &self.pool,
849 sqlx::query_as::<Postgres, T>(&query).bind(id),
850 )
851 .await?;
852 Ok(result)
853 }
854
855 async fn find_all(&self) -> anyhow::Result<Vec<T>> {
856 let query = format!("SELECT * FROM {}", self.table_name);
857 let results = crate::company_scope::fetch_all_scoped(
858 &self.pool,
859 sqlx::query_as::<Postgres, T>(&query),
860 )
861 .await?;
862 Ok(results)
863 }
864
865 async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>> {
866 let json_value = serde_json::to_value(entity)?;
868
869 let json_obj = match json_value {
870 Value::Object(obj) => obj,
871 _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
872 };
873
874 let update_columns: Vec<&String> = json_obj.keys()
876 .filter(|k| *k != "id")
877 .collect();
878
879 let column_names = update_columns.iter()
880 .map(|k| quote_ident(k))
881 .collect::<Vec<_>>()
882 .join(", ");
883
884 let json_str = serde_json::to_string(&json_obj)?;
885
886 let query = format!(
888 r#"
889 WITH new_row AS (
890 SELECT (jsonb_populate_record(NULL::{table}, $1::jsonb)).*
891 )
892 UPDATE {table} AS t
893 SET ({columns}) = (SELECT {columns} FROM new_row)
894 WHERE t.id = $2::uuid
895 RETURNING t.*
896 "#,
897 table = self.table_name,
898 columns = column_names
899 );
900
901 let result = crate::company_scope::fetch_optional_scoped(
902 &self.pool,
903 sqlx::query_as::<_, T>(&query).bind(&json_str).bind(id),
904 )
905 .await?;
906
907 Ok(result)
908 }
909
910 async fn delete(&self, id: &str) -> anyhow::Result<bool> {
911 let query = format!("DELETE FROM {} WHERE id = $1::uuid", self.table_name);
912 let result = crate::company_scope::execute_scoped(
913 &self.pool,
914 sqlx::query(&query).bind(id),
915 )
916 .await?;
917 Ok(result.rows_affected() > 0)
918 }
919
920 async fn count(&self) -> anyhow::Result<u64> {
921 let query = format!("SELECT COUNT(*) FROM {}", self.table_name);
922 let count = crate::company_scope::fetch_one_scalar_scoped(
923 &self.pool,
924 sqlx::query_scalar::<_, i64>(&query),
925 )
926 .await? as u64;
927 Ok(count)
928 }
929
930 async fn exists(&self, id: &str) -> anyhow::Result<bool> {
931 let query = format!("SELECT 1 FROM {} WHERE id = $1::uuid LIMIT 1", self.table_name);
932 let result = crate::company_scope::fetch_optional_scalar_scoped(
933 &self.pool,
934 sqlx::query_scalar::<_, i32>(&query).bind(id),
935 )
936 .await?;
937 Ok(result.is_some())
938 }
939
940 async fn execute_query(&self, query: &str) -> anyhow::Result<u64> {
941 let result = crate::company_scope::execute_scoped(
942 &self.pool,
943 sqlx::query(query),
944 )
945 .await?;
946 Ok(result.rows_affected())
947 }
948}
949
950#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
954pub enum AggregateFn {
955 Sum,
956 Avg,
957 Min,
958 Max,
959}
960
961impl AggregateFn {
962 fn sql(self) -> &'static str {
963 match self {
964 AggregateFn::Sum => "SUM",
965 AggregateFn::Avg => "AVG",
966 AggregateFn::Min => "MIN",
967 AggregateFn::Max => "MAX",
968 }
969 }
970
971 fn label(self) -> &'static str {
972 match self {
973 AggregateFn::Sum => "sum",
974 AggregateFn::Avg => "avg",
975 AggregateFn::Min => "min",
976 AggregateFn::Max => "max",
977 }
978 }
979
980 fn requires_numeric(self) -> bool {
983 matches!(self, AggregateFn::Sum | AggregateFn::Avg)
984 }
985}
986
987#[derive(Debug, Clone, Default)]
989pub struct AggregateSpec {
990 pub group_by: Option<String>,
992 pub reductions: Vec<(AggregateFn, String)>,
994 pub group_limit: usize,
996 pub label_field: Option<String>,
1000 pub label_relation: Option<(String, String)>,
1004}
1005
1006pub const DEFAULT_GROUP_LIMIT: usize = 200;
1014
1015#[derive(Debug, Clone, Serialize, Deserialize)]
1018pub struct AggregateGroup {
1019 pub key: Option<String>,
1020 pub label: Option<String>,
1023 pub count: u64,
1024 pub values: HashMap<String, Option<String>>,
1032}
1033
1034#[derive(Debug, Clone, Serialize, Deserialize)]
1036pub struct AggregateResult {
1037 pub groups: Vec<AggregateGroup>,
1038 pub total: AggregateGroup,
1039 pub truncated: bool,
1041}
1042
1043#[derive(Debug, Clone)]
1046pub struct AggregateFieldError(pub String);
1047
1048impl std::fmt::Display for AggregateFieldError {
1049 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1050 f.write_str(&self.0)
1051 }
1052}
1053
1054impl std::error::Error for AggregateFieldError {}
1055
1056fn is_numeric_pg_type(pg_type: &str) -> bool {
1058 let t = pg_type.trim().to_ascii_lowercase();
1059 let t = t.split('(').next().unwrap_or(&t).trim();
1060 matches!(
1061 t,
1062 "numeric" | "decimal" | "money"
1063 | "smallint" | "int2" | "integer" | "int" | "int4" | "bigint" | "int8"
1064 | "real" | "float4" | "double precision" | "float8"
1065 | "smallserial" | "serial" | "bigserial"
1066 )
1067}
1068
1069async fn catalog_columns(
1078 pool: &PgPool,
1079 qualified_table: &str,
1080) -> anyhow::Result<HashMap<String, String>> {
1081 let (schema, table) = match qualified_table.split_once('.') {
1082 Some((s, t)) => (s.to_string(), t.to_string()),
1083 None => ("public".to_string(), qualified_table.to_string()),
1084 };
1085 let q = sqlx::query_as::<Postgres, (String, String)>(
1086 "SELECT column_name, data_type FROM information_schema.columns
1087 WHERE table_schema = $1 AND table_name = $2",
1088 )
1089 .bind(schema)
1090 .bind(table);
1091 let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
1092 Ok(rows.into_iter().collect())
1093}
1094
1095fn resolve_column<'a>(
1104 name: &str,
1105 column_types: &'a HashMap<String, String>,
1106) -> Result<(&'a str, &'a str), AggregateFieldError> {
1107 column_types
1108 .get_key_value(name)
1109 .map(|(k, v)| (k.as_str(), v.as_str()))
1110 .ok_or_else(|| {
1111 AggregateFieldError(format!("unknown field `{name}` — not a column of this entity"))
1112 })
1113}
1114
1115impl<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> PostgresRepository<T> {
1116 pub async fn aggregate_with_filters(
1124 &self,
1125 spec: &AggregateSpec,
1126 filters: &HashMap<String, String>,
1127 column_types: &HashMap<String, String>,
1128 search_fields: &[&str],
1129 ) -> anyhow::Result<AggregateResult> {
1130 let mut query_filter = self.parse_typed_filters(filters, column_types, None).await?;
1131 if !search_fields.is_empty() {
1132 query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
1133 }
1134 query_filter.limit = None;
1137 query_filter.offset = None;
1138 let (where_clause, filter_params) = query_filter.build_where_clause();
1139
1140 let columns = catalog_columns(&self.pool, &self.table_name).await?;
1142 let mut selects: Vec<String> = Vec::new();
1143 let mut value_keys: Vec<String> = Vec::new();
1144 for (func, field) in &spec.reductions {
1145 let (column, pg_type) = resolve_column(field, &columns)?;
1146 if func.requires_numeric() && !is_numeric_pg_type(pg_type) {
1147 return Err(AggregateFieldError(format!(
1148 "cannot {} `{}`: its type is {} — {} needs a numeric column",
1149 func.label(),
1150 column,
1151 pg_type,
1152 func.label()
1153 ))
1154 .into());
1155 }
1156 let key = format!("{}:{}", func.label(), column);
1157 selects.push(format!("{}({})::text AS \"{}\"", func.sql(), column, key));
1160 value_keys.push(key);
1161 }
1162
1163 let group_limit = if spec.group_limit == 0 { DEFAULT_GROUP_LIMIT } else { spec.group_limit };
1164 let reductions = if selects.is_empty() { String::new() } else { format!(", {}", selects.join(", ")) };
1165
1166 let mut label_select = String::new();
1172 let mut label_join = String::new();
1173 if let (Some(field), Some(label_field), Some((rel_table, base_fk))) = (
1174 &spec.group_by,
1175 &spec.label_field,
1176 &spec.label_relation,
1177 ) {
1178 let _ = field;
1179 let qualified = qualify_relation_table(&self.table_name, rel_table);
1180 let rel_columns = catalog_columns(&self.pool, &qualified).await?;
1181 let (label_col, _) = resolve_column(label_field, &rel_columns).map_err(|_| {
1182 AggregateFieldError(format!(
1183 "cannot label groups by `{label_field}`: the related table `{qualified}` has no such column"
1184 ))
1185 })?;
1186 label_select = format!(", (label_rel.{label_col})::text AS __group_label");
1187 label_join = format!(
1188 " LEFT JOIN {qualified} AS label_rel ON label_rel.id IS NOT DISTINCT FROM {base_fk}"
1189 );
1190 }
1191
1192 let sql = match &spec.group_by {
1193 Some(field) => {
1194 let (column, _) = resolve_column(field, &columns)?;
1195 format!(
1196 "SELECT GROUPING({column}) AS __is_total, ({column})::text AS __group_key{label_select}, \
1197 COUNT(*) AS __count{reductions} \
1198 FROM {table}{label_join}{where_clause} \
1199 GROUP BY GROUPING SETS (({column}), ()) \
1200 ORDER BY __is_total DESC, __count DESC \
1201 LIMIT {limit}",
1202 column = column,
1203 label_select = label_select,
1204 label_join = label_join,
1205 reductions = reductions,
1206 table = self.table_name,
1207 where_clause = where_clause,
1208 limit = group_limit + 2,
1211 )
1212 }
1213 None => format!(
1214 "SELECT 1 AS __is_total, NULL::text AS __group_key, COUNT(*) AS __count{reductions} \
1215 FROM {table}{where_clause}",
1216 reductions = reductions,
1217 table = self.table_name,
1218 where_clause = where_clause,
1219 ),
1220 };
1221
1222 let mut builder = sqlx::query(&sql);
1223 for param in &filter_params {
1224 builder = builder.bind(param);
1225 }
1226 let rows = crate::company_scope::fetch_all_rows_scoped(&self.pool, builder).await?;
1227
1228 let read_group = |row: &PgRow| -> AggregateGroup {
1229 use sqlx::Row as _;
1230 let mut values = HashMap::with_capacity(value_keys.len());
1231 for key in &value_keys {
1232 values.insert(key.clone(), row.try_get::<Option<String>, _>(key.as_str()).ok().flatten());
1233 }
1234 AggregateGroup {
1235 key: row.try_get::<Option<String>, _>("__group_key").ok().flatten(),
1236 label: row.try_get::<Option<String>, _>("__group_label").ok().flatten(),
1237 count: row.try_get::<i64, _>("__count").unwrap_or(0).max(0) as u64,
1238 values,
1239 }
1240 };
1241
1242 use sqlx::Row as _;
1243 let mut total: Option<AggregateGroup> = None;
1244 let mut groups: Vec<AggregateGroup> = Vec::new();
1245 for row in &rows {
1246 let is_total = row.try_get::<i32, _>("__is_total").unwrap_or(0) == 1;
1247 if is_total {
1248 total = Some(read_group(row));
1250 } else {
1251 groups.push(read_group(row));
1252 }
1253 }
1254
1255 let truncated = groups.len() > group_limit;
1256 groups.truncate(group_limit);
1257
1258 let total = total.unwrap_or_else(|| AggregateGroup {
1261 key: None,
1262 label: None,
1263 count: 0,
1264 values: value_keys.iter().map(|k| (k.clone(), None)).collect(),
1265 });
1266
1267 Ok(AggregateResult { groups, total, truncated })
1268 }
1269}
1270
1271#[cfg(test)]
1272mod aggregate_field_tests {
1273 use super::*;
1274
1275 fn columns() -> HashMap<String, String> {
1276 [
1277 ("status", "text"),
1278 ("total", "numeric"),
1279 ("qty", "integer"),
1280 ("notes", "text"),
1281 ]
1282 .iter()
1283 .map(|(k, v)| (k.to_string(), v.to_string()))
1284 .collect()
1285 }
1286
1287 #[test]
1288 fn resolves_only_declared_columns() {
1289 let cols = columns();
1290 assert_eq!(resolve_column("total", &cols).unwrap().0, "total");
1291 assert!(resolve_column("password_hash", &cols).is_err());
1292 }
1293
1294 #[test]
1297 fn rejects_injection_attempts_rather_than_escaping_them() {
1298 let cols = columns();
1299 for probe in [
1300 "total) FROM selling.sales_orders; DROP TABLE users --",
1301 "status\"",
1302 "1=1",
1303 "total, (SELECT password FROM users)",
1304 "",
1305 ] {
1306 assert!(
1307 resolve_column(probe, &cols).is_err(),
1308 "`{probe}` must be refused, never escaped into the query"
1309 );
1310 }
1311 }
1312
1313 #[test]
1316 fn returns_the_declared_key_not_the_callers_string() {
1317 let cols = columns();
1318 let (name, _) = resolve_column("total", &cols).unwrap();
1319 assert!(std::ptr::eq(name, cols.get_key_value("total").unwrap().0.as_str()));
1320 }
1321
1322 #[test]
1323 fn sum_and_avg_require_a_numeric_type() {
1324 assert!(AggregateFn::Sum.requires_numeric());
1325 assert!(AggregateFn::Avg.requires_numeric());
1326 assert!(!AggregateFn::Min.requires_numeric());
1328 assert!(!AggregateFn::Max.requires_numeric());
1329 }
1330
1331 #[test]
1332 fn recognises_the_numeric_postgres_types() {
1333 for t in ["numeric", "NUMERIC(14,2)", "integer", "bigint", "double precision", "money"] {
1334 assert!(is_numeric_pg_type(t), "{t} should count as numeric");
1335 }
1336 for t in ["text", "uuid", "timestamptz", "boolean", "jsonb", "USER-DEFINED"] {
1337 assert!(!is_numeric_pg_type(t), "{t} must not accept a SUM");
1338 }
1339 }
1340}
1341
1342#[cfg(test)]
1343mod filter_cast_tests {
1344 use super::*;
1345
1346 #[test]
1347 fn typed_base_columns_get_their_own_cast() {
1348 for (t, want) in [
1349 ("boolean", "boolean"),
1350 ("integer", "integer"),
1351 ("bigint", "bigint"),
1352 ("smallint", "smallint"),
1353 ("numeric", "numeric"),
1354 ("double precision", "double precision"),
1355 ("uuid", "uuid"),
1356 ("date", "date"),
1357 ("time without time zone", "time"),
1358 ("timestamp with time zone", "timestamptz"),
1359 ("timestamp without time zone", "timestamp"),
1360 ] {
1361 assert_eq!(filter_cast_for(t, "b").as_deref(), Some(want), "{t}");
1362 }
1363 }
1364
1365 #[test]
1366 fn text_like_and_composite_columns_keep_the_text_bind() {
1367 for t in ["text", "character varying", "character", "jsonb", "json", "bytea", "text[]", "uuid[]"] {
1368 assert_eq!(filter_cast_for(t, "b"), None, "{t}");
1369 }
1370 assert_eq!(filter_cast_for("approvals.money_amount", "d"), None);
1372 assert_eq!(filter_cast_for("approvals.address", "c"), None);
1373 }
1374
1375 #[test]
1376 fn an_enum_casts_to_its_catalog_name() {
1377 assert_eq!(filter_cast_for("approval_status", "e").as_deref(), Some("approval_status"));
1378 assert_eq!(
1379 filter_cast_for("recruitment.stage_kind", "e").as_deref(),
1380 Some("recruitment.stage_kind")
1381 );
1382 }
1383
1384 #[test]
1385 fn a_generated_hint_wins_over_the_catalog_for_its_column() {
1386 let hints: HashMap<String, String> =
1387 [("id", "uuid"), ("status", "approval_status")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1388 let catalog: HashMap<String, String> = [
1389 ("id", "uuid"),
1390 ("status", "approvals.approval_status"),
1391 ("folded", "boolean"),
1392 ("requested_by", "uuid"),
1393 ]
1394 .iter()
1395 .map(|(k, v)| (k.to_string(), v.to_string()))
1396 .collect();
1397 let merged = merge_filter_casts(&hints, catalog);
1398 assert_eq!(merged["status"], "approval_status");
1399 assert_eq!(merged["folded"], "boolean");
1400 assert_eq!(merged["requested_by"], "uuid");
1401 assert_eq!(merged.len(), 4);
1402 }
1403
1404 #[test]
1405 fn only_a_filter_with_an_uncast_comparison_reads_the_catalog() {
1406 let hints: HashMap<String, String> =
1407 [("id", "uuid")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1408 let parse = |pairs: &[(&str, &str)]| {
1409 let f: HashMap<String, String> =
1410 pairs.iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1411 parse_query_filter(&f, &hints, None).unwrap().has_uncast_value_conditions()
1412 };
1413 assert!(!parse(&[("id[in]", "a,b"), ("name[contain]", "x"), ("limit", "5")]));
1414 assert!(!parse(&[("deleted_by[isnull]", "1")]));
1415 assert!(parse(&[("folded[eq]", "false")]));
1416 assert!(parse(&[("folded", "false")]));
1417 assert!(parse(&[("scheduled_at[between]", "2026-10-01,2026-10-03")]));
1418 assert!(parse(&[("sequence[or]", "3")]));
1419 }
1420}