Skip to main content

backbone_orm/
repository.rs

1//! Repository implementations for PostgreSQL with comprehensive CRUD operations
2
3use 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
15/// Generic entity trait that all repository entities must implement
16pub trait Entity {
17    /// Get the entity's ID
18    fn id(&self) -> Option<&str>;
19
20    /// Get the table name for this entity
21    fn table_name() -> &'static str where Self: Sized;
22
23    /// Check if entity is soft deleted
24    fn is_deleted(&self) -> bool { false }
25
26    /// Get creation timestamp
27    fn created_at(&self) -> Option<NaiveDateTime> { None }
28
29    /// Get update timestamp
30    fn updated_at(&self) -> Option<NaiveDateTime> { None }
31}
32
33/// Pagination parameters
34#[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), // Limit to 1-100 per page
45        }
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/// Sorting parameters
58#[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/// Filter parameters
72#[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/// Paginated result wrapper
90#[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    /// Keyset paging: the position of this page's last row, for the next
103    /// page's `after=`. None when the page is empty or no more rows follow.
104    #[serde(default, skip_serializing_if = "Option::is_none")]
105    pub next_cursor: Option<String>,
106    /// Keyset paging: the position of this page's first row, for the
107    /// previous page's `before=`.
108    #[serde(default, skip_serializing_if = "Option::is_none")]
109    pub prev_cursor: Option<String>,
110    /// Whether another page follows (fetched with limit+1, so it is known
111    /// without a count).
112    #[serde(default, skip_serializing_if = "Option::is_none")]
113    pub has_more: Option<bool>,
114    /// How `total` came to be: "exact" (counted), "estimate" (the planner's
115    /// row estimate — `estimate=1` asked for it), or "none" (a cursor walk
116    /// without an estimate; the exact figure is the separate count call).
117    #[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/// Database operations trait - requires Serialize for write operations
142#[async_trait]
143pub trait DatabaseOperations<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> {
144    /// Create a new entity
145    async fn create(&self, entity: &T) -> anyhow::Result<T>;
146
147    /// Find entity by ID
148    async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>>;
149
150    /// Find all entities
151    async fn find_all(&self) -> anyhow::Result<Vec<T>>;
152
153    /// Update an existing entity
154    async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>>;
155
156    /// Delete an entity
157    async fn delete(&self, id: &str) -> anyhow::Result<bool>;
158
159    /// Count all entities
160    async fn count(&self) -> anyhow::Result<u64>;
161
162    /// Check if entity exists
163    async fn exists(&self, id: &str) -> anyhow::Result<bool>;
164
165    /// Execute custom query
166    async fn execute_query(&self, query: &str) -> anyhow::Result<u64>;
167}
168
169/// PostgreSQL repository implementation with JSON-based dynamic queries
170pub 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    /// List entities with pagination and advanced filtering
194    ///
195    /// This method provides comprehensive filtering capabilities similar to Laravel's Filter Query String.
196    ///
197    /// # Supported Filter Operators
198    ///
199    /// - `field[eq]=value` - Equal
200    /// - `field[notEq]=value` - Not equal
201    /// - `field[gt]=value` - Greater than
202    /// - `field[gte]=value` - Greater than or equal
203    /// - `field[lt]=value` - Less than
204    /// - `field[lte]=value` - Less than or equal
205    /// - `field[like]=value` - LIKE (case-sensitive)
206    /// - `field[ilike]=value` - ILIKE (case-insensitive)
207    /// - `field[notlike]=value` - NOT LIKE
208    /// - `field[contain]=value` - Contains (%value%)
209    /// - `field[notcontain]=value` - Does not contain
210    /// - `field[startwith]=value` - Starts with (value%)
211    /// - `field[endwith]=value` - Ends with (%value)
212    /// - `field[in]=val1,val2` - IN array
213    /// - `field[notin]=val1,val2` - NOT IN array
214    /// - `field[between]=val1,val2` - BETWEEN
215    /// - `field[notbetween]=val1,val2` - NOT BETWEEN
216    /// - `field[isnull]` - IS NULL
217    /// - `field[isnotnull]` - IS NOT NULL
218    ///
219    /// # Special Parameters
220    ///
221    /// - `search=value&searchFields=field1,field2` - Search in multiple fields
222    /// - `orderby=field` or `orderby[field]=asc` - Sort results
223    /// - `limit=10` - Limit results
224    /// - `page=1` - Page number
225    ///
226    /// # Column Type Casting
227    ///
228    /// The `column_types` HashMap maps field names to their PostgreSQL types for proper casting.
229    /// For example, `{"status": "user_status"}` will cast the status parameter to `user_status` enum type.
230    ///
231    /// # Example
232    ///
233    /// ```ignore
234    /// let mut filters = HashMap::new();
235    /// filters.insert("username[contain]".to_string(), "john".to_string());
236    /// filters.insert("age[gt]".to_string(), "18".to_string());
237    ///
238    /// let mut column_types = HashMap::new();
239    /// column_types.insert("status".to_string(), "user_status".to_string());
240    ///
241    /// let result = repo.list_with_filters(
242    ///     PaginationParams::new(1, 10),
243    ///     &filters,
244    ///     &column_types,
245    ///     &["username", "email"]  // search fields
246    /// ).await?;
247    /// ```
248    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        // Parse filters from HashMap (no field allow-list by default for backward compatibility)
259        let mut query_filter = parse_query_filter(filters, column_types, None)?;
260
261        // Set up search fields if provided
262        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    /// List entities with pagination, filtering, and field whitelist enforcement
270    ///
271    /// Similar to `list_with_filters` but accepts an optional set of allowed field names.
272    /// When provided, only filter conditions on whitelisted fields are applied;
273    /// conditions on unknown fields are silently dropped.
274    ///
275    /// This prevents clients from filtering on internal or sensitive columns
276    /// (e.g., `password_hash`, `internal_notes`).
277    ///
278    /// # Arguments
279    ///
280    /// * `pagination` - Page and limit parameters
281    /// * `filters` - HTTP query parameters (e.g., `field[operator]=value`)
282    /// * `column_types` - PostgreSQL type mappings for enum casting
283    /// * `search_fields` - Fields to search when `search` parameter is present
284    /// * `allowed_fields` - Optional whitelist of field names; `None` allows all fields
285    ///
286    /// # Example
287    ///
288    /// ```ignore
289    /// let allowed: HashSet<String> = ["username", "email", "status"]
290    ///     .iter().map(|s| s.to_string()).collect();
291    ///
292    /// let result = repo.list_with_filters_whitelisted(
293    ///     PaginationParams::new(1, 10),
294    ///     &filters,
295    ///     &column_types,
296    ///     &["username", "email"],
297    ///     Some(&allowed),
298    /// ).await?;
299    /// ```
300    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        // Parse filters with optional field whitelist
312        let mut query_filter = parse_query_filter(filters, column_types, allowed_fields)?;
313
314        // Set up search fields if provided
315        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    /// The shared list execution: filters, order, paging, and the total.
323    ///
324    /// Three shapes, chosen by the request:
325    ///
326    /// * **page mode** (today's behaviour, unchanged): exact COUNT, then
327    ///   `LIMIT l OFFSET o` in the requested order;
328    /// * **cursor mode** (`after=`/`before=`): the keyset predicate replaces
329    ///   the offset, the order always ends on the `id` tiebreaker, and one
330    ///   extra row is fetched so `has_more` is known without a count — the
331    ///   page costs the same at any depth, which is the point;
332    /// * **estimate** (`estimate=1`, either mode): the exact COUNT (a scan
333    ///   of the whole filtered set) is replaced by the planner's row
334    ///   estimate; the exact figure stays the separate count call.
335    #[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        // The deterministic order a cursor walks in (cursor mode only).
349        let mut boundary_sorts: Vec<(String, FilterSortDirection)> = Vec::new();
350        // Cast suffixes for the sort columns, for the keyset binds.
351        let mut boundary_casts: Vec<Option<String>> = Vec::new();
352
353        // The deterministic order: a cursor walks one, and a page-mode list
354        // that carries a sort gets the same treatment so it can HAND OUT a
355        // cursor to start a keyset walk from. Appending the id tiebreaker
356        // only reorders rows that tied — ties had no order to preserve.
357        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        // Casts whenever there is a deterministic order to key: cursor mode
371        // walks on them, page mode encodes the handed-out cursor from them.
372        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            // Page mode with a sort: the caller's order, made deterministic
411            // by the same id tiebreaker so the cursor it hands out is real.
412            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        // The total: exact in page mode (today's behaviour), the planner's
429        // estimate when asked, nothing on a cursor walk that did not ask.
430        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        // The page: limit+1 rows so has_more is known without a count. The
451        // extra row is truncated away before returning.
452        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        // One fetch path for both modes: untyped rows, decoded per-entity
473        // through FromRow exactly as query_as would (decimals keep their
474        // scale), with the boundary values read for the cursors off the same
475        // rows — no second query, no JSON round-trip.
476        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        // Cursors from the boundary rows' own values, whenever the order is
494        // deterministic. A NULL in a sort column cannot key a position, so
495        // that side's cursor is omitted.
496        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                    // Page mode never resolved casts; the encode does not
501                    // need them, only the walk does.
502                    &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    /// The planner's row estimate for a filtered read, from EXPLAIN. The
528    /// exact figure stays the separate count endpoint; this exists because
529    /// counting a hot filtered set is a scan of all of it.
530    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    /// The SQL cast suffix for each sort column, from the table's real
553    /// columns (never a cached hint list — the aggregate lesson). The
554    /// placeholder is cast to the column's type so the bind compares
555    /// against the column without coercing the column itself.
556    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    /// One boundary row's cursor: its values for the sort columns (typed
582    /// exactly as sqlx decodes them, so a decimal keeps its scale) plus its
583    /// id. NULL in any sort column yields None: a null has no position.
584    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            // The ORM appends the id tiebreaker itself and page mode never
596            // resolves casts — without this, the uuid id decodes as text,
597            // the read fails, and the whole cursor silently vanishes.
598            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
641/// The placeholder cast for a column type, or None when a bare text bind
642/// compares correctly.
643fn 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        // An enum: bind text, cast to the enum's own name so the comparison
655        // runs in the enum's ordering.
656        "USER-DEFINED" if !udt_name.is_empty() => Some(udt_name.to_string()),
657        _ => None,
658    }
659}
660
661/// Turn "column ... does not exist" into a sentence that names the cause.
662///
663/// The insert names the columns the entity serializes. A field that is serialized but is not a
664/// column of the table used to vanish quietly, because selecting every column of the row type threw
665/// unknown keys away; now it fails, and the bare Postgres error does not say why. Serialized field
666/// and table column are meant to be the same set — the update path has always assumed it — so this
667/// points at the mismatch rather than leaving someone to guess.
668fn 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
680/// Quote a column name as a SQL identifier.
681///
682/// Column names here come from serializing the caller's entity, so they are Rust field names in
683/// practice — but they are interpolated into DDL/DML, where Postgres has no bind parameter for an
684/// identifier. Doubling an embedded quote is the identifier escape, so a name can never end the
685/// quoted section early.
686fn 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        // Serialize entity to JSON to extract field names and values
697        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        // Build dynamic INSERT query using jsonb_populate_record
705        // This approach handles all PostgreSQL types correctly including ENUMs and booleans
706        let json_str = serde_json::to_string(&json_obj)?;
707
708        // Name only the columns the payload actually carries.
709        //
710        // `SELECT (jsonb_populate_record(...)).*` emits EVERY column of the row type, so a column
711        // the entity does not know about arrived as an explicit NULL — and an explicit NULL is not
712        // an absent value: it overrides the column DEFAULT. That is invisible until a table's
713        // correctness depends on a default, which is exactly what composition-installed tenancy
714        // does (a scoped table defaults `org_unit_id` from the acting unit), so generic creates
715        // over such a table wrote NULL and were refused by the write-path guard.
716        //
717        // Listing the payload's own keys leaves every other column unmentioned, so its default
718        // applies. A key that is present with a JSON null is still written as NULL, which is
719        // right: the caller said so. This mirrors the update path below, which has always built
720        // its column list from these same keys.
721        let insert_columns: Vec<String> = json_obj.keys().map(|k| quote_ident(k)).collect();
722
723        let query = if insert_columns.is_empty() {
724            // Nothing supplied at all: let every column take its default rather than emitting
725            // `INSERT INTO t () SELECT`, which is not valid SQL.
726            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        // The DEFAULT VALUES form takes no bind; every other form binds the payload.
741        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        // Cast text to UUID for PostgreSQL UUID columns
755        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        // Serialize entity to JSON
776        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        // Build column list for the update (excluding 'id')
784        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        // Use jsonb_populate_record with CTE to get properly typed values
796        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// ─── Aggregation ──────────────────────────────────────────────────────────────
860
861/// Which reduction to apply to a column.
862#[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    /// `SUM`/`AVG` on a text column is a type error, not a zero. `MIN`/`MAX`
890    /// order any comparable type, so they carry no such restriction.
891    fn requires_numeric(self) -> bool {
892        matches!(self, AggregateFn::Sum | AggregateFn::Avg)
893    }
894}
895
896/// What to group by and what to reduce — the parsed form of the query string.
897#[derive(Debug, Clone, Default)]
898pub struct AggregateSpec {
899    /// Column whose distinct values become groups. `None` asks for one total.
900    pub group_by: Option<String>,
901    /// `(function, column)` pairs, in the order the caller asked for them.
902    pub reductions: Vec<(AggregateFn, String)>,
903    /// Most groups to return before reporting the answer as truncated.
904    pub group_limit: usize,
905    /// Carry a label per group: the column on the group column's RELATED
906    /// table (resolved through the entity's relation metadata) to show
907    /// instead of a bare uuid key. `None` = keys stay as they are.
908    pub label_field: Option<String>,
909    /// The resolved relation behind the group column — `(target table, the
910    /// BASE table's FK column, snake)`, filled by the generic layer from
911    /// the entity's `relations()` metadata. Callers never set this.
912    pub label_relation: Option<(String, String)>,
913}
914
915/// The default ceiling on distinct groups.
916///
917/// A `group_by` on a uuid or a timestamp yields one group per row, which is a
918/// table scan wearing a chart's clothes. Rather than refusing those columns —
919/// a list that would be wrong for some schema sooner or later — the answer is
920/// capped and the cap is *reported*, so a caller can tell a complete picture
921/// from a partial one instead of quietly drawing the wrong one.
922pub const DEFAULT_GROUP_LIMIT: usize = 200;
923
924/// One group's numbers. `key` is the group's value; `None` is a real answer —
925/// the rows whose group column is null — and is distinct from "no rows".
926#[derive(Debug, Clone, Serialize, Deserialize)]
927pub struct AggregateGroup {
928    pub key: Option<String>,
929    /// The group's display label (the related row's label column), when the
930    /// caller asked for one and the group column is a relation FK.
931    pub label: Option<String>,
932    pub count: u64,
933    /// Reduction results keyed `"sum:amount"`, carried as strings.
934    ///
935    /// Postgres `numeric` holds more precision than an IEEE double, and money
936    /// columns are exactly where that bites: a tenant large enough for the
937    /// total to matter is a tenant large enough to round it. The string is the
938    /// exact value Postgres computed; the caller decides how to parse it.
939    /// `None` is SQL NULL — no rows contributed — which is not zero.
940    pub values: HashMap<String, Option<String>>,
941}
942
943/// Groups plus the overall total, computed together.
944#[derive(Debug, Clone, Serialize, Deserialize)]
945pub struct AggregateResult {
946    pub groups: Vec<AggregateGroup>,
947    pub total: AggregateGroup,
948    /// True when more distinct groups exist than `group_limit` allowed.
949    pub truncated: bool,
950}
951
952/// A column name rejected by the allow-list, or a reduction that its type
953/// cannot answer.
954#[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
965/// True for the Postgres types `SUM`/`AVG` accept.
966fn 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
978/// Read a table's real columns and their types from the catalog.
979///
980/// `EntityRepoMeta::column_types()` looks like the natural allow-list and is
981/// not one: it carries only the columns the filter parser must CAST — uuids and
982/// enums — so every numeric column is absent from it, which is exactly the set
983/// `sum` and `avg` exist for. The catalog is the only complete and current
984/// answer, and it cannot drift from the table the query will actually run
985/// against.
986async 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
1004/// Resolve a caller-supplied column name against the entity's real columns.
1005///
1006/// This is the whole defence for the aggregate path. Unlike a filter *value*,
1007/// which is bound as a parameter, a `group_by` or `sum` column is spliced into
1008/// the SQL as an identifier — binding cannot protect it. So the name is never
1009/// escaped or quoted into safety; it is *replaced* by the matching key already
1010/// present in the entity's declared column map, and a name with no match is
1011/// refused. Nothing a caller types can reach the query text.
1012fn 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    /// Group and reduce rows in one statement, under the same filters, the same
1026    /// soft-delete convention and the same tenancy fence as the list endpoint.
1027    ///
1028    /// Groups and the overall total come back from a single `GROUPING SETS`
1029    /// query, which is what lets a caller draw a chart and its headline from
1030    /// one reply, and what keeps the two numbers consistent — a separate total
1031    /// query could observe a different set of rows.
1032    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        // Grouping replaces row output entirely: paging and ordering describe a
1044        // page of rows, and there are none.
1045        query_filter.limit = None;
1046        query_filter.offset = None;
1047        let (where_clause, filter_params) = query_filter.build_where_clause();
1048
1049        // Every identifier below comes from the catalog, never from the caller.
1050        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            // Cast to text in SQL so the exact value Postgres computed is what
1067            // crosses the wire — see `AggregateGroup::values`.
1068            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        // The group's label: when the caller names one (group_label) and the
1076        // group column is a relation FK of THIS entity, LEFT JOIN the
1077        // related table and carry its label column beside the key — the
1078        // aggregate's equivalent of `?include=`, which has no row to
1079        // hydrate otherwise.
1080        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                    // One total row, the groups themselves, and one more to
1118                    // detect that a further group existed.
1119                    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                // Ordered first, so it survives the cap.
1158                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        // No rows at all means no total row either: an empty result is a real
1168        // answer of zero, not a missing one.
1169        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    /// The allow-list is the entire defence, because a group/sum column is
1204    /// spliced into SQL as an identifier and cannot be bound as a parameter.
1205    #[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    /// The returned name is the map's own key, not the caller's string, so no
1223    /// caller-controlled bytes can reach the SQL even on a match.
1224    #[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        // Ordering works on any comparable column, so these stay open.
1236        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}