Adds 415 response for any non-empty Accept header different from application/json or text/csv. Uses apropriate Content-Type header when sending CSV format.

This commit is contained in:
Diogo Biazus
2015-08-20 20:47:36 -07:00
committed by Joe Nelson
parent cb7d00b839
commit 3b017dfdf6
6 changed files with 81 additions and 43 deletions
+40 -23
View File
@@ -1,19 +1,19 @@
{-# LANGUAGE FlexibleContexts #-} {-# LANGUAGE FlexibleContexts #-}
module PostgREST.App (app, sqlError, isSqlError) where module PostgREST.App (app, sqlError, isSqlError, contentTypeForAccept) where
import Control.Monad (join) import Control.Monad (join)
import Control.Arrow ((***), second) import Control.Arrow ((***), second)
import Control.Applicative import Control.Applicative
import Data.Text hiding (map) import Data.Text hiding (map, find)
import Data.Maybe (fromMaybe, mapMaybe) import Data.Maybe (fromMaybe, mapMaybe, isJust, isNothing)
import Text.Regex.TDFA ((=~)) import Text.Regex.TDFA ((=~))
import Data.Ord (comparing) import Data.Ord (comparing)
import Data.Ranged.Ranges (emptyRange) import Data.Ranged.Ranges (emptyRange)
import qualified Data.HashMap.Strict as M import qualified Data.HashMap.Strict as M
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.CaseInsensitive (original) import Data.CaseInsensitive (original)
import Data.List (sortBy) import Data.List (sortBy, find)
import Data.Functor.Identity import Data.Functor.Identity
import qualified Data.Set as S import qualified Data.Set as S
import qualified Data.ByteString.Lazy as BL import qualified Data.ByteString.Lazy as BL
@@ -26,6 +26,7 @@ import Network.HTTP.Types.Header
import Network.HTTP.Types.URI (parseSimpleQuery) import Network.HTTP.Types.URI (parseSimpleQuery)
import Network.HTTP.Base (urlEncodeVars) import Network.HTTP.Base (urlEncodeVars)
import Network.Wai import Network.Wai
import Network.Wai.Parse (parseHttpAccept)
import Network.Wai.Internal (Response(..)) import Network.Wai.Internal (Response(..))
import Data.Aeson import Data.Aeson
@@ -67,7 +68,7 @@ app conf reqBody req =
parentheticT ( parentheticT (
whereT qt qq $ countRows qt whereT qt qq $ countRows qt
) <> commaq <> ( ) <> commaq <> (
bodyForAccept accept qt bodyForAccept contentType qt
. limitT range . limitT range
. orderT (orderParse qq) . orderT (orderParse qq)
. whereT qt qq . whereT qt qq
@@ -85,7 +86,7 @@ app conf reqBody req =
. parseSimpleQuery . parseSimpleQuery
$ rawQueryString req $ rawQueryString req
return $ responseLBS status return $ responseLBS status
[if accept == Just "text/csv" then csvH else jsonH, contentRange, [contentTypeH, contentRange,
("Content-Location", ("Content-Location",
"/" <> cs table <> "/" <> cs table <>
if Prelude.null canonical then "" else "?" <> cs canonical if Prelude.null canonical then "" else "?" <> cs canonical
@@ -131,7 +132,7 @@ app conf reqBody req =
let qt = qualify table let qt = qualify table
echoRequested = lookupHeader "Prefer" == Just "return=representation" echoRequested = lookupHeader "Prefer" == Just "return=representation"
parsed :: Either String (V.Vector Text, V.Vector (V.Vector Value)) parsed :: Either String (V.Vector Text, V.Vector (V.Vector Value))
parsed = if lookupHeader "Content-Type" == Just "text/csv" parsed = if lookupHeader "Content-Type" == Just csvMT
then do then do
rows <- CSV.decode CSV.NoHeader reqBody rows <- CSV.decode CSV.NoHeader reqBody
if V.null rows then Left "CSV requires header" if V.null rows then Left "CSV requires header"
@@ -227,13 +228,15 @@ app conf reqBody req =
qq = queryString req qq = queryString req
qualify = QualifiedTable schema qualify = QualifiedTable schema
hdrs = requestHeaders req hdrs = requestHeaders req
schema = requestedSchema (cs $ configV1Schema conf) hdrs lookupHeader = flip lookup hdrs
accept = lookupHeader hAccept
schema = requestedSchema (cs $ configV1Schema conf) accept
authenticator = cs $ configDbUser conf authenticator = cs $ configDbUser conf
jwtSecret = cs $ configJwtSecret conf jwtSecret = cs $ configJwtSecret conf
range = rangeRequested hdrs range = rangeRequested hdrs
allOrigins = ("Access-Control-Allow-Origin", "*") :: Header allOrigins = ("Access-Control-Allow-Origin", "*") :: Header
lookupHeader = flip lookup hdrs contentType = fromMaybe "application/json" $ contentTypeForAccept accept
accept = lookupHeader hAccept contentTypeH = (hContentType, contentType)
sqlError :: t sqlError :: t
sqlError = undefined sqlError = undefined
@@ -247,12 +250,6 @@ rangeStatus from to total
| (1 + to - from) < total = status206 | (1 + to - from) < total = status206
| otherwise = status200 | 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 :: Int -> Int -> Int -> Header
contentRangeH from to total = contentRangeH from to total =
("Content-Range", ("Content-Range",
@@ -263,21 +260,41 @@ contentRangeH from to total =
<> cs (show total) <> cs (show total)
) )
requestedSchema :: Text -> RequestHeaders -> Text requestedSchema :: Text -> Maybe BS.ByteString -> Text
requestedSchema v1schema hdrs = requestedSchema v1schema accept =
case verStr of case verStr of
Just [[_, ver]] -> if ver == "1" then v1schema else cs ver Just [[_, ver]] -> if ver == "1" then v1schema else cs ver
_ -> v1schema _ -> v1schema
where verRegex = "version[ ]*=[ ]*([0-9]+)" :: BS.ByteString where verRegex = "version[ ]*=[ ]*([0-9]+)" :: BS.ByteString
accept = cs <$> lookup hAccept hdrs :: Maybe BS.ByteString
verStr = (=~ verRegex) <$> accept :: Maybe [[BS.ByteString]] verStr = (=~ verRegex) <$> accept :: Maybe [[BS.ByteString]]
jsonH :: Header
jsonH = (hContentType, "application/json")
csvH :: Header jsonMT :: BS.ByteString
csvH = (hContentType, "text/csv") jsonMT = "application/json"
csvMT :: BS.ByteString
csvMT = "text/csv"
jsonH :: Header
jsonH = (hContentType, jsonMT)
contentTypeForAccept :: Maybe BS.ByteString -> Maybe BS.ByteString
contentTypeForAccept accept
| isNothing accept || hasJson = Just jsonMT
| hasCsv = Just csvMT
| otherwise = Nothing
where
Just acceptH = accept
findInAccept = flip find $ parseHttpAccept acceptH
hasJson = isJust $ findInAccept $ BS.isPrefixOf jsonMT
hasCsv = isJust $ findInAccept $ BS.isPrefixOf csvMT
bodyForAccept :: BS.ByteString -> QualifiedTable -> StatementT
bodyForAccept contentType table
| contentType == csvMT = asCsvWithCount table
| otherwise = asJsonWithCount -- defaults to JSON
handleJsonObj :: BL.ByteString -> (Object -> H.Tx P.Postgres s Response) handleJsonObj :: BL.ByteString -> (Object -> H.Tx P.Postgres s Response)
-> H.Tx P.Postgres s Response -> H.Tx P.Postgres s Response
+2 -8
View File
@@ -11,10 +11,7 @@ import Control.Monad (unless)
import Control.Monad.IO.Class (liftIO) import Control.Monad.IO.Class (liftIO)
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Network.Wai (strictRequestBody) import Network.Wai (strictRequestBody)
import Network.Wai.Middleware.Cors (cors)
import Network.Wai.Handler.Warp hiding (Connection) import Network.Wai.Handler.Warp hiding (Connection)
import Network.Wai.Middleware.Gzip (gzip, def)
import Network.Wai.Middleware.Static (staticPolicy, only)
import Network.Wai.Middleware.RequestLogger (logStdout) import Network.Wai.Middleware.RequestLogger (logStdout)
import Data.List (intercalate) import Data.List (intercalate)
import Data.Version (versionBranch) import Data.Version (versionBranch)
@@ -26,7 +23,7 @@ import Options.Applicative hiding (columns)
import System.IO (stderr, stdin, stdout, hSetBuffering, BufferMode(..)) import System.IO (stderr, stdin, stdout, hSetBuffering, BufferMode(..))
import PostgREST.Config (AppConfig(..), argParser, corsPolicy) import PostgREST.Config (AppConfig(..), argParser)
isServerVersionSupported = do isServerVersionSupported = do
Identity (row :: Text) <- H.tx Nothing $ H.singleEx $ [H.stmt|SHOW server_version_num|] Identity (row :: Text) <- H.tx Nothing $ H.singleEx $ [H.stmt|SHOW server_version_num|]
@@ -64,10 +61,7 @@ main = do
appSettings = setPort port appSettings = setPort port
. setServerName (cs $ "postgrest/" <> prettyVersion) . setServerName (cs $ "postgrest/" <> prettyVersion)
$ defaultSettings $ defaultSettings
middle = logStdout middle = logStdout . defaultMiddle (configSecure conf)
. (if configSecure conf then redirectInsecure else id)
. gzip def . cors corsPolicy
. staticPolicy (only [("favicon.ico", "static/favicon.ico")])
poolSettings <- maybe (fail "Improper session settings") return $ poolSettings <- maybe (fail "Improper session settings") return $
H.poolSettings (fromIntegral $ configPool conf) 30 H.poolSettings (fromIntegral $ configPool conf) 30
+22 -4
View File
@@ -3,7 +3,7 @@
module PostgREST.Middleware where module PostgREST.Middleware where
import Data.Maybe (fromMaybe) import Data.Maybe (fromMaybe, isNothing)
import Data.Monoid import Data.Monoid
import Data.Text import Data.Text
-- import Data.Pool(withResource, Pool) -- import Data.Pool(withResource, Pool)
@@ -12,15 +12,19 @@ import qualified Hasql as H
import qualified Hasql.Postgres as P import qualified Hasql.Postgres as P
import Data.String.Conversions(cs) import Data.String.Conversions(cs)
import Network.HTTP.Types.Header (hLocation, hAuthorization) import Network.HTTP.Types.Header (hLocation, hAuthorization, hAccept)
import Network.HTTP.Types (RequestHeaders) import Network.HTTP.Types (RequestHeaders)
import Network.HTTP.Types.Status (status400, status401, status301) import Network.HTTP.Types.Status (status400, status401, status301, status415)
import Network.Wai (Application, requestHeaders, responseLBS, rawPathInfo, import Network.Wai (Application, requestHeaders, responseLBS, rawPathInfo,
rawQueryString, isSecure, Request(..), Response) rawQueryString, isSecure, Request(..), Response)
import Network.Wai.Middleware.Gzip (gzip, def)
import Network.Wai.Middleware.Cors (cors)
import Network.Wai.Middleware.Static (staticPolicy, only)
import Network.URI (URI(..), parseURI) import Network.URI (URI(..), parseURI)
import PostgREST.Config (AppConfig(..)) import PostgREST.Config (AppConfig(..), corsPolicy)
import PostgREST.Auth (LoginAttempt(..), signInRole, signInWithJWT, setRole, resetRole, setUserId, resetUserId) import PostgREST.Auth (LoginAttempt(..), signInRole, signInWithJWT, setRole, resetRole, setUserId, resetUserId)
import PostgREST.App (contentTypeForAccept)
import Codec.Binary.Base64.String (decode) import Codec.Binary.Base64.String (decode)
import Prelude import Prelude
@@ -84,3 +88,17 @@ redirectInsecure app req respond = do
Nothing -> Nothing ->
respond $ responseLBS status400 [] "SSL is required" respond $ responseLBS status400 [] "SSL is required"
else app req respond else app req respond
unsupportedAccept :: Application -> Application
unsupportedAccept app req respond = do
let
accept = lookup hAccept $ requestHeaders req
if isNothing $ contentTypeForAccept accept
then respond $ responseLBS status415 [] "Unsupported Accept header, try: application/json"
else app req respond
defaultMiddle :: Bool -> Application -> Application
defaultMiddle secure = (if secure then redirectInsecure else id)
. gzip def . cors corsPolicy
. staticPolicy (only [("favicon.ico", "static/favicon.ico")])
. unsupportedAccept
+1 -1
View File
@@ -22,7 +22,7 @@ spec = around withApp $ describe "CORS" $ do
("Host", "localhost:3000"), ("Host", "localhost:3000"),
("User-Agent", "Mozilla/5.0 (Macintosh; Intel Mac OS X 10.9; rv:32.0) Gecko/20100101 Firefox/32.0"), ("User-Agent", "Mozilla/5.0 (Macintosh; Intel Mac OS X 10.9; rv:32.0) Gecko/20100101 Firefox/32.0"),
("Origin", "http://localhost:8000"), ("Origin", "http://localhost:8000"),
("Accept", "text/plain, */*; q=0.01"), ("Accept", "text/csv, */*; q=0.01"),
("Accept-Language", "en-US,en;q=0.5"), ("Accept-Language", "en-US,en;q=0.5"),
("Accept-Encoding", "gzip, deflate"), ("Accept-Encoding", "gzip, deflate"),
("Referer", "http://localhost:8000/"), ("Referer", "http://localhost:8000/"),
+14 -4
View File
@@ -130,13 +130,23 @@ spec =
it "without other constraints" $ it "without other constraints" $
get "/items?order=asc.id" `shouldRespondWith` 200 get "/items?order=asc.id" `shouldRespondWith` 200
describe "Accept headers" $ describe "Accept headers" $ do
it "should respond an unknown accept type with 415" $
request methodGet "/simple_pk"
(acceptHdrs "text/unknowntype") ""
`shouldRespondWith` 415
it "should respond correctly to multiple types in accept header" $
request methodGet "/simple_pk"
(acceptHdrs "text/unknowntype, text/csv") ""
`shouldRespondWith` 200
it "should respond with CSV to 'text/csv' request" $ it "should respond with CSV to 'text/csv' request" $
request methodGet "/simple_pk" request methodGet "/simple_pk"
(acceptHdrs "text/csv") "" (acceptHdrs "text/csv; version=1") ""
`shouldRespondWith` ResponseMatcher { `shouldRespondWith` ResponseMatcher {
matchBody = Just "k,extra\rxyyx,u\rxYYx,v" matchBody = Just "k,extra\rxyyx,u\rxYYx,v"
, matchStatus = 200 , matchStatus = 200
, matchHeaders = ["Content-Type" <:> "text/csv"] , matchHeaders = ["Content-Type" <:> "text/csv"]
} }
+2 -3
View File
@@ -21,13 +21,12 @@ import Data.CaseInsensitive (CI(..))
import Data.Maybe (fromMaybe) import Data.Maybe (fromMaybe)
import Text.Regex.TDFA ((=~)) import Text.Regex.TDFA ((=~))
import qualified Data.ByteString.Char8 as BS import qualified Data.ByteString.Char8 as BS
import Network.Wai.Middleware.Cors (cors)
import System.Process (readProcess) import System.Process (readProcess)
import qualified Data.Aeson.Types as J import qualified Data.Aeson.Types as J
import PostgREST.App (app) import PostgREST.App (app)
import PostgREST.Config (AppConfig(..), corsPolicy) import PostgREST.Config (AppConfig(..))
import PostgREST.Middleware import PostgREST.Middleware
import PostgREST.Error(errResponse) import PostgREST.Error(errResponse)
@@ -59,7 +58,7 @@ withApp perform = do
$ authenticated cfg (app cfg body) req $ authenticated cfg (app cfg body) req
either (resp . errResponse) resp result either (resp . errResponse) resp result
where middle = cors corsPolicy where middle = defaultMiddle False
resetDb :: IO () resetDb :: IO ()