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 mut 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_sorts: Vec<(String, FilterSortDirection)> = Vec::new();
380 let mut boundary_casts: Vec<Option<String>> = Vec::new();
382
383 let mut sorts: Vec<(String, FilterSortDirection)> = query_filter
388 .sorts
389 .iter()
390 .map(|s| (s.field.clone(), s.direction.clone()))
391 .collect();
392 if cursor_walk || !sorts.is_empty() {
393 if sorts.is_empty() {
394 sorts.push(("id".into(), FilterSortDirection::Asc));
395 } else if sorts.last().map(|(f, _)| f != "id").unwrap_or(true) {
396 sorts.push(("id".into(), FilterSortDirection::Asc));
397 }
398 }
399
400 if !sorts.is_empty() {
403 boundary_casts = self.sort_column_casts(&sorts).await?;
404 }
405
406 if cursor_walk {
407 let opaque = if backwards {
408 query_filter.cursor_before.clone().unwrap()
409 } else {
410 query_filter.cursor_after.clone().unwrap()
411 };
412 let payload = crate::filter::cursor::decode_cursor(&opaque, &sorts)
413 .map_err(|e| anyhow::anyhow!("cursor refused: {e}"))?;
414 let mut idx = filter_params.len() + 1;
415 let (keyset_sql, keyset_params) = crate::filter::cursor::build_keyset_predicate(
416 &payload,
417 &mut idx,
418 &boundary_casts,
419 backwards,
420 );
421 if where_clause.is_empty() {
422 where_clause = format!(" WHERE {}", keyset_sql);
423 } else {
424 where_clause = format!("{} AND ({})", where_clause, keyset_sql);
425 }
426 filter_params.extend(keyset_params);
427 let parts: Vec<String> = sorts
428 .iter()
429 .map(|(f, d)| {
430 let dir = if (*d == FilterSortDirection::Desc) != backwards {
431 "DESC"
432 } else {
433 "ASC"
434 };
435 format!("{} {}", f, dir)
436 })
437 .collect();
438 order_clause = format!(" ORDER BY {}", parts.join(", "));
439 } else {
440 if sorts.is_empty() {
443 order_clause = query_filter.build_order_by_clause();
444 } else {
445 let parts: Vec<String> = sorts
446 .iter()
447 .map(|(f, d)| {
448 let dir =
449 if *d == FilterSortDirection::Desc { "DESC" } else { "ASC" };
450 format!("{} {}", f, dir)
451 })
452 .collect();
453 order_clause = format!(" ORDER BY {}", parts.join(", "));
454 }
455 }
456 boundary_sorts = sorts;
457
458 let (total, count_mode) = if query_filter.estimate_total {
461 (
462 self.estimate_filtered_rows(&where_clause, &filter_params).await?,
463 "estimate",
464 )
465 } else if cursor_walk {
466 (0u64, "none")
467 } else {
468 let count_query = format!("SELECT COUNT(*) FROM {}{}", self.table_name, where_clause);
469 let mut count_query_builder = sqlx::query_scalar::<_, i64>(&count_query);
470 for param in &filter_params {
471 count_query_builder = count_query_builder.bind(param);
472 }
473 (
474 crate::company_scope::fetch_one_scalar_scoped(&self.pool, count_query_builder)
475 .await? as u64,
476 "exact",
477 )
478 };
479
480 let fetch = limit + 1;
483 let data_query = if cursor_walk {
484 format!(
485 "SELECT * FROM {}{}{} LIMIT {}",
486 self.table_name, where_clause, order_clause, fetch
487 )
488 } else {
489 format!(
490 "SELECT * FROM {}{}{} LIMIT {} OFFSET {}",
491 self.table_name,
492 where_clause,
493 order_clause,
494 fetch,
495 pagination.offset()
496 )
497 };
498
499 let mut pagination_info = PaginationInfo::new(pagination.page, pagination.per_page, total);
500 pagination_info.count_mode = count_mode.to_string();
501
502 let mut rows_query = sqlx::query(&data_query);
507 for param in &filter_params {
508 rows_query = rows_query.bind(param);
509 }
510 let rows: Vec<PgRow> =
511 crate::company_scope::fetch_all_rows_scoped(&self.pool, rows_query).await?;
512 let has_more = rows.len() as i64 > limit;
513 let mut page: Vec<PgRow> = rows.into_iter().take(limit as usize).collect();
514 if backwards {
515 page.reverse();
516 }
517 let data: anyhow::Result<Vec<T>> = page
518 .iter()
519 .map(|row| T::from_row(row).map_err(|e| anyhow::anyhow!("decode row: {e}")))
520 .collect();
521 let data = data?;
522
523 let deterministic = !boundary_sorts.is_empty();
527 let next_cursor = if (has_more || backwards) && deterministic {
528 page.last().and_then(|r| {
529 let casts = if boundary_casts.is_empty() {
530 &boundary_casts
533 } else {
534 &boundary_casts
535 };
536 self.row_cursor(r, &boundary_sorts, casts)
537 })
538 } else {
539 None
540 };
541 let prev_cursor = if !page.is_empty() && deterministic {
542 page.first().and_then(|r| self.row_cursor(r, &boundary_sorts, &boundary_casts))
543 } else {
544 None
545 };
546
547 pagination_info.has_more = Some(has_more);
548 pagination_info.next_cursor = next_cursor;
549 pagination_info.prev_cursor = prev_cursor;
550
551 Ok(PaginatedResult {
552 data,
553 pagination: pagination_info,
554 })
555 }
556
557 async fn estimate_filtered_rows(
561 &self,
562 where_clause: &str,
563 filter_params: &[String],
564 ) -> anyhow::Result<u64> {
565 let explain = format!("EXPLAIN (FORMAT JSON) SELECT 1 FROM {}{}", self.table_name, where_clause);
566 let mut builder = sqlx::query_scalar::<_, serde_json::Value>(&explain);
567 for param in filter_params {
568 builder = builder.bind(param);
569 }
570 let plan: serde_json::Value =
571 crate::company_scope::fetch_one_scalar_scoped(&self.pool, builder).await?;
572 let rows = plan
573 .as_array()
574 .and_then(|a| a.first())
575 .and_then(|top| top.get("Plan"))
576 .and_then(|p| p.get("Plan Rows"))
577 .and_then(|r| r.as_i64())
578 .unwrap_or(0);
579 Ok(rows.max(0) as u64)
580 }
581
582 async fn sort_column_casts(
587 &self,
588 sorts: &[(String, FilterSortDirection)],
589 ) -> anyhow::Result<Vec<Option<String>>> {
590 let (schema, table) = match self.table_name.rsplit_once('.') {
591 Some((s, t)) => (s.to_string(), t.to_string()),
592 None => ("public".to_string(), self.table_name.clone()),
593 };
594 let mut casts: Vec<Option<String>> = Vec::with_capacity(sorts.len());
595 for (field, _) in sorts {
596 let row: Option<(String, String)> = sqlx::query_as(
597 "SELECT data_type, coalesce(udt_name, '') FROM information_schema.columns \
598 WHERE table_schema = $1 AND table_name = $2 AND column_name = $3",
599 )
600 .bind(&schema)
601 .bind(&table)
602 .bind(field)
603 .fetch_optional(&self.pool)
604 .await?;
605 let cast = row.map(|(data_type, udt)| cast_suffix(&data_type, &udt)).flatten();
606 casts.push(cast);
607 }
608 Ok(casts)
609 }
610
611 fn row_cursor(
615 &self,
616 row: &PgRow,
617 sorts: &[(String, FilterSortDirection)],
618 casts: &[Option<String>],
619 ) -> Option<String> {
620 let id: uuid::Uuid = row.try_get("id").ok()?;
621 let mut values: Vec<serde_json::Value> = Vec::with_capacity(sorts.len());
622 for (i, (field, _)) in sorts.iter().enumerate() {
623 let field = field.as_str();
624 let mut data_type = casts.get(i).and_then(|c| c.as_deref()).unwrap_or("");
625 if field == "id" && data_type.is_empty() {
629 data_type = "uuid";
630 }
631 let text: Option<String> = match data_type {
632 "numeric" => row
633 .try_get::<Option<sqlx::types::Decimal>, _>(field)
634 .ok()?
635 .map(|d| d.to_string()),
636 "uuid" => row
637 .try_get::<Option<uuid::Uuid>, _>(field)
638 .ok()?
639 .map(|u| u.to_string()),
640 "timestamptz" => row
641 .try_get::<Option<chrono::DateTime<chrono::Utc>>, _>(field)
642 .ok()?
643 .map(|t| t.to_rfc3339()),
644 "integer" | "smallint" => row
645 .try_get::<Option<i32>, _>(field)
646 .ok()?
647 .map(|n| n.to_string()),
648 "bigint" => row
649 .try_get::<Option<i64>, _>(field)
650 .ok()?
651 .map(|n| n.to_string()),
652 "boolean" => row
653 .try_get::<Option<bool>, _>(field)
654 .ok()?
655 .map(|b| b.to_string()),
656 "date" => row
657 .try_get::<Option<chrono::NaiveDate>, _>(field)
658 .ok()?
659 .map(|d| d.to_string()),
660 _ => row
661 .try_get::<Option<String>, _>(field)
662 .ok()?
663 .filter(|s| !s.is_empty() || data_type.is_empty()),
664 };
665 values.push(serde_json::Value::String(text?));
666 }
667 crate::filter::cursor::encode_cursor(sorts, &values, &id.to_string()).ok()
668 }
669}
670
671fn cast_suffix(data_type: &str, udt_name: &str) -> Option<String> {
674 match data_type {
675 "uuid" => Some("uuid".into()),
676 "numeric" => Some("numeric".into()),
677 "integer" => Some("integer".into()),
678 "smallint" => Some("smallint".into()),
679 "bigint" => Some("bigint".into()),
680 "boolean" => Some("boolean".into()),
681 "date" => Some("date".into()),
682 "timestamp with time zone" => Some("timestamptz".into()),
683 "timestamp without time zone" => Some("timestamp".into()),
684 "USER-DEFINED" if !udt_name.is_empty() => Some(udt_name.to_string()),
687 _ => None,
688 }
689}
690
691fn filter_cast_for(type_name: &str, typtype: &str) -> Option<String> {
696 match typtype {
697 "e" => Some(type_name.to_string()),
700 "b" => match type_name {
701 "uuid" | "boolean" | "smallint" | "integer" | "bigint" | "numeric" | "real"
702 | "double precision" | "date" | "interval" | "inet" | "cidr" | "macaddr" => {
703 Some(type_name.to_string())
704 }
705 "time without time zone" => Some("time".into()),
706 "time with time zone" => Some("timetz".into()),
707 "timestamp with time zone" => Some("timestamptz".into()),
708 "timestamp without time zone" => Some("timestamp".into()),
709 _ => None,
710 },
711 _ => None,
712 }
713}
714
715async fn catalog_filter_casts(
722 pool: &PgPool,
723 qualified_table: &str,
724) -> anyhow::Result<HashMap<String, String>> {
725 let q = sqlx::query_as::<Postgres, (String, String, String)>(
726 "SELECT a.attname::text, format_type(a.atttypid, NULL), t.typtype::text
727 FROM pg_catalog.pg_attribute a
728 JOIN pg_catalog.pg_type t ON t.oid = a.atttypid
729 WHERE a.attrelid = to_regclass($1) AND a.attnum > 0 AND NOT a.attisdropped",
730 )
731 .bind(qualified_table.to_string());
732 let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
733 Ok(rows
734 .into_iter()
735 .filter_map(|(column, type_name, typtype)| {
736 filter_cast_for(&type_name, &typtype).map(|cast| (column, cast))
737 })
738 .collect())
739}
740
741fn merge_filter_casts(
744 hints: &HashMap<String, String>,
745 mut catalog: HashMap<String, String>,
746) -> HashMap<String, String> {
747 for (column, cast) in hints {
748 catalog.insert(column.clone(), cast.clone());
749 }
750 catalog
751}
752
753fn explain_unknown_column(error: sqlx::Error, table: &str) -> anyhow::Error {
761 let text = error.to_string();
762 if text.contains("does not exist") && text.contains("column") {
763 return anyhow::Error::new(error).context(format!(
764 "insert into {table} named a column that does not exist: the entity serializes a field \
765 with no matching column. Every serialized field must be a column of the table (rename \
766 it, map it with #[serde(rename)], or skip it with #[serde(skip)])"
767 ));
768 }
769 anyhow::Error::new(error)
770}
771
772fn quote_ident(name: &str) -> String {
779 format!("\"{}\"", name.replace('"', "\"\""))
780}
781
782#[async_trait]
783impl<T> DatabaseOperations<T> for PostgresRepository<T>
784where
785 T: for<'a> FromRow<'a, PgRow> + Send + Sync + Unpin + Serialize,
786{
787 async fn create(&self, entity: &T) -> anyhow::Result<T> {
788 let json_value = serde_json::to_value(entity)?;
790
791 let json_obj = match json_value {
792 Value::Object(obj) => obj,
793 _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
794 };
795
796 let json_str = serde_json::to_string(&json_obj)?;
799
800 let insert_columns: Vec<String> = json_obj.keys().map(|k| quote_ident(k)).collect();
814
815 let query = if insert_columns.is_empty() {
816 format!("INSERT INTO {table} DEFAULT VALUES RETURNING *", table = self.table_name)
819 } else {
820 let columns = insert_columns.join(", ");
821 format!(
822 r#"
823 INSERT INTO {table} ({columns})
824 SELECT {columns} FROM jsonb_populate_record(NULL::{table}, $1::jsonb)
825 RETURNING *
826 "#,
827 table = self.table_name,
828 columns = columns
829 )
830 };
831
832 let statement = if insert_columns.is_empty() {
834 sqlx::query_as::<_, T>(&query)
835 } else {
836 sqlx::query_as::<_, T>(&query).bind(&json_str)
837 };
838 let result = crate::company_scope::fetch_one_scoped(&self.pool, statement)
839 .await
840 .map_err(|e| explain_unknown_column(e, &self.table_name))?;
841
842 Ok(result)
843 }
844
845 async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>> {
846 let query = format!("SELECT * FROM {} WHERE id = $1::uuid", self.table_name);
848 let result = crate::company_scope::fetch_optional_scoped(
849 &self.pool,
850 sqlx::query_as::<Postgres, T>(&query).bind(id),
851 )
852 .await?;
853 Ok(result)
854 }
855
856 async fn find_all(&self) -> anyhow::Result<Vec<T>> {
857 let query = format!("SELECT * FROM {}", self.table_name);
858 let results = crate::company_scope::fetch_all_scoped(
859 &self.pool,
860 sqlx::query_as::<Postgres, T>(&query),
861 )
862 .await?;
863 Ok(results)
864 }
865
866 async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>> {
867 let json_value = serde_json::to_value(entity)?;
869
870 let json_obj = match json_value {
871 Value::Object(obj) => obj,
872 _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
873 };
874
875 let update_columns: Vec<&String> = json_obj.keys()
877 .filter(|k| *k != "id")
878 .collect();
879
880 let column_names = update_columns.iter()
881 .map(|k| quote_ident(k))
882 .collect::<Vec<_>>()
883 .join(", ");
884
885 let json_str = serde_json::to_string(&json_obj)?;
886
887 let query = format!(
889 r#"
890 WITH new_row AS (
891 SELECT (jsonb_populate_record(NULL::{table}, $1::jsonb)).*
892 )
893 UPDATE {table} AS t
894 SET ({columns}) = (SELECT {columns} FROM new_row)
895 WHERE t.id = $2::uuid
896 RETURNING t.*
897 "#,
898 table = self.table_name,
899 columns = column_names
900 );
901
902 let result = crate::company_scope::fetch_optional_scoped(
903 &self.pool,
904 sqlx::query_as::<_, T>(&query).bind(&json_str).bind(id),
905 )
906 .await?;
907
908 Ok(result)
909 }
910
911 async fn delete(&self, id: &str) -> anyhow::Result<bool> {
912 let query = format!("DELETE FROM {} WHERE id = $1::uuid", self.table_name);
913 let result = crate::company_scope::execute_scoped(
914 &self.pool,
915 sqlx::query(&query).bind(id),
916 )
917 .await?;
918 Ok(result.rows_affected() > 0)
919 }
920
921 async fn count(&self) -> anyhow::Result<u64> {
922 let query = format!("SELECT COUNT(*) FROM {}", self.table_name);
923 let count = crate::company_scope::fetch_one_scalar_scoped(
924 &self.pool,
925 sqlx::query_scalar::<_, i64>(&query),
926 )
927 .await? as u64;
928 Ok(count)
929 }
930
931 async fn exists(&self, id: &str) -> anyhow::Result<bool> {
932 let query = format!("SELECT 1 FROM {} WHERE id = $1::uuid LIMIT 1", self.table_name);
933 let result = crate::company_scope::fetch_optional_scalar_scoped(
934 &self.pool,
935 sqlx::query_scalar::<_, i32>(&query).bind(id),
936 )
937 .await?;
938 Ok(result.is_some())
939 }
940
941 async fn execute_query(&self, query: &str) -> anyhow::Result<u64> {
942 let result = crate::company_scope::execute_scoped(
943 &self.pool,
944 sqlx::query(query),
945 )
946 .await?;
947 Ok(result.rows_affected())
948 }
949}
950
951#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
955pub enum AggregateFn {
956 Sum,
957 Avg,
958 Min,
959 Max,
960}
961
962impl AggregateFn {
963 fn sql(self) -> &'static str {
964 match self {
965 AggregateFn::Sum => "SUM",
966 AggregateFn::Avg => "AVG",
967 AggregateFn::Min => "MIN",
968 AggregateFn::Max => "MAX",
969 }
970 }
971
972 fn label(self) -> &'static str {
973 match self {
974 AggregateFn::Sum => "sum",
975 AggregateFn::Avg => "avg",
976 AggregateFn::Min => "min",
977 AggregateFn::Max => "max",
978 }
979 }
980
981 fn requires_numeric(self) -> bool {
984 matches!(self, AggregateFn::Sum | AggregateFn::Avg)
985 }
986}
987
988#[derive(Debug, Clone, Default)]
990pub struct AggregateSpec {
991 pub group_by: Option<String>,
993 pub reductions: Vec<(AggregateFn, String)>,
995 pub group_limit: usize,
997 pub label_field: Option<String>,
1001 pub label_relation: Option<(String, String)>,
1005}
1006
1007pub const DEFAULT_GROUP_LIMIT: usize = 200;
1015
1016#[derive(Debug, Clone, Serialize, Deserialize)]
1019pub struct AggregateGroup {
1020 pub key: Option<String>,
1021 pub label: Option<String>,
1024 pub count: u64,
1025 pub values: HashMap<String, Option<String>>,
1033}
1034
1035#[derive(Debug, Clone, Serialize, Deserialize)]
1037pub struct AggregateResult {
1038 pub groups: Vec<AggregateGroup>,
1039 pub total: AggregateGroup,
1040 pub truncated: bool,
1042}
1043
1044#[derive(Debug, Clone)]
1047pub struct AggregateFieldError(pub String);
1048
1049impl std::fmt::Display for AggregateFieldError {
1050 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1051 f.write_str(&self.0)
1052 }
1053}
1054
1055impl std::error::Error for AggregateFieldError {}
1056
1057fn is_numeric_pg_type(pg_type: &str) -> bool {
1059 let t = pg_type.trim().to_ascii_lowercase();
1060 let t = t.split('(').next().unwrap_or(&t).trim();
1061 matches!(
1062 t,
1063 "numeric" | "decimal" | "money"
1064 | "smallint" | "int2" | "integer" | "int" | "int4" | "bigint" | "int8"
1065 | "real" | "float4" | "double precision" | "float8"
1066 | "smallserial" | "serial" | "bigserial"
1067 )
1068}
1069
1070async fn catalog_columns(
1079 pool: &PgPool,
1080 qualified_table: &str,
1081) -> anyhow::Result<HashMap<String, String>> {
1082 let (schema, table) = match qualified_table.split_once('.') {
1083 Some((s, t)) => (s.to_string(), t.to_string()),
1084 None => ("public".to_string(), qualified_table.to_string()),
1085 };
1086 let q = sqlx::query_as::<Postgres, (String, String)>(
1087 "SELECT column_name, data_type FROM information_schema.columns
1088 WHERE table_schema = $1 AND table_name = $2",
1089 )
1090 .bind(schema)
1091 .bind(table);
1092 let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
1093 Ok(rows.into_iter().collect())
1094}
1095
1096fn resolve_column<'a>(
1105 name: &str,
1106 column_types: &'a HashMap<String, String>,
1107) -> Result<(&'a str, &'a str), AggregateFieldError> {
1108 column_types
1109 .get_key_value(name)
1110 .map(|(k, v)| (k.as_str(), v.as_str()))
1111 .ok_or_else(|| {
1112 AggregateFieldError(format!("unknown field `{name}` — not a column of this entity"))
1113 })
1114}
1115
1116impl<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> PostgresRepository<T> {
1117 pub async fn aggregate_with_filters(
1125 &self,
1126 spec: &AggregateSpec,
1127 filters: &HashMap<String, String>,
1128 column_types: &HashMap<String, String>,
1129 search_fields: &[&str],
1130 ) -> anyhow::Result<AggregateResult> {
1131 let mut query_filter = self.parse_typed_filters(filters, column_types, None).await?;
1132 if !search_fields.is_empty() {
1133 query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
1134 }
1135 query_filter.limit = None;
1138 query_filter.offset = None;
1139 let (where_clause, filter_params) = query_filter.build_where_clause();
1140
1141 let columns = catalog_columns(&self.pool, &self.table_name).await?;
1143 let mut selects: Vec<String> = Vec::new();
1144 let mut value_keys: Vec<String> = Vec::new();
1145 for (func, field) in &spec.reductions {
1146 let (column, pg_type) = resolve_column(field, &columns)?;
1147 if func.requires_numeric() && !is_numeric_pg_type(pg_type) {
1148 return Err(AggregateFieldError(format!(
1149 "cannot {} `{}`: its type is {} — {} needs a numeric column",
1150 func.label(),
1151 column,
1152 pg_type,
1153 func.label()
1154 ))
1155 .into());
1156 }
1157 let key = format!("{}:{}", func.label(), column);
1158 selects.push(format!("{}({})::text AS \"{}\"", func.sql(), column, key));
1161 value_keys.push(key);
1162 }
1163
1164 let group_limit = if spec.group_limit == 0 { DEFAULT_GROUP_LIMIT } else { spec.group_limit };
1165 let reductions = if selects.is_empty() { String::new() } else { format!(", {}", selects.join(", ")) };
1166
1167 let mut label_select = String::new();
1173 let mut label_join = String::new();
1174 if let (Some(field), Some(label_field), Some((rel_table, base_fk))) = (
1175 &spec.group_by,
1176 &spec.label_field,
1177 &spec.label_relation,
1178 ) {
1179 let _ = field;
1180 let qualified = qualify_relation_table(&self.table_name, rel_table);
1181 let rel_columns = catalog_columns(&self.pool, &qualified).await?;
1182 let (label_col, _) = resolve_column(label_field, &rel_columns).map_err(|_| {
1183 AggregateFieldError(format!(
1184 "cannot label groups by `{label_field}`: the related table `{qualified}` has no such column"
1185 ))
1186 })?;
1187 label_select = format!(", (label_rel.{label_col})::text AS __group_label");
1188 label_join = format!(
1189 " LEFT JOIN {qualified} AS label_rel ON label_rel.id IS NOT DISTINCT FROM {base_fk}"
1190 );
1191 }
1192
1193 let sql = match &spec.group_by {
1194 Some(field) => {
1195 let (column, _) = resolve_column(field, &columns)?;
1196 format!(
1197 "SELECT GROUPING({column}) AS __is_total, ({column})::text AS __group_key{label_select}, \
1198 COUNT(*) AS __count{reductions} \
1199 FROM {table}{label_join}{where_clause} \
1200 GROUP BY GROUPING SETS (({column}), ()) \
1201 ORDER BY __is_total DESC, __count DESC \
1202 LIMIT {limit}",
1203 column = column,
1204 label_select = label_select,
1205 label_join = label_join,
1206 reductions = reductions,
1207 table = self.table_name,
1208 where_clause = where_clause,
1209 limit = group_limit + 2,
1212 )
1213 }
1214 None => format!(
1215 "SELECT 1 AS __is_total, NULL::text AS __group_key, COUNT(*) AS __count{reductions} \
1216 FROM {table}{where_clause}",
1217 reductions = reductions,
1218 table = self.table_name,
1219 where_clause = where_clause,
1220 ),
1221 };
1222
1223 let mut builder = sqlx::query(&sql);
1224 for param in &filter_params {
1225 builder = builder.bind(param);
1226 }
1227 let rows = crate::company_scope::fetch_all_rows_scoped(&self.pool, builder).await?;
1228
1229 let read_group = |row: &PgRow| -> AggregateGroup {
1230 use sqlx::Row as _;
1231 let mut values = HashMap::with_capacity(value_keys.len());
1232 for key in &value_keys {
1233 values.insert(key.clone(), row.try_get::<Option<String>, _>(key.as_str()).ok().flatten());
1234 }
1235 AggregateGroup {
1236 key: row.try_get::<Option<String>, _>("__group_key").ok().flatten(),
1237 label: row.try_get::<Option<String>, _>("__group_label").ok().flatten(),
1238 count: row.try_get::<i64, _>("__count").unwrap_or(0).max(0) as u64,
1239 values,
1240 }
1241 };
1242
1243 use sqlx::Row as _;
1244 let mut total: Option<AggregateGroup> = None;
1245 let mut groups: Vec<AggregateGroup> = Vec::new();
1246 for row in &rows {
1247 let is_total = row.try_get::<i32, _>("__is_total").unwrap_or(0) == 1;
1248 if is_total {
1249 total = Some(read_group(row));
1251 } else {
1252 groups.push(read_group(row));
1253 }
1254 }
1255
1256 let truncated = groups.len() > group_limit;
1257 groups.truncate(group_limit);
1258
1259 let total = total.unwrap_or_else(|| AggregateGroup {
1262 key: None,
1263 label: None,
1264 count: 0,
1265 values: value_keys.iter().map(|k| (k.clone(), None)).collect(),
1266 });
1267
1268 Ok(AggregateResult { groups, total, truncated })
1269 }
1270}
1271
1272#[cfg(test)]
1273mod aggregate_field_tests {
1274 use super::*;
1275
1276 fn columns() -> HashMap<String, String> {
1277 [
1278 ("status", "text"),
1279 ("total", "numeric"),
1280 ("qty", "integer"),
1281 ("notes", "text"),
1282 ]
1283 .iter()
1284 .map(|(k, v)| (k.to_string(), v.to_string()))
1285 .collect()
1286 }
1287
1288 #[test]
1289 fn resolves_only_declared_columns() {
1290 let cols = columns();
1291 assert_eq!(resolve_column("total", &cols).unwrap().0, "total");
1292 assert!(resolve_column("password_hash", &cols).is_err());
1293 }
1294
1295 #[test]
1298 fn rejects_injection_attempts_rather_than_escaping_them() {
1299 let cols = columns();
1300 for probe in [
1301 "total) FROM selling.sales_orders; DROP TABLE users --",
1302 "status\"",
1303 "1=1",
1304 "total, (SELECT password FROM users)",
1305 "",
1306 ] {
1307 assert!(
1308 resolve_column(probe, &cols).is_err(),
1309 "`{probe}` must be refused, never escaped into the query"
1310 );
1311 }
1312 }
1313
1314 #[test]
1317 fn returns_the_declared_key_not_the_callers_string() {
1318 let cols = columns();
1319 let (name, _) = resolve_column("total", &cols).unwrap();
1320 assert!(std::ptr::eq(name, cols.get_key_value("total").unwrap().0.as_str()));
1321 }
1322
1323 #[test]
1324 fn sum_and_avg_require_a_numeric_type() {
1325 assert!(AggregateFn::Sum.requires_numeric());
1326 assert!(AggregateFn::Avg.requires_numeric());
1327 assert!(!AggregateFn::Min.requires_numeric());
1329 assert!(!AggregateFn::Max.requires_numeric());
1330 }
1331
1332 #[test]
1333 fn recognises_the_numeric_postgres_types() {
1334 for t in ["numeric", "NUMERIC(14,2)", "integer", "bigint", "double precision", "money"] {
1335 assert!(is_numeric_pg_type(t), "{t} should count as numeric");
1336 }
1337 for t in ["text", "uuid", "timestamptz", "boolean", "jsonb", "USER-DEFINED"] {
1338 assert!(!is_numeric_pg_type(t), "{t} must not accept a SUM");
1339 }
1340 }
1341}
1342
1343#[cfg(test)]
1344mod filter_cast_tests {
1345 use super::*;
1346
1347 #[test]
1348 fn typed_base_columns_get_their_own_cast() {
1349 for (t, want) in [
1350 ("boolean", "boolean"),
1351 ("integer", "integer"),
1352 ("bigint", "bigint"),
1353 ("smallint", "smallint"),
1354 ("numeric", "numeric"),
1355 ("double precision", "double precision"),
1356 ("uuid", "uuid"),
1357 ("date", "date"),
1358 ("time without time zone", "time"),
1359 ("timestamp with time zone", "timestamptz"),
1360 ("timestamp without time zone", "timestamp"),
1361 ] {
1362 assert_eq!(filter_cast_for(t, "b").as_deref(), Some(want), "{t}");
1363 }
1364 }
1365
1366 #[test]
1367 fn text_like_and_composite_columns_keep_the_text_bind() {
1368 for t in ["text", "character varying", "character", "jsonb", "json", "bytea", "text[]", "uuid[]"] {
1369 assert_eq!(filter_cast_for(t, "b"), None, "{t}");
1370 }
1371 assert_eq!(filter_cast_for("approvals.money_amount", "d"), None);
1373 assert_eq!(filter_cast_for("approvals.address", "c"), None);
1374 }
1375
1376 #[test]
1377 fn an_enum_casts_to_its_catalog_name() {
1378 assert_eq!(filter_cast_for("approval_status", "e").as_deref(), Some("approval_status"));
1379 assert_eq!(
1380 filter_cast_for("recruitment.stage_kind", "e").as_deref(),
1381 Some("recruitment.stage_kind")
1382 );
1383 }
1384
1385 #[test]
1386 fn a_generated_hint_wins_over_the_catalog_for_its_column() {
1387 let hints: HashMap<String, String> =
1388 [("id", "uuid"), ("status", "approval_status")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1389 let catalog: HashMap<String, String> = [
1390 ("id", "uuid"),
1391 ("status", "approvals.approval_status"),
1392 ("folded", "boolean"),
1393 ("requested_by", "uuid"),
1394 ]
1395 .iter()
1396 .map(|(k, v)| (k.to_string(), v.to_string()))
1397 .collect();
1398 let merged = merge_filter_casts(&hints, catalog);
1399 assert_eq!(merged["status"], "approval_status");
1400 assert_eq!(merged["folded"], "boolean");
1401 assert_eq!(merged["requested_by"], "uuid");
1402 assert_eq!(merged.len(), 4);
1403 }
1404
1405 #[test]
1406 fn only_a_filter_with_an_uncast_comparison_reads_the_catalog() {
1407 let hints: HashMap<String, String> =
1408 [("id", "uuid")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1409 let parse = |pairs: &[(&str, &str)]| {
1410 let f: HashMap<String, String> =
1411 pairs.iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1412 parse_query_filter(&f, &hints, None).unwrap().has_uncast_value_conditions()
1413 };
1414 assert!(!parse(&[("id[in]", "a,b"), ("name[contain]", "x"), ("limit", "5")]));
1415 assert!(!parse(&[("deleted_by[isnull]", "1")]));
1416 assert!(parse(&[("folded[eq]", "false")]));
1417 assert!(parse(&[("folded", "false")]));
1418 assert!(parse(&[("scheduled_at[between]", "2026-10-01,2026-10-03")]));
1419 assert!(parse(&[("sequence[or]", "3")]));
1420 }
1421}