diff --git a/src/PostgREST/App.hs b/src/PostgREST/App.hs index 3cc092b83..5c392b499 100644 --- a/src/PostgREST/App.hs +++ b/src/PostgREST/App.hs @@ -1,19 +1,19 @@ {-# LANGUAGE FlexibleContexts #-} -module PostgREST.App (app, sqlError, isSqlError) where +module PostgREST.App (app, sqlError, isSqlError, contentTypeForAccept) where import Control.Monad (join) import Control.Arrow ((***), second) import Control.Applicative -import Data.Text hiding (map) -import Data.Maybe (fromMaybe, mapMaybe) +import Data.Text hiding (map, find) +import Data.Maybe (fromMaybe, mapMaybe, isJust, isNothing) import Text.Regex.TDFA ((=~)) import Data.Ord (comparing) import Data.Ranged.Ranges (emptyRange) import qualified Data.HashMap.Strict as M import Data.String.Conversions (cs) import Data.CaseInsensitive (original) -import Data.List (sortBy) +import Data.List (sortBy, find) import Data.Functor.Identity import qualified Data.Set as S 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.Base (urlEncodeVars) import Network.Wai +import Network.Wai.Parse (parseHttpAccept) import Network.Wai.Internal (Response(..)) import Data.Aeson @@ -67,7 +68,7 @@ app conf reqBody req = parentheticT ( whereT qt qq $ countRows qt ) <> commaq <> ( - bodyForAccept accept qt + bodyForAccept contentType qt . limitT range . orderT (orderParse qq) . whereT qt qq @@ -85,7 +86,7 @@ app conf reqBody req = . parseSimpleQuery $ rawQueryString req return $ responseLBS status - [if accept == Just "text/csv" then csvH else jsonH, contentRange, + [contentTypeH, contentRange, ("Content-Location", "/" <> cs table <> if Prelude.null canonical then "" else "?" <> cs canonical @@ -131,7 +132,7 @@ app conf reqBody req = let qt = qualify table echoRequested = lookupHeader "Prefer" == Just "return=representation" 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 rows <- CSV.decode CSV.NoHeader reqBody if V.null rows then Left "CSV requires header" @@ -227,13 +228,15 @@ app conf reqBody req = qq = queryString req qualify = QualifiedTable schema 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 jwtSecret = cs $ configJwtSecret conf range = rangeRequested hdrs allOrigins = ("Access-Control-Allow-Origin", "*") :: Header - lookupHeader = flip lookup hdrs - accept = lookupHeader hAccept + contentType = fromMaybe "application/json" $ contentTypeForAccept accept + contentTypeH = (hContentType, contentType) sqlError :: t sqlError = undefined @@ -247,12 +250,6 @@ 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", @@ -263,21 +260,41 @@ contentRangeH from to total = <> cs (show total) ) -requestedSchema :: Text -> RequestHeaders -> Text -requestedSchema v1schema hdrs = +requestedSchema :: Text -> Maybe BS.ByteString -> Text +requestedSchema v1schema accept = case verStr of Just [[_, ver]] -> if ver == "1" then v1schema else cs ver _ -> v1schema where verRegex = "version[ ]*=[ ]*([0-9]+)" :: BS.ByteString - accept = cs <$> lookup hAccept hdrs :: Maybe BS.ByteString verStr = (=~ verRegex) <$> accept :: Maybe [[BS.ByteString]] -jsonH :: Header -jsonH = (hContentType, "application/json") -csvH :: Header -csvH = (hContentType, "text/csv") +jsonMT :: BS.ByteString +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) -> H.Tx P.Postgres s Response diff --git a/src/PostgREST/Main.hs b/src/PostgREST/Main.hs index 8592f5538..acae88bd6 100644 --- a/src/PostgREST/Main.hs +++ b/src/PostgREST/Main.hs @@ -11,10 +11,7 @@ import Control.Monad (unless) import Control.Monad.IO.Class (liftIO) import Data.String.Conversions (cs) import Network.Wai (strictRequestBody) -import Network.Wai.Middleware.Cors (cors) 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 Data.List (intercalate) import Data.Version (versionBranch) @@ -26,7 +23,7 @@ import Options.Applicative hiding (columns) import System.IO (stderr, stdin, stdout, hSetBuffering, BufferMode(..)) -import PostgREST.Config (AppConfig(..), argParser, corsPolicy) +import PostgREST.Config (AppConfig(..), argParser) isServerVersionSupported = do Identity (row :: Text) <- H.tx Nothing $ H.singleEx $ [H.stmt|SHOW server_version_num|] @@ -64,10 +61,7 @@ main = do appSettings = setPort port . setServerName (cs $ "postgrest/" <> prettyVersion) $ defaultSettings - middle = logStdout - . (if configSecure conf then redirectInsecure else id) - . gzip def . cors corsPolicy - . staticPolicy (only [("favicon.ico", "static/favicon.ico")]) + middle = logStdout . defaultMiddle (configSecure conf) poolSettings <- maybe (fail "Improper session settings") return $ H.poolSettings (fromIntegral $ configPool conf) 30 diff --git a/src/PostgREST/Middleware.hs b/src/PostgREST/Middleware.hs index 568cfff9a..b14bf5005 100644 --- a/src/PostgREST/Middleware.hs +++ b/src/PostgREST/Middleware.hs @@ -3,7 +3,7 @@ module PostgREST.Middleware where -import Data.Maybe (fromMaybe) +import Data.Maybe (fromMaybe, isNothing) import Data.Monoid import Data.Text -- import Data.Pool(withResource, Pool) @@ -12,15 +12,19 @@ import qualified Hasql as H import qualified Hasql.Postgres as P 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.Status (status400, status401, status301) +import Network.HTTP.Types.Status (status400, status401, status301, status415) import Network.Wai (Application, requestHeaders, responseLBS, rawPathInfo, 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 PostgREST.Config (AppConfig(..)) +import PostgREST.Config (AppConfig(..), corsPolicy) import PostgREST.Auth (LoginAttempt(..), signInRole, signInWithJWT, setRole, resetRole, setUserId, resetUserId) +import PostgREST.App (contentTypeForAccept) import Codec.Binary.Base64.String (decode) import Prelude @@ -84,3 +88,17 @@ redirectInsecure app req respond = do Nothing -> respond $ responseLBS status400 [] "SSL is required" 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 diff --git a/test/Feature/CorsSpec.hs b/test/Feature/CorsSpec.hs index 2e108e8f3..fa005dc22 100644 --- a/test/Feature/CorsSpec.hs +++ b/test/Feature/CorsSpec.hs @@ -22,7 +22,7 @@ spec = around withApp $ describe "CORS" $ do ("Host", "localhost:3000"), ("User-Agent", "Mozilla/5.0 (Macintosh; Intel Mac OS X 10.9; rv:32.0) Gecko/20100101 Firefox/32.0"), ("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-Encoding", "gzip, deflate"), ("Referer", "http://localhost:8000/"), diff --git a/test/Feature/QuerySpec.hs b/test/Feature/QuerySpec.hs index ccebcf20d..e3f83c6e3 100644 --- a/test/Feature/QuerySpec.hs +++ b/test/Feature/QuerySpec.hs @@ -130,13 +130,23 @@ spec = it "without other constraints" $ 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" $ request methodGet "/simple_pk" - (acceptHdrs "text/csv") "" + (acceptHdrs "text/csv; version=1") "" `shouldRespondWith` ResponseMatcher { - matchBody = Just "k,extra\rxyyx,u\rxYYx,v" - , matchStatus = 200 + matchBody = Just "k,extra\rxyyx,u\rxYYx,v" + , matchStatus = 200 , matchHeaders = ["Content-Type" <:> "text/csv"] } diff --git a/test/SpecHelper.hs b/test/SpecHelper.hs index 345a9b561..a9c5df075 100644 --- a/test/SpecHelper.hs +++ b/test/SpecHelper.hs @@ -21,13 +21,12 @@ import Data.CaseInsensitive (CI(..)) import Data.Maybe (fromMaybe) import Text.Regex.TDFA ((=~)) import qualified Data.ByteString.Char8 as BS -import Network.Wai.Middleware.Cors (cors) import System.Process (readProcess) import qualified Data.Aeson.Types as J import PostgREST.App (app) -import PostgREST.Config (AppConfig(..), corsPolicy) +import PostgREST.Config (AppConfig(..)) import PostgREST.Middleware import PostgREST.Error(errResponse) @@ -59,7 +58,7 @@ withApp perform = do $ authenticated cfg (app cfg body) req either (resp . errResponse) resp result - where middle = cors corsPolicy + where middle = defaultMiddle False resetDb :: IO ()