diff --git a/CHANGELOG.md b/CHANGELOG.md index 1fb97db65..9f901f0fa 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -7,6 +7,8 @@ This project adheres to [Semantic Versioning](http://semver.org/). ### Fixed +- Prevent query error from infecting later connection - @begriffs, @ruslantalpa, @nikita-volkov, @jwiegley + ## [0.3.0.4] - 2016-02-12 ### Fixed diff --git a/postgrest.cabal b/postgrest.cabal index 30b477626..9fdb9cda3 100644 --- a/postgrest.cabal +++ b/postgrest.cabal @@ -23,7 +23,7 @@ Flag CI executable postgrest main-is: PostgREST/Main.hs - default-extensions: OverloadedStrings, ScopedTypeVariables, QuasiQuotes, LambdaCase + default-extensions: OverloadedStrings, ScopedTypeVariables, QuasiQuotes ghc-options: -threaded -rtsopts -with-rtsopts=-N default-language: Haskell2010 build-depends: aeson >= 0.8 && < 0.10 @@ -34,15 +34,17 @@ executable postgrest , containers , contravariant , 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 , interpolatedstring-perl6 , jwt + , mtl , optparse-applicative >= 0.11 && < 0.13 , parsec , postgrest , regex-tdfa - , resource-pool , safe >= 0.3 && < 0.4 , scientific , string-conversions @@ -86,9 +88,12 @@ library , contravariant , errors , hasql + , hasql-transaction + , hasql-pool , http-types , interpolatedstring-perl6 , jwt + , mtl , optparse-applicative , parsec , regex-tdfa @@ -123,10 +128,12 @@ library Test-Suite spec Type: exitcode-stdio-1.0 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 Main-Is: Main.hs Other-Modules: Feature.AuthSpec + , Feature.ConcurrentSpec , Feature.CorsSpec , Feature.DeleteSpec , Feature.InsertSpec @@ -138,6 +145,7 @@ Test-Suite spec , PostgREST.Auth , PostgREST.Config , PostgREST.Error + , PostgREST.Main , PostgREST.Middleware , PostgREST.Parsers , PostgREST.DbStructure @@ -148,6 +156,7 @@ Test-Suite spec , SpecHelper , TestTypes Build-Depends: aeson + , async , base , base64-string , bytestring @@ -157,6 +166,8 @@ Test-Suite spec , contravariant , errors , hasql + , hasql-pool + , hasql-transaction , heredoc , hspec == 2.2.* , hspec-wai @@ -164,6 +175,8 @@ Test-Suite spec , http-types , interpolatedstring-perl6 , jwt + , monad-control + , mtl , optparse-applicative , parsec , process @@ -173,11 +186,15 @@ Test-Suite spec , string-conversions , text , time + , transformers + , transformers-base , unordered-containers + , unix , vector , wai , wai-cors , wai-extra , wai-middleware-static + , warp , HTTP , Ranged-sets diff --git a/src/PostgREST/App.hs b/src/PostgREST/App.hs index 418a84f55..000d66544 100644 --- a/src/PostgREST/App.hs +++ b/src/PostgREST/App.hs @@ -31,7 +31,7 @@ import Data.Aeson import Data.Aeson.Types (emptyArray) import Data.Monoid import qualified Data.Vector as V -import qualified Hasql.Session as H +import qualified Hasql.Transaction as H import PostgREST.Config (AppConfig (..)) import PostgREST.Parsers @@ -58,7 +58,7 @@ import PostgREST.QueryBuilder ( callProc import Prelude -app :: DbStructure -> AppConfig -> RequestBody -> Request -> H.Session Response +app :: DbStructure -> AppConfig -> RequestBody -> Request -> H.Transaction Response app dbStructure conf reqBody req = let -- TODO: blow up for Left values (there is a middleware that checks the headers) diff --git a/src/PostgREST/Config.hs b/src/PostgREST/Config.hs index 5aaccf9a7..5cddcec12 100644 --- a/src/PostgREST/Config.hs +++ b/src/PostgREST/Config.hs @@ -43,6 +43,7 @@ data AppConfig = AppConfig { , configJwtSecret :: Secret , configPool :: Int , configMaxRows :: Maybe Integer + , configQuiet :: Bool } 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)) <*> 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)) + <*> pure False defaultCorsPolicy :: CorsResourcePolicy defaultCorsPolicy = CorsResourcePolicy Nothing diff --git a/src/PostgREST/Error.hs b/src/PostgREST/Error.hs index e6629f5dd..e7893aa0b 100644 --- a/src/PostgREST/Error.hs +++ b/src/PostgREST/Error.hs @@ -7,11 +7,13 @@ module PostgREST.Error (pgErrResponse, errResponse) where import Data.Aeson ((.=)) import qualified Data.Aeson as JSON +import Data.Maybe (fromMaybe) import Data.Monoid ((<>)) import Data.String.Conversions (cs) import Data.Text (Text) import qualified Data.Text as T import qualified Hasql.Session as H +import qualified Hasql.Pool as P import Network.HTTP.Types.Header import qualified Network.HTTP.Types.Status as HT import Network.Wai (Response, responseLBS) @@ -19,10 +21,17 @@ import Network.Wai (Response, responseLBS) errResponse :: HT.Status -> Text -> Response 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) [(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 toJSON (H.ResultError (H.ServerError c m d h)) = JSON.object [ "code" .= (cs c::T.Text), @@ -51,15 +60,17 @@ instance JSON.ToJSON H.Error where "message" .= ("Database client error"::String), "details" .= (fmap cs d::Maybe T.Text)] -httpStatus :: H.Error -> HT.Status -httpStatus (H.ResultError (H.ServerError c _ _ _)) = +httpStatus :: P.UsageError -> HT.Status +httpStatus (P.ConnectionError _) = + HT.status500 +httpStatus (P.SessionError (H.ResultError (H.ServerError c _ _ _))) = case cs c of '0':'8':_ -> HT.status503 -- pg connection err '0':'9':_ -> HT.status500 -- triggered action exception '0':'L':_ -> HT.status403 -- invalid grantor '0':'P':_ -> HT.status403 -- invalid role specification - "23503" -> HT.status409 -- foreign_key_violation - "23505" -> HT.status409 -- unique_violation + "23503" -> HT.status409 -- foreign_key_violation + "23505" -> HT.status409 -- unique_violation '2':'5':_ -> HT.status500 -- invalid tx state '2':'8':_ -> HT.status403 -- invalid auth specification '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 'P':'0':_ -> HT.status500 -- PL/pgSQL Error 'X':'X':_ -> HT.status500 -- internal Error - "42P01" -> HT.status404 -- undefined table - "42501" -> HT.status404 -- insufficient privilege - _ -> HT.status400 -httpStatus (H.ResultError _) = HT.status500 -httpStatus (H.ClientError _) = HT.status503 + "42P01" -> HT.status404 -- undefined table + "42501" -> HT.status404 -- insufficient privilege + _ -> HT.status400 +httpStatus (P.SessionError (H.ResultError _)) = HT.status500 +httpStatus (P.SessionError (H.ClientError _)) = HT.status503 diff --git a/src/PostgREST/Main.hs b/src/PostgREST/Main.hs index 3701e36a7..ae840249d 100644 --- a/src/PostgREST/Main.hs +++ b/src/PostgREST/Main.hs @@ -1,6 +1,6 @@ {-# LANGUAGE CPP #-} -module Main where +module PostgREST.Main where import PostgREST.App @@ -9,21 +9,20 @@ import PostgREST.Config (AppConfig (..), prettyVersion, readOptions) import PostgREST.DbStructure -import PostgREST.Error (errResponse, pgErrResponse) +import PostgREST.Error (pgErrResponse) import PostgREST.Middleware -import PostgREST.QueryBuilder (inTransaction, Isolation(..)) +import PostgREST.Types (DbStructure) import Control.Monad import Data.Monoid ((<>)) -import Data.Pool import Data.String.Conversions (cs) import Data.Time.Clock.POSIX (getPOSIXTime) import qualified Hasql.Query as H -import qualified Hasql.Connection as H import qualified Hasql.Session as H +import qualified Hasql.Transaction as HT import qualified Hasql.Decoders as HD 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.Handler.Warp import Network.Wai.Middleware.RequestLogger (logStdout) @@ -55,50 +54,46 @@ main = do conf <- readOptions let port = configPort conf + pgSettings = cs (configDatabase conf) + appSettings = setPort port + . setServerName (cs $ "postgrest/" <> prettyVersion) + $ defaultSettings unless (secret "secret" /= configJwtSecret conf) $ putStrLn "WARNING, running in insecure mode, JWT secret is the default value" Prelude.putStrLn $ "Listening on port " ++ (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) - (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 + pool <- P.acquire (configPool conf, 10, pgSettings) #ifndef mingw32_HOST_OS tid <- myThreadId void $ installHandler keyboardSignal (Catch $ do - destroyAllResources pool + P.release pool throwTo tid UserInterrupt ) Nothing #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 body <- strictRequestBody req - let handleReq = H.run $ inTransaction ReadCommitted - (runWithClaims conf time (app dbStructure conf body) req) - withResource pool $ \case - Left err -> respond $ errResponse HT.status500 (cs . show $ err) - Right c -> do - resOrError <- handleReq c - either (respond . pgErrResponse) respond resOrError + + let handleReq = runWithClaims conf time (app dbStructure conf body) req + resp <- either pgErrResponse id <$> P.use pool + (HT.run handleReq HT.ReadCommitted HT.Write) + respond resp diff --git a/src/PostgREST/Middleware.hs b/src/PostgREST/Middleware.hs index 4292e261c..e68ddbc6f 100644 --- a/src/PostgREST/Middleware.hs +++ b/src/PostgREST/Middleware.hs @@ -7,7 +7,7 @@ import Data.Maybe (fromMaybe) import Data.Text import Data.String.Conversions (cs) 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.Status (status415, status400) @@ -27,8 +27,8 @@ import Prelude hiding(concat) import qualified Data.Map.Lazy as M runWithClaims :: AppConfig -> NominalDiffTime -> - (Request -> H.Session Response) -> - Request -> H.Session Response + (Request -> H.Transaction Response) -> + Request -> H.Transaction Response runWithClaims conf time app req = do H.sql setAnon case split (== ' ') (cs auth) of diff --git a/src/PostgREST/QueryBuilder.hs b/src/PostgREST/QueryBuilder.hs index 0da811869..046040909 100644 --- a/src/PostgREST/QueryBuilder.hs +++ b/src/PostgREST/QueryBuilder.hs @@ -18,7 +18,6 @@ module PostgREST.QueryBuilder ( , callProc , createReadStatement , createWriteStatement - , inTransaction , operators , pgFmtIdent , pgFmtLit @@ -27,11 +26,9 @@ module PostgREST.QueryBuilder ( , sourceCTEName , unquoted , ResultsWithCount - , Isolation(..) ) where import qualified Hasql.Query as H -import qualified Hasql.Session as H import qualified Hasql.Encoders as HE import qualified Hasql.Decoders as HD @@ -504,20 +501,3 @@ pgFmtAsJsonPath (Just xx) = " AS " <> last xx trimNullChars :: Text -> Text 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" diff --git a/stack.yaml b/stack.yaml index 1774ffffa..74f1079af 100644 --- a/stack.yaml +++ b/stack.yaml @@ -1,10 +1,15 @@ resolver: lts-5.0 extra-deps: - - hasql-0.19.3.3 - 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 + - postgresql-error-codes-1 + - postgresql-binary-0.8.1 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: - '.' diff --git a/test/Feature/AuthSpec.hs b/test/Feature/AuthSpec.hs index f678ae269..3ae39724f 100644 --- a/test/Feature/AuthSpec.hs +++ b/test/Feature/AuthSpec.hs @@ -5,15 +5,13 @@ import Test.Hspec import Test.Hspec.Wai import Test.Hspec.Wai.JSON import Network.HTTP.Types -import qualified Hasql.Connection as H import SpecHelper -import PostgREST.Types (DbStructure(..)) +import Network.Wai (Application) -- }}} -spec :: DbStructure -> H.Connection -> Spec -spec struct c = around (withApp cfgDefault struct c) - $ describe "authorization" $ do +spec :: SpecWith Application +spec = describe "authorization" $ do it "hides tables that anonymous does not own" $ get "/authors_only" `shouldRespondWith` 404 diff --git a/test/Feature/ConcurrentSpec.hs b/test/Feature/ConcurrentSpec.hs new file mode 100644 index 000000000..be154375d --- /dev/null +++ b/test/Feature/ConcurrentSpec.hs @@ -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 diff --git a/test/Feature/CorsSpec.hs b/test/Feature/CorsSpec.hs index 811af712a..c639f7309 100644 --- a/test/Feature/CorsSpec.hs +++ b/test/Feature/CorsSpec.hs @@ -5,16 +5,16 @@ import Test.Hspec import Test.Hspec.Wai import Network.Wai.Test (SResponse(simpleHeaders, simpleBody)) import qualified Data.ByteString.Lazy as BL -import qualified Hasql.Connection as H import SpecHelper -import PostgREST.Types (DbStructure(..)) import Network.HTTP.Types +import Network.Wai (Application) -- }}} -spec :: DbStructure -> H.Connection -> Spec -spec struct c = around (withApp cfgDefault struct c) $ describe "CORS" $ do +spec :: SpecWith Application +spec = + describe "CORS" $ do let preflightHeaders = [ ("Accept", "*/*"), ("Origin", "http://example.com"), diff --git a/test/Feature/DeleteSpec.hs b/test/Feature/DeleteSpec.hs index ba9c60eea..76d1fcc64 100644 --- a/test/Feature/DeleteSpec.hs +++ b/test/Feature/DeleteSpec.hs @@ -4,15 +4,11 @@ import Test.Hspec import Test.Hspec.Wai import Text.Heredoc -import SpecHelper -import PostgREST.Types (DbStructure(..)) -import qualified Hasql.Connection as H - import Network.HTTP.Types +import Network.Wai (Application) -spec :: DbStructure -> H.Connection -> Spec -spec struct c = beforeAll resetDb - . around (withApp cfgDefault struct c) $ +spec :: SpecWith Application +spec = describe "Deleting" $ do context "existing record" $ do it "succeeds with 204 and deletion count" $ diff --git a/test/Feature/InsertSpec.hs b/test/Feature/InsertSpec.hs index e3030fe5a..7b3d41bde 100644 --- a/test/Feature/InsertSpec.hs +++ b/test/Feature/InsertSpec.hs @@ -6,7 +6,6 @@ import Test.Hspec.Wai.JSON import Network.Wai.Test (SResponse(simpleBody,simpleHeaders,simpleStatus)) import SpecHelper -import PostgREST.Types (DbStructure(..)) import qualified Data.Aeson as JSON import Data.Maybe (fromJust) @@ -14,12 +13,12 @@ import Text.Heredoc import Network.HTTP.Types.Header import Network.HTTP.Types import Control.Monad (replicateM_) -import qualified Hasql.Connection as H import TestTypes(IncPK(..), CompoundPK(..)) +import Network.Wai (Application) -spec :: DbStructure -> H.Connection -> Spec -spec struct c = beforeAll_ resetDb $ around (withApp cfgDefault struct c) $ do +spec :: SpecWith Application +spec = do describe "Posting new record" $ do context "disparate json types" $ do it "accepts disparate json types" $ do diff --git a/test/Feature/QueryLimitedSpec.hs b/test/Feature/QueryLimitedSpec.hs index 71e6aa714..95d8b9223 100644 --- a/test/Feature/QueryLimitedSpec.hs +++ b/test/Feature/QueryLimitedSpec.hs @@ -5,15 +5,12 @@ import Test.Hspec.Wai import Test.Hspec.Wai.JSON import Network.HTTP.Types import Network.Wai.Test (SResponse(simpleHeaders, simpleStatus)) -import qualified Hasql.Connection as H import SpecHelper -import PostgREST.Types (DbStructure(..)) +import Network.Wai (Application) -spec :: DbStructure -> H.Connection -> Spec -spec struct c = - beforeAll resetDb - . around (withApp (cfgLimitRows 3) struct c) $ +spec :: SpecWith Application +spec = describe "Requesting many items with server limits enabled" $ do it "restricts results" $ get "/items" diff --git a/test/Feature/QuerySpec.hs b/test/Feature/QuerySpec.hs index 3c3062db2..6ae4194e7 100644 --- a/test/Feature/QuerySpec.hs +++ b/test/Feature/QuerySpec.hs @@ -5,14 +5,13 @@ import Test.Hspec.Wai import Test.Hspec.Wai.JSON import Network.HTTP.Types import Network.Wai.Test (SResponse(simpleHeaders)) -import qualified Hasql.Connection as H import SpecHelper -import PostgREST.Types (DbStructure(..)) import Text.Heredoc +import Network.Wai (Application) -spec :: DbStructure -> H.Connection -> Spec -spec struct c = around (withApp cfgDefault struct c) $ do +spec :: SpecWith Application +spec = do describe "Querying a table with a column called count" $ it "should not confuse count column with pg_catalog.count aggregate" $ diff --git a/test/Feature/RangeSpec.hs b/test/Feature/RangeSpec.hs index 0ab998712..1b20c8d5a 100644 --- a/test/Feature/RangeSpec.hs +++ b/test/Feature/RangeSpec.hs @@ -5,14 +5,13 @@ import Test.Hspec.Wai import Test.Hspec.Wai.JSON import Network.HTTP.Types import Network.Wai.Test (SResponse(simpleHeaders,simpleStatus)) -import qualified Hasql.Connection as H 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 context "without range headers" $ do diff --git a/test/Feature/StructureSpec.hs b/test/Feature/StructureSpec.hs index a441f00c9..3ab2ac82b 100644 --- a/test/Feature/StructureSpec.hs +++ b/test/Feature/StructureSpec.hs @@ -3,15 +3,15 @@ module Feature.StructureSpec where import Test.Hspec hiding (pendingWith) import Test.Hspec.Wai import Test.Hspec.Wai.JSON -import qualified Hasql.Connection as H import SpecHelper -import PostgREST.Types (DbStructure(..)) 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 it "lists views in schema" $ request methodGet "/" [] "" diff --git a/test/Main.hs b/test/Main.hs index 6256f5d98..cc52c0fb4 100644 --- a/test/Main.hs +++ b/test/Main.hs @@ -3,13 +3,14 @@ module Main where import Test.Hspec import SpecHelper -import qualified Hasql.Session as H -import qualified Hasql.Connection as H +import qualified Hasql.Pool as P import PostgREST.DbStructure (getDbStructure) +import PostgREST.Main (postgrest) import Data.String.Conversions (cs) import qualified Feature.AuthSpec +import qualified Feature.ConcurrentSpec import qualified Feature.CorsSpec import qualified Feature.DeleteSpec import qualified Feature.InsertSpec @@ -22,22 +23,28 @@ main :: IO () main = do setupDb - H.acquire (cs dbString) >>= \case - Left err -> error $ show err - Right c -> do - dbOrErr <- H.run (getDbStructure "test") c - -- Not using hspec-discover because we want to precompute - -- the db structure and pass it to specs for speed - either (error.show) (hspec . specs c) dbOrErr - H.release c + pool <- P.acquire (3, 10, cs testDbConn) + + result <- P.use pool $ getDbStructure "test" + let dbStructure = either (error.show) id result + withApp = return $ postgrest testCfg dbStructure pool + ltdApp = return $ postgrest testLtdRowsCfg dbStructure pool + + 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 - specs conn dbStructure = do - describe "Feature.AuthSpec" $ Feature.AuthSpec.spec dbStructure conn - describe "Feature.CorsSpec" $ Feature.CorsSpec.spec dbStructure conn - describe "Feature.DeleteSpec" $ Feature.DeleteSpec.spec dbStructure conn - describe "Feature.InsertSpec" $ Feature.InsertSpec.spec dbStructure conn - describe "Feature.QueryLimitedSpec" $ Feature.QueryLimitedSpec.spec dbStructure conn - describe "Feature.QuerySpec" $ Feature.QuerySpec.spec dbStructure conn - describe "Feature.RangeSpec" $ Feature.RangeSpec.spec dbStructure conn - describe "Feature.StructureSpec" $ Feature.StructureSpec.spec dbStructure conn + specs = map (uncurry describe) [ + ("Feature.AuthSpec" , Feature.AuthSpec.spec) + , ("Feature.ConcurrentSpec" , Feature.ConcurrentSpec.spec) + , ("Feature.CorsSpec" , Feature.CorsSpec.spec) + , ("Feature.DeleteSpec" , Feature.DeleteSpec.spec) + , ("Feature.InsertSpec" , Feature.InsertSpec.spec) + , ("Feature.QuerySpec" , Feature.QuerySpec.spec) + , ("Feature.RangeSpec" , Feature.RangeSpec.spec) + , ("Feature.StructureSpec" , Feature.StructureSpec.spec) + ] diff --git a/test/SpecHelper.hs b/test/SpecHelper.hs index 22b80ee3c..87abae532 100644 --- a/test/SpecHelper.hs +++ b/test/SpecHelper.hs @@ -1,10 +1,6 @@ module SpecHelper where -import Network.Wai -import Test.Hspec - import Data.String.Conversions (cs) -import Data.Time.Clock.POSIX (getPOSIXTime) import Control.Monad (void) import Network.HTTP.Types.Header (Header, ByteRange, renderByteRange, @@ -16,42 +12,18 @@ import qualified Data.ByteString.Char8 as BS import System.Process (readProcess) 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.Middleware -import PostgREST.Error(pgErrResponse) -import PostgREST.Types -import PostgREST.QueryBuilder (inTransaction, Isolation(..)) -dbString :: String -dbString = "postgres://postgrest_test_authenticator@localhost:5432/postgrest_test" +testDbConn :: String +testDbConn = "postgres://postgrest_test_authenticator@localhost:5432/postgrest_test" -cfg :: String -> Maybe Integer -> AppConfig -cfg conStr = AppConfig conStr "postgrest_test_anonymous" "test" 3000 (secret "safe") 10 +testCfg :: AppConfig +testCfg = + AppConfig testDbConn "postgrest_test_anonymous" "test" 3000 (secret "safe") 10 Nothing True -cfgDefault :: AppConfig -cfgDefault = cfg dbString Nothing - -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 +testLtdRowsCfg :: AppConfig +testLtdRowsCfg = + AppConfig testDbConn "postgrest_test_anonymous" "test" 3000 (secret "safe") 10 (Just 3) True setupDb :: IO () setupDb = do