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
121pub const EXACT_COUNT_BELOW: u64 = 10_000;
124
125fn default_count_mode() -> String {
126 "exact".to_string()
127}
128
129impl PaginationInfo {
130 pub fn new(page: u32, per_page: u32, total: u64) -> Self {
131 let total_pages = ((total as f64) / (per_page as f64)).ceil() as u32;
132 Self {
133 page,
134 per_page,
135 total,
136 total_pages,
137 next_cursor: None,
138 prev_cursor: None,
139 has_more: None,
140 count_mode: default_count_mode(),
141 }
142 }
143}
144
145#[async_trait]
147pub trait DatabaseOperations<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> {
148 async fn create(&self, entity: &T) -> anyhow::Result<T>;
150
151 async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>>;
153
154 async fn find_all(&self) -> anyhow::Result<Vec<T>>;
156
157 async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>>;
159
160 async fn delete(&self, id: &str) -> anyhow::Result<bool>;
162
163 async fn count(&self) -> anyhow::Result<u64>;
165
166 async fn exists(&self, id: &str) -> anyhow::Result<bool>;
168
169 async fn execute_query(&self, query: &str) -> anyhow::Result<u64>;
171}
172
173pub struct PostgresRepository<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> {
175 pool: PgPool,
176 table_name: String,
177 _phantom: std::marker::PhantomData<T>,
178}
179
180impl<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> PostgresRepository<T> {
181 pub fn new(pool: PgPool, table_name: &str) -> Self {
182 Self {
183 pool,
184 table_name: table_name.to_string(),
185 _phantom: std::marker::PhantomData,
186 }
187 }
188
189 pub fn pool(&self) -> &PgPool {
190 &self.pool
191 }
192
193 pub fn table_name(&self) -> &str {
194 &self.table_name
195 }
196
197 pub async fn list_with_filters(
253 &self,
254 pagination: PaginationParams,
255 filters: &HashMap<String, String>,
256 column_types: &HashMap<String, String>,
257 search_fields: &[&str],
258 ) -> anyhow::Result<PaginatedResult<T>>
259 where
260 T: Send + Sync,
261 {
262 let mut query_filter = self.parse_typed_filters(filters, column_types, None).await?;
264
265 if !search_fields.is_empty() {
267 query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
268 }
269
270 self.execute_list(pagination, query_filter).await
271 }
272
273 pub async fn list_with_filters_whitelisted(
305 &self,
306 pagination: PaginationParams,
307 filters: &HashMap<String, String>,
308 column_types: &HashMap<String, String>,
309 search_fields: &[&str],
310 allowed_fields: Option<&HashSet<String>>,
311 ) -> anyhow::Result<PaginatedResult<T>>
312 where
313 T: Send + Sync,
314 {
315 let mut query_filter =
317 self.parse_typed_filters(filters, column_types, allowed_fields).await?;
318
319 if !search_fields.is_empty() {
321 query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
322 }
323
324 self.execute_list(pagination, query_filter).await
325 }
326
327 async fn parse_typed_filters(
342 &self,
343 filters: &HashMap<String, String>,
344 column_types: &HashMap<String, String>,
345 allowed_fields: Option<&HashSet<String>>,
346 ) -> anyhow::Result<crate::QueryFilter> {
347 let query_filter = parse_query_filter(filters, column_types, allowed_fields)?;
348 if !query_filter.needs_catalog_types() {
349 return Ok(query_filter);
350 }
351 let catalog = catalog_filter_casts(&self.pool, &self.table_name).await?;
352 let merged = merge_filter_casts(column_types, catalog);
353 parse_query_filter(filters, &merged, allowed_fields)
354 }
355
356 #[allow(clippy::type_complexity)]
370 async fn execute_list(
371 &self,
372 pagination: PaginationParams,
373 query_filter: crate::QueryFilter,
374 ) -> anyhow::Result<PaginatedResult<T>> {
375 let limit = pagination.limit() as i64;
376 let backwards =
377 query_filter.cursor_before.is_some() && query_filter.cursor_after.is_none();
378 let cursor_walk = query_filter.cursor_after.is_some() || backwards;
379
380 let (mut where_clause, mut filter_params) = query_filter.build_where_clause();
381 let order_clause;
382 let mut boundary_casts: Vec<Option<String>> = Vec::new();
384
385 let mut sorts: Vec<(String, FilterSortDirection)> = query_filter
390 .sorts
391 .iter()
392 .map(|s| (s.field.clone(), s.direction.clone()))
393 .collect();
394 if cursor_walk || !sorts.is_empty() {
395 if sorts.is_empty() {
396 sorts.push(("id".into(), FilterSortDirection::Asc));
397 } else if sorts.last().map(|(f, _)| f != "id").unwrap_or(true) {
398 sorts.push(("id".into(), FilterSortDirection::Asc));
399 }
400 }
401
402 if !sorts.is_empty() {
405 boundary_casts = self.sort_column_casts(&sorts).await?;
406 }
407
408 if cursor_walk {
409 let opaque = if backwards {
410 query_filter.cursor_before.clone().unwrap()
411 } else {
412 query_filter.cursor_after.clone().unwrap()
413 };
414 let payload = crate::filter::cursor::decode_cursor(&opaque, &sorts)
415 .map_err(|e| anyhow::anyhow!("cursor refused: {e}"))?;
416 let mut idx = filter_params.len() + 1;
417 let (keyset_sql, keyset_params) = crate::filter::cursor::build_keyset_predicate(
418 &payload,
419 &mut idx,
420 &boundary_casts,
421 backwards,
422 );
423 if where_clause.is_empty() {
424 where_clause = format!(" WHERE {}", keyset_sql);
425 } else {
426 where_clause = format!("{} AND ({})", where_clause, keyset_sql);
427 }
428 filter_params.extend(keyset_params);
429 let parts: Vec<String> = sorts
430 .iter()
431 .map(|(f, d)| {
432 let dir = if (*d == FilterSortDirection::Desc) != backwards {
433 "DESC"
434 } else {
435 "ASC"
436 };
437 format!("{} {}", f, dir)
438 })
439 .collect();
440 order_clause = format!(" ORDER BY {}", parts.join(", "));
441 } else {
442 if sorts.is_empty() {
445 order_clause = query_filter.build_order_by_clause();
446 } else {
447 let parts: Vec<String> = sorts
448 .iter()
449 .map(|(f, d)| {
450 let dir =
451 if *d == FilterSortDirection::Desc { "DESC" } else { "ASC" };
452 format!("{} {}", f, dir)
453 })
454 .collect();
455 order_clause = format!(" ORDER BY {}", parts.join(", "));
456 }
457 }
458 let boundary_sorts: Vec<(String, FilterSortDirection)> = sorts;
460
461 let (total, count_mode) = if query_filter.estimate_total {
467 let estimate = self.estimate_filtered_rows(&where_clause, &filter_params).await?;
468 if estimate <= EXACT_COUNT_BELOW {
469 (self.count_filtered_rows(&where_clause, &filter_params).await?, "exact")
470 } else {
471 (estimate, "estimate")
472 }
473 } else if cursor_walk {
474 (0u64, "none")
475 } else {
476 (self.count_filtered_rows(&where_clause, &filter_params).await?, "exact")
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 count_filtered_rows(
561 &self,
562 where_clause: &str,
563 filter_params: &[String],
564 ) -> anyhow::Result<u64> {
565 let count_query = format!("SELECT COUNT(*) FROM {}{}", self.table_name, where_clause);
566 let mut count_query_builder = sqlx::query_scalar::<_, i64>(&count_query);
567 for param in filter_params {
568 count_query_builder = count_query_builder.bind(param);
569 }
570 Ok(crate::company_scope::fetch_one_scalar_scoped(&self.pool, count_query_builder).await?
571 as u64)
572 }
573
574 async fn estimate_filtered_rows(
575 &self,
576 where_clause: &str,
577 filter_params: &[String],
578 ) -> anyhow::Result<u64> {
579 let explain = format!("EXPLAIN (FORMAT JSON) SELECT 1 FROM {}{}", self.table_name, where_clause);
580 let mut builder = sqlx::query_scalar::<_, serde_json::Value>(&explain);
581 for param in filter_params {
582 builder = builder.bind(param);
583 }
584 let plan: serde_json::Value =
585 crate::company_scope::fetch_one_scalar_scoped(&self.pool, builder).await?;
586 let rows = plan
587 .as_array()
588 .and_then(|a| a.first())
589 .and_then(|top| top.get("Plan"))
590 .and_then(|p| p.get("Plan Rows"))
591 .and_then(|r| r.as_i64())
592 .unwrap_or(0);
593 Ok(rows.max(0) as u64)
594 }
595
596 async fn sort_column_casts(
601 &self,
602 sorts: &[(String, FilterSortDirection)],
603 ) -> anyhow::Result<Vec<Option<String>>> {
604 let (schema, table) = match self.table_name.rsplit_once('.') {
605 Some((s, t)) => (s.to_string(), t.to_string()),
606 None => ("public".to_string(), self.table_name.clone()),
607 };
608 let mut casts: Vec<Option<String>> = Vec::with_capacity(sorts.len());
609 for (field, _) in sorts {
610 let row: Option<(String, String, String)> = sqlx::query_as(
611 "SELECT data_type, coalesce(udt_name, ''), coalesce(udt_schema, '') \
612 FROM information_schema.columns \
613 WHERE table_schema = $1 AND table_name = $2 AND column_name = $3",
614 )
615 .bind(&schema)
616 .bind(&table)
617 .bind(field)
618 .fetch_optional(&self.pool)
619 .await?;
620 let cast = row
621 .map(|(data_type, udt, udt_schema)| cast_suffix(&data_type, &udt, &udt_schema))
622 .flatten();
623 casts.push(cast);
624 }
625 Ok(casts)
626 }
627
628 fn row_cursor(
632 &self,
633 row: &PgRow,
634 sorts: &[(String, FilterSortDirection)],
635 casts: &[Option<String>],
636 ) -> Option<String> {
637 let id: uuid::Uuid = row.try_get("id").ok()?;
638 let mut values: Vec<serde_json::Value> = Vec::with_capacity(sorts.len());
639 for (i, (field, _)) in sorts.iter().enumerate() {
640 let field = field.as_str();
641 let mut data_type = casts.get(i).and_then(|c| c.as_deref()).unwrap_or("");
642 if field == "id" && data_type.is_empty() {
646 data_type = "uuid";
647 }
648 let text: Option<String> = match data_type {
649 "numeric" => row
650 .try_get::<Option<sqlx::types::Decimal>, _>(field)
651 .ok()?
652 .map(|d| d.to_string()),
653 "uuid" => row
654 .try_get::<Option<uuid::Uuid>, _>(field)
655 .ok()?
656 .map(|u| u.to_string()),
657 "timestamptz" => row
658 .try_get::<Option<chrono::DateTime<chrono::Utc>>, _>(field)
659 .ok()?
660 .map(|t| t.to_rfc3339()),
661 "integer" | "smallint" => row
662 .try_get::<Option<i32>, _>(field)
663 .ok()?
664 .map(|n| n.to_string()),
665 "bigint" => row
666 .try_get::<Option<i64>, _>(field)
667 .ok()?
668 .map(|n| n.to_string()),
669 "boolean" => row
670 .try_get::<Option<bool>, _>(field)
671 .ok()?
672 .map(|b| b.to_string()),
673 "date" => row
674 .try_get::<Option<chrono::NaiveDate>, _>(field)
675 .ok()?
676 .map(|d| d.to_string()),
677 _ => row
678 .try_get::<Option<String>, _>(field)
679 .ok()?
680 .filter(|s| !s.is_empty() || data_type.is_empty()),
681 };
682 values.push(serde_json::Value::String(text?));
683 }
684 crate::filter::cursor::encode_cursor(sorts, &values, &id.to_string()).ok()
685 }
686}
687
688fn cast_suffix(data_type: &str, udt_name: &str, udt_schema: &str) -> Option<String> {
691 match data_type {
692 "uuid" => Some("uuid".into()),
693 "numeric" => Some("numeric".into()),
694 "integer" => Some("integer".into()),
695 "smallint" => Some("smallint".into()),
696 "bigint" => Some("bigint".into()),
697 "boolean" => Some("boolean".into()),
698 "date" => Some("date".into()),
699 "timestamp with time zone" => Some("timestamptz".into()),
700 "timestamp without time zone" => Some("timestamp".into()),
701 "USER-DEFINED" if !udt_name.is_empty() => Some(match udt_schema {
705 "" | "public" | "pg_catalog" => udt_name.to_string(),
706 schema => format!("{schema}.{udt_name}"),
707 }),
708 _ => None,
709 }
710}
711
712fn filter_cast_for(type_name: &str, typtype: &str) -> Option<String> {
717 match typtype {
718 "e" => Some(type_name.to_string()),
721 "b" => match type_name {
722 "uuid" | "boolean" | "smallint" | "integer" | "bigint" | "numeric" | "real"
723 | "double precision" | "date" | "interval" | "inet" | "cidr" | "macaddr" => {
724 Some(type_name.to_string())
725 }
726 "time without time zone" => Some("time".into()),
727 "time with time zone" => Some("timetz".into()),
728 "timestamp with time zone" => Some("timestamptz".into()),
729 "timestamp without time zone" => Some("timestamp".into()),
730 _ => None,
731 },
732 _ => None,
733 }
734}
735
736async fn catalog_filter_casts(
743 pool: &PgPool,
744 qualified_table: &str,
745) -> anyhow::Result<HashMap<String, String>> {
746 let q = sqlx::query_as::<Postgres, (String, String, String)>(
747 "SELECT a.attname::text, format_type(a.atttypid, NULL), t.typtype::text
748 FROM pg_catalog.pg_attribute a
749 JOIN pg_catalog.pg_type t ON t.oid = a.atttypid
750 WHERE a.attrelid = to_regclass($1) AND a.attnum > 0 AND NOT a.attisdropped",
751 )
752 .bind(qualified_table.to_string());
753 let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
754 Ok(rows
755 .into_iter()
756 .filter_map(|(column, type_name, typtype)| {
757 filter_cast_for(&type_name, &typtype).map(|cast| (column, cast))
758 })
759 .collect())
760}
761
762fn merge_filter_casts(
767 hints: &HashMap<String, String>,
768 mut catalog: HashMap<String, String>,
769) -> HashMap<String, String> {
770 for (column, cast) in hints {
771 let qualified_same_type = catalog
772 .get(column)
773 .and_then(|c| c.rsplit_once('.'))
774 .is_some_and(|(_, name)| name.trim_matches('"') == cast.as_str());
775 if !qualified_same_type {
776 catalog.insert(column.clone(), cast.clone());
777 }
778 }
779 catalog
780}
781
782fn explain_unknown_column(error: sqlx::Error, table: &str) -> anyhow::Error {
790 let text = error.to_string();
791 if text.contains("does not exist") && text.contains("column") {
792 return anyhow::Error::new(error).context(format!(
793 "insert into {table} named a column that does not exist: the entity serializes a field \
794 with no matching column. Every serialized field must be a column of the table (rename \
795 it, map it with #[serde(rename)], or skip it with #[serde(skip)])"
796 ));
797 }
798 anyhow::Error::new(error)
799}
800
801fn quote_ident(name: &str) -> String {
808 format!("\"{}\"", name.replace('"', "\"\""))
809}
810
811#[async_trait]
812impl<T> DatabaseOperations<T> for PostgresRepository<T>
813where
814 T: for<'a> FromRow<'a, PgRow> + Send + Sync + Unpin + Serialize,
815{
816 async fn create(&self, entity: &T) -> anyhow::Result<T> {
817 let json_value = serde_json::to_value(entity)?;
819
820 let json_obj = match json_value {
821 Value::Object(obj) => obj,
822 _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
823 };
824
825 let json_str = serde_json::to_string(&json_obj)?;
828
829 let insert_columns: Vec<String> = json_obj.keys().map(|k| quote_ident(k)).collect();
843
844 let query = if insert_columns.is_empty() {
845 format!("INSERT INTO {table} DEFAULT VALUES RETURNING *", table = self.table_name)
848 } else {
849 let columns = insert_columns.join(", ");
850 format!(
851 r#"
852 INSERT INTO {table} ({columns})
853 SELECT {columns} FROM jsonb_populate_record(NULL::{table}, $1::jsonb)
854 RETURNING *
855 "#,
856 table = self.table_name,
857 columns = columns
858 )
859 };
860
861 let statement = if insert_columns.is_empty() {
863 sqlx::query_as::<_, T>(&query)
864 } else {
865 sqlx::query_as::<_, T>(&query).bind(&json_str)
866 };
867 let result = crate::company_scope::fetch_one_scoped(&self.pool, statement)
868 .await
869 .map_err(|e| explain_unknown_column(e, &self.table_name))?;
870
871 Ok(result)
872 }
873
874 async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>> {
875 let query = format!("SELECT * FROM {} WHERE id = $1::uuid", self.table_name);
877 let result = crate::company_scope::fetch_optional_scoped(
878 &self.pool,
879 sqlx::query_as::<Postgres, T>(&query).bind(id),
880 )
881 .await?;
882 Ok(result)
883 }
884
885 async fn find_all(&self) -> anyhow::Result<Vec<T>> {
886 let query = format!("SELECT * FROM {}", self.table_name);
887 let results = crate::company_scope::fetch_all_scoped(
888 &self.pool,
889 sqlx::query_as::<Postgres, T>(&query),
890 )
891 .await?;
892 Ok(results)
893 }
894
895 async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>> {
896 let json_value = serde_json::to_value(entity)?;
898
899 let json_obj = match json_value {
900 Value::Object(obj) => obj,
901 _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
902 };
903
904 let update_columns: Vec<&String> = json_obj.keys()
906 .filter(|k| *k != "id")
907 .collect();
908
909 let column_names = update_columns.iter()
910 .map(|k| quote_ident(k))
911 .collect::<Vec<_>>()
912 .join(", ");
913
914 let json_str = serde_json::to_string(&json_obj)?;
915
916 let query = format!(
918 r#"
919 WITH new_row AS (
920 SELECT (jsonb_populate_record(NULL::{table}, $1::jsonb)).*
921 )
922 UPDATE {table} AS t
923 SET ({columns}) = (SELECT {columns} FROM new_row)
924 WHERE t.id = $2::uuid
925 RETURNING t.*
926 "#,
927 table = self.table_name,
928 columns = column_names
929 );
930
931 let result = crate::company_scope::fetch_optional_scoped(
932 &self.pool,
933 sqlx::query_as::<_, T>(&query).bind(&json_str).bind(id),
934 )
935 .await?;
936
937 Ok(result)
938 }
939
940 async fn delete(&self, id: &str) -> anyhow::Result<bool> {
941 let query = format!("DELETE FROM {} WHERE id = $1::uuid", self.table_name);
942 let result = crate::company_scope::execute_scoped(
943 &self.pool,
944 sqlx::query(&query).bind(id),
945 )
946 .await?;
947 Ok(result.rows_affected() > 0)
948 }
949
950 async fn count(&self) -> anyhow::Result<u64> {
951 let query = format!("SELECT COUNT(*) FROM {}", self.table_name);
952 let count = crate::company_scope::fetch_one_scalar_scoped(
953 &self.pool,
954 sqlx::query_scalar::<_, i64>(&query),
955 )
956 .await? as u64;
957 Ok(count)
958 }
959
960 async fn exists(&self, id: &str) -> anyhow::Result<bool> {
961 let query = format!("SELECT 1 FROM {} WHERE id = $1::uuid LIMIT 1", self.table_name);
962 let result = crate::company_scope::fetch_optional_scalar_scoped(
963 &self.pool,
964 sqlx::query_scalar::<_, i32>(&query).bind(id),
965 )
966 .await?;
967 Ok(result.is_some())
968 }
969
970 async fn execute_query(&self, query: &str) -> anyhow::Result<u64> {
971 let result = crate::company_scope::execute_scoped(
972 &self.pool,
973 sqlx::query(query),
974 )
975 .await?;
976 Ok(result.rows_affected())
977 }
978}
979
980#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
984pub enum AggregateFn {
985 Sum,
986 Avg,
987 Min,
988 Max,
989}
990
991impl AggregateFn {
992 fn sql(self) -> &'static str {
993 match self {
994 AggregateFn::Sum => "SUM",
995 AggregateFn::Avg => "AVG",
996 AggregateFn::Min => "MIN",
997 AggregateFn::Max => "MAX",
998 }
999 }
1000
1001 fn label(self) -> &'static str {
1002 match self {
1003 AggregateFn::Sum => "sum",
1004 AggregateFn::Avg => "avg",
1005 AggregateFn::Min => "min",
1006 AggregateFn::Max => "max",
1007 }
1008 }
1009
1010 fn requires_numeric(self) -> bool {
1013 matches!(self, AggregateFn::Sum | AggregateFn::Avg)
1014 }
1015}
1016
1017#[derive(Debug, Clone, Default)]
1019pub struct AggregateSpec {
1020 pub group_by: Option<String>,
1022 pub reductions: Vec<(AggregateFn, String)>,
1024 pub group_limit: usize,
1026 pub label_field: Option<String>,
1030 pub label_relation: Option<(String, String)>,
1034}
1035
1036pub const DEFAULT_GROUP_LIMIT: usize = 200;
1044
1045#[derive(Debug, Clone, Serialize, Deserialize)]
1048pub struct AggregateGroup {
1049 pub key: Option<String>,
1050 pub label: Option<String>,
1053 pub count: u64,
1054 pub values: HashMap<String, Option<String>>,
1062}
1063
1064#[derive(Debug, Clone, Serialize, Deserialize)]
1066pub struct AggregateResult {
1067 pub groups: Vec<AggregateGroup>,
1068 pub total: AggregateGroup,
1069 pub truncated: bool,
1071}
1072
1073#[derive(Debug, Clone)]
1076pub struct AggregateFieldError(pub String);
1077
1078impl std::fmt::Display for AggregateFieldError {
1079 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1080 f.write_str(&self.0)
1081 }
1082}
1083
1084impl std::error::Error for AggregateFieldError {}
1085
1086fn is_numeric_pg_type(pg_type: &str) -> bool {
1088 let t = pg_type.trim().to_ascii_lowercase();
1089 let t = t.split('(').next().unwrap_or(&t).trim();
1090 matches!(
1091 t,
1092 "numeric" | "decimal" | "money"
1093 | "smallint" | "int2" | "integer" | "int" | "int4" | "bigint" | "int8"
1094 | "real" | "float4" | "double precision" | "float8"
1095 | "smallserial" | "serial" | "bigserial"
1096 )
1097}
1098
1099async fn catalog_columns(
1108 pool: &PgPool,
1109 qualified_table: &str,
1110) -> anyhow::Result<HashMap<String, String>> {
1111 let (schema, table) = match qualified_table.split_once('.') {
1112 Some((s, t)) => (s.to_string(), t.to_string()),
1113 None => ("public".to_string(), qualified_table.to_string()),
1114 };
1115 let q = sqlx::query_as::<Postgres, (String, String)>(
1116 "SELECT column_name, data_type FROM information_schema.columns
1117 WHERE table_schema = $1 AND table_name = $2",
1118 )
1119 .bind(schema)
1120 .bind(table);
1121 let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
1122 Ok(rows.into_iter().collect())
1123}
1124
1125fn resolve_column<'a>(
1134 name: &str,
1135 column_types: &'a HashMap<String, String>,
1136) -> Result<(&'a str, &'a str), AggregateFieldError> {
1137 column_types
1138 .get_key_value(name)
1139 .map(|(k, v)| (k.as_str(), v.as_str()))
1140 .ok_or_else(|| {
1141 AggregateFieldError(format!("unknown field `{name}` — not a column of this entity"))
1142 })
1143}
1144
1145impl<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> PostgresRepository<T> {
1146 pub async fn aggregate_with_filters(
1154 &self,
1155 spec: &AggregateSpec,
1156 filters: &HashMap<String, String>,
1157 column_types: &HashMap<String, String>,
1158 search_fields: &[&str],
1159 ) -> anyhow::Result<AggregateResult> {
1160 let mut query_filter = self.parse_typed_filters(filters, column_types, None).await?;
1161 if !search_fields.is_empty() {
1162 query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
1163 }
1164 query_filter.limit = None;
1167 query_filter.offset = None;
1168 let (where_clause, filter_params) = query_filter.build_where_clause();
1169
1170 let columns = catalog_columns(&self.pool, &self.table_name).await?;
1172 let mut selects: Vec<String> = Vec::new();
1173 let mut value_keys: Vec<String> = Vec::new();
1174 for (func, field) in &spec.reductions {
1175 let (column, pg_type) = resolve_column(field, &columns)?;
1176 if func.requires_numeric() && !is_numeric_pg_type(pg_type) {
1177 return Err(AggregateFieldError(format!(
1178 "cannot {} `{}`: its type is {} — {} needs a numeric column",
1179 func.label(),
1180 column,
1181 pg_type,
1182 func.label()
1183 ))
1184 .into());
1185 }
1186 let key = format!("{}:{}", func.label(), column);
1187 selects.push(format!("{}({})::text AS \"{}\"", func.sql(), column, key));
1190 value_keys.push(key);
1191 }
1192
1193 let group_limit = if spec.group_limit == 0 { DEFAULT_GROUP_LIMIT } else { spec.group_limit };
1194 let reductions = if selects.is_empty() { String::new() } else { format!(", {}", selects.join(", ")) };
1195
1196 let mut label_select = String::new();
1202 let mut label_join = String::new();
1203 if let (Some(field), Some(label_field), Some((rel_table, base_fk))) = (
1204 &spec.group_by,
1205 &spec.label_field,
1206 &spec.label_relation,
1207 ) {
1208 let _ = field;
1209 let qualified = qualify_relation_table(&self.table_name, rel_table);
1210 let rel_columns = catalog_columns(&self.pool, &qualified).await?;
1211 let (label_col, _) = resolve_column(label_field, &rel_columns).map_err(|_| {
1212 AggregateFieldError(format!(
1213 "cannot label groups by `{label_field}`: the related table `{qualified}` has no such column"
1214 ))
1215 })?;
1216 label_select = format!(", (label_rel.{label_col})::text AS __group_label");
1217 label_join = format!(
1218 " LEFT JOIN {qualified} AS label_rel ON label_rel.id IS NOT DISTINCT FROM {base_fk}"
1219 );
1220 }
1221
1222 let sql = match &spec.group_by {
1223 Some(field) => {
1224 let (column, _) = resolve_column(field, &columns)?;
1225 format!(
1226 "SELECT GROUPING({column}) AS __is_total, ({column})::text AS __group_key{label_select}, \
1227 COUNT(*) AS __count{reductions} \
1228 FROM {table}{label_join}{where_clause} \
1229 GROUP BY GROUPING SETS (({column}), ()) \
1230 ORDER BY __is_total DESC, __count DESC \
1231 LIMIT {limit}",
1232 column = column,
1233 label_select = label_select,
1234 label_join = label_join,
1235 reductions = reductions,
1236 table = self.table_name,
1237 where_clause = where_clause,
1238 limit = group_limit + 2,
1241 )
1242 }
1243 None => format!(
1244 "SELECT 1 AS __is_total, NULL::text AS __group_key, COUNT(*) AS __count{reductions} \
1245 FROM {table}{where_clause}",
1246 reductions = reductions,
1247 table = self.table_name,
1248 where_clause = where_clause,
1249 ),
1250 };
1251
1252 let mut builder = sqlx::query(&sql);
1253 for param in &filter_params {
1254 builder = builder.bind(param);
1255 }
1256 let rows = crate::company_scope::fetch_all_rows_scoped(&self.pool, builder).await?;
1257
1258 let read_group = |row: &PgRow| -> AggregateGroup {
1259 use sqlx::Row as _;
1260 let mut values = HashMap::with_capacity(value_keys.len());
1261 for key in &value_keys {
1262 values.insert(key.clone(), row.try_get::<Option<String>, _>(key.as_str()).ok().flatten());
1263 }
1264 AggregateGroup {
1265 key: row.try_get::<Option<String>, _>("__group_key").ok().flatten(),
1266 label: row.try_get::<Option<String>, _>("__group_label").ok().flatten(),
1267 count: row.try_get::<i64, _>("__count").unwrap_or(0).max(0) as u64,
1268 values,
1269 }
1270 };
1271
1272 use sqlx::Row as _;
1273 let mut total: Option<AggregateGroup> = None;
1274 let mut groups: Vec<AggregateGroup> = Vec::new();
1275 for row in &rows {
1276 let is_total = row.try_get::<i32, _>("__is_total").unwrap_or(0) == 1;
1277 if is_total {
1278 total = Some(read_group(row));
1280 } else {
1281 groups.push(read_group(row));
1282 }
1283 }
1284
1285 let truncated = groups.len() > group_limit;
1286 groups.truncate(group_limit);
1287
1288 let total = total.unwrap_or_else(|| AggregateGroup {
1291 key: None,
1292 label: None,
1293 count: 0,
1294 values: value_keys.iter().map(|k| (k.clone(), None)).collect(),
1295 });
1296
1297 Ok(AggregateResult { groups, total, truncated })
1298 }
1299}
1300
1301#[cfg(test)]
1302mod aggregate_field_tests {
1303 use super::*;
1304
1305 fn columns() -> HashMap<String, String> {
1306 [
1307 ("status", "text"),
1308 ("total", "numeric"),
1309 ("qty", "integer"),
1310 ("notes", "text"),
1311 ]
1312 .iter()
1313 .map(|(k, v)| (k.to_string(), v.to_string()))
1314 .collect()
1315 }
1316
1317 #[test]
1318 fn resolves_only_declared_columns() {
1319 let cols = columns();
1320 assert_eq!(resolve_column("total", &cols).unwrap().0, "total");
1321 assert!(resolve_column("password_hash", &cols).is_err());
1322 }
1323
1324 #[test]
1327 fn rejects_injection_attempts_rather_than_escaping_them() {
1328 let cols = columns();
1329 for probe in [
1330 "total) FROM selling.sales_orders; DROP TABLE users --",
1331 "status\"",
1332 "1=1",
1333 "total, (SELECT password FROM users)",
1334 "",
1335 ] {
1336 assert!(
1337 resolve_column(probe, &cols).is_err(),
1338 "`{probe}` must be refused, never escaped into the query"
1339 );
1340 }
1341 }
1342
1343 #[test]
1346 fn returns_the_declared_key_not_the_callers_string() {
1347 let cols = columns();
1348 let (name, _) = resolve_column("total", &cols).unwrap();
1349 assert!(std::ptr::eq(name, cols.get_key_value("total").unwrap().0.as_str()));
1350 }
1351
1352 #[test]
1353 fn sum_and_avg_require_a_numeric_type() {
1354 assert!(AggregateFn::Sum.requires_numeric());
1355 assert!(AggregateFn::Avg.requires_numeric());
1356 assert!(!AggregateFn::Min.requires_numeric());
1358 assert!(!AggregateFn::Max.requires_numeric());
1359 }
1360
1361 #[test]
1362 fn recognises_the_numeric_postgres_types() {
1363 for t in ["numeric", "NUMERIC(14,2)", "integer", "bigint", "double precision", "money"] {
1364 assert!(is_numeric_pg_type(t), "{t} should count as numeric");
1365 }
1366 for t in ["text", "uuid", "timestamptz", "boolean", "jsonb", "USER-DEFINED"] {
1367 assert!(!is_numeric_pg_type(t), "{t} must not accept a SUM");
1368 }
1369 }
1370}
1371
1372#[cfg(test)]
1373mod filter_cast_tests {
1374 use super::*;
1375
1376 #[test]
1377 fn typed_base_columns_get_their_own_cast() {
1378 for (t, want) in [
1379 ("boolean", "boolean"),
1380 ("integer", "integer"),
1381 ("bigint", "bigint"),
1382 ("smallint", "smallint"),
1383 ("numeric", "numeric"),
1384 ("double precision", "double precision"),
1385 ("uuid", "uuid"),
1386 ("date", "date"),
1387 ("time without time zone", "time"),
1388 ("timestamp with time zone", "timestamptz"),
1389 ("timestamp without time zone", "timestamp"),
1390 ] {
1391 assert_eq!(filter_cast_for(t, "b").as_deref(), Some(want), "{t}");
1392 }
1393 }
1394
1395 #[test]
1396 fn text_like_and_composite_columns_keep_the_text_bind() {
1397 for t in ["text", "character varying", "character", "jsonb", "json", "bytea", "text[]", "uuid[]"] {
1398 assert_eq!(filter_cast_for(t, "b"), None, "{t}");
1399 }
1400 assert_eq!(filter_cast_for("approvals.money_amount", "d"), None);
1402 assert_eq!(filter_cast_for("approvals.address", "c"), None);
1403 }
1404
1405 #[test]
1406 fn an_enum_casts_to_its_catalog_name() {
1407 assert_eq!(filter_cast_for("approval_status", "e").as_deref(), Some("approval_status"));
1408 assert_eq!(
1409 filter_cast_for("recruitment.stage_kind", "e").as_deref(),
1410 Some("recruitment.stage_kind")
1411 );
1412 }
1413
1414 #[test]
1415 fn a_generated_hint_wins_over_the_catalog_for_its_column() {
1416 let hints: HashMap<String, String> =
1417 [("id", "uuid"), ("status", "approval_status")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1418 let catalog: HashMap<String, String> = [
1419 ("id", "uuid"),
1420 ("status", "approvals.approval_status"),
1421 ("folded", "boolean"),
1422 ("requested_by", "uuid"),
1423 ]
1424 .iter()
1425 .map(|(k, v)| (k.to_string(), v.to_string()))
1426 .collect();
1427 let merged = merge_filter_casts(&hints, catalog);
1428 assert_eq!(merged["status"], "approvals.approval_status");
1431 assert_eq!(merged["id"], "uuid");
1432 assert_eq!(merged["folded"], "boolean");
1433 assert_eq!(merged["requested_by"], "uuid");
1434 assert_eq!(merged.len(), 4);
1435 }
1436
1437 #[test]
1438 fn a_hint_for_another_type_still_decides_its_column() {
1439 let hints: HashMap<String, String> =
1440 [("amount", "numeric")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1441 let catalog: HashMap<String, String> =
1442 [("amount", "real")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1443 assert_eq!(merge_filter_casts(&hints, catalog)["amount"], "numeric");
1444 }
1445
1446 #[test]
1447 fn a_filter_on_a_bare_enum_hint_reads_the_catalog_to_qualify_it() {
1448 let hints: HashMap<String, String> = [("id", "uuid"), ("status", "task_status")]
1449 .iter()
1450 .map(|(k, v)| (k.to_string(), v.to_string()))
1451 .collect();
1452 let needs = |pairs: &[(&str, &str)]| {
1453 let f: HashMap<String, String> =
1454 pairs.iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1455 parse_query_filter(&f, &hints, None).unwrap().needs_catalog_types()
1456 };
1457 assert!(needs(&[("status[in]", "done,skipped")]));
1458 assert!(needs(&[("status[eq]", "done")]));
1459 assert!(!needs(&[("id[in]", "a,b")]));
1461 }
1462
1463 #[test]
1464 fn a_sort_on_an_enum_casts_to_its_schema_qualified_name() {
1465 assert_eq!(
1466 cast_suffix("USER-DEFINED", "task_status", "lifecycle").as_deref(),
1467 Some("lifecycle.task_status")
1468 );
1469 assert_eq!(cast_suffix("USER-DEFINED", "mood", "public").as_deref(), Some("mood"));
1470 assert_eq!(cast_suffix("uuid", "uuid", "pg_catalog").as_deref(), Some("uuid"));
1471 }
1472
1473 #[test]
1474 fn only_a_filter_with_an_uncast_comparison_reads_the_catalog() {
1475 let hints: HashMap<String, String> =
1476 [("id", "uuid")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1477 let parse = |pairs: &[(&str, &str)]| {
1478 let f: HashMap<String, String> =
1479 pairs.iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1480 parse_query_filter(&f, &hints, None).unwrap().has_uncast_value_conditions()
1481 };
1482 assert!(!parse(&[("id[in]", "a,b"), ("name[contain]", "x"), ("limit", "5")]));
1483 assert!(!parse(&[("deleted_by[isnull]", "1")]));
1484 assert!(parse(&[("folded[eq]", "false")]));
1485 assert!(parse(&[("folded", "false")]));
1486 assert!(parse(&[("scheduled_at[between]", "2026-10-01,2026-10-03")]));
1487 assert!(parse(&[("sequence[or]", "3")]));
1488 }
1489}