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