From 76fdf29f5c60c7796ac2a5b1df26b0479e4ef479 Mon Sep 17 00:00:00 2001 From: Wolfgang Walther Date: Wed, 5 Jan 2022 17:11:28 +0100 Subject: [PATCH] refactor: Move each Wai Middleware to a separate file --- postgrest.cabal | 2 + src/PostgREST/App.hs | 7 +++- src/PostgREST/Cors.hs | 42 ++++++++++++++++++++ src/PostgREST/Logger.hs | 28 ++++++++++++++ src/PostgREST/Middleware.hs | 77 ++++++------------------------------- 5 files changed, 89 insertions(+), 67 deletions(-) create mode 100644 src/PostgREST/Cors.hs create mode 100644 src/PostgREST/Logger.hs diff --git a/postgrest.cabal b/postgrest.cabal index 88518a19c..31a958702 100644 --- a/postgrest.cabal +++ b/postgrest.cabal @@ -44,6 +44,7 @@ library PostgREST.Config.PgVersion PostgREST.Config.Proxy PostgREST.ContentType + PostgREST.Cors PostgREST.DbStructure PostgREST.DbStructure.Identifiers PostgREST.DbStructure.Proc @@ -51,6 +52,7 @@ library PostgREST.DbStructure.Table PostgREST.Error PostgREST.GucHeader + PostgREST.Logger PostgREST.Middleware PostgREST.OpenAPI PostgREST.Query.QueryBuilder diff --git a/src/PostgREST/App.hs b/src/PostgREST/App.hs index d70ac55ef..a2d84fe22 100644 --- a/src/PostgREST/App.hs +++ b/src/PostgREST/App.hs @@ -42,8 +42,10 @@ import qualified Network.Wai.Handler.Warp as Warp import qualified PostgREST.AppState as AppState import qualified PostgREST.Auth as Auth +import qualified PostgREST.Cors as Cors import qualified PostgREST.DbStructure as DbStructure import qualified PostgREST.Error as Error +import qualified PostgREST.Logger as Logger import qualified PostgREST.Middleware as Middleware import qualified PostgREST.OpenAPI as OpenAPI import qualified PostgREST.Query.QueryBuilder as QueryBuilder @@ -137,8 +139,9 @@ serverSettings AppConfig{..} = -- | PostgREST application postgrest :: LogLevel -> AppState.AppState -> IO () -> Wai.Application -postgrest logLev appState connWorker = - Middleware.pgrstMiddleware logLev $ +postgrest logLevel appState connWorker = + Logger.middleware logLevel . + Cors.middleware $ \req respond -> do time <- AppState.getTime appState conf <- AppState.getConfig appState diff --git a/src/PostgREST/Cors.hs b/src/PostgREST/Cors.hs new file mode 100644 index 000000000..df38f1d90 --- /dev/null +++ b/src/PostgREST/Cors.hs @@ -0,0 +1,42 @@ +{-| +Module : PostgREST.Cors +Description : Wai Middleware to set cors policy. +-} +module PostgREST.Cors (middleware) where + +import qualified Data.ByteString.Char8 as BS +import qualified Data.CaseInsensitive as CI +import qualified Network.Wai as Wai +import qualified Network.Wai.Middleware.Cors as Wai + +import Data.List (lookup) + +import Protolude + +middleware :: Wai.Middleware +middleware = Wai.cors corsPolicy + +-- | CORS policy to be used in by Wai Cors middleware +corsPolicy :: Wai.Request -> Maybe Wai.CorsResourcePolicy +corsPolicy req = case lookup "origin" headers of + Just origin -> + Just Wai.CorsResourcePolicy + { Wai.corsOrigins = Just ([origin], True) + , Wai.corsMethods = ["GET", "POST", "PATCH", "PUT", "DELETE", "OPTIONS"] + , Wai.corsRequestHeaders = "Authorization" : accHeaders + , Wai.corsExposedHeaders = Just + [ "Content-Encoding", "Content-Location", "Content-Range", "Content-Type" + , "Date", "Location", "Server", "Transfer-Encoding", "Range-Unit"] + , Wai.corsMaxAge = Just $ 60*60*24 + , Wai.corsVaryOrigin = False + , Wai.corsRequireOrigin = False + , Wai.corsIgnoreFailures = True + } + Nothing -> Nothing + where + headers = Wai.requestHeaders req + accHeaders = case lookup "access-control-request-headers" headers of + Just hdrs -> map (CI.mk . BS.strip) $ BS.split ',' hdrs + -- Impossible case, Middleware.Cors will not evaluate this when + -- the Access-Control-Request-Headers header is not set. + Nothing -> [] diff --git a/src/PostgREST/Logger.hs b/src/PostgREST/Logger.hs new file mode 100644 index 000000000..ba2645d3a --- /dev/null +++ b/src/PostgREST/Logger.hs @@ -0,0 +1,28 @@ +{-| +Module : PostgREST.Logger +Description : Wai Middleware to log requests to stdout. +-} +module PostgREST.Logger (middleware) where + +import qualified Network.Wai as Wai +import qualified Network.Wai.Middleware.RequestLogger as Wai + +import Network.HTTP.Types.Status (status400, status500) +import System.IO.Unsafe (unsafePerformIO) + +import PostgREST.Config (LogLevel (..)) + +import Protolude + +middleware :: LogLevel -> Wai.Middleware +middleware logLevel = case logLevel of + LogInfo -> requestLogger (const True) + LogWarn -> requestLogger (>= status400) + LogError -> requestLogger (>= status500) + LogCrit -> requestLogger (const False) + where + requestLogger filterStatus = unsafePerformIO $ Wai.mkRequestLogger Wai.defaultRequestLoggerSettings + { Wai.outputFormat = Wai.ApacheWithSettings $ + Wai.defaultApacheSettings + & Wai.setApacheRequestFilter (\_ res -> filterStatus $ Wai.responseStatus res) + } diff --git a/src/PostgREST/Middleware.hs b/src/PostgREST/Middleware.hs index aeafe5365..ef48dc381 100644 --- a/src/PostgREST/Middleware.hs +++ b/src/PostgREST/Middleware.hs @@ -6,35 +6,25 @@ Description : Sets CORS policy. Also the PostgreSQL GUCs, role, search_path and {-# LANGUAGE RecordWildCards #-} module PostgREST.Middleware ( runPgLocals - , pgrstMiddleware , optionalRollback ) where -import qualified Data.Aeson as JSON -import qualified Data.ByteString.Char8 as BS -import qualified Data.ByteString.Lazy.Char8 as LBS -import qualified Data.CaseInsensitive as CI -import qualified Data.HashMap.Strict as M -import qualified Data.Text as T -import qualified Data.Text.Encoding as T -import qualified Hasql.Decoders as HD -import qualified Hasql.DynamicStatements.Snippet as SQL hiding - (sql) -import qualified Hasql.DynamicStatements.Statement as SQL -import qualified Hasql.Transaction as SQL -import qualified Network.Wai as Wai -import qualified Network.Wai.Middleware.Cors as Wai -import qualified Network.Wai.Middleware.RequestLogger as Wai +import qualified Data.Aeson as JSON +import qualified Data.ByteString.Lazy.Char8 as LBS +import qualified Data.HashMap.Strict as M +import qualified Data.Text as T +import qualified Data.Text.Encoding as T +import qualified Hasql.Decoders as HD +import qualified Hasql.DynamicStatements.Snippet as SQL hiding (sql) +import qualified Hasql.DynamicStatements.Statement as SQL +import qualified Hasql.Transaction as SQL +import qualified Network.Wai as Wai import Control.Arrow ((***)) -import Data.List (lookup) -import Data.Scientific (FPFormat (..), formatScientific, - isInteger) -import Network.HTTP.Types.Status (status400, status500) -import System.IO.Unsafe (unsafePerformIO) +import Data.Scientific (FPFormat (..), formatScientific, isInteger) -import PostgREST.Config (AppConfig (..), LogLevel (..)) +import PostgREST.Config (AppConfig (..)) import PostgREST.Config.PgVersion (PgVersion (..), pgVersion140) import PostgREST.Error (Error, errorResponseFor) import PostgREST.GucHeader (addHeadersIfNotIncluded) @@ -89,49 +79,6 @@ runPgLocals conf claims app req jsonDbS actualPgVersion = do unquoted (JSON.Bool b) = show b unquoted v = T.decodeUtf8 . LBS.toStrict $ JSON.encode v -pgrstMiddleware :: LogLevel -> Wai.Middleware -pgrstMiddleware logLevel = - logger logLevel - . Wai.cors corsPolicy - -logger :: LogLevel -> Wai.Middleware -logger logLevel = case logLevel of - LogInfo -> requestLogger (const True) - LogWarn -> requestLogger (>= status400) - LogError -> requestLogger (>= status500) - LogCrit -> requestLogger (const False) - where - requestLogger filterStatus = unsafePerformIO $ Wai.mkRequestLogger Wai.defaultRequestLoggerSettings - { Wai.outputFormat = Wai.ApacheWithSettings $ - Wai.defaultApacheSettings - & Wai.setApacheRequestFilter (\_ res -> filterStatus $ Wai.responseStatus res) - } - --- | CORS policy to be used in by Wai Cors middleware -corsPolicy :: Wai.Request -> Maybe Wai.CorsResourcePolicy -corsPolicy req = case lookup "origin" headers of - Just origin -> - Just Wai.CorsResourcePolicy - { Wai.corsOrigins = Just ([origin], True) - , Wai.corsMethods = ["GET", "POST", "PATCH", "PUT", "DELETE", "OPTIONS"] - , Wai.corsRequestHeaders = "Authorization" : accHeaders - , Wai.corsExposedHeaders = Just - [ "Content-Encoding", "Content-Location", "Content-Range", "Content-Type" - , "Date", "Location", "Server", "Transfer-Encoding", "Range-Unit"] - , Wai.corsMaxAge = Just $ 60*60*24 - , Wai.corsVaryOrigin = False - , Wai.corsRequireOrigin = False - , Wai.corsIgnoreFailures = True - } - Nothing -> Nothing - where - headers = Wai.requestHeaders req - accHeaders = case lookup "access-control-request-headers" headers of - Just hdrs -> map (CI.mk . BS.strip) $ BS.split ',' hdrs - -- Impossible case, Middleware.Cors will not evaluate this when - -- the Access-Control-Request-Headers header is not set. - Nothing -> [] - -- | Set a transaction to eventually roll back if requested and set respective -- headers on the response. optionalRollback