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 = parse_query_filter(filters, column_types, None)?;
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 = parse_query_filter(filters, column_types, allowed_fields)?;
313
314 if !search_fields.is_empty() {
316 query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
317 }
318
319 self.execute_list(pagination, query_filter).await
320 }
321
322 #[allow(clippy::type_complexity)]
336 async fn execute_list(
337 &self,
338 pagination: PaginationParams,
339 mut query_filter: crate::QueryFilter,
340 ) -> anyhow::Result<PaginatedResult<T>> {
341 let limit = pagination.limit() as i64;
342 let backwards =
343 query_filter.cursor_before.is_some() && query_filter.cursor_after.is_none();
344 let cursor_walk = query_filter.cursor_after.is_some() || backwards;
345
346 let (mut where_clause, mut filter_params) = query_filter.build_where_clause();
347 let order_clause;
348 let mut boundary_sorts: Vec<(String, FilterSortDirection)> = Vec::new();
350 let mut boundary_casts: Vec<Option<String>> = Vec::new();
352
353 let mut sorts: Vec<(String, FilterSortDirection)> = query_filter
358 .sorts
359 .iter()
360 .map(|s| (s.field.clone(), s.direction.clone()))
361 .collect();
362 if cursor_walk || !sorts.is_empty() {
363 if sorts.is_empty() {
364 sorts.push(("id".into(), FilterSortDirection::Asc));
365 } else if sorts.last().map(|(f, _)| f != "id").unwrap_or(true) {
366 sorts.push(("id".into(), FilterSortDirection::Asc));
367 }
368 }
369
370 if !sorts.is_empty() {
373 boundary_casts = self.sort_column_casts(&sorts).await?;
374 }
375
376 if cursor_walk {
377 let opaque = if backwards {
378 query_filter.cursor_before.clone().unwrap()
379 } else {
380 query_filter.cursor_after.clone().unwrap()
381 };
382 let payload = crate::filter::cursor::decode_cursor(&opaque, &sorts)
383 .map_err(|e| anyhow::anyhow!("cursor refused: {e}"))?;
384 let mut idx = filter_params.len() + 1;
385 let (keyset_sql, keyset_params) = crate::filter::cursor::build_keyset_predicate(
386 &payload,
387 &mut idx,
388 &boundary_casts,
389 backwards,
390 );
391 if where_clause.is_empty() {
392 where_clause = format!(" WHERE {}", keyset_sql);
393 } else {
394 where_clause = format!("{} AND ({})", where_clause, keyset_sql);
395 }
396 filter_params.extend(keyset_params);
397 let parts: Vec<String> = sorts
398 .iter()
399 .map(|(f, d)| {
400 let dir = if (*d == FilterSortDirection::Desc) != backwards {
401 "DESC"
402 } else {
403 "ASC"
404 };
405 format!("{} {}", f, dir)
406 })
407 .collect();
408 order_clause = format!(" ORDER BY {}", parts.join(", "));
409 } else {
410 if sorts.is_empty() {
413 order_clause = query_filter.build_order_by_clause();
414 } else {
415 let parts: Vec<String> = sorts
416 .iter()
417 .map(|(f, d)| {
418 let dir =
419 if *d == FilterSortDirection::Desc { "DESC" } else { "ASC" };
420 format!("{} {}", f, dir)
421 })
422 .collect();
423 order_clause = format!(" ORDER BY {}", parts.join(", "));
424 }
425 }
426 boundary_sorts = sorts;
427
428 let (total, count_mode) = if query_filter.estimate_total {
431 (
432 self.estimate_filtered_rows(&where_clause, &filter_params).await?,
433 "estimate",
434 )
435 } else if cursor_walk {
436 (0u64, "none")
437 } else {
438 let count_query = format!("SELECT COUNT(*) FROM {}{}", self.table_name, where_clause);
439 let mut count_query_builder = sqlx::query_scalar::<_, i64>(&count_query);
440 for param in &filter_params {
441 count_query_builder = count_query_builder.bind(param);
442 }
443 (
444 crate::company_scope::fetch_one_scalar_scoped(&self.pool, count_query_builder)
445 .await? as u64,
446 "exact",
447 )
448 };
449
450 let fetch = limit + 1;
453 let data_query = if cursor_walk {
454 format!(
455 "SELECT * FROM {}{}{} LIMIT {}",
456 self.table_name, where_clause, order_clause, fetch
457 )
458 } else {
459 format!(
460 "SELECT * FROM {}{}{} LIMIT {} OFFSET {}",
461 self.table_name,
462 where_clause,
463 order_clause,
464 fetch,
465 pagination.offset()
466 )
467 };
468
469 let mut pagination_info = PaginationInfo::new(pagination.page, pagination.per_page, total);
470 pagination_info.count_mode = count_mode.to_string();
471
472 let mut rows_query = sqlx::query(&data_query);
477 for param in &filter_params {
478 rows_query = rows_query.bind(param);
479 }
480 let rows: Vec<PgRow> =
481 crate::company_scope::fetch_all_rows_scoped(&self.pool, rows_query).await?;
482 let has_more = rows.len() as i64 > limit;
483 let mut page: Vec<PgRow> = rows.into_iter().take(limit as usize).collect();
484 if backwards {
485 page.reverse();
486 }
487 let data: anyhow::Result<Vec<T>> = page
488 .iter()
489 .map(|row| T::from_row(row).map_err(|e| anyhow::anyhow!("decode row: {e}")))
490 .collect();
491 let data = data?;
492
493 let deterministic = !boundary_sorts.is_empty();
497 let next_cursor = if (has_more || backwards) && deterministic {
498 page.last().and_then(|r| {
499 let casts = if boundary_casts.is_empty() {
500 &boundary_casts
503 } else {
504 &boundary_casts
505 };
506 self.row_cursor(r, &boundary_sorts, casts)
507 })
508 } else {
509 None
510 };
511 let prev_cursor = if !page.is_empty() && deterministic {
512 page.first().and_then(|r| self.row_cursor(r, &boundary_sorts, &boundary_casts))
513 } else {
514 None
515 };
516
517 pagination_info.has_more = Some(has_more);
518 pagination_info.next_cursor = next_cursor;
519 pagination_info.prev_cursor = prev_cursor;
520
521 Ok(PaginatedResult {
522 data,
523 pagination: pagination_info,
524 })
525 }
526
527 async fn estimate_filtered_rows(
531 &self,
532 where_clause: &str,
533 filter_params: &[String],
534 ) -> anyhow::Result<u64> {
535 let explain = format!("EXPLAIN (FORMAT JSON) SELECT 1 FROM {}{}", self.table_name, where_clause);
536 let mut builder = sqlx::query_scalar::<_, serde_json::Value>(&explain);
537 for param in filter_params {
538 builder = builder.bind(param);
539 }
540 let plan: serde_json::Value =
541 crate::company_scope::fetch_one_scalar_scoped(&self.pool, builder).await?;
542 let rows = plan
543 .as_array()
544 .and_then(|a| a.first())
545 .and_then(|top| top.get("Plan"))
546 .and_then(|p| p.get("Plan Rows"))
547 .and_then(|r| r.as_i64())
548 .unwrap_or(0);
549 Ok(rows.max(0) as u64)
550 }
551
552 async fn sort_column_casts(
557 &self,
558 sorts: &[(String, FilterSortDirection)],
559 ) -> anyhow::Result<Vec<Option<String>>> {
560 let (schema, table) = match self.table_name.rsplit_once('.') {
561 Some((s, t)) => (s.to_string(), t.to_string()),
562 None => ("public".to_string(), self.table_name.clone()),
563 };
564 let mut casts: Vec<Option<String>> = Vec::with_capacity(sorts.len());
565 for (field, _) in sorts {
566 let row: Option<(String, String)> = sqlx::query_as(
567 "SELECT data_type, coalesce(udt_name, '') FROM information_schema.columns \
568 WHERE table_schema = $1 AND table_name = $2 AND column_name = $3",
569 )
570 .bind(&schema)
571 .bind(&table)
572 .bind(field)
573 .fetch_optional(&self.pool)
574 .await?;
575 let cast = row.map(|(data_type, udt)| cast_suffix(&data_type, &udt)).flatten();
576 casts.push(cast);
577 }
578 Ok(casts)
579 }
580
581 fn row_cursor(
585 &self,
586 row: &PgRow,
587 sorts: &[(String, FilterSortDirection)],
588 casts: &[Option<String>],
589 ) -> Option<String> {
590 let id: uuid::Uuid = row.try_get("id").ok()?;
591 let mut values: Vec<serde_json::Value> = Vec::with_capacity(sorts.len());
592 for (i, (field, _)) in sorts.iter().enumerate() {
593 let field = field.as_str();
594 let mut data_type = casts.get(i).and_then(|c| c.as_deref()).unwrap_or("");
595 if field == "id" && data_type.is_empty() {
599 data_type = "uuid";
600 }
601 let text: Option<String> = match data_type {
602 "numeric" => row
603 .try_get::<Option<sqlx::types::Decimal>, _>(field)
604 .ok()?
605 .map(|d| d.to_string()),
606 "uuid" => row
607 .try_get::<Option<uuid::Uuid>, _>(field)
608 .ok()?
609 .map(|u| u.to_string()),
610 "timestamptz" => row
611 .try_get::<Option<chrono::DateTime<chrono::Utc>>, _>(field)
612 .ok()?
613 .map(|t| t.to_rfc3339()),
614 "integer" | "smallint" => row
615 .try_get::<Option<i32>, _>(field)
616 .ok()?
617 .map(|n| n.to_string()),
618 "bigint" => row
619 .try_get::<Option<i64>, _>(field)
620 .ok()?
621 .map(|n| n.to_string()),
622 "boolean" => row
623 .try_get::<Option<bool>, _>(field)
624 .ok()?
625 .map(|b| b.to_string()),
626 "date" => row
627 .try_get::<Option<chrono::NaiveDate>, _>(field)
628 .ok()?
629 .map(|d| d.to_string()),
630 _ => row
631 .try_get::<Option<String>, _>(field)
632 .ok()?
633 .filter(|s| !s.is_empty() || data_type.is_empty()),
634 };
635 values.push(serde_json::Value::String(text?));
636 }
637 crate::filter::cursor::encode_cursor(sorts, &values, &id.to_string()).ok()
638 }
639}
640
641fn cast_suffix(data_type: &str, udt_name: &str) -> Option<String> {
644 match data_type {
645 "uuid" => Some("uuid".into()),
646 "numeric" => Some("numeric".into()),
647 "integer" => Some("integer".into()),
648 "smallint" => Some("smallint".into()),
649 "bigint" => Some("bigint".into()),
650 "boolean" => Some("boolean".into()),
651 "date" => Some("date".into()),
652 "timestamp with time zone" => Some("timestamptz".into()),
653 "timestamp without time zone" => Some("timestamp".into()),
654 "USER-DEFINED" if !udt_name.is_empty() => Some(udt_name.to_string()),
657 _ => None,
658 }
659}
660
661fn explain_unknown_column(error: sqlx::Error, table: &str) -> anyhow::Error {
669 let text = error.to_string();
670 if text.contains("does not exist") && text.contains("column") {
671 return anyhow::Error::new(error).context(format!(
672 "insert into {table} named a column that does not exist: the entity serializes a field \
673 with no matching column. Every serialized field must be a column of the table (rename \
674 it, map it with #[serde(rename)], or skip it with #[serde(skip)])"
675 ));
676 }
677 anyhow::Error::new(error)
678}
679
680fn quote_ident(name: &str) -> String {
687 format!("\"{}\"", name.replace('"', "\"\""))
688}
689
690#[async_trait]
691impl<T> DatabaseOperations<T> for PostgresRepository<T>
692where
693 T: for<'a> FromRow<'a, PgRow> + Send + Sync + Unpin + Serialize,
694{
695 async fn create(&self, entity: &T) -> anyhow::Result<T> {
696 let json_value = serde_json::to_value(entity)?;
698
699 let json_obj = match json_value {
700 Value::Object(obj) => obj,
701 _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
702 };
703
704 let json_str = serde_json::to_string(&json_obj)?;
707
708 let insert_columns: Vec<String> = json_obj.keys().map(|k| quote_ident(k)).collect();
722
723 let query = if insert_columns.is_empty() {
724 format!("INSERT INTO {table} DEFAULT VALUES RETURNING *", table = self.table_name)
727 } else {
728 let columns = insert_columns.join(", ");
729 format!(
730 r#"
731 INSERT INTO {table} ({columns})
732 SELECT {columns} FROM jsonb_populate_record(NULL::{table}, $1::jsonb)
733 RETURNING *
734 "#,
735 table = self.table_name,
736 columns = columns
737 )
738 };
739
740 let statement = if insert_columns.is_empty() {
742 sqlx::query_as::<_, T>(&query)
743 } else {
744 sqlx::query_as::<_, T>(&query).bind(&json_str)
745 };
746 let result = crate::company_scope::fetch_one_scoped(&self.pool, statement)
747 .await
748 .map_err(|e| explain_unknown_column(e, &self.table_name))?;
749
750 Ok(result)
751 }
752
753 async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>> {
754 let query = format!("SELECT * FROM {} WHERE id = $1::uuid", self.table_name);
756 let result = crate::company_scope::fetch_optional_scoped(
757 &self.pool,
758 sqlx::query_as::<Postgres, T>(&query).bind(id),
759 )
760 .await?;
761 Ok(result)
762 }
763
764 async fn find_all(&self) -> anyhow::Result<Vec<T>> {
765 let query = format!("SELECT * FROM {}", self.table_name);
766 let results = crate::company_scope::fetch_all_scoped(
767 &self.pool,
768 sqlx::query_as::<Postgres, T>(&query),
769 )
770 .await?;
771 Ok(results)
772 }
773
774 async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>> {
775 let json_value = serde_json::to_value(entity)?;
777
778 let json_obj = match json_value {
779 Value::Object(obj) => obj,
780 _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
781 };
782
783 let update_columns: Vec<&String> = json_obj.keys()
785 .filter(|k| *k != "id")
786 .collect();
787
788 let column_names = update_columns.iter()
789 .map(|k| quote_ident(k))
790 .collect::<Vec<_>>()
791 .join(", ");
792
793 let json_str = serde_json::to_string(&json_obj)?;
794
795 let query = format!(
797 r#"
798 WITH new_row AS (
799 SELECT (jsonb_populate_record(NULL::{table}, $1::jsonb)).*
800 )
801 UPDATE {table} AS t
802 SET ({columns}) = (SELECT {columns} FROM new_row)
803 WHERE t.id = $2::uuid
804 RETURNING t.*
805 "#,
806 table = self.table_name,
807 columns = column_names
808 );
809
810 let result = crate::company_scope::fetch_optional_scoped(
811 &self.pool,
812 sqlx::query_as::<_, T>(&query).bind(&json_str).bind(id),
813 )
814 .await?;
815
816 Ok(result)
817 }
818
819 async fn delete(&self, id: &str) -> anyhow::Result<bool> {
820 let query = format!("DELETE FROM {} WHERE id = $1::uuid", self.table_name);
821 let result = crate::company_scope::execute_scoped(
822 &self.pool,
823 sqlx::query(&query).bind(id),
824 )
825 .await?;
826 Ok(result.rows_affected() > 0)
827 }
828
829 async fn count(&self) -> anyhow::Result<u64> {
830 let query = format!("SELECT COUNT(*) FROM {}", self.table_name);
831 let count = crate::company_scope::fetch_one_scalar_scoped(
832 &self.pool,
833 sqlx::query_scalar::<_, i64>(&query),
834 )
835 .await? as u64;
836 Ok(count)
837 }
838
839 async fn exists(&self, id: &str) -> anyhow::Result<bool> {
840 let query = format!("SELECT 1 FROM {} WHERE id = $1::uuid LIMIT 1", self.table_name);
841 let result = crate::company_scope::fetch_optional_scalar_scoped(
842 &self.pool,
843 sqlx::query_scalar::<_, i32>(&query).bind(id),
844 )
845 .await?;
846 Ok(result.is_some())
847 }
848
849 async fn execute_query(&self, query: &str) -> anyhow::Result<u64> {
850 let result = crate::company_scope::execute_scoped(
851 &self.pool,
852 sqlx::query(query),
853 )
854 .await?;
855 Ok(result.rows_affected())
856 }
857}
858
859#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
863pub enum AggregateFn {
864 Sum,
865 Avg,
866 Min,
867 Max,
868}
869
870impl AggregateFn {
871 fn sql(self) -> &'static str {
872 match self {
873 AggregateFn::Sum => "SUM",
874 AggregateFn::Avg => "AVG",
875 AggregateFn::Min => "MIN",
876 AggregateFn::Max => "MAX",
877 }
878 }
879
880 fn label(self) -> &'static str {
881 match self {
882 AggregateFn::Sum => "sum",
883 AggregateFn::Avg => "avg",
884 AggregateFn::Min => "min",
885 AggregateFn::Max => "max",
886 }
887 }
888
889 fn requires_numeric(self) -> bool {
892 matches!(self, AggregateFn::Sum | AggregateFn::Avg)
893 }
894}
895
896#[derive(Debug, Clone, Default)]
898pub struct AggregateSpec {
899 pub group_by: Option<String>,
901 pub reductions: Vec<(AggregateFn, String)>,
903 pub group_limit: usize,
905 pub label_field: Option<String>,
909 pub label_relation: Option<(String, String)>,
913}
914
915pub const DEFAULT_GROUP_LIMIT: usize = 200;
923
924#[derive(Debug, Clone, Serialize, Deserialize)]
927pub struct AggregateGroup {
928 pub key: Option<String>,
929 pub label: Option<String>,
932 pub count: u64,
933 pub values: HashMap<String, Option<String>>,
941}
942
943#[derive(Debug, Clone, Serialize, Deserialize)]
945pub struct AggregateResult {
946 pub groups: Vec<AggregateGroup>,
947 pub total: AggregateGroup,
948 pub truncated: bool,
950}
951
952#[derive(Debug, Clone)]
955pub struct AggregateFieldError(pub String);
956
957impl std::fmt::Display for AggregateFieldError {
958 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
959 f.write_str(&self.0)
960 }
961}
962
963impl std::error::Error for AggregateFieldError {}
964
965fn is_numeric_pg_type(pg_type: &str) -> bool {
967 let t = pg_type.trim().to_ascii_lowercase();
968 let t = t.split('(').next().unwrap_or(&t).trim();
969 matches!(
970 t,
971 "numeric" | "decimal" | "money"
972 | "smallint" | "int2" | "integer" | "int" | "int4" | "bigint" | "int8"
973 | "real" | "float4" | "double precision" | "float8"
974 | "smallserial" | "serial" | "bigserial"
975 )
976}
977
978async fn catalog_columns(
987 pool: &PgPool,
988 qualified_table: &str,
989) -> anyhow::Result<HashMap<String, String>> {
990 let (schema, table) = match qualified_table.split_once('.') {
991 Some((s, t)) => (s.to_string(), t.to_string()),
992 None => ("public".to_string(), qualified_table.to_string()),
993 };
994 let q = sqlx::query_as::<Postgres, (String, String)>(
995 "SELECT column_name, data_type FROM information_schema.columns
996 WHERE table_schema = $1 AND table_name = $2",
997 )
998 .bind(schema)
999 .bind(table);
1000 let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
1001 Ok(rows.into_iter().collect())
1002}
1003
1004fn resolve_column<'a>(
1013 name: &str,
1014 column_types: &'a HashMap<String, String>,
1015) -> Result<(&'a str, &'a str), AggregateFieldError> {
1016 column_types
1017 .get_key_value(name)
1018 .map(|(k, v)| (k.as_str(), v.as_str()))
1019 .ok_or_else(|| {
1020 AggregateFieldError(format!("unknown field `{name}` — not a column of this entity"))
1021 })
1022}
1023
1024impl<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> PostgresRepository<T> {
1025 pub async fn aggregate_with_filters(
1033 &self,
1034 spec: &AggregateSpec,
1035 filters: &HashMap<String, String>,
1036 column_types: &HashMap<String, String>,
1037 search_fields: &[&str],
1038 ) -> anyhow::Result<AggregateResult> {
1039 let mut query_filter = parse_query_filter(filters, column_types, None)?;
1040 if !search_fields.is_empty() {
1041 query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
1042 }
1043 query_filter.limit = None;
1046 query_filter.offset = None;
1047 let (where_clause, filter_params) = query_filter.build_where_clause();
1048
1049 let columns = catalog_columns(&self.pool, &self.table_name).await?;
1051 let mut selects: Vec<String> = Vec::new();
1052 let mut value_keys: Vec<String> = Vec::new();
1053 for (func, field) in &spec.reductions {
1054 let (column, pg_type) = resolve_column(field, &columns)?;
1055 if func.requires_numeric() && !is_numeric_pg_type(pg_type) {
1056 return Err(AggregateFieldError(format!(
1057 "cannot {} `{}`: its type is {} — {} needs a numeric column",
1058 func.label(),
1059 column,
1060 pg_type,
1061 func.label()
1062 ))
1063 .into());
1064 }
1065 let key = format!("{}:{}", func.label(), column);
1066 selects.push(format!("{}({})::text AS \"{}\"", func.sql(), column, key));
1069 value_keys.push(key);
1070 }
1071
1072 let group_limit = if spec.group_limit == 0 { DEFAULT_GROUP_LIMIT } else { spec.group_limit };
1073 let reductions = if selects.is_empty() { String::new() } else { format!(", {}", selects.join(", ")) };
1074
1075 let mut label_select = String::new();
1081 let mut label_join = String::new();
1082 if let (Some(field), Some(label_field), Some((rel_table, base_fk))) = (
1083 &spec.group_by,
1084 &spec.label_field,
1085 &spec.label_relation,
1086 ) {
1087 let _ = field;
1088 let qualified = qualify_relation_table(&self.table_name, rel_table);
1089 let rel_columns = catalog_columns(&self.pool, &qualified).await?;
1090 let (label_col, _) = resolve_column(label_field, &rel_columns).map_err(|_| {
1091 AggregateFieldError(format!(
1092 "cannot label groups by `{label_field}`: the related table `{qualified}` has no such column"
1093 ))
1094 })?;
1095 label_select = format!(", (label_rel.{label_col})::text AS __group_label");
1096 label_join = format!(
1097 " LEFT JOIN {qualified} AS label_rel ON label_rel.id IS NOT DISTINCT FROM {base_fk}"
1098 );
1099 }
1100
1101 let sql = match &spec.group_by {
1102 Some(field) => {
1103 let (column, _) = resolve_column(field, &columns)?;
1104 format!(
1105 "SELECT GROUPING({column}) AS __is_total, ({column})::text AS __group_key{label_select}, \
1106 COUNT(*) AS __count{reductions} \
1107 FROM {table}{label_join}{where_clause} \
1108 GROUP BY GROUPING SETS (({column}), ()) \
1109 ORDER BY __is_total DESC, __count DESC \
1110 LIMIT {limit}",
1111 column = column,
1112 label_select = label_select,
1113 label_join = label_join,
1114 reductions = reductions,
1115 table = self.table_name,
1116 where_clause = where_clause,
1117 limit = group_limit + 2,
1120 )
1121 }
1122 None => format!(
1123 "SELECT 1 AS __is_total, NULL::text AS __group_key, COUNT(*) AS __count{reductions} \
1124 FROM {table}{where_clause}",
1125 reductions = reductions,
1126 table = self.table_name,
1127 where_clause = where_clause,
1128 ),
1129 };
1130
1131 let mut builder = sqlx::query(&sql);
1132 for param in &filter_params {
1133 builder = builder.bind(param);
1134 }
1135 let rows = crate::company_scope::fetch_all_rows_scoped(&self.pool, builder).await?;
1136
1137 let read_group = |row: &PgRow| -> AggregateGroup {
1138 use sqlx::Row as _;
1139 let mut values = HashMap::with_capacity(value_keys.len());
1140 for key in &value_keys {
1141 values.insert(key.clone(), row.try_get::<Option<String>, _>(key.as_str()).ok().flatten());
1142 }
1143 AggregateGroup {
1144 key: row.try_get::<Option<String>, _>("__group_key").ok().flatten(),
1145 label: row.try_get::<Option<String>, _>("__group_label").ok().flatten(),
1146 count: row.try_get::<i64, _>("__count").unwrap_or(0).max(0) as u64,
1147 values,
1148 }
1149 };
1150
1151 use sqlx::Row as _;
1152 let mut total: Option<AggregateGroup> = None;
1153 let mut groups: Vec<AggregateGroup> = Vec::new();
1154 for row in &rows {
1155 let is_total = row.try_get::<i32, _>("__is_total").unwrap_or(0) == 1;
1156 if is_total {
1157 total = Some(read_group(row));
1159 } else {
1160 groups.push(read_group(row));
1161 }
1162 }
1163
1164 let truncated = groups.len() > group_limit;
1165 groups.truncate(group_limit);
1166
1167 let total = total.unwrap_or_else(|| AggregateGroup {
1170 key: None,
1171 label: None,
1172 count: 0,
1173 values: value_keys.iter().map(|k| (k.clone(), None)).collect(),
1174 });
1175
1176 Ok(AggregateResult { groups, total, truncated })
1177 }
1178}
1179
1180#[cfg(test)]
1181mod aggregate_field_tests {
1182 use super::*;
1183
1184 fn columns() -> HashMap<String, String> {
1185 [
1186 ("status", "text"),
1187 ("total", "numeric"),
1188 ("qty", "integer"),
1189 ("notes", "text"),
1190 ]
1191 .iter()
1192 .map(|(k, v)| (k.to_string(), v.to_string()))
1193 .collect()
1194 }
1195
1196 #[test]
1197 fn resolves_only_declared_columns() {
1198 let cols = columns();
1199 assert_eq!(resolve_column("total", &cols).unwrap().0, "total");
1200 assert!(resolve_column("password_hash", &cols).is_err());
1201 }
1202
1203 #[test]
1206 fn rejects_injection_attempts_rather_than_escaping_them() {
1207 let cols = columns();
1208 for probe in [
1209 "total) FROM selling.sales_orders; DROP TABLE users --",
1210 "status\"",
1211 "1=1",
1212 "total, (SELECT password FROM users)",
1213 "",
1214 ] {
1215 assert!(
1216 resolve_column(probe, &cols).is_err(),
1217 "`{probe}` must be refused, never escaped into the query"
1218 );
1219 }
1220 }
1221
1222 #[test]
1225 fn returns_the_declared_key_not_the_callers_string() {
1226 let cols = columns();
1227 let (name, _) = resolve_column("total", &cols).unwrap();
1228 assert!(std::ptr::eq(name, cols.get_key_value("total").unwrap().0.as_str()));
1229 }
1230
1231 #[test]
1232 fn sum_and_avg_require_a_numeric_type() {
1233 assert!(AggregateFn::Sum.requires_numeric());
1234 assert!(AggregateFn::Avg.requires_numeric());
1235 assert!(!AggregateFn::Min.requires_numeric());
1237 assert!(!AggregateFn::Max.requires_numeric());
1238 }
1239
1240 #[test]
1241 fn recognises_the_numeric_postgres_types() {
1242 for t in ["numeric", "NUMERIC(14,2)", "integer", "bigint", "double precision", "money"] {
1243 assert!(is_numeric_pg_type(t), "{t} should count as numeric");
1244 }
1245 for t in ["text", "uuid", "timestamptz", "boolean", "jsonb", "USER-DEFINED"] {
1246 assert!(!is_numeric_pg_type(t), "{t} must not accept a SUM");
1247 }
1248 }
1249}