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 = self.parse_typed_filters(filters, column_types, None).await?;
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 =
313            self.parse_typed_filters(filters, column_types, allowed_fields).await?;
314
315        // Set up search fields if provided
316        if !search_fields.is_empty() {
317            query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
318        }
319
320        self.execute_list(pagination, query_filter).await
321    }
322
323    /// Parse the wire filters with a cast for every typed column they compare.
324    ///
325    /// Filter values arrive as text and are bound as text; PostgreSQL has no implicit comparison
326    /// between text and a boolean, number, uuid, date, time or timestamp column, so each such
327    /// comparison needs a cast on its placeholder. The entity's generated `column_types()` hints
328    /// supply it where they exist, and stay the answer for the columns they name. They were never
329    /// a complete list — no booleans or numbers, only uuids named `id`/`*_id`, and temporal
330    /// columns only in modules generated after the generator learned them — so whenever a
331    /// comparison is left without a cast, the table's real column types are read from the
332    /// catalog and fill the gaps. The catalog cannot drift from the table the query runs
333    /// against, which is the same reason the aggregate and sort paths read it.
334    ///
335    /// A filter on text columns only, or with hints for every compared column, costs nothing
336    /// extra beyond the parse; otherwise one catalog lookup on the request's own connection.
337    async fn parse_typed_filters(
338        &self,
339        filters: &HashMap<String, String>,
340        column_types: &HashMap<String, String>,
341        allowed_fields: Option<&HashSet<String>>,
342    ) -> anyhow::Result<crate::QueryFilter> {
343        let query_filter = parse_query_filter(filters, column_types, allowed_fields)?;
344        if !query_filter.has_uncast_value_conditions() {
345            return Ok(query_filter);
346        }
347        let catalog = catalog_filter_casts(&self.pool, &self.table_name).await?;
348        let merged = merge_filter_casts(column_types, catalog);
349        parse_query_filter(filters, &merged, allowed_fields)
350    }
351
352    /// The shared list execution: filters, order, paging, and the total.
353    ///
354    /// Three shapes, chosen by the request:
355    ///
356    /// * **page mode** (today's behaviour, unchanged): exact COUNT, then
357    ///   `LIMIT l OFFSET o` in the requested order;
358    /// * **cursor mode** (`after=`/`before=`): the keyset predicate replaces
359    ///   the offset, the order always ends on the `id` tiebreaker, and one
360    ///   extra row is fetched so `has_more` is known without a count — the
361    ///   page costs the same at any depth, which is the point;
362    /// * **estimate** (`estimate=1`, either mode): the exact COUNT (a scan
363    ///   of the whole filtered set) is replaced by the planner's row
364    ///   estimate; the exact figure stays the separate count call.
365    #[allow(clippy::type_complexity)]
366    async fn execute_list(
367        &self,
368        pagination: PaginationParams,
369        query_filter: crate::QueryFilter,
370    ) -> anyhow::Result<PaginatedResult<T>> {
371        let limit = pagination.limit() as i64;
372        let backwards =
373            query_filter.cursor_before.is_some() && query_filter.cursor_after.is_none();
374        let cursor_walk = query_filter.cursor_after.is_some() || backwards;
375
376        let (mut where_clause, mut filter_params) = query_filter.build_where_clause();
377        let order_clause;
378        // Cast suffixes for the sort columns, for the keyset binds.
379        let mut boundary_casts: Vec<Option<String>> = Vec::new();
380
381        // The deterministic order: a cursor walks one, and a page-mode list
382        // that carries a sort gets the same treatment so it can HAND OUT a
383        // cursor to start a keyset walk from. Appending the id tiebreaker
384        // only reorders rows that tied — ties had no order to preserve.
385        let mut sorts: Vec<(String, FilterSortDirection)> = query_filter
386            .sorts
387            .iter()
388            .map(|s| (s.field.clone(), s.direction.clone()))
389            .collect();
390        if cursor_walk || !sorts.is_empty() {
391            if sorts.is_empty() {
392                sorts.push(("id".into(), FilterSortDirection::Asc));
393            } else if sorts.last().map(|(f, _)| f != "id").unwrap_or(true) {
394                sorts.push(("id".into(), FilterSortDirection::Asc));
395            }
396        }
397
398        // Casts whenever there is a deterministic order to key: cursor mode
399        // walks on them, page mode encodes the handed-out cursor from them.
400        if !sorts.is_empty() {
401            boundary_casts = self.sort_column_casts(&sorts).await?;
402        }
403
404        if cursor_walk {
405            let opaque = if backwards {
406                query_filter.cursor_before.clone().unwrap()
407            } else {
408                query_filter.cursor_after.clone().unwrap()
409            };
410            let payload = crate::filter::cursor::decode_cursor(&opaque, &sorts)
411                .map_err(|e| anyhow::anyhow!("cursor refused: {e}"))?;
412            let mut idx = filter_params.len() + 1;
413            let (keyset_sql, keyset_params) = crate::filter::cursor::build_keyset_predicate(
414                &payload,
415                &mut idx,
416                &boundary_casts,
417                backwards,
418            );
419            if where_clause.is_empty() {
420                where_clause = format!(" WHERE {}", keyset_sql);
421            } else {
422                where_clause = format!("{} AND ({})", where_clause, keyset_sql);
423            }
424            filter_params.extend(keyset_params);
425            let parts: Vec<String> = sorts
426                .iter()
427                .map(|(f, d)| {
428                    let dir = if (*d == FilterSortDirection::Desc) != backwards {
429                        "DESC"
430                    } else {
431                        "ASC"
432                    };
433                    format!("{} {}", f, dir)
434                })
435                .collect();
436            order_clause = format!(" ORDER BY {}", parts.join(", "));
437        } else {
438            // Page mode with a sort: the caller's order, made deterministic
439            // by the same id tiebreaker so the cursor it hands out is real.
440            if sorts.is_empty() {
441                order_clause = query_filter.build_order_by_clause();
442            } else {
443                let parts: Vec<String> = sorts
444                    .iter()
445                    .map(|(f, d)| {
446                        let dir =
447                            if *d == FilterSortDirection::Desc { "DESC" } else { "ASC" };
448                        format!("{} {}", f, dir)
449                    })
450                    .collect();
451                order_clause = format!(" ORDER BY {}", parts.join(", "));
452            }
453        }
454        // The deterministic order a cursor walks in (cursor mode only).
455        let boundary_sorts: Vec<(String, FilterSortDirection)> = sorts;
456
457        // The total: exact in page mode (today's behaviour), the planner's
458        // estimate when asked, nothing on a cursor walk that did not ask.
459        let (total, count_mode) = if query_filter.estimate_total {
460            (
461                self.estimate_filtered_rows(&where_clause, &filter_params).await?,
462                "estimate",
463            )
464        } else if cursor_walk {
465            (0u64, "none")
466        } else {
467            let count_query = format!("SELECT COUNT(*) FROM {}{}", self.table_name, where_clause);
468            let mut count_query_builder = sqlx::query_scalar::<_, i64>(&count_query);
469            for param in &filter_params {
470                count_query_builder = count_query_builder.bind(param);
471            }
472            (
473                crate::company_scope::fetch_one_scalar_scoped(&self.pool, count_query_builder)
474                    .await? as u64,
475                "exact",
476            )
477        };
478
479        // The page: limit+1 rows so has_more is known without a count. The
480        // extra row is truncated away before returning.
481        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        // One fetch path for both modes: untyped rows, decoded per-entity
502        // through FromRow exactly as query_as would (decimals keep their
503        // scale), with the boundary values read for the cursors off the same
504        // rows — no second query, no JSON round-trip.
505        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        // Cursors from the boundary rows' own values, whenever the order is
523        // deterministic. A NULL in a sort column cannot key a position, so
524        // that side's cursor is omitted.
525        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                    // Page mode never resolved casts; the encode does not
530                    // need them, only the walk does.
531                    &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    /// The planner's row estimate for a filtered read, from EXPLAIN. The
557    /// exact figure stays the separate count endpoint; this exists because
558    /// counting a hot filtered set is a scan of all of it.
559    async fn estimate_filtered_rows(
560        &self,
561        where_clause: &str,
562        filter_params: &[String],
563    ) -> anyhow::Result<u64> {
564        let explain = format!("EXPLAIN (FORMAT JSON) SELECT 1 FROM {}{}", self.table_name, where_clause);
565        let mut builder = sqlx::query_scalar::<_, serde_json::Value>(&explain);
566        for param in filter_params {
567            builder = builder.bind(param);
568        }
569        let plan: serde_json::Value =
570            crate::company_scope::fetch_one_scalar_scoped(&self.pool, builder).await?;
571        let rows = plan
572            .as_array()
573            .and_then(|a| a.first())
574            .and_then(|top| top.get("Plan"))
575            .and_then(|p| p.get("Plan Rows"))
576            .and_then(|r| r.as_i64())
577            .unwrap_or(0);
578        Ok(rows.max(0) as u64)
579    }
580
581    /// The SQL cast suffix for each sort column, from the table's real
582    /// columns (never a cached hint list — the aggregate lesson). The
583    /// placeholder is cast to the column's type so the bind compares
584    /// against the column without coercing the column itself.
585    async fn sort_column_casts(
586        &self,
587        sorts: &[(String, FilterSortDirection)],
588    ) -> anyhow::Result<Vec<Option<String>>> {
589        let (schema, table) = match self.table_name.rsplit_once('.') {
590            Some((s, t)) => (s.to_string(), t.to_string()),
591            None => ("public".to_string(), self.table_name.clone()),
592        };
593        let mut casts: Vec<Option<String>> = Vec::with_capacity(sorts.len());
594        for (field, _) in sorts {
595            let row: Option<(String, String)> = sqlx::query_as(
596                "SELECT data_type, coalesce(udt_name, '') FROM information_schema.columns \
597                  WHERE table_schema = $1 AND table_name = $2 AND column_name = $3",
598            )
599            .bind(&schema)
600            .bind(&table)
601            .bind(field)
602            .fetch_optional(&self.pool)
603            .await?;
604            let cast = row.map(|(data_type, udt)| cast_suffix(&data_type, &udt)).flatten();
605            casts.push(cast);
606        }
607        Ok(casts)
608    }
609
610    /// One boundary row's cursor: its values for the sort columns (typed
611    /// exactly as sqlx decodes them, so a decimal keeps its scale) plus its
612    /// id. NULL in any sort column yields None: a null has no position.
613    fn row_cursor(
614        &self,
615        row: &PgRow,
616        sorts: &[(String, FilterSortDirection)],
617        casts: &[Option<String>],
618    ) -> Option<String> {
619        let id: uuid::Uuid = row.try_get("id").ok()?;
620        let mut values: Vec<serde_json::Value> = Vec::with_capacity(sorts.len());
621        for (i, (field, _)) in sorts.iter().enumerate() {
622            let field = field.as_str();
623            let mut data_type = casts.get(i).and_then(|c| c.as_deref()).unwrap_or("");
624            // The ORM appends the id tiebreaker itself and page mode never
625            // resolves casts — without this, the uuid id decodes as text,
626            // the read fails, and the whole cursor silently vanishes.
627            if field == "id" && data_type.is_empty() {
628                data_type = "uuid";
629            }
630            let text: Option<String> = match data_type {
631                "numeric" => row
632                    .try_get::<Option<sqlx::types::Decimal>, _>(field)
633                    .ok()?
634                    .map(|d| d.to_string()),
635                "uuid" => row
636                    .try_get::<Option<uuid::Uuid>, _>(field)
637                    .ok()?
638                    .map(|u| u.to_string()),
639                "timestamptz" => row
640                    .try_get::<Option<chrono::DateTime<chrono::Utc>>, _>(field)
641                    .ok()?
642                    .map(|t| t.to_rfc3339()),
643                "integer" | "smallint" => row
644                    .try_get::<Option<i32>, _>(field)
645                    .ok()?
646                    .map(|n| n.to_string()),
647                "bigint" => row
648                    .try_get::<Option<i64>, _>(field)
649                    .ok()?
650                    .map(|n| n.to_string()),
651                "boolean" => row
652                    .try_get::<Option<bool>, _>(field)
653                    .ok()?
654                    .map(|b| b.to_string()),
655                "date" => row
656                    .try_get::<Option<chrono::NaiveDate>, _>(field)
657                    .ok()?
658                    .map(|d| d.to_string()),
659                _ => row
660                    .try_get::<Option<String>, _>(field)
661                    .ok()?
662                    .filter(|s| !s.is_empty() || data_type.is_empty()),
663            };
664            values.push(serde_json::Value::String(text?));
665        }
666        crate::filter::cursor::encode_cursor(sorts, &values, &id.to_string()).ok()
667    }
668}
669
670/// The placeholder cast for a column type, or None when a bare text bind
671/// compares correctly.
672fn cast_suffix(data_type: &str, udt_name: &str) -> Option<String> {
673    match data_type {
674        "uuid" => Some("uuid".into()),
675        "numeric" => Some("numeric".into()),
676        "integer" => Some("integer".into()),
677        "smallint" => Some("smallint".into()),
678        "bigint" => Some("bigint".into()),
679        "boolean" => Some("boolean".into()),
680        "date" => Some("date".into()),
681        "timestamp with time zone" => Some("timestamptz".into()),
682        "timestamp without time zone" => Some("timestamp".into()),
683        // An enum: bind text, cast to the enum's own name so the comparison
684        // runs in the enum's ordering.
685        "USER-DEFINED" if !udt_name.is_empty() => Some(udt_name.to_string()),
686        _ => None,
687    }
688}
689
690/// The cast a filter placeholder needs to compare against a column, from the column's catalog
691/// type: `format_type(atttypid, NULL)` and `pg_type.typtype`. None when a bare text bind already
692/// compares correctly (text-like columns) or when no single-value cast fits (arrays, json,
693/// composite and domain types keep today's text bind).
694fn filter_cast_for(type_name: &str, typtype: &str) -> Option<String> {
695    match typtype {
696        // An enum: cast to the enum itself so the comparison runs in the enum's ordering. The
697        // name comes from `format_type`, schema-qualified when the type is not on the search path.
698        "e" => Some(type_name.to_string()),
699        "b" => match type_name {
700            "uuid" | "boolean" | "smallint" | "integer" | "bigint" | "numeric" | "real"
701            | "double precision" | "date" | "interval" | "inet" | "cidr" | "macaddr" => {
702                Some(type_name.to_string())
703            }
704            "time without time zone" => Some("time".into()),
705            "time with time zone" => Some("timetz".into()),
706            "timestamp with time zone" => Some("timestamptz".into()),
707            "timestamp without time zone" => Some("timestamp".into()),
708            _ => None,
709        },
710        _ => None,
711    }
712}
713
714/// Every column of `qualified_table` whose filter placeholder needs a cast, with that cast.
715///
716/// Read from `pg_catalog` rather than `information_schema`: it is a direct lookup on the
717/// relation, it says whether a user-defined type is an enum (a domain or composite must not be
718/// treated as one), and `to_regclass` resolves the name exactly as the query will. Runs through
719/// the company-scoped helper so it uses the request's own connection when one is held.
720async fn catalog_filter_casts(
721    pool: &PgPool,
722    qualified_table: &str,
723) -> anyhow::Result<HashMap<String, String>> {
724    let q = sqlx::query_as::<Postgres, (String, String, String)>(
725        "SELECT a.attname::text, format_type(a.atttypid, NULL), t.typtype::text
726           FROM pg_catalog.pg_attribute a
727           JOIN pg_catalog.pg_type t ON t.oid = a.atttypid
728          WHERE a.attrelid = to_regclass($1) AND a.attnum > 0 AND NOT a.attisdropped",
729    )
730    .bind(qualified_table.to_string());
731    let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
732    Ok(rows
733        .into_iter()
734        .filter_map(|(column, type_name, typtype)| {
735            filter_cast_for(&type_name, &typtype).map(|cast| (column, cast))
736        })
737        .collect())
738}
739
740/// The generated hints, with the catalog's casts filling every column they do not name. A hint
741/// keeps deciding its own column, so nothing that compares correctly today changes.
742fn merge_filter_casts(
743    hints: &HashMap<String, String>,
744    mut catalog: HashMap<String, String>,
745) -> HashMap<String, String> {
746    for (column, cast) in hints {
747        catalog.insert(column.clone(), cast.clone());
748    }
749    catalog
750}
751
752/// Turn "column ... does not exist" into a sentence that names the cause.
753///
754/// The insert names the columns the entity serializes. A field that is serialized but is not a
755/// column of the table used to vanish quietly, because selecting every column of the row type threw
756/// unknown keys away; now it fails, and the bare Postgres error does not say why. Serialized field
757/// and table column are meant to be the same set — the update path has always assumed it — so this
758/// points at the mismatch rather than leaving someone to guess.
759fn explain_unknown_column(error: sqlx::Error, table: &str) -> anyhow::Error {
760    let text = error.to_string();
761    if text.contains("does not exist") && text.contains("column") {
762        return anyhow::Error::new(error).context(format!(
763            "insert into {table} named a column that does not exist: the entity serializes a field \
764             with no matching column. Every serialized field must be a column of the table (rename \
765             it, map it with #[serde(rename)], or skip it with #[serde(skip)])"
766        ));
767    }
768    anyhow::Error::new(error)
769}
770
771/// Quote a column name as a SQL identifier.
772///
773/// Column names here come from serializing the caller's entity, so they are Rust field names in
774/// practice — but they are interpolated into DDL/DML, where Postgres has no bind parameter for an
775/// identifier. Doubling an embedded quote is the identifier escape, so a name can never end the
776/// quoted section early.
777fn quote_ident(name: &str) -> String {
778    format!("\"{}\"", name.replace('"', "\"\""))
779}
780
781#[async_trait]
782impl<T> DatabaseOperations<T> for PostgresRepository<T>
783where
784    T: for<'a> FromRow<'a, PgRow> + Send + Sync + Unpin + Serialize,
785{
786    async fn create(&self, entity: &T) -> anyhow::Result<T> {
787        // Serialize entity to JSON to extract field names and values
788        let json_value = serde_json::to_value(entity)?;
789
790        let json_obj = match json_value {
791            Value::Object(obj) => obj,
792            _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
793        };
794
795        // Build dynamic INSERT query using jsonb_populate_record
796        // This approach handles all PostgreSQL types correctly including ENUMs and booleans
797        let json_str = serde_json::to_string(&json_obj)?;
798
799        // Name only the columns the payload actually carries.
800        //
801        // `SELECT (jsonb_populate_record(...)).*` emits EVERY column of the row type, so a column
802        // the entity does not know about arrived as an explicit NULL — and an explicit NULL is not
803        // an absent value: it overrides the column DEFAULT. That is invisible until a table's
804        // correctness depends on a default, which is exactly what composition-installed tenancy
805        // does (a scoped table defaults `org_unit_id` from the acting unit), so generic creates
806        // over such a table wrote NULL and were refused by the write-path guard.
807        //
808        // Listing the payload's own keys leaves every other column unmentioned, so its default
809        // applies. A key that is present with a JSON null is still written as NULL, which is
810        // right: the caller said so. This mirrors the update path below, which has always built
811        // its column list from these same keys.
812        let insert_columns: Vec<String> = json_obj.keys().map(|k| quote_ident(k)).collect();
813
814        let query = if insert_columns.is_empty() {
815            // Nothing supplied at all: let every column take its default rather than emitting
816            // `INSERT INTO t () SELECT`, which is not valid SQL.
817            format!("INSERT INTO {table} DEFAULT VALUES RETURNING *", table = self.table_name)
818        } else {
819            let columns = insert_columns.join(", ");
820            format!(
821                r#"
822            INSERT INTO {table} ({columns})
823            SELECT {columns} FROM jsonb_populate_record(NULL::{table}, $1::jsonb)
824            RETURNING *
825            "#,
826                table = self.table_name,
827                columns = columns
828            )
829        };
830
831        // The DEFAULT VALUES form takes no bind; every other form binds the payload.
832        let statement = if insert_columns.is_empty() {
833            sqlx::query_as::<_, T>(&query)
834        } else {
835            sqlx::query_as::<_, T>(&query).bind(&json_str)
836        };
837        let result = crate::company_scope::fetch_one_scoped(&self.pool, statement)
838            .await
839            .map_err(|e| explain_unknown_column(e, &self.table_name))?;
840
841        Ok(result)
842    }
843
844    async fn find_by_id(&self, id: &str) -> anyhow::Result<Option<T>> {
845        // Cast text to UUID for PostgreSQL UUID columns
846        let query = format!("SELECT * FROM {} WHERE id = $1::uuid", self.table_name);
847        let result = crate::company_scope::fetch_optional_scoped(
848            &self.pool,
849            sqlx::query_as::<Postgres, T>(&query).bind(id),
850        )
851        .await?;
852        Ok(result)
853    }
854
855    async fn find_all(&self) -> anyhow::Result<Vec<T>> {
856        let query = format!("SELECT * FROM {}", self.table_name);
857        let results = crate::company_scope::fetch_all_scoped(
858            &self.pool,
859            sqlx::query_as::<Postgres, T>(&query),
860        )
861        .await?;
862        Ok(results)
863    }
864
865    async fn update(&self, id: &str, entity: &T) -> anyhow::Result<Option<T>> {
866        // Serialize entity to JSON
867        let json_value = serde_json::to_value(entity)?;
868
869        let json_obj = match json_value {
870            Value::Object(obj) => obj,
871            _ => return Err(anyhow::anyhow!("Entity must serialize to a JSON object")),
872        };
873
874        // Build column list for the update (excluding 'id')
875        let update_columns: Vec<&String> = json_obj.keys()
876            .filter(|k| *k != "id")
877            .collect();
878
879        let column_names = update_columns.iter()
880            .map(|k| quote_ident(k))
881            .collect::<Vec<_>>()
882            .join(", ");
883
884        let json_str = serde_json::to_string(&json_obj)?;
885
886        // Use jsonb_populate_record with CTE to get properly typed values
887        let query = format!(
888            r#"
889            WITH new_row AS (
890                SELECT (jsonb_populate_record(NULL::{table}, $1::jsonb)).*
891            )
892            UPDATE {table} AS t
893            SET ({columns}) = (SELECT {columns} FROM new_row)
894            WHERE t.id = $2::uuid
895            RETURNING t.*
896            "#,
897            table = self.table_name,
898            columns = column_names
899        );
900
901        let result = crate::company_scope::fetch_optional_scoped(
902            &self.pool,
903            sqlx::query_as::<_, T>(&query).bind(&json_str).bind(id),
904        )
905        .await?;
906
907        Ok(result)
908    }
909
910    async fn delete(&self, id: &str) -> anyhow::Result<bool> {
911        let query = format!("DELETE FROM {} WHERE id = $1::uuid", self.table_name);
912        let result = crate::company_scope::execute_scoped(
913            &self.pool,
914            sqlx::query(&query).bind(id),
915        )
916        .await?;
917        Ok(result.rows_affected() > 0)
918    }
919
920    async fn count(&self) -> anyhow::Result<u64> {
921        let query = format!("SELECT COUNT(*) FROM {}", self.table_name);
922        let count = crate::company_scope::fetch_one_scalar_scoped(
923            &self.pool,
924            sqlx::query_scalar::<_, i64>(&query),
925        )
926        .await? as u64;
927        Ok(count)
928    }
929
930    async fn exists(&self, id: &str) -> anyhow::Result<bool> {
931        let query = format!("SELECT 1 FROM {} WHERE id = $1::uuid LIMIT 1", self.table_name);
932        let result = crate::company_scope::fetch_optional_scalar_scoped(
933            &self.pool,
934            sqlx::query_scalar::<_, i32>(&query).bind(id),
935        )
936        .await?;
937        Ok(result.is_some())
938    }
939
940    async fn execute_query(&self, query: &str) -> anyhow::Result<u64> {
941        let result = crate::company_scope::execute_scoped(
942            &self.pool,
943            sqlx::query(query),
944        )
945        .await?;
946        Ok(result.rows_affected())
947    }
948}
949
950// ─── Aggregation ──────────────────────────────────────────────────────────────
951
952/// Which reduction to apply to a column.
953#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
954pub enum AggregateFn {
955    Sum,
956    Avg,
957    Min,
958    Max,
959}
960
961impl AggregateFn {
962    fn sql(self) -> &'static str {
963        match self {
964            AggregateFn::Sum => "SUM",
965            AggregateFn::Avg => "AVG",
966            AggregateFn::Min => "MIN",
967            AggregateFn::Max => "MAX",
968        }
969    }
970
971    fn label(self) -> &'static str {
972        match self {
973            AggregateFn::Sum => "sum",
974            AggregateFn::Avg => "avg",
975            AggregateFn::Min => "min",
976            AggregateFn::Max => "max",
977        }
978    }
979
980    /// `SUM`/`AVG` on a text column is a type error, not a zero. `MIN`/`MAX`
981    /// order any comparable type, so they carry no such restriction.
982    fn requires_numeric(self) -> bool {
983        matches!(self, AggregateFn::Sum | AggregateFn::Avg)
984    }
985}
986
987/// What to group by and what to reduce — the parsed form of the query string.
988#[derive(Debug, Clone, Default)]
989pub struct AggregateSpec {
990    /// Column whose distinct values become groups. `None` asks for one total.
991    pub group_by: Option<String>,
992    /// `(function, column)` pairs, in the order the caller asked for them.
993    pub reductions: Vec<(AggregateFn, String)>,
994    /// Most groups to return before reporting the answer as truncated.
995    pub group_limit: usize,
996    /// Carry a label per group: the column on the group column's RELATED
997    /// table (resolved through the entity's relation metadata) to show
998    /// instead of a bare uuid key. `None` = keys stay as they are.
999    pub label_field: Option<String>,
1000    /// The resolved relation behind the group column — `(target table, the
1001    /// BASE table's FK column, snake)`, filled by the generic layer from
1002    /// the entity's `relations()` metadata. Callers never set this.
1003    pub label_relation: Option<(String, String)>,
1004}
1005
1006/// The default ceiling on distinct groups.
1007///
1008/// A `group_by` on a uuid or a timestamp yields one group per row, which is a
1009/// table scan wearing a chart's clothes. Rather than refusing those columns —
1010/// a list that would be wrong for some schema sooner or later — the answer is
1011/// capped and the cap is *reported*, so a caller can tell a complete picture
1012/// from a partial one instead of quietly drawing the wrong one.
1013pub const DEFAULT_GROUP_LIMIT: usize = 200;
1014
1015/// One group's numbers. `key` is the group's value; `None` is a real answer —
1016/// the rows whose group column is null — and is distinct from "no rows".
1017#[derive(Debug, Clone, Serialize, Deserialize)]
1018pub struct AggregateGroup {
1019    pub key: Option<String>,
1020    /// The group's display label (the related row's label column), when the
1021    /// caller asked for one and the group column is a relation FK.
1022    pub label: Option<String>,
1023    pub count: u64,
1024    /// Reduction results keyed `"sum:amount"`, carried as strings.
1025    ///
1026    /// Postgres `numeric` holds more precision than an IEEE double, and money
1027    /// columns are exactly where that bites: a tenant large enough for the
1028    /// total to matter is a tenant large enough to round it. The string is the
1029    /// exact value Postgres computed; the caller decides how to parse it.
1030    /// `None` is SQL NULL — no rows contributed — which is not zero.
1031    pub values: HashMap<String, Option<String>>,
1032}
1033
1034/// Groups plus the overall total, computed together.
1035#[derive(Debug, Clone, Serialize, Deserialize)]
1036pub struct AggregateResult {
1037    pub groups: Vec<AggregateGroup>,
1038    pub total: AggregateGroup,
1039    /// True when more distinct groups exist than `group_limit` allowed.
1040    pub truncated: bool,
1041}
1042
1043/// A column name rejected by the allow-list, or a reduction that its type
1044/// cannot answer.
1045#[derive(Debug, Clone)]
1046pub struct AggregateFieldError(pub String);
1047
1048impl std::fmt::Display for AggregateFieldError {
1049    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1050        f.write_str(&self.0)
1051    }
1052}
1053
1054impl std::error::Error for AggregateFieldError {}
1055
1056/// True for the Postgres types `SUM`/`AVG` accept.
1057fn is_numeric_pg_type(pg_type: &str) -> bool {
1058    let t = pg_type.trim().to_ascii_lowercase();
1059    let t = t.split('(').next().unwrap_or(&t).trim();
1060    matches!(
1061        t,
1062        "numeric" | "decimal" | "money"
1063            | "smallint" | "int2" | "integer" | "int" | "int4" | "bigint" | "int8"
1064            | "real" | "float4" | "double precision" | "float8"
1065            | "smallserial" | "serial" | "bigserial"
1066    )
1067}
1068
1069/// Read a table's real columns and their types from the catalog.
1070///
1071/// `EntityRepoMeta::column_types()` looks like the natural allow-list and is
1072/// not one: it carries only the columns the filter parser must CAST — uuids and
1073/// enums — so every numeric column is absent from it, which is exactly the set
1074/// `sum` and `avg` exist for. The catalog is the only complete and current
1075/// answer, and it cannot drift from the table the query will actually run
1076/// against.
1077async fn catalog_columns(
1078    pool: &PgPool,
1079    qualified_table: &str,
1080) -> anyhow::Result<HashMap<String, String>> {
1081    let (schema, table) = match qualified_table.split_once('.') {
1082        Some((s, t)) => (s.to_string(), t.to_string()),
1083        None => ("public".to_string(), qualified_table.to_string()),
1084    };
1085    let q = sqlx::query_as::<Postgres, (String, String)>(
1086        "SELECT column_name, data_type FROM information_schema.columns
1087          WHERE table_schema = $1 AND table_name = $2",
1088    )
1089    .bind(schema)
1090    .bind(table);
1091    let rows = crate::company_scope::fetch_all_scoped(pool, q).await?;
1092    Ok(rows.into_iter().collect())
1093}
1094
1095/// Resolve a caller-supplied column name against the entity's real columns.
1096///
1097/// This is the whole defence for the aggregate path. Unlike a filter *value*,
1098/// which is bound as a parameter, a `group_by` or `sum` column is spliced into
1099/// the SQL as an identifier — binding cannot protect it. So the name is never
1100/// escaped or quoted into safety; it is *replaced* by the matching key already
1101/// present in the entity's declared column map, and a name with no match is
1102/// refused. Nothing a caller types can reach the query text.
1103fn resolve_column<'a>(
1104    name: &str,
1105    column_types: &'a HashMap<String, String>,
1106) -> Result<(&'a str, &'a str), AggregateFieldError> {
1107    column_types
1108        .get_key_value(name)
1109        .map(|(k, v)| (k.as_str(), v.as_str()))
1110        .ok_or_else(|| {
1111            AggregateFieldError(format!("unknown field `{name}` — not a column of this entity"))
1112        })
1113}
1114
1115impl<T: for<'a> FromRow<'a, PgRow> + Send + Unpin> PostgresRepository<T> {
1116    /// Group and reduce rows in one statement, under the same filters, the same
1117    /// soft-delete convention and the same tenancy fence as the list endpoint.
1118    ///
1119    /// Groups and the overall total come back from a single `GROUPING SETS`
1120    /// query, which is what lets a caller draw a chart and its headline from
1121    /// one reply, and what keeps the two numbers consistent — a separate total
1122    /// query could observe a different set of rows.
1123    pub async fn aggregate_with_filters(
1124        &self,
1125        spec: &AggregateSpec,
1126        filters: &HashMap<String, String>,
1127        column_types: &HashMap<String, String>,
1128        search_fields: &[&str],
1129    ) -> anyhow::Result<AggregateResult> {
1130        let mut query_filter = self.parse_typed_filters(filters, column_types, None).await?;
1131        if !search_fields.is_empty() {
1132            query_filter.search_fields = search_fields.iter().map(|s| s.to_string()).collect();
1133        }
1134        // Grouping replaces row output entirely: paging and ordering describe a
1135        // page of rows, and there are none.
1136        query_filter.limit = None;
1137        query_filter.offset = None;
1138        let (where_clause, filter_params) = query_filter.build_where_clause();
1139
1140        // Every identifier below comes from the catalog, never from the caller.
1141        let columns = catalog_columns(&self.pool, &self.table_name).await?;
1142        let mut selects: Vec<String> = Vec::new();
1143        let mut value_keys: Vec<String> = Vec::new();
1144        for (func, field) in &spec.reductions {
1145            let (column, pg_type) = resolve_column(field, &columns)?;
1146            if func.requires_numeric() && !is_numeric_pg_type(pg_type) {
1147                return Err(AggregateFieldError(format!(
1148                    "cannot {} `{}`: its type is {} — {} needs a numeric column",
1149                    func.label(),
1150                    column,
1151                    pg_type,
1152                    func.label()
1153                ))
1154                .into());
1155            }
1156            let key = format!("{}:{}", func.label(), column);
1157            // Cast to text in SQL so the exact value Postgres computed is what
1158            // crosses the wire — see `AggregateGroup::values`.
1159            selects.push(format!("{}({})::text AS \"{}\"", func.sql(), column, key));
1160            value_keys.push(key);
1161        }
1162
1163        let group_limit = if spec.group_limit == 0 { DEFAULT_GROUP_LIMIT } else { spec.group_limit };
1164        let reductions = if selects.is_empty() { String::new() } else { format!(", {}", selects.join(", ")) };
1165
1166        // The group's label: when the caller names one (group_label) and the
1167        // group column is a relation FK of THIS entity, LEFT JOIN the
1168        // related table and carry its label column beside the key — the
1169        // aggregate's equivalent of `?include=`, which has no row to
1170        // hydrate otherwise.
1171        let mut label_select = String::new();
1172        let mut label_join = String::new();
1173        if let (Some(field), Some(label_field), Some((rel_table, base_fk))) = (
1174            &spec.group_by,
1175            &spec.label_field,
1176            &spec.label_relation,
1177        ) {
1178            let _ = field;
1179            let qualified = qualify_relation_table(&self.table_name, rel_table);
1180            let rel_columns = catalog_columns(&self.pool, &qualified).await?;
1181            let (label_col, _) = resolve_column(label_field, &rel_columns).map_err(|_| {
1182                AggregateFieldError(format!(
1183                    "cannot label groups by `{label_field}`: the related table `{qualified}` has no such column"
1184                ))
1185            })?;
1186            label_select = format!(", (label_rel.{label_col})::text AS __group_label");
1187            label_join = format!(
1188                " LEFT JOIN {qualified} AS label_rel ON label_rel.id IS NOT DISTINCT FROM {base_fk}"
1189            );
1190        }
1191
1192        let sql = match &spec.group_by {
1193            Some(field) => {
1194                let (column, _) = resolve_column(field, &columns)?;
1195                format!(
1196                    "SELECT GROUPING({column}) AS __is_total, ({column})::text AS __group_key{label_select}, \
1197                     COUNT(*) AS __count{reductions} \
1198                     FROM {table}{label_join}{where_clause} \
1199                     GROUP BY GROUPING SETS (({column}), ()) \
1200                     ORDER BY __is_total DESC, __count DESC \
1201                     LIMIT {limit}",
1202                    column = column,
1203                    label_select = label_select,
1204                    label_join = label_join,
1205                    reductions = reductions,
1206                    table = self.table_name,
1207                    where_clause = where_clause,
1208                    // One total row, the groups themselves, and one more to
1209                    // detect that a further group existed.
1210                    limit = group_limit + 2,
1211                )
1212            }
1213            None => format!(
1214                "SELECT 1 AS __is_total, NULL::text AS __group_key, COUNT(*) AS __count{reductions} \
1215                 FROM {table}{where_clause}",
1216                reductions = reductions,
1217                table = self.table_name,
1218                where_clause = where_clause,
1219            ),
1220        };
1221
1222        let mut builder = sqlx::query(&sql);
1223        for param in &filter_params {
1224            builder = builder.bind(param);
1225        }
1226        let rows = crate::company_scope::fetch_all_rows_scoped(&self.pool, builder).await?;
1227
1228        let read_group = |row: &PgRow| -> AggregateGroup {
1229            use sqlx::Row as _;
1230            let mut values = HashMap::with_capacity(value_keys.len());
1231            for key in &value_keys {
1232                values.insert(key.clone(), row.try_get::<Option<String>, _>(key.as_str()).ok().flatten());
1233            }
1234            AggregateGroup {
1235                key: row.try_get::<Option<String>, _>("__group_key").ok().flatten(),
1236                label: row.try_get::<Option<String>, _>("__group_label").ok().flatten(),
1237                count: row.try_get::<i64, _>("__count").unwrap_or(0).max(0) as u64,
1238                values,
1239            }
1240        };
1241
1242        use sqlx::Row as _;
1243        let mut total: Option<AggregateGroup> = None;
1244        let mut groups: Vec<AggregateGroup> = Vec::new();
1245        for row in &rows {
1246            let is_total = row.try_get::<i32, _>("__is_total").unwrap_or(0) == 1;
1247            if is_total {
1248                // Ordered first, so it survives the cap.
1249                total = Some(read_group(row));
1250            } else {
1251                groups.push(read_group(row));
1252            }
1253        }
1254
1255        let truncated = groups.len() > group_limit;
1256        groups.truncate(group_limit);
1257
1258        // No rows at all means no total row either: an empty result is a real
1259        // answer of zero, not a missing one.
1260        let total = total.unwrap_or_else(|| AggregateGroup {
1261            key: None,
1262            label: None,
1263            count: 0,
1264            values: value_keys.iter().map(|k| (k.clone(), None)).collect(),
1265        });
1266
1267        Ok(AggregateResult { groups, total, truncated })
1268    }
1269}
1270
1271#[cfg(test)]
1272mod aggregate_field_tests {
1273    use super::*;
1274
1275    fn columns() -> HashMap<String, String> {
1276        [
1277            ("status", "text"),
1278            ("total", "numeric"),
1279            ("qty", "integer"),
1280            ("notes", "text"),
1281        ]
1282        .iter()
1283        .map(|(k, v)| (k.to_string(), v.to_string()))
1284        .collect()
1285    }
1286
1287    #[test]
1288    fn resolves_only_declared_columns() {
1289        let cols = columns();
1290        assert_eq!(resolve_column("total", &cols).unwrap().0, "total");
1291        assert!(resolve_column("password_hash", &cols).is_err());
1292    }
1293
1294    /// The allow-list is the entire defence, because a group/sum column is
1295    /// spliced into SQL as an identifier and cannot be bound as a parameter.
1296    #[test]
1297    fn rejects_injection_attempts_rather_than_escaping_them() {
1298        let cols = columns();
1299        for probe in [
1300            "total) FROM selling.sales_orders; DROP TABLE users --",
1301            "status\"",
1302            "1=1",
1303            "total, (SELECT password FROM users)",
1304            "",
1305        ] {
1306            assert!(
1307                resolve_column(probe, &cols).is_err(),
1308                "`{probe}` must be refused, never escaped into the query"
1309            );
1310        }
1311    }
1312
1313    /// The returned name is the map's own key, not the caller's string, so no
1314    /// caller-controlled bytes can reach the SQL even on a match.
1315    #[test]
1316    fn returns_the_declared_key_not_the_callers_string() {
1317        let cols = columns();
1318        let (name, _) = resolve_column("total", &cols).unwrap();
1319        assert!(std::ptr::eq(name, cols.get_key_value("total").unwrap().0.as_str()));
1320    }
1321
1322    #[test]
1323    fn sum_and_avg_require_a_numeric_type() {
1324        assert!(AggregateFn::Sum.requires_numeric());
1325        assert!(AggregateFn::Avg.requires_numeric());
1326        // Ordering works on any comparable column, so these stay open.
1327        assert!(!AggregateFn::Min.requires_numeric());
1328        assert!(!AggregateFn::Max.requires_numeric());
1329    }
1330
1331    #[test]
1332    fn recognises_the_numeric_postgres_types() {
1333        for t in ["numeric", "NUMERIC(14,2)", "integer", "bigint", "double precision", "money"] {
1334            assert!(is_numeric_pg_type(t), "{t} should count as numeric");
1335        }
1336        for t in ["text", "uuid", "timestamptz", "boolean", "jsonb", "USER-DEFINED"] {
1337            assert!(!is_numeric_pg_type(t), "{t} must not accept a SUM");
1338        }
1339    }
1340}
1341
1342#[cfg(test)]
1343mod filter_cast_tests {
1344    use super::*;
1345
1346    #[test]
1347    fn typed_base_columns_get_their_own_cast() {
1348        for (t, want) in [
1349            ("boolean", "boolean"),
1350            ("integer", "integer"),
1351            ("bigint", "bigint"),
1352            ("smallint", "smallint"),
1353            ("numeric", "numeric"),
1354            ("double precision", "double precision"),
1355            ("uuid", "uuid"),
1356            ("date", "date"),
1357            ("time without time zone", "time"),
1358            ("timestamp with time zone", "timestamptz"),
1359            ("timestamp without time zone", "timestamp"),
1360        ] {
1361            assert_eq!(filter_cast_for(t, "b").as_deref(), Some(want), "{t}");
1362        }
1363    }
1364
1365    #[test]
1366    fn text_like_and_composite_columns_keep_the_text_bind() {
1367        for t in ["text", "character varying", "character", "jsonb", "json", "bytea", "text[]", "uuid[]"] {
1368            assert_eq!(filter_cast_for(t, "b"), None, "{t}");
1369        }
1370        // A domain or composite is not an enum, even though both are user-defined.
1371        assert_eq!(filter_cast_for("approvals.money_amount", "d"), None);
1372        assert_eq!(filter_cast_for("approvals.address", "c"), None);
1373    }
1374
1375    #[test]
1376    fn an_enum_casts_to_its_catalog_name() {
1377        assert_eq!(filter_cast_for("approval_status", "e").as_deref(), Some("approval_status"));
1378        assert_eq!(
1379            filter_cast_for("recruitment.stage_kind", "e").as_deref(),
1380            Some("recruitment.stage_kind")
1381        );
1382    }
1383
1384    #[test]
1385    fn a_generated_hint_wins_over_the_catalog_for_its_column() {
1386        let hints: HashMap<String, String> =
1387            [("id", "uuid"), ("status", "approval_status")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1388        let catalog: HashMap<String, String> = [
1389            ("id", "uuid"),
1390            ("status", "approvals.approval_status"),
1391            ("folded", "boolean"),
1392            ("requested_by", "uuid"),
1393        ]
1394        .iter()
1395        .map(|(k, v)| (k.to_string(), v.to_string()))
1396        .collect();
1397        let merged = merge_filter_casts(&hints, catalog);
1398        assert_eq!(merged["status"], "approval_status");
1399        assert_eq!(merged["folded"], "boolean");
1400        assert_eq!(merged["requested_by"], "uuid");
1401        assert_eq!(merged.len(), 4);
1402    }
1403
1404    #[test]
1405    fn only_a_filter_with_an_uncast_comparison_reads_the_catalog() {
1406        let hints: HashMap<String, String> =
1407            [("id", "uuid")].iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1408        let parse = |pairs: &[(&str, &str)]| {
1409            let f: HashMap<String, String> =
1410                pairs.iter().map(|(k, v)| (k.to_string(), v.to_string())).collect();
1411            parse_query_filter(&f, &hints, None).unwrap().has_uncast_value_conditions()
1412        };
1413        assert!(!parse(&[("id[in]", "a,b"), ("name[contain]", "x"), ("limit", "5")]));
1414        assert!(!parse(&[("deleted_by[isnull]", "1")]));
1415        assert!(parse(&[("folded[eq]", "false")]));
1416        assert!(parse(&[("folded", "false")]));
1417        assert!(parse(&[("scheduled_at[between]", "2026-10-01,2026-10-03")]));
1418        assert!(parse(&[("sequence[or]", "3")]));
1419    }
1420}