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