From d031bb2df55ef56ebb0a4594f7ab00b2701d53d3 Mon Sep 17 00:00:00 2001 From: Laurence Isla Date: Sat, 20 Dec 2025 16:17:11 -0500 Subject: [PATCH] refactor: use a single function to get the page_total count --- src/PostgREST/Query/SqlFragment.hs | 8 ++++++++ src/PostgREST/Query/Statements.hs | 9 +++------ 2 files changed, 11 insertions(+), 6 deletions(-) diff --git a/src/PostgREST/Query/SqlFragment.hs b/src/PostgREST/Query/SqlFragment.hs index 461d4e827..cdd14bdd5 100644 --- a/src/PostgREST/Query/SqlFragment.hs +++ b/src/PostgREST/Query/SqlFragment.hs @@ -23,6 +23,7 @@ module PostgREST.Query.SqlFragment , locationF , noLocationF , orderF + , pageCountSelectF , pgFmtColumn , pgFmtFilter , pgFmtIdent @@ -96,6 +97,7 @@ import PostgREST.SchemaCache.Routine (MediaHandler (..), Routine (..), funcReturnsScalar, funcReturnsSetOfScalar, + funcReturnsSingle, funcReturnsSingleComposite) import Protolude hiding (Sum, cast) @@ -495,6 +497,12 @@ countF countQuery shouldCount = mempty , "null::bigint") +pageCountSelectF :: Maybe Routine -> SQL.Snippet +pageCountSelectF rout = + if maybe False funcReturnsSingle rout + then "1" + else "pg_catalog.count(_postgrest_t)" + returningF :: QualifiedIdentifier -> [FieldName] -> SQL.Snippet returningF qi returnings = if null returnings diff --git a/src/PostgREST/Query/Statements.hs b/src/PostgREST/Query/Statements.hs index e90611aea..1805ea0ba 100644 --- a/src/PostgREST/Query/Statements.hs +++ b/src/PostgREST/Query/Statements.hs @@ -20,8 +20,7 @@ import PostgREST.Plan.MutatePlan as MTPlan import PostgREST.Plan.ReadPlan import PostgREST.Query.QueryBuilder import PostgREST.Query.SqlFragment -import PostgREST.SchemaCache.Routine (MediaHandler (..), Routine, - funcReturnsSingle) +import PostgREST.SchemaCache.Routine (MediaHandler (..), Routine) import Protolude @@ -72,7 +71,7 @@ mainRead rPlan countQuery pCount maxRows mt handler = mtSnippet mt snippet countCTEF <> " " <> "SELECT " <> countResultF <> " AS total_result_set, " <> - "pg_catalog.count(_postgrest_t) AS page_total, " <> + pageCountSelectF Nothing <> " AS page_total, " <> handlerF Nothing handler <> " AS body, " <> responseHeadersF <> " AS response_headers, " <> responseStatusF <> " AS response_status, " <> @@ -97,9 +96,7 @@ mainCall rout cPlan rPlan pCount mt handler = mtSnippet mt snippet countCTEF <> "SELECT " <> countResultF <> " AS total_result_set, " <> - (if funcReturnsSingle rout - then "1" - else "pg_catalog.count(_postgrest_t)") <> " AS page_total, " <> + pageCountSelectF (Just rout) <> " AS page_total, " <> handlerF (Just rout) handler <> " AS body, " <> responseHeadersF <> " AS response_headers, " <> responseStatusF <> " AS response_status, " <>