Merge pull request #503 from begriffs/one-tx-per-client

Reduces pool resource locking (2)
This commit is contained in:
Joe Nelson
2016-02-26 10:11:47 -08:00
20 changed files with 201 additions and 171 deletions
+2
View File
@@ -7,6 +7,8 @@ This project adheres to [Semantic Versioning](http://semver.org/).
### Fixed ### Fixed
- Prevent query error from infecting later connection - @begriffs, @ruslantalpa, @nikita-volkov, @jwiegley
## [0.3.0.4] - 2016-02-12 ## [0.3.0.4] - 2016-02-12
### Fixed ### Fixed
+21 -4
View File
@@ -23,7 +23,7 @@ Flag CI
executable postgrest executable postgrest
main-is: PostgREST/Main.hs main-is: PostgREST/Main.hs
default-extensions: OverloadedStrings, ScopedTypeVariables, QuasiQuotes, LambdaCase default-extensions: OverloadedStrings, ScopedTypeVariables, QuasiQuotes
ghc-options: -threaded -rtsopts -with-rtsopts=-N ghc-options: -threaded -rtsopts -with-rtsopts=-N
default-language: Haskell2010 default-language: Haskell2010
build-depends: aeson >= 0.8 && < 0.10 build-depends: aeson >= 0.8 && < 0.10
@@ -34,15 +34,17 @@ executable postgrest
, containers , containers
, contravariant , contravariant
, errors , errors
, hasql >= 0.19.3.3 && < 0.20 , hasql >= 0.19.9 && < 0.20
, hasql-pool >= 0.4 && < 0.5
, hasql-transaction >= 0.4.3 && < 0.5
, http-types , http-types
, interpolatedstring-perl6 , interpolatedstring-perl6
, jwt , jwt
, mtl
, optparse-applicative >= 0.11 && < 0.13 , optparse-applicative >= 0.11 && < 0.13
, parsec , parsec
, postgrest , postgrest
, regex-tdfa , regex-tdfa
, resource-pool
, safe >= 0.3 && < 0.4 , safe >= 0.3 && < 0.4
, scientific , scientific
, string-conversions , string-conversions
@@ -86,9 +88,12 @@ library
, contravariant , contravariant
, errors , errors
, hasql , hasql
, hasql-transaction
, hasql-pool
, http-types , http-types
, interpolatedstring-perl6 , interpolatedstring-perl6
, jwt , jwt
, mtl
, optparse-applicative , optparse-applicative
, parsec , parsec
, regex-tdfa , regex-tdfa
@@ -123,10 +128,12 @@ library
Test-Suite spec Test-Suite spec
Type: exitcode-stdio-1.0 Type: exitcode-stdio-1.0
Default-Language: Haskell2010 Default-Language: Haskell2010
default-extensions: OverloadedStrings, ScopedTypeVariables, QuasiQuotes, LambdaCase default-extensions: OverloadedStrings, ScopedTypeVariables, QuasiQuotes
ghc-options: -threaded -rtsopts -with-rtsopts=-N
Hs-Source-Dirs: test, src Hs-Source-Dirs: test, src
Main-Is: Main.hs Main-Is: Main.hs
Other-Modules: Feature.AuthSpec Other-Modules: Feature.AuthSpec
, Feature.ConcurrentSpec
, Feature.CorsSpec , Feature.CorsSpec
, Feature.DeleteSpec , Feature.DeleteSpec
, Feature.InsertSpec , Feature.InsertSpec
@@ -138,6 +145,7 @@ Test-Suite spec
, PostgREST.Auth , PostgREST.Auth
, PostgREST.Config , PostgREST.Config
, PostgREST.Error , PostgREST.Error
, PostgREST.Main
, PostgREST.Middleware , PostgREST.Middleware
, PostgREST.Parsers , PostgREST.Parsers
, PostgREST.DbStructure , PostgREST.DbStructure
@@ -148,6 +156,7 @@ Test-Suite spec
, SpecHelper , SpecHelper
, TestTypes , TestTypes
Build-Depends: aeson Build-Depends: aeson
, async
, base , base
, base64-string , base64-string
, bytestring , bytestring
@@ -157,6 +166,8 @@ Test-Suite spec
, contravariant , contravariant
, errors , errors
, hasql , hasql
, hasql-pool
, hasql-transaction
, heredoc , heredoc
, hspec == 2.2.* , hspec == 2.2.*
, hspec-wai , hspec-wai
@@ -164,6 +175,8 @@ Test-Suite spec
, http-types , http-types
, interpolatedstring-perl6 , interpolatedstring-perl6
, jwt , jwt
, monad-control
, mtl
, optparse-applicative , optparse-applicative
, parsec , parsec
, process , process
@@ -173,11 +186,15 @@ Test-Suite spec
, string-conversions , string-conversions
, text , text
, time , time
, transformers
, transformers-base
, unordered-containers , unordered-containers
, unix
, vector , vector
, wai , wai
, wai-cors , wai-cors
, wai-extra , wai-extra
, wai-middleware-static , wai-middleware-static
, warp
, HTTP , HTTP
, Ranged-sets , Ranged-sets
+2 -2
View File
@@ -31,7 +31,7 @@ import Data.Aeson
import Data.Aeson.Types (emptyArray) import Data.Aeson.Types (emptyArray)
import Data.Monoid import Data.Monoid
import qualified Data.Vector as V import qualified Data.Vector as V
import qualified Hasql.Session as H import qualified Hasql.Transaction as H
import PostgREST.Config (AppConfig (..)) import PostgREST.Config (AppConfig (..))
import PostgREST.Parsers import PostgREST.Parsers
@@ -58,7 +58,7 @@ import PostgREST.QueryBuilder ( callProc
import Prelude import Prelude
app :: DbStructure -> AppConfig -> RequestBody -> Request -> H.Session Response app :: DbStructure -> AppConfig -> RequestBody -> Request -> H.Transaction Response
app dbStructure conf reqBody req = app dbStructure conf reqBody req =
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)
+2
View File
@@ -43,6 +43,7 @@ data AppConfig = AppConfig {
, configJwtSecret :: Secret , configJwtSecret :: Secret
, configPool :: Int , configPool :: Int
, configMaxRows :: Maybe Integer , configMaxRows :: Maybe Integer
, configQuiet :: Bool
} }
argParser :: Parser AppConfig argParser :: Parser AppConfig
@@ -55,6 +56,7 @@ argParser = AppConfig
strOption (long "jwt-secret" <> short 'j' <> help "secret used to encrypt and decrypt JWT tokens" <> metavar "SECRET" <> value "secret" <> showDefault)) strOption (long "jwt-secret" <> short 'j' <> help "secret used to encrypt and decrypt JWT tokens" <> metavar "SECRET" <> value "secret" <> showDefault))
<*> option auto (long "pool" <> short 'o' <> help "max connections in database pool" <> metavar "COUNT" <> value 10 <> showDefault) <*> option auto (long "pool" <> short 'o' <> help "max connections in database pool" <> metavar "COUNT" <> value 10 <> showDefault)
<*> (readMay <$> strOption (long "max-rows" <> short 'm' <> help "max rows in response" <> metavar "COUNT" <> value "infinity" <> showDefault)) <*> (readMay <$> strOption (long "max-rows" <> short 'm' <> help "max rows in response" <> metavar "COUNT" <> value "infinity" <> showDefault))
<*> pure False
defaultCorsPolicy :: CorsResourcePolicy defaultCorsPolicy :: CorsResourcePolicy
defaultCorsPolicy = CorsResourcePolicy Nothing defaultCorsPolicy = CorsResourcePolicy Nothing
+21 -10
View File
@@ -7,11 +7,13 @@ module PostgREST.Error (pgErrResponse, errResponse) where
import Data.Aeson ((.=)) import Data.Aeson ((.=))
import qualified Data.Aeson as JSON import qualified Data.Aeson as JSON
import Data.Maybe (fromMaybe)
import Data.Monoid ((<>)) import Data.Monoid ((<>))
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.Text (Text) import Data.Text (Text)
import qualified Data.Text as T import qualified Data.Text as T
import qualified Hasql.Session as H import qualified Hasql.Session as H
import qualified Hasql.Pool as P
import Network.HTTP.Types.Header import Network.HTTP.Types.Header
import qualified Network.HTTP.Types.Status as HT import qualified Network.HTTP.Types.Status as HT
import Network.Wai (Response, responseLBS) import Network.Wai (Response, responseLBS)
@@ -19,10 +21,17 @@ import Network.Wai (Response, responseLBS)
errResponse :: HT.Status -> Text -> Response errResponse :: HT.Status -> Text -> Response
errResponse status message = responseLBS status [(hContentType, "application/json")] (cs $ T.concat ["{\"message\":\"",message,"\"}"]) errResponse status message = responseLBS status [(hContentType, "application/json")] (cs $ T.concat ["{\"message\":\"",message,"\"}"])
pgErrResponse :: H.Error -> Response pgErrResponse :: P.UsageError -> Response
pgErrResponse e = responseLBS (httpStatus e) pgErrResponse e = responseLBS (httpStatus e)
[(hContentType, "application/json")] (JSON.encode e) [(hContentType, "application/json")] (JSON.encode e)
instance JSON.ToJSON P.UsageError where
toJSON (P.ConnectionError e) = JSON.object [
"code" .= ("" :: T.Text),
"message" .= ("Connection error" :: T.Text),
"details" .= (cs (fromMaybe "" e) :: T.Text)]
toJSON (P.SessionError e) = JSON.toJSON e -- H.Error
instance JSON.ToJSON H.Error where instance JSON.ToJSON H.Error where
toJSON (H.ResultError (H.ServerError c m d h)) = JSON.object [ toJSON (H.ResultError (H.ServerError c m d h)) = JSON.object [
"code" .= (cs c::T.Text), "code" .= (cs c::T.Text),
@@ -51,15 +60,17 @@ instance JSON.ToJSON H.Error where
"message" .= ("Database client error"::String), "message" .= ("Database client error"::String),
"details" .= (fmap cs d::Maybe T.Text)] "details" .= (fmap cs d::Maybe T.Text)]
httpStatus :: H.Error -> HT.Status httpStatus :: P.UsageError -> HT.Status
httpStatus (H.ResultError (H.ServerError c _ _ _)) = httpStatus (P.ConnectionError _) =
HT.status500
httpStatus (P.SessionError (H.ResultError (H.ServerError c _ _ _))) =
case cs c of case cs c of
'0':'8':_ -> HT.status503 -- pg connection err '0':'8':_ -> HT.status503 -- pg connection err
'0':'9':_ -> HT.status500 -- triggered action exception '0':'9':_ -> HT.status500 -- triggered action exception
'0':'L':_ -> HT.status403 -- invalid grantor '0':'L':_ -> HT.status403 -- invalid grantor
'0':'P':_ -> HT.status403 -- invalid role specification '0':'P':_ -> HT.status403 -- invalid role specification
"23503" -> HT.status409 -- foreign_key_violation "23503" -> HT.status409 -- foreign_key_violation
"23505" -> HT.status409 -- unique_violation "23505" -> HT.status409 -- unique_violation
'2':'5':_ -> HT.status500 -- invalid tx state '2':'5':_ -> HT.status500 -- invalid tx state
'2':'8':_ -> HT.status403 -- invalid auth specification '2':'8':_ -> HT.status403 -- invalid auth specification
'2':'D':_ -> HT.status500 -- invalid tx termination '2':'D':_ -> HT.status500 -- invalid tx termination
@@ -76,8 +87,8 @@ httpStatus (H.ResultError (H.ServerError c _ _ _)) =
'H':'V':_ -> HT.status500 -- foreign data wrapper error 'H':'V':_ -> HT.status500 -- foreign data wrapper error
'P':'0':_ -> HT.status500 -- PL/pgSQL Error 'P':'0':_ -> HT.status500 -- PL/pgSQL Error
'X':'X':_ -> HT.status500 -- internal Error 'X':'X':_ -> HT.status500 -- internal Error
"42P01" -> HT.status404 -- undefined table "42P01" -> HT.status404 -- undefined table
"42501" -> HT.status404 -- insufficient privilege "42501" -> HT.status404 -- insufficient privilege
_ -> HT.status400 _ -> HT.status400
httpStatus (H.ResultError _) = HT.status500 httpStatus (P.SessionError (H.ResultError _)) = HT.status500
httpStatus (H.ClientError _) = HT.status503 httpStatus (P.SessionError (H.ClientError _)) = HT.status503
+31 -36
View File
@@ -1,6 +1,6 @@
{-# LANGUAGE CPP #-} {-# LANGUAGE CPP #-}
module Main where module PostgREST.Main where
import PostgREST.App import PostgREST.App
@@ -9,21 +9,20 @@ import PostgREST.Config (AppConfig (..),
prettyVersion, prettyVersion,
readOptions) readOptions)
import PostgREST.DbStructure import PostgREST.DbStructure
import PostgREST.Error (errResponse, pgErrResponse) import PostgREST.Error (pgErrResponse)
import PostgREST.Middleware import PostgREST.Middleware
import PostgREST.QueryBuilder (inTransaction, Isolation(..)) import PostgREST.Types (DbStructure)
import Control.Monad import Control.Monad
import Data.Monoid ((<>)) import Data.Monoid ((<>))
import Data.Pool
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.Time.Clock.POSIX (getPOSIXTime) import Data.Time.Clock.POSIX (getPOSIXTime)
import qualified Hasql.Query as H import qualified Hasql.Query as H
import qualified Hasql.Connection as H
import qualified Hasql.Session as H import qualified Hasql.Session as H
import qualified Hasql.Transaction as HT
import qualified Hasql.Decoders as HD import qualified Hasql.Decoders as HD
import qualified Hasql.Encoders as HE import qualified Hasql.Encoders as HE
import qualified Network.HTTP.Types.Status as HT import qualified Hasql.Pool as P
import Network.Wai import Network.Wai
import Network.Wai.Handler.Warp import Network.Wai.Handler.Warp
import Network.Wai.Middleware.RequestLogger (logStdout) import Network.Wai.Middleware.RequestLogger (logStdout)
@@ -55,50 +54,46 @@ main = do
conf <- readOptions conf <- readOptions
let port = configPort conf let port = configPort conf
pgSettings = cs (configDatabase conf)
appSettings = setPort port
. setServerName (cs $ "postgrest/" <> prettyVersion)
$ defaultSettings
unless (secret "secret" /= configJwtSecret conf) $ unless (secret "secret" /= configJwtSecret conf) $
putStrLn "WARNING, running in insecure mode, JWT secret is the default value" putStrLn "WARNING, running in insecure mode, JWT secret is the default value"
Prelude.putStrLn $ "Listening on port " ++ Prelude.putStrLn $ "Listening on port " ++
(show $ configPort conf :: String) (show $ configPort conf :: String)
let pgSettings = cs (configDatabase conf)
appSettings = setPort port
. setServerName (cs $ "postgrest/" <> prettyVersion)
$ defaultSettings
middle = logStdout . defaultMiddle
pool <- createPool (H.acquire pgSettings) pool <- P.acquire (configPool conf, 10, pgSettings)
(either (const $ return ()) H.release) 1 1 (configPool conf)
dbStructure <- withResource pool $ \case
Left err -> error $ show err
Right c -> do
supported <- H.run isServerVersionSupported c
case supported of
Left e -> error $ show e
Right good -> unless good $
error (
"Cannot run in this PostgreSQL version, PostgREST needs at least "
<> show minimumPgVersion)
dbOrError <- H.run (getDbStructure (cs $ configSchema conf)) c
either (error . show) return dbOrError
#ifndef mingw32_HOST_OS #ifndef mingw32_HOST_OS
tid <- myThreadId tid <- myThreadId
void $ installHandler keyboardSignal (Catch $ do void $ installHandler keyboardSignal (Catch $ do
destroyAllResources pool P.release pool
throwTo tid UserInterrupt throwTo tid UserInterrupt
) Nothing ) Nothing
#endif #endif
runSettings appSettings $ middle $ \ req respond -> do result <- P.use pool $ do
supported <- isServerVersionSupported
unless supported $ error (
"Cannot run in this PostgreSQL version, PostgREST needs at least "
<> show minimumPgVersion)
getDbStructure (cs $ configSchema conf)
let dbStructure = either (error.show) id result
runSettings appSettings $ postgrest conf dbStructure pool
postgrest :: AppConfig -> DbStructure -> P.Pool -> Application
postgrest conf dbStructure pool =
let middle = (if configQuiet conf then id else logStdout) . defaultMiddle in
middle $ \ req respond -> do
time <- getPOSIXTime time <- getPOSIXTime
body <- strictRequestBody req body <- strictRequestBody req
let handleReq = H.run $ inTransaction ReadCommitted
(runWithClaims conf time (app dbStructure conf body) req) let handleReq = runWithClaims conf time (app dbStructure conf body) req
withResource pool $ \case resp <- either pgErrResponse id <$> P.use pool
Left err -> respond $ errResponse HT.status500 (cs . show $ err) (HT.run handleReq HT.ReadCommitted HT.Write)
Right c -> do respond resp
resOrError <- handleReq c
either (respond . pgErrResponse) respond resOrError
+3 -3
View File
@@ -7,7 +7,7 @@ import Data.Maybe (fromMaybe)
import Data.Text import Data.Text
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.Time.Clock (NominalDiffTime) import Data.Time.Clock (NominalDiffTime)
import qualified Hasql.Session as H import qualified Hasql.Transaction as H
import Network.HTTP.Types.Header (hAccept, hAuthorization) import Network.HTTP.Types.Header (hAccept, hAuthorization)
import Network.HTTP.Types.Status (status415, status400) import Network.HTTP.Types.Status (status415, status400)
@@ -27,8 +27,8 @@ import Prelude hiding(concat)
import qualified Data.Map.Lazy as M import qualified Data.Map.Lazy as M
runWithClaims :: AppConfig -> NominalDiffTime -> runWithClaims :: AppConfig -> NominalDiffTime ->
(Request -> H.Session Response) -> (Request -> H.Transaction Response) ->
Request -> H.Session Response Request -> H.Transaction Response
runWithClaims conf time app req = do runWithClaims conf time app req = do
H.sql setAnon H.sql setAnon
case split (== ' ') (cs auth) of case split (== ' ') (cs auth) of
-20
View File
@@ -18,7 +18,6 @@ module PostgREST.QueryBuilder (
, callProc , callProc
, createReadStatement , createReadStatement
, createWriteStatement , createWriteStatement
, inTransaction
, operators , operators
, pgFmtIdent , pgFmtIdent
, pgFmtLit , pgFmtLit
@@ -27,11 +26,9 @@ module PostgREST.QueryBuilder (
, sourceCTEName , sourceCTEName
, unquoted , unquoted
, ResultsWithCount , ResultsWithCount
, Isolation(..)
) where ) where
import qualified Hasql.Query as H import qualified Hasql.Query as H
import qualified Hasql.Session as H
import qualified Hasql.Encoders as HE import qualified Hasql.Encoders as HE
import qualified Hasql.Decoders as HD import qualified Hasql.Decoders as HD
@@ -504,20 +501,3 @@ pgFmtAsJsonPath (Just xx) = " AS " <> last xx
trimNullChars :: Text -> Text trimNullChars :: Text -> Text
trimNullChars = T.takeWhile (/= '\x0') trimNullChars = T.takeWhile (/= '\x0')
data Isolation = ReadCommitted | RepeatableRead | Serializable
{- |
Wrap a session in a transaction of desired isolation level
-}
inTransaction :: Isolation -> H.Session a -> H.Session a
inTransaction lvl f = do
H.sql $ "begin " <> isolate <> ";"
r <- f
H.sql "commit;"
return r
where
isolate = case lvl of
ReadCommitted -> "ISOLATION LEVEL READ COMMITTED"
RepeatableRead -> "ISOLATION LEVEL REPEATABLE READ"
Serializable -> "ISOLATION LEVEL SERIALIZABLE"
+7 -2
View File
@@ -1,10 +1,15 @@
resolver: lts-5.0 resolver: lts-5.0
extra-deps: extra-deps:
- hasql-0.19.3.3
- Ranged-sets-0.3.0 - Ranged-sets-0.3.0
- bytestring-tree-builder-0.2.5
- hasql-0.19.9
- hasql-pool-0.4
- hasql-transaction-0.4.3
- packdeps-0.4.2.1 - packdeps-0.4.2.1
- postgresql-error-codes-1
- postgresql-binary-0.8.1
ghc-options: ghc-options:
postgrest: -O2 -Werror -Wall -fwarn-monomorphism-restriction -fwarn-missing-exported-sigs -fwarn-identities postgrest: -O1 -Werror -Wall -fwarn-monomorphism-restriction -fwarn-missing-exported-sigs -fwarn-identities
packages: packages:
- '.' - '.'
+3 -5
View File
@@ -5,15 +5,13 @@ import Test.Hspec
import Test.Hspec.Wai import Test.Hspec.Wai
import Test.Hspec.Wai.JSON import Test.Hspec.Wai.JSON
import Network.HTTP.Types import Network.HTTP.Types
import qualified Hasql.Connection as H
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import Network.Wai (Application)
-- }}} -- }}}
spec :: DbStructure -> H.Connection -> Spec spec :: SpecWith Application
spec struct c = around (withApp cfgDefault struct c) spec = describe "authorization" $ do
$ describe "authorization" $ do
it "hides tables that anonymous does not own" $ it "hides tables that anonymous does not own" $
get "/authors_only" `shouldRespondWith` 404 get "/authors_only" `shouldRespondWith` 404
+51
View File
@@ -0,0 +1,51 @@
{-# LANGUAGE MultiParamTypeClasses, TypeFamilies, UndecidableInstances #-}
{-# OPTIONS_GHC -fno-warn-orphans #-}
module Feature.ConcurrentSpec where
import Control.Monad (void)
import Control.Monad.Base
import Control.Monad.Trans.Control
import Control.Concurrent.Async (mapConcurrently)
import Test.Hspec hiding (pendingWith)
import Test.Hspec.Wai.Internal
import Test.Hspec.Wai
import Test.Hspec.Wai.JSON
import Network.Wai.Test (Session)
import Network.Wai (Application)
spec :: SpecWith Application
spec =
describe "Queryiny in parallel" $
it "should not raise 'transaction in progress' error" $
raceTest 10 $
get "/fakefake"
`shouldRespondWith` ResponseMatcher {
matchBody = Just [json|
{ "hint": null,
"details":null,
"code":"42P01",
"message":"relation \"test.fakefake\" does not exist"
} |]
, matchStatus = 404
, matchHeaders = []
}
raceTest :: Int -> WaiExpectation -> WaiExpectation
raceTest times = liftBaseDiscard go
where
go test = void $ mapConcurrently (const test) [1..times]
instance MonadBaseControl IO WaiSession where
type StM WaiSession a = StM Session a
liftBaseWith f = WaiSession $
liftBaseWith $ \runInBase ->
f $ \k -> runInBase (unWaiSession k)
restoreM = WaiSession . restoreM
{-# INLINE liftBaseWith #-}
{-# INLINE restoreM #-}
instance MonadBase IO WaiSession where
liftBase = liftIO
+4 -4
View File
@@ -5,16 +5,16 @@ import Test.Hspec
import Test.Hspec.Wai import Test.Hspec.Wai
import Network.Wai.Test (SResponse(simpleHeaders, simpleBody)) import Network.Wai.Test (SResponse(simpleHeaders, simpleBody))
import qualified Data.ByteString.Lazy as BL import qualified Data.ByteString.Lazy as BL
import qualified Hasql.Connection as H
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..))
import Network.HTTP.Types import Network.HTTP.Types
import Network.Wai (Application)
-- }}} -- }}}
spec :: DbStructure -> H.Connection -> Spec spec :: SpecWith Application
spec struct c = around (withApp cfgDefault struct c) $ describe "CORS" $ do spec =
describe "CORS" $ do
let preflightHeaders = [ let preflightHeaders = [
("Accept", "*/*"), ("Accept", "*/*"),
("Origin", "http://example.com"), ("Origin", "http://example.com"),
+3 -7
View File
@@ -4,15 +4,11 @@ import Test.Hspec
import Test.Hspec.Wai import Test.Hspec.Wai
import Text.Heredoc import Text.Heredoc
import SpecHelper
import PostgREST.Types (DbStructure(..))
import qualified Hasql.Connection as H
import Network.HTTP.Types import Network.HTTP.Types
import Network.Wai (Application)
spec :: DbStructure -> H.Connection -> Spec spec :: SpecWith Application
spec struct c = beforeAll resetDb spec =
. around (withApp cfgDefault struct c) $
describe "Deleting" $ do describe "Deleting" $ do
context "existing record" $ do context "existing record" $ do
it "succeeds with 204 and deletion count" $ it "succeeds with 204 and deletion count" $
+3 -4
View File
@@ -6,7 +6,6 @@ import Test.Hspec.Wai.JSON
import Network.Wai.Test (SResponse(simpleBody,simpleHeaders,simpleStatus)) import Network.Wai.Test (SResponse(simpleBody,simpleHeaders,simpleStatus))
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..))
import qualified Data.Aeson as JSON import qualified Data.Aeson as JSON
import Data.Maybe (fromJust) import Data.Maybe (fromJust)
@@ -14,12 +13,12 @@ import Text.Heredoc
import Network.HTTP.Types.Header import Network.HTTP.Types.Header
import Network.HTTP.Types import Network.HTTP.Types
import Control.Monad (replicateM_) import Control.Monad (replicateM_)
import qualified Hasql.Connection as H
import TestTypes(IncPK(..), CompoundPK(..)) import TestTypes(IncPK(..), CompoundPK(..))
import Network.Wai (Application)
spec :: DbStructure -> H.Connection -> Spec spec :: SpecWith Application
spec struct c = beforeAll_ resetDb $ around (withApp cfgDefault struct c) $ do spec = do
describe "Posting new record" $ do describe "Posting new record" $ do
context "disparate json types" $ do context "disparate json types" $ do
it "accepts disparate json types" $ do it "accepts disparate json types" $ do
+3 -6
View File
@@ -5,15 +5,12 @@ import Test.Hspec.Wai
import Test.Hspec.Wai.JSON import Test.Hspec.Wai.JSON
import Network.HTTP.Types import Network.HTTP.Types
import Network.Wai.Test (SResponse(simpleHeaders, simpleStatus)) import Network.Wai.Test (SResponse(simpleHeaders, simpleStatus))
import qualified Hasql.Connection as H
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import Network.Wai (Application)
spec :: DbStructure -> H.Connection -> Spec spec :: SpecWith Application
spec struct c = spec =
beforeAll resetDb
. around (withApp (cfgLimitRows 3) struct c) $
describe "Requesting many items with server limits enabled" $ do describe "Requesting many items with server limits enabled" $ do
it "restricts results" $ it "restricts results" $
get "/items" get "/items"
+3 -4
View File
@@ -5,14 +5,13 @@ import Test.Hspec.Wai
import Test.Hspec.Wai.JSON import Test.Hspec.Wai.JSON
import Network.HTTP.Types import Network.HTTP.Types
import Network.Wai.Test (SResponse(simpleHeaders)) import Network.Wai.Test (SResponse(simpleHeaders))
import qualified Hasql.Connection as H
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..))
import Text.Heredoc import Text.Heredoc
import Network.Wai (Application)
spec :: DbStructure -> H.Connection -> Spec spec :: SpecWith Application
spec struct c = around (withApp cfgDefault struct c) $ do spec = do
describe "Querying a table with a column called count" $ describe "Querying a table with a column called count" $
it "should not confuse count column with pg_catalog.count aggregate" $ it "should not confuse count column with pg_catalog.count aggregate" $
+4 -5
View File
@@ -5,14 +5,13 @@ import Test.Hspec.Wai
import Test.Hspec.Wai.JSON import Test.Hspec.Wai.JSON
import Network.HTTP.Types import Network.HTTP.Types
import Network.Wai.Test (SResponse(simpleHeaders,simpleStatus)) import Network.Wai.Test (SResponse(simpleHeaders,simpleStatus))
import qualified Hasql.Connection as H
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..)) import Network.Wai (Application)
spec :: SpecWith Application
spec =
spec :: DbStructure -> H.Connection -> Spec
spec struct c = beforeAll resetDb
. around (withApp cfgDefault struct c) $
describe "GET /items" $ do describe "GET /items" $ do
context "without range headers" $ do context "without range headers" $ do
+4 -4
View File
@@ -3,15 +3,15 @@ module Feature.StructureSpec where
import Test.Hspec hiding (pendingWith) import Test.Hspec hiding (pendingWith)
import Test.Hspec.Wai import Test.Hspec.Wai
import Test.Hspec.Wai.JSON import Test.Hspec.Wai.JSON
import qualified Hasql.Connection as H
import SpecHelper import SpecHelper
import PostgREST.Types (DbStructure(..))
import Network.HTTP.Types import Network.HTTP.Types
import Network.Wai (Application)
spec :: SpecWith Application
spec = do
spec :: DbStructure -> H.Connection -> Spec
spec struct c = around (withApp cfgDefault struct c) $ do
describe "GET /" $ do describe "GET /" $ do
it "lists views in schema" $ it "lists views in schema" $
request methodGet "/" [] "" request methodGet "/" [] ""
+26 -19
View File
@@ -3,13 +3,14 @@ module Main where
import Test.Hspec import Test.Hspec
import SpecHelper import SpecHelper
import qualified Hasql.Session as H import qualified Hasql.Pool as P
import qualified Hasql.Connection as H
import PostgREST.DbStructure (getDbStructure) import PostgREST.DbStructure (getDbStructure)
import PostgREST.Main (postgrest)
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import qualified Feature.AuthSpec import qualified Feature.AuthSpec
import qualified Feature.ConcurrentSpec
import qualified Feature.CorsSpec import qualified Feature.CorsSpec
import qualified Feature.DeleteSpec import qualified Feature.DeleteSpec
import qualified Feature.InsertSpec import qualified Feature.InsertSpec
@@ -22,22 +23,28 @@ main :: IO ()
main = do main = do
setupDb setupDb
H.acquire (cs dbString) >>= \case pool <- P.acquire (3, 10, cs testDbConn)
Left err -> error $ show err
Right c -> do result <- P.use pool $ getDbStructure "test"
dbOrErr <- H.run (getDbStructure "test") c let dbStructure = either (error.show) id result
-- Not using hspec-discover because we want to precompute withApp = return $ postgrest testCfg dbStructure pool
-- the db structure and pass it to specs for speed ltdApp = return $ postgrest testLtdRowsCfg dbStructure pool
either (error.show) (hspec . specs c) dbOrErr
H.release c hspec $ do
mapM_ (beforeAll_ resetDb . before withApp) specs
-- this test runs with a different server flag
beforeAll_ resetDb . before ltdApp $
describe "Feature.QueryLimitedSpec" Feature.QueryLimitedSpec.spec
where where
specs conn dbStructure = do specs = map (uncurry describe) [
describe "Feature.AuthSpec" $ Feature.AuthSpec.spec dbStructure conn ("Feature.AuthSpec" , Feature.AuthSpec.spec)
describe "Feature.CorsSpec" $ Feature.CorsSpec.spec dbStructure conn , ("Feature.ConcurrentSpec" , Feature.ConcurrentSpec.spec)
describe "Feature.DeleteSpec" $ Feature.DeleteSpec.spec dbStructure conn , ("Feature.CorsSpec" , Feature.CorsSpec.spec)
describe "Feature.InsertSpec" $ Feature.InsertSpec.spec dbStructure conn , ("Feature.DeleteSpec" , Feature.DeleteSpec.spec)
describe "Feature.QueryLimitedSpec" $ Feature.QueryLimitedSpec.spec dbStructure conn , ("Feature.InsertSpec" , Feature.InsertSpec.spec)
describe "Feature.QuerySpec" $ Feature.QuerySpec.spec dbStructure conn , ("Feature.QuerySpec" , Feature.QuerySpec.spec)
describe "Feature.RangeSpec" $ Feature.RangeSpec.spec dbStructure conn , ("Feature.RangeSpec" , Feature.RangeSpec.spec)
describe "Feature.StructureSpec" $ Feature.StructureSpec.spec dbStructure conn , ("Feature.StructureSpec" , Feature.StructureSpec.spec)
]
+8 -36
View File
@@ -1,10 +1,6 @@
module SpecHelper where module SpecHelper where
import Network.Wai
import Test.Hspec
import Data.String.Conversions (cs) import Data.String.Conversions (cs)
import Data.Time.Clock.POSIX (getPOSIXTime)
import Control.Monad (void) import Control.Monad (void)
import Network.HTTP.Types.Header (Header, ByteRange, renderByteRange, import Network.HTTP.Types.Header (Header, ByteRange, renderByteRange,
@@ -16,42 +12,18 @@ import qualified Data.ByteString.Char8 as BS
import System.Process (readProcess) import System.Process (readProcess)
import Web.JWT (secret) import Web.JWT (secret)
import qualified Hasql.Connection as H
import qualified Hasql.Session as H
import PostgREST.App (app)
import PostgREST.Config (AppConfig(..)) import PostgREST.Config (AppConfig(..))
import PostgREST.Middleware
import PostgREST.Error(pgErrResponse)
import PostgREST.Types
import PostgREST.QueryBuilder (inTransaction, Isolation(..))
dbString :: String testDbConn :: String
dbString = "postgres://postgrest_test_authenticator@localhost:5432/postgrest_test" testDbConn = "postgres://postgrest_test_authenticator@localhost:5432/postgrest_test"
cfg :: String -> Maybe Integer -> AppConfig testCfg :: AppConfig
cfg conStr = AppConfig conStr "postgrest_test_anonymous" "test" 3000 (secret "safe") 10 testCfg =
AppConfig testDbConn "postgrest_test_anonymous" "test" 3000 (secret "safe") 10 Nothing True
cfgDefault :: AppConfig testLtdRowsCfg :: AppConfig
cfgDefault = cfg dbString Nothing testLtdRowsCfg =
AppConfig testDbConn "postgrest_test_anonymous" "test" 3000 (secret "safe") 10 (Just 3) True
cfgLimitRows :: Integer -> AppConfig
cfgLimitRows = cfg dbString . Just
withApp :: AppConfig -> DbStructure -> H.Connection
-> ActionWith Application -> IO ()
withApp config dbStructure c perform =
perform $ defaultMiddle $ \req resp -> do
time <- getPOSIXTime
body <- strictRequestBody req
let handleReq = H.run $ inTransaction ReadCommitted
(runWithClaims config time (app dbStructure config body) req)
handleReq c >>= \case
Left err -> do
void $ H.run (H.sql "rollback;") c
resp $ pgErrResponse err
Right res -> resp res
setupDb :: IO () setupDb :: IO ()
setupDb = do setupDb = do