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:
+40
-23
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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/"),
|
||||||
|
|||||||
@@ -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
@@ -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 ()
|
||||||
|
|||||||
Reference in New Issue
Block a user