diff --git a/src/PostgREST/App.hs b/src/PostgREST/App.hs index 4744fd33d..b2b6b6326 100644 --- a/src/PostgREST/App.hs +++ b/src/PostgREST/App.hs @@ -62,11 +62,12 @@ app conf reqBody req = then return $ responseLBS status416 [] "HTTP Range error" else do let qt = qualify table + from = fromMaybe 0 $ rangeOffset <$> range select = B.Stmt "select " V.empty True <> parentheticT ( whereT qt qq $ countRows qt ) <> commaq <> ( - asJsonWithCount + bodyForAccept accept qt . limitT range . orderT (orderParse qq) . whereT qt qq @@ -75,7 +76,6 @@ app conf reqBody req = row <- H.maybeEx select let (tableTotal, queryTotal, body) = fromMaybe (0, 0, Just "" :: Maybe Text) row - from = fromMaybe 0 $ rangeOffset <$> range to = from+queryTotal-1 contentRange = contentRangeH from to tableTotal status = rangeStatus from to tableTotal @@ -233,6 +233,7 @@ app conf reqBody req = range = rangeRequested hdrs allOrigins = ("Access-Control-Allow-Origin", "*") :: Header lookupHeader = flip lookup hdrs + accept = lookupHeader hAccept sqlError :: t sqlError = undefined @@ -246,6 +247,12 @@ rangeStatus from to total | (1 + to - from) < total = status206 | otherwise = status200 +bodyForAccept :: Maybe BS.ByteString -> QualifiedTable -> StatementT +bodyForAccept accept table = + case accept of + Just "text/csv" -> asCsvWithCount table + _ -> asJsonWithCount -- defaults to JSON + contentRangeH :: Int -> Int -> Int -> Header contentRangeH from to total = ("Content-Range", diff --git a/src/PostgREST/PgQuery.hs b/src/PostgREST/PgQuery.hs index b9dc71210..c84c1fab3 100644 --- a/src/PostgREST/PgQuery.hs +++ b/src/PostgREST/PgQuery.hs @@ -100,10 +100,27 @@ countT s = countRows :: QualifiedTable -> PStmt countRows t = B.Stmt ("select pg_catalog.count(1) from " <> fromQt t) empty True +asCsvWithCount :: QualifiedTable -> StatementT +asCsvWithCount table = withCount . asCsv table + +asCsv :: QualifiedTable -> StatementT +asCsv table s = s { B.stmtTemplate = + "(select string_agg(quote_ident(column_name::text), ',') from " + <> "(select column_name from information_schema.columns where quote_ident(table_schema) || '.' || table_name = '" + <> fromQt table <> "' order by ordinal_position) h) || '\r' || " + <> "coalesce(string_agg(substring(t::text, 2, length(t::text) - 2), '\r'), '') from (" + <> B.stmtTemplate s <> ") t" } + asJsonWithCount :: StatementT -asJsonWithCount s = s { B.stmtTemplate = - "pg_catalog.count(t), array_to_json(array_agg(row_to_json(t)))::character varying from (" - <> B.stmtTemplate s <> ") t" } +asJsonWithCount = withCount . asJson + +asJson :: StatementT +asJson s = s { B.stmtTemplate = + "array_to_json(array_agg(row_to_json(t)))::character varying from (" + <> B.stmtTemplate s <> ") t" } + +withCount :: StatementT +withCount s = s { B.stmtTemplate = "pg_catalog.count(t), " <> B.stmtTemplate s } asJsonRow :: StatementT asJsonRow s = s { B.stmtTemplate = "row_to_json(t) from (" <> B.stmtTemplate s <> ") t" } diff --git a/test/Feature/QuerySpec.hs b/test/Feature/QuerySpec.hs index 7cadff625..98f036e98 100644 --- a/test/Feature/QuerySpec.hs +++ b/test/Feature/QuerySpec.hs @@ -3,6 +3,7 @@ module Feature.QuerySpec where import Test.Hspec import Test.Hspec.Wai import Test.Hspec.Wai.JSON +import Network.HTTP.Types import Network.Wai.Test (SResponse(simpleHeaders)) import SpecHelper @@ -121,6 +122,12 @@ spec = it "without other constraints" $ get "/items?order=asc.id" `shouldRespondWith` 200 + describe "Accept headers" $ + it "should respond with CSV to 'text/csv' request" $ + request methodGet "/simple_pk" + (acceptHdrs "text/csv") "" + `shouldRespondWith` "k,extra\rxyyx,u\rxYYx,v" + describe "Canonical location" $ do it "Sets Content-Location with alphabetized params" $ get "/no_pk?b=eq.1&a=eq.1" diff --git a/test/SpecHelper.hs b/test/SpecHelper.hs index a88da793e..345a9b561 100644 --- a/test/SpecHelper.hs +++ b/test/SpecHelper.hs @@ -15,7 +15,7 @@ import qualified Data.Vector as V import Control.Monad (void) import Network.HTTP.Types.Header (Header, ByteRange, renderByteRange, - hRange, hAuthorization) + hRange, hAuthorization, hAccept) import Codec.Binary.Base64.String (encode) import Data.CaseInsensitive (CI(..)) import Data.Maybe (fromMaybe) @@ -84,6 +84,9 @@ loadFixture name = rangeHdrs :: ByteRange -> [Header] rangeHdrs r = [rangeUnit, (hRange, renderByteRange r)] +acceptHdrs :: BS.ByteString -> [Header] +acceptHdrs mime = [(hAccept, mime)] + rangeUnit :: Header rangeUnit = ("Range-Unit" :: CI BS.ByteString, "items")