Use read-only transaction mode for read requests (#561)

* Make middleware use ApiRequest rather than Request

* Fix outdated comments

* Use read-only transaction mode for read requests

This allows API requests against read replicas
This commit is contained in:
Joe Nelson
2016-04-15 12:26:40 -07:00
parent eae5857d0e
commit 5aadfba84b
4 changed files with 65 additions and 52 deletions
+1
View File
@@ -8,6 +8,7 @@ This project adheres to [Semantic Versioning](http://semver.org/).
### Fixed ### Fixed
- Prevent role from being changed twice - @begriffs - Prevent role from being changed twice - @begriffs
- Use read-only transaction for read requests - @ruslantalpa
## [0.3.1.1] - 2016-03-28 ## [0.3.1.1] - 2016-03-28
+23 -4
View File
@@ -4,16 +4,21 @@ import qualified Data.Aeson as JSON
import qualified Data.ByteString as BS import qualified Data.ByteString as BS
import qualified Data.ByteString.Lazy as BL import qualified Data.ByteString.Lazy as BL
import qualified Data.Csv as CSV import qualified Data.Csv as CSV
import Data.List (find) import Data.List (find, sortBy)
import qualified Data.HashMap.Strict as M import qualified Data.HashMap.Strict as M
import qualified Data.Set as S import qualified Data.Set as S
import Data.Maybe (fromMaybe, isJust, isNothing, import Data.Maybe (fromMaybe, isJust, isNothing,
listToMaybe, fromJust) listToMaybe, fromJust)
import Control.Arrow ((***))
import Control.Monad (join) import Control.Monad (join)
import Data.Monoid ((<>)) import Data.Monoid ((<>))
import Data.Ord (comparing)
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import qualified Data.Text as T import qualified Data.Text as T
import qualified Data.Vector as V import qualified Data.Vector as V
import Network.HTTP.Base (urlEncodeVars)
import Network.HTTP.Types.Header (hAuthorization)
import Network.HTTP.Types.URI (parseSimpleQuery)
import Network.Wai (Request (..)) import Network.Wai (Request (..))
import Network.Wai.Parse (parseHttpAccept) import Network.Wai.Parse (parseHttpAccept)
import PostgREST.RangeQuery (NonnegRange, rangeRequested) import PostgREST.RangeQuery (NonnegRange, rangeRequested)
@@ -52,11 +57,11 @@ instance Show ContentType where
if it is an action we are able to perform. if it is an action we are able to perform.
-} -}
data ApiRequest = ApiRequest { data ApiRequest = ApiRequest {
-- | Set to Nothing for unknown HTTP verbs -- | Similar but not identical to HTTP verb, e.g. Create/Invoke both POST
iAction :: Action iAction :: Action
-- | Set to Nothing for malformed range -- | Requested range of rows within response
, iRange :: NonnegRange , iRange :: NonnegRange
-- | Set to Nothing for strangely nested urls -- | The target, be it calling a proc or accessing a table
, iTarget :: Target , iTarget :: Target
-- | The content type the client most desires (or JSON if undecided) -- | The content type the client most desires (or JSON if undecided)
, iAccepts :: Either BS.ByteString ContentType , iAccepts :: Either BS.ByteString ContentType
@@ -74,6 +79,10 @@ data ApiRequest = ApiRequest {
, iSelect :: String , iSelect :: String
-- | &order parameter -- | &order parameter
, iOrder :: Maybe String , iOrder :: Maybe String
-- | Alphabetized (canonical) request query string for response URLs
, iCanonicalQS :: String
-- | JSON Web Token
, iJWT :: T.Text
} }
-- | Examines HTTP request and translates it into user intent. -- | Examines HTTP request and translates it into user intent.
@@ -136,6 +145,12 @@ userApiRequest schema req reqBody =
then "*" then "*"
else fromMaybe "*" $ fromMaybe (Just "*") $ lookup "select" qParams else fromMaybe "*" $ fromMaybe (Just "*") $ lookup "select" qParams
, iOrder = join $ lookup "order" qParams , iOrder = join $ lookup "order" qParams
, iCanonicalQS = urlEncodeVars
. sortBy (comparing fst)
. map (join (***) cs)
. parseSimpleQuery
$ rawQueryString req
, iJWT = tokenStr
} }
where where
@@ -155,6 +170,10 @@ userApiRequest schema req reqBody =
| hasPrefer "return=representation" = Full | hasPrefer "return=representation" = Full
| hasPrefer "return=minimal" = None | hasPrefer "return=minimal" = None
| otherwise = HeadersOnly | otherwise = HeadersOnly
auth = fromMaybe "" $ lookupHeader hAuthorization
tokenStr = case T.split (== ' ') (cs auth) of
("Bearer" : t : _) -> t
_ -> ""
-- PRIVATE --------------------------------------------------------------- -- PRIVATE ---------------------------------------------------------------
+15 -16
View File
@@ -7,12 +7,9 @@ module PostgREST.App (
) where ) where
import Control.Applicative import Control.Applicative
import Control.Arrow ((***))
import Control.Monad (join)
import Data.Bifunctor (first) import Data.Bifunctor (first)
import Data.List (find, sortBy, delete) import Data.List (find, delete)
import Data.Maybe (isJust, fromMaybe, fromJust, mapMaybe) import Data.Maybe (isJust, fromMaybe, fromJust, mapMaybe)
import Data.Ord (comparing)
import Data.Ranged.Ranges (emptyRange) import Data.Ranged.Ranges (emptyRange)
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.Text (Text, replace, strip) import Data.Text (Text, replace, strip)
@@ -24,10 +21,8 @@ import qualified Hasql.Transaction as HT
import Text.Parsec.Error import Text.Parsec.Error
import Text.ParserCombinators.Parsec (parse) import Text.ParserCombinators.Parsec (parse)
import Network.HTTP.Base (urlEncodeVars)
import Network.HTTP.Types.Header import Network.HTTP.Types.Header
import Network.HTTP.Types.Status import Network.HTTP.Types.Status
import Network.HTTP.Types.URI (parseSimpleQuery)
import Network.Wai import Network.Wai
import Network.Wai.Middleware.RequestLogger (logStdout) import Network.Wai.Middleware.RequestLogger (logStdout)
@@ -72,13 +67,22 @@ postgrest conf dbStructure pool =
time <- getPOSIXTime time <- getPOSIXTime
body <- strictRequestBody req body <- strictRequestBody req
let handleReq = runWithClaims conf time (app dbStructure conf body) req let schema = cs $ configSchema conf
apiRequest = userApiRequest schema req body
handleReq = runWithClaims conf time (app dbStructure conf) apiRequest
txMode = transactionMode $ iAction apiRequest
resp <- either pgErrResponse id <$> P.use pool resp <- either pgErrResponse id <$> P.use pool
(HT.run handleReq HT.ReadCommitted HT.Write) (HT.run handleReq HT.ReadCommitted txMode)
respond resp respond resp
app :: DbStructure -> AppConfig -> RequestBody -> Request -> H.Transaction Response transactionMode :: Action -> H.Mode
app dbStructure conf reqBody req = transactionMode ActionRead = HT.Read
transactionMode ActionInfo = HT.Read
transactionMode _ = HT.Write
app :: DbStructure -> AppConfig -> ApiRequest -> H.Transaction Response
app dbStructure conf apiRequest =
let let
-- TODO: blow up for Left values (there is a middleware that checks the headers) -- TODO: blow up for Left values (there is a middleware that checks the headers)
contentType = either (const ApplicationJSON) id (iAccepts apiRequest) contentType = either (const ApplicationJSON) id (iAccepts apiRequest)
@@ -102,11 +106,7 @@ app dbStructure conf reqBody req =
else responseLBS status200 [contentTypeH] (cs body) else responseLBS status200 [contentTypeH] (cs body)
else do else do
let (status, contentRange) = rangeHeader queryTotal tableTotal let (status, contentRange) = rangeHeader queryTotal tableTotal
canonical = urlEncodeVars -- should this be moved to the dbStructure (location)? canonical = iCanonicalQS apiRequest
. sortBy (comparing fst)
. map (join (***) cs)
. parseSimpleQuery
$ rawQueryString req
return $ responseLBS status return $ responseLBS status
[contentTypeH, contentRange, [contentTypeH, contentRange,
("Content-Location", ("Content-Location",
@@ -210,7 +210,6 @@ app dbStructure conf reqBody req =
allPrKeys = dbPrimaryKeys dbStructure allPrKeys = dbPrimaryKeys dbStructure
allOrigins = ("Access-Control-Allow-Origin", "*") :: Header allOrigins = ("Access-Control-Allow-Origin", "*") :: Header
schema = cs $ configSchema conf schema = cs $ configSchema conf
apiRequest = userApiRequest schema req reqBody
shouldCount = iPreferCount apiRequest shouldCount = iPreferCount apiRequest
range = restrictRange (configMaxRows conf) $ iRange apiRequest range = restrictRange (configMaxRows conf) $ iRange apiRequest
readDbRequest = DbRead <$> buildReadRequest (dbRelations dbStructure) apiRequest readDbRequest = DbRead <$> buildReadRequest (dbRelations dbStructure) apiRequest
+6 -12
View File
@@ -5,13 +5,12 @@ module PostgREST.Middleware where
import Data.Aeson (Value (..)) import Data.Aeson (Value (..))
import qualified Data.HashMap.Strict as M import qualified Data.HashMap.Strict as M
import Data.Maybe (fromMaybe)
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.Text import Data.Text
import Data.Time.Clock (NominalDiffTime) import Data.Time.Clock (NominalDiffTime)
import qualified Hasql.Transaction as H import qualified Hasql.Transaction as H
import Network.HTTP.Types.Header (hAccept, hAuthorization) import Network.HTTP.Types.Header (hAccept)
import Network.HTTP.Types.Status (status400, status415) import Network.HTTP.Types.Status (status400, status415)
import Network.Wai (Application, Request (..), import Network.Wai (Application, Request (..),
Response, requestHeaders) Response, requestHeaders)
@@ -19,7 +18,7 @@ import Network.Wai.Middleware.Cors (cors)
import Network.Wai.Middleware.Gzip (def, gzip) import Network.Wai.Middleware.Gzip (def, gzip)
import Network.Wai.Middleware.Static (only, staticPolicy) import Network.Wai.Middleware.Static (only, staticPolicy)
import PostgREST.ApiRequest (pickContentType) import PostgREST.ApiRequest (ApiRequest(..), pickContentType)
import PostgREST.Auth (jwtClaims, claimsToSQL) import PostgREST.Auth (jwtClaims, claimsToSQL)
import PostgREST.Config (AppConfig (..), corsPolicy) import PostgREST.Config (AppConfig (..), corsPolicy)
import PostgREST.Error (errResponse) import PostgREST.Error (errResponse)
@@ -27,26 +26,21 @@ import PostgREST.Error (errResponse)
import Prelude hiding (concat, null) import Prelude hiding (concat, null)
runWithClaims :: AppConfig -> NominalDiffTime -> runWithClaims :: AppConfig -> NominalDiffTime ->
(Request -> H.Transaction Response) -> (ApiRequest -> H.Transaction Response) ->
Request -> H.Transaction Response ApiRequest -> H.Transaction Response
runWithClaims conf time app req = do runWithClaims conf time app req = do
let tokenStr = case split (== ' ') (cs auth) of let eClaims = jwtClaims jwtSecret (iJWT req) time
("Bearer" : t : _) -> t
_ -> ""
eClaims = jwtClaims jwtSecret tokenStr time
case eClaims of case eClaims of
Left e -> clientErr e Left e -> clientErr e
Right claims -> Right claims ->
if M.null claims && not (null tokenStr) if M.null claims && not (null $ iJWT req)
then clientErr "Invalid JWT" then clientErr "Invalid JWT"
else do else do
-- role claim defaults to anon if not specified in jwt -- role claim defaults to anon if not specified in jwt
H.sql . mconcat . claimsToSQL $ M.union claims (M.singleton "role" anon) H.sql . mconcat . claimsToSQL $ M.union claims (M.singleton "role" anon)
app req app req
where where
hdrs = requestHeaders req
jwtSecret = configJwtSecret conf jwtSecret = configJwtSecret conf
auth = fromMaybe "" $ lookup hAuthorization hdrs
anon = String . cs $ configAnonRole conf anon = String . cs $ configAnonRole conf
clientErr = return . errResponse status400 clientErr = return . errResponse status400